mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-09-30 12:01:59 +00:00
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:
co-authored by
Claude Sonnet 5.5
parent
bc8bff7955
commit
da1af1487e
@@ -576,11 +576,32 @@ Centralized management of multiple database connections with support for Postgre
|
|||||||
- Multiple named database connections
|
- Multiple named database connections
|
||||||
- Multi-ORM access (Bun, GORM, Native SQL) sharing the same connection pool
|
- Multi-ORM access (Bun, GORM, Native SQL) sharing the same connection pool
|
||||||
- Automatic SQLite schema translation (`schema.table` → `schema_table`)
|
- 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
|
- Prometheus metrics for monitoring
|
||||||
- Configuration-driven via YAML
|
- Configuration-driven via YAML
|
||||||
- Per-connection statistics and management
|
- 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).
|
For documentation, see [pkg/dbmanager/README.md](pkg/dbmanager/README.md).
|
||||||
|
|
||||||
#### Cache
|
#### Cache
|
||||||
|
|||||||
@@ -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).
|
||||||
@@ -50,7 +50,6 @@ dbmanager:
|
|||||||
|
|
||||||
# Health checks
|
# Health checks
|
||||||
health_check_interval: 30s
|
health_check_interval: 30s
|
||||||
enable_auto_reconnect: true
|
|
||||||
|
|
||||||
connections:
|
connections:
|
||||||
# Primary PostgreSQL connection
|
# Primary PostgreSQL connection
|
||||||
@@ -256,7 +255,7 @@ db, _ := mgr.GetDefaultDatabase()
|
|||||||
| `retry_delay` | duration | 1s | Initial retry delay |
|
| `retry_delay` | duration | 1s | Initial retry delay |
|
||||||
| `retry_max_delay` | duration | 10s | Maximum retry delay |
|
| `retry_max_delay` | duration | 10s | Maximum retry delay |
|
||||||
| `health_check_interval` | duration | 30s | Interval between health checks |
|
| `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
|
### Connection Configuration
|
||||||
|
|
||||||
@@ -451,7 +450,6 @@ db.NewSelect().Model(&User{}).Scan(ctx)
|
|||||||
3. **Enable Health Checks**: Catch connection issues early
|
3. **Enable Health Checks**: Catch connection issues early
|
||||||
```yaml
|
```yaml
|
||||||
health_check_interval: 30s
|
health_check_interval: 30s
|
||||||
enable_auto_reconnect: true
|
|
||||||
```
|
```
|
||||||
|
|
||||||
4. **Use Appropriate ORM**: Choose based on your needs
|
4. **Use Appropriate ORM**: Choose based on your needs
|
||||||
|
|||||||
+106
-71
@@ -2,6 +2,10 @@ package dbmanager
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/url"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/config"
|
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||||
@@ -57,9 +61,15 @@ type ManagerConfig struct {
|
|||||||
RetryDelay time.Duration `mapstructure:"retry_delay"`
|
RetryDelay time.Duration `mapstructure:"retry_delay"`
|
||||||
RetryMaxDelay time.Duration `mapstructure:"retry_max_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"`
|
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
|
// ConnectionConfig defines configuration for a single database connection
|
||||||
@@ -103,6 +113,11 @@ type ConnectionConfig struct {
|
|||||||
ConnectTimeout time.Duration `mapstructure:"connect_timeout"`
|
ConnectTimeout time.Duration `mapstructure:"connect_timeout"`
|
||||||
QueryTimeout time.Duration `mapstructure:"query_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
|
// Features
|
||||||
EnableTracing bool `mapstructure:"enable_tracing"`
|
EnableTracing bool `mapstructure:"enable_tracing"`
|
||||||
EnableMetrics bool `mapstructure:"enable_metrics"`
|
EnableMetrics bool `mapstructure:"enable_metrics"`
|
||||||
@@ -129,7 +144,6 @@ func DefaultManagerConfig() ManagerConfig {
|
|||||||
RetryDelay: 1 * time.Second,
|
RetryDelay: 1 * time.Second,
|
||||||
RetryMaxDelay: 10 * time.Second,
|
RetryMaxDelay: 10 * time.Second,
|
||||||
HealthCheckInterval: 15 * time.Second,
|
HealthCheckInterval: 15 * time.Second,
|
||||||
EnableAutoReconnect: true,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -161,11 +175,6 @@ func (c *ManagerConfig) ApplyDefaults() {
|
|||||||
if c.HealthCheckInterval == 0 {
|
if c.HealthCheckInterval == 0 {
|
||||||
c.HealthCheckInterval = defaults.HealthCheckInterval
|
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
|
// Validate validates the manager configuration
|
||||||
@@ -222,9 +231,18 @@ func (cc *ConnectionConfig) ApplyDefaults(global *ManagerConfig) {
|
|||||||
}
|
}
|
||||||
if cc.QueryTimeout == 0 {
|
if cc.QueryTimeout == 0 {
|
||||||
cc.QueryTimeout = 2 * time.Minute // Default to 2 minutes
|
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
|
// 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 {
|
func (cc *ConnectionConfig) buildPostgresDSN() string {
|
||||||
dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s",
|
q := url.Values{}
|
||||||
cc.Host, cc.Port, cc.User, cc.Password, cc.Database)
|
|
||||||
|
|
||||||
if cc.SSLMode != "" {
|
if cc.SSLMode != "" {
|
||||||
dsn += fmt.Sprintf(" sslmode=%s", cc.SSLMode)
|
q.Set("sslmode", cc.SSLMode)
|
||||||
} else {
|
} 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 != "" {
|
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)
|
u := url.URL{
|
||||||
if cc.QueryTimeout > 0 {
|
Scheme: "postgres",
|
||||||
timeoutMs := int(cc.QueryTimeout.Milliseconds())
|
Host: hostPort(cc.Host, cc.Port),
|
||||||
dsn += fmt.Sprintf(" statement_timeout=%d", timeoutMs)
|
Path: "/" + cc.Database,
|
||||||
|
RawQuery: q.Encode(),
|
||||||
}
|
}
|
||||||
|
if cc.User != "" || cc.Password != "" {
|
||||||
return dsn
|
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 {
|
func (cc *ConnectionConfig) buildSQLiteDSN() string {
|
||||||
filepath := cc.FilePath
|
filepath := cc.FilePath
|
||||||
if filepath == "" {
|
if filepath == "" {
|
||||||
filepath = ":memory:"
|
filepath = ":memory:"
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add query parameters for timeouts
|
var pragmas []string
|
||||||
// Note: SQLite driver supports _timeout parameter (in milliseconds)
|
|
||||||
if cc.QueryTimeout > 0 {
|
if cc.QueryTimeout > 0 {
|
||||||
timeoutMs := int(cc.QueryTimeout.Milliseconds())
|
pragmas = append(pragmas, fmt.Sprintf("busy_timeout(%d)", cc.QueryTimeout.Milliseconds()))
|
||||||
filepath += fmt.Sprintf("?_timeout=%d", timeoutMs)
|
}
|
||||||
|
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 {
|
func (cc *ConnectionConfig) buildMSSQLDSN() string {
|
||||||
// Format: sqlserver://username:password@host:port?database=dbname
|
// Format: sqlserver://username:password@host:port?database=dbname
|
||||||
dsn := fmt.Sprintf("sqlserver://%s:%s@%s:%d?database=%s",
|
q := url.Values{}
|
||||||
cc.User, cc.Password, cc.Host, cc.Port, cc.Database)
|
q.Set("database", cc.Database)
|
||||||
|
|
||||||
if cc.Schema != "" {
|
if cc.Schema != "" {
|
||||||
dsn += fmt.Sprintf("&schema=%s", cc.Schema)
|
q.Set("schema", cc.Schema)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add connection timeout (in seconds)
|
|
||||||
if cc.ConnectTimeout > 0 {
|
if cc.ConnectTimeout > 0 {
|
||||||
timeoutSec := int(cc.ConnectTimeout.Seconds())
|
sec := strconv.Itoa(int(cc.ConnectTimeout.Seconds()))
|
||||||
dsn += fmt.Sprintf("&connection timeout=%d", timeoutSec)
|
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 {
|
if cc.QueryTimeout > 0 {
|
||||||
readTimeoutSec := int(cc.QueryTimeout.Seconds())
|
q.Set("read timeout", strconv.Itoa(int(cc.QueryTimeout.Seconds())))
|
||||||
dsn += fmt.Sprintf("&read timeout=%d", readTimeoutSec)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
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 {
|
func (cc *ConnectionConfig) buildMongoDSN() string {
|
||||||
// Format: mongodb://username:password@host:port/database?authSource=admin
|
// Format: mongodb://username:password@host:port/database?authSource=admin
|
||||||
var dsn string
|
q := url.Values{}
|
||||||
|
|
||||||
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 := ""
|
|
||||||
if cc.AuthSource != "" {
|
if cc.AuthSource != "" {
|
||||||
params += fmt.Sprintf("authSource=%s", cc.AuthSource)
|
q.Set("authSource", cc.AuthSource)
|
||||||
}
|
}
|
||||||
if cc.ReplicaSet != "" {
|
if cc.ReplicaSet != "" {
|
||||||
if params != "" {
|
q.Set("replicaSet", cc.ReplicaSet)
|
||||||
params += "&"
|
|
||||||
}
|
|
||||||
params += fmt.Sprintf("replicaSet=%s", cc.ReplicaSet)
|
|
||||||
}
|
}
|
||||||
if cc.ReadPreference != "" {
|
if cc.ReadPreference != "" {
|
||||||
if params != "" {
|
q.Set("readPreference", cc.ReadPreference)
|
||||||
params += "&"
|
|
||||||
}
|
|
||||||
params += fmt.Sprintf("readPreference=%s", cc.ReadPreference)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if params != "" {
|
u := url.URL{
|
||||||
dsn += "?" + params
|
Scheme: "mongodb",
|
||||||
|
Host: hostPort(cc.Host, cc.Port),
|
||||||
|
Path: "/" + cc.Database,
|
||||||
|
RawQuery: q.Encode(),
|
||||||
}
|
}
|
||||||
|
if cc.User != "" && cc.Password != "" {
|
||||||
return dsn
|
u.User = url.UserPassword(cc.User, cc.Password)
|
||||||
|
}
|
||||||
|
return u.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
// FromConfig converts config.DBManagerConfig to internal ManagerConfig
|
// 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) GetQueryTimeout() time.Duration { return cc.QueryTimeout }
|
||||||
func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics }
|
func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics }
|
||||||
func (cc *ConnectionConfig) GetReadPreference() string { return cc.ReadPreference }
|
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 }
|
||||||
|
|||||||
@@ -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
@@ -3,6 +3,7 @@ package dbmanager
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -14,6 +15,7 @@ import (
|
|||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Connection represents a single named database connection
|
// Connection represents a single named database connection
|
||||||
@@ -82,6 +84,9 @@ type sqlConnection struct {
|
|||||||
// State
|
// State
|
||||||
connected bool
|
connected bool
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
|
// lifecycleMu serialises Connect/Close/Reconnect against health-check pings.
|
||||||
|
// Lock order: lifecycleMu before mu.
|
||||||
|
lifecycleMu sync.RWMutex
|
||||||
|
|
||||||
// Health check
|
// Health check
|
||||||
lastHealthCheck time.Time
|
lastHealthCheck time.Time
|
||||||
@@ -110,9 +115,16 @@ func (c *sqlConnection) Type() DatabaseType {
|
|||||||
|
|
||||||
// Connect establishes the database connection
|
// Connect establishes the database connection
|
||||||
func (c *sqlConnection) Connect(ctx context.Context) error {
|
func (c *sqlConnection) Connect(ctx context.Context) error {
|
||||||
|
c.lifecycleMu.Lock()
|
||||||
|
defer c.lifecycleMu.Unlock()
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
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 {
|
if c.connected {
|
||||||
return ErrAlreadyConnected
|
return ErrAlreadyConnected
|
||||||
}
|
}
|
||||||
@@ -127,17 +139,29 @@ func (c *sqlConnection) Connect(ctx context.Context) error {
|
|||||||
|
|
||||||
// Close closes the database connection and all ORM instances
|
// Close closes the database connection and all ORM instances
|
||||||
func (c *sqlConnection) Close() error {
|
func (c *sqlConnection) Close() error {
|
||||||
|
c.lifecycleMu.Lock()
|
||||||
|
defer c.lifecycleMu.Unlock()
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
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 {
|
if !c.connected {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close Bun if initialized
|
var errs []error
|
||||||
if c.bunDB != nil {
|
|
||||||
|
// 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 {
|
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)
|
// Close the provider (which closes the underlying sql.DB)
|
||||||
if err := c.provider.Close(); err != nil {
|
if err := c.provider.Close(); err != nil {
|
||||||
return NewConnectionError(c.name, "close", err)
|
errs = append(errs, NewConnectionError(c.name, "close", err))
|
||||||
}
|
}
|
||||||
|
|
||||||
c.connected = false
|
c.connected = false
|
||||||
@@ -156,39 +180,75 @@ func (c *sqlConnection) Close() error {
|
|||||||
c.gormAdapter = nil
|
c.gormAdapter = nil
|
||||||
c.nativeAdapter = 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 {
|
func (c *sqlConnection) HealthCheck(ctx context.Context) error {
|
||||||
if c == nil {
|
if c == nil {
|
||||||
return fmt.Errorf("connection is 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.mu.RLock()
|
||||||
c.healthCheckStatus = "disconnected"
|
connected := c.connected
|
||||||
|
provider := c.provider
|
||||||
|
c.mu.RUnlock()
|
||||||
|
|
||||||
|
if !connected {
|
||||||
|
c.setHealth("disconnected")
|
||||||
return ErrConnectionClosed
|
return ErrConnectionClosed
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.provider.HealthCheck(ctx); err != nil {
|
if err := provider.HealthCheck(ctx); err != nil {
|
||||||
c.healthCheckStatus = "unhealthy: " + err.Error()
|
c.setHealth("unhealthy: " + err.Error())
|
||||||
return NewConnectionError(c.name, "health check", err)
|
return NewConnectionError(c.name, "health check", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
c.healthCheckStatus = "healthy"
|
c.setHealth("healthy")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reconnect closes and re-establishes the connection
|
func (c *sqlConnection) setHealth(status string) {
|
||||||
func (c *sqlConnection) Reconnect(ctx context.Context) error {
|
c.mu.Lock()
|
||||||
if err := c.Close(); err != nil {
|
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 err
|
||||||
}
|
}
|
||||||
return c.Connect(ctx)
|
return c.connectLocked(ctx)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Native returns the native *sql.DB connection
|
// Native returns the native *sql.DB connection
|
||||||
@@ -250,6 +310,10 @@ func (c *sqlConnection) Bun() (*bun.DB, error) {
|
|||||||
return c.bunDB, nil
|
return c.bunDB, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !c.connected {
|
||||||
|
return nil, ErrConnectionClosed
|
||||||
|
}
|
||||||
|
|
||||||
// Get native connection first
|
// Get native connection first
|
||||||
native, err := c.provider.GetNative()
|
native, err := c.provider.GetNative()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -283,6 +347,10 @@ func (c *sqlConnection) GORM() (*gorm.DB, error) {
|
|||||||
return c.gormDB, nil
|
return c.gormDB, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !c.connected {
|
||||||
|
return nil, ErrConnectionClosed
|
||||||
|
}
|
||||||
|
|
||||||
// Get native connection first
|
// Get native connection first
|
||||||
native, err := c.provider.GetNative()
|
native, err := c.provider.GetNative()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -359,39 +427,18 @@ func (c *sqlConnection) Stats() *ConnectionStats {
|
|||||||
return stats
|
return stats
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *sqlConnection) reconnectForAdapter() error {
|
// The adapter factories only re-fetch the current handle. They must not close
|
||||||
timeout := c.config.ConnectTimeout
|
// the shared pool: *sql.DB discards bad connections on its own, and closing it
|
||||||
if timeout <= 0 {
|
// here would break every other holder of the pool.
|
||||||
timeout = 10 * time.Second
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
return c.Reconnect(ctx)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *sqlConnection) reopenNativeForAdapter() (*sql.DB, error) {
|
func (c *sqlConnection) reopenNativeForAdapter() (*sql.DB, error) {
|
||||||
if err := c.reconnectForAdapter(); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return c.Native()
|
return c.Native()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *sqlConnection) reopenBunForAdapter() (*bun.DB, error) {
|
func (c *sqlConnection) reopenBunForAdapter() (*bun.DB, error) {
|
||||||
if err := c.reconnectForAdapter(); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return c.Bun()
|
return c.Bun()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *sqlConnection) reopenGORMForAdapter() (*gorm.DB, error) {
|
func (c *sqlConnection) reopenGORMForAdapter() (*gorm.DB, error) {
|
||||||
if err := c.reconnectForAdapter(); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return c.GORM()
|
return c.GORM()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -512,15 +559,8 @@ func (c *sqlConnection) getNativeAdapter() (common.Database, error) {
|
|||||||
|
|
||||||
// Create a native adapter based on database type
|
// Create a native adapter based on database type
|
||||||
switch c.dbType {
|
switch c.dbType {
|
||||||
case DatabaseTypePostgreSQL:
|
case DatabaseTypePostgreSQL, DatabaseTypeSQLite, DatabaseTypeMSSQL:
|
||||||
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
|
// The adapter takes the driver name so it can adjust its dialect.
|
||||||
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:
|
|
||||||
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
|
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
|
||||||
WithDBFactory(c.reopenNativeForAdapter).
|
WithDBFactory(c.reopenNativeForAdapter).
|
||||||
SetMetricsEnabled(c.config.EnableMetrics)
|
SetMetricsEnabled(c.config.EnableMetrics)
|
||||||
@@ -572,8 +612,9 @@ type mongoConnection struct {
|
|||||||
client *mongo.Client
|
client *mongo.Client
|
||||||
|
|
||||||
// State
|
// State
|
||||||
connected bool
|
connected bool
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
|
lifecycleMu sync.RWMutex // see sqlConnection.lifecycleMu
|
||||||
|
|
||||||
// Health check
|
// Health check
|
||||||
lastHealthCheck time.Time
|
lastHealthCheck time.Time
|
||||||
@@ -601,9 +642,16 @@ func (c *mongoConnection) Type() DatabaseType {
|
|||||||
|
|
||||||
// Connect establishes the MongoDB connection
|
// Connect establishes the MongoDB connection
|
||||||
func (c *mongoConnection) Connect(ctx context.Context) error {
|
func (c *mongoConnection) Connect(ctx context.Context) error {
|
||||||
|
c.lifecycleMu.Lock()
|
||||||
|
defer c.lifecycleMu.Unlock()
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
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 {
|
if c.connected {
|
||||||
return ErrAlreadyConnected
|
return ErrAlreadyConnected
|
||||||
}
|
}
|
||||||
@@ -615,6 +663,7 @@ func (c *mongoConnection) Connect(ctx context.Context) error {
|
|||||||
// Get the mongo client
|
// Get the mongo client
|
||||||
client, err := c.provider.GetMongo()
|
client, err := c.provider.GetMongo()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
_ = c.provider.Close()
|
||||||
return NewConnectionError(c.name, "get mongo client", err)
|
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
|
// Close closes the MongoDB connection
|
||||||
func (c *mongoConnection) Close() error {
|
func (c *mongoConnection) Close() error {
|
||||||
|
c.lifecycleMu.Lock()
|
||||||
|
defer c.lifecycleMu.Unlock()
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
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 {
|
if !c.connected {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.provider.Close(); err != nil {
|
err := c.provider.Close()
|
||||||
return NewConnectionError(c.name, "close", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
c.connected = false
|
c.connected = false
|
||||||
c.client = nil
|
c.client = nil
|
||||||
|
if err != nil {
|
||||||
|
return NewConnectionError(c.name, "close", err)
|
||||||
|
}
|
||||||
return nil
|
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 {
|
func (c *mongoConnection) HealthCheck(ctx context.Context) error {
|
||||||
c.mu.Lock()
|
c.lifecycleMu.RLock()
|
||||||
defer c.mu.Unlock()
|
defer c.lifecycleMu.RUnlock()
|
||||||
|
|
||||||
c.lastHealthCheck = time.Now()
|
c.mu.RLock()
|
||||||
|
connected := c.connected
|
||||||
|
c.mu.RUnlock()
|
||||||
|
|
||||||
if !c.connected {
|
if !connected {
|
||||||
c.healthCheckStatus = "disconnected"
|
c.setHealth("disconnected")
|
||||||
return ErrConnectionClosed
|
return ErrConnectionClosed
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.provider.HealthCheck(ctx); err != nil {
|
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)
|
return NewConnectionError(c.name, "health check", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
c.healthCheckStatus = "healthy"
|
c.setHealth("healthy")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reconnect closes and re-establishes the MongoDB connection
|
func (c *mongoConnection) setHealth(status string) {
|
||||||
func (c *mongoConnection) Reconnect(ctx context.Context) error {
|
c.mu.Lock()
|
||||||
if err := c.Close(); err != nil {
|
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 err
|
||||||
}
|
}
|
||||||
return c.Connect(ctx)
|
return c.connectLocked(ctx)
|
||||||
}
|
}
|
||||||
|
|
||||||
// MongoDB returns the MongoDB client
|
// MongoDB returns the MongoDB client
|
||||||
|
|||||||
@@ -4,13 +4,8 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
|
||||||
|
|
||||||
_ "github.com/mattn/go-sqlite3"
|
_ "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) {
|
func TestNewConnectionFromDB(t *testing.T) {
|
||||||
@@ -213,157 +208,3 @@ func TestNewConnectionFromDB_PostgreSQL(t *testing.T) {
|
|||||||
t.Errorf("Expected type DatabaseTypePostgreSQL, got '%s'", conn.Type())
|
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")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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
@@ -2,9 +2,7 @@ package dbmanager
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -49,6 +47,7 @@ type connectionManager struct {
|
|||||||
// Background health check
|
// Background health check
|
||||||
healthTicker *time.Ticker
|
healthTicker *time.Ticker
|
||||||
stopChan chan struct{}
|
stopChan chan struct{}
|
||||||
|
healthMu sync.Mutex // guards healthTicker and stopChan
|
||||||
wg sync.WaitGroup
|
wg sync.WaitGroup
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -100,7 +99,9 @@ func ResetInstance() {
|
|||||||
defer instanceMu.Unlock()
|
defer instanceMu.Unlock()
|
||||||
|
|
||||||
if instance != nil {
|
if instance != nil {
|
||||||
_ = instance.Close()
|
if err := instance.Close(); err != nil {
|
||||||
|
logger.Error("Failed to close manager during reset: %v", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
instance = nil
|
instance = nil
|
||||||
}
|
}
|
||||||
@@ -116,7 +117,6 @@ func NewManager(cfg ManagerConfig) (Manager, error) {
|
|||||||
mgr := &connectionManager{
|
mgr := &connectionManager{
|
||||||
connections: make(map[string]Connection),
|
connections: make(map[string]Connection),
|
||||||
config: cfg,
|
config: cfg,
|
||||||
stopChan: make(chan struct{}),
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return mgr, nil
|
return mgr, nil
|
||||||
@@ -195,11 +195,26 @@ func (m *connectionManager) SetDefaultDatabase(name string) error {
|
|||||||
|
|
||||||
// Connect establishes all configured database connections
|
// Connect establishes all configured database connections
|
||||||
func (m *connectionManager) Connect(ctx context.Context) error {
|
func (m *connectionManager) Connect(ctx context.Context) error {
|
||||||
m.mu.Lock()
|
// Dial outside m.mu so a slow connect never blocks Get/Stats/HealthCheck.
|
||||||
defer m.mu.Unlock()
|
m.mu.RLock()
|
||||||
|
names := make([]string, 0, len(m.config.Connections))
|
||||||
// Create connections from configuration
|
|
||||||
for name := range 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
|
// Get a copy of the connection config
|
||||||
connCfg := m.config.Connections[name]
|
connCfg := m.config.Connections[name]
|
||||||
// Apply global defaults to connection config
|
// Apply global defaults to connection config
|
||||||
@@ -209,25 +224,39 @@ func (m *connectionManager) Connect(ctx context.Context) error {
|
|||||||
// Create connection using factory
|
// Create connection using factory
|
||||||
conn, err := createConnection(connCfg)
|
conn, err := createConnection(connCfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
closeOpened()
|
||||||
return fmt.Errorf("failed to create connection '%s': %w", name, err)
|
return fmt.Errorf("failed to create connection '%s': %w", name, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Connect
|
// Connect
|
||||||
if err := conn.Connect(ctx); err != nil {
|
if err := conn.Connect(ctx); err != nil {
|
||||||
|
closeOpened()
|
||||||
return fmt.Errorf("failed to connect '%s': %w", name, err)
|
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)
|
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
|
// Always start background health checks
|
||||||
if m.config.HealthCheckInterval > 0 {
|
if m.config.HealthCheckInterval > 0 {
|
||||||
m.startHealthChecker()
|
m.startHealthChecker()
|
||||||
logger.Info("Background health checker started: interval=%v", m.config.HealthCheckInterval)
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -246,7 +275,7 @@ func (m *connectionManager) Close() error {
|
|||||||
for name, conn := range m.connections {
|
for name, conn := range m.connections {
|
||||||
if err := conn.Close(); err != nil {
|
if err := conn.Close(); err != nil {
|
||||||
errors = append(errors, fmt.Errorf("failed to close connection '%s': %w", name, err))
|
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 {
|
} else {
|
||||||
logger.Info("Connection closed: name=%s", name)
|
logger.Info("Connection closed: name=%s", name)
|
||||||
}
|
}
|
||||||
@@ -311,11 +340,17 @@ func (m *connectionManager) Stats() *ManagerStats {
|
|||||||
|
|
||||||
// startHealthChecker starts background health checking
|
// startHealthChecker starts background health checking
|
||||||
func (m *connectionManager) startHealthChecker() {
|
func (m *connectionManager) startHealthChecker() {
|
||||||
|
m.healthMu.Lock()
|
||||||
|
defer m.healthMu.Unlock()
|
||||||
|
|
||||||
if m.healthTicker != nil {
|
if m.healthTicker != nil {
|
||||||
return // Already running
|
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)
|
m.wg.Add(1)
|
||||||
go func() {
|
go func() {
|
||||||
@@ -324,9 +359,9 @@ func (m *connectionManager) startHealthChecker() {
|
|||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-m.healthTicker.C:
|
case <-ticker.C:
|
||||||
m.performHealthCheck()
|
m.performHealthCheck()
|
||||||
case <-m.stopChan:
|
case <-stop:
|
||||||
logger.Info("Health checker stopped")
|
logger.Info("Health checker stopped")
|
||||||
return
|
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() {
|
func (m *connectionManager) stopHealthChecker() {
|
||||||
if m.healthTicker != nil {
|
m.healthMu.Lock()
|
||||||
m.healthTicker.Stop()
|
defer m.healthMu.Unlock()
|
||||||
close(m.stopChan)
|
|
||||||
m.wg.Wait()
|
if m.healthTicker == nil {
|
||||||
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
|
// performHealthCheck performs a health check on all connections
|
||||||
@@ -362,40 +402,14 @@ func (m *connectionManager) performHealthCheck() {
|
|||||||
}
|
}
|
||||||
m.mu.RUnlock()
|
m.mu.RUnlock()
|
||||||
|
|
||||||
|
defer m.PublishMetrics()
|
||||||
|
|
||||||
for _, item := range connections {
|
for _, item := range connections {
|
||||||
if err := item.conn.HealthCheck(ctx); err != nil {
|
if err := item.conn.HealthCheck(ctx); err != nil {
|
||||||
logger.Warn("Health check failed",
|
// Do not reconnect here: *sql.DB discards bad connections and dials
|
||||||
"connection", item.name,
|
// new ones by itself, while Reconnect closes the pool and breaks
|
||||||
"error", err)
|
// every handle already handed out. Reconnect is operator-only.
|
||||||
|
logger.Warn("Health check failed: connection=%s, error=%v", item.name, 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)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
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")
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -21,19 +21,30 @@ type healthCheckStubConnection struct {
|
|||||||
reconnectCalls int
|
reconnectCalls int
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *healthCheckStubConnection) Name() string { return "stub" }
|
func (c *healthCheckStubConnection) Name() string { return "stub" }
|
||||||
func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL }
|
func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL }
|
||||||
func (c *healthCheckStubConnection) Bun() (*bun.DB, error) { return nil, fmt.Errorf("not implemented") }
|
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) GORM() (*gorm.DB, error) {
|
||||||
func (c *healthCheckStubConnection) Native() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") }
|
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) Native() (*sql.DB, error) {
|
||||||
func (c *healthCheckStubConnection) MongoDB() (*mongo.Client, error) { return nil, fmt.Errorf("not implemented") }
|
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) DB() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") }
|
||||||
func (c *healthCheckStubConnection) HealthCheck(ctx context.Context) error { return c.healthErr }
|
func (c *healthCheckStubConnection) Database() (common.Database, error) {
|
||||||
func (c *healthCheckStubConnection) Reconnect(ctx context.Context) error { c.reconnectCalls++; return nil }
|
return nil, fmt.Errorf("not implemented")
|
||||||
func (c *healthCheckStubConnection) Stats() *ConnectionStats { return &ConnectionStats{} }
|
}
|
||||||
|
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) {
|
func TestBackgroundHealthChecker(t *testing.T) {
|
||||||
// Create a SQLite in-memory database
|
// 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",
|
t.Errorf("Expected default health check interval to be %v, got %v",
|
||||||
expectedInterval, defaults.HealthCheckInterval)
|
expectedInterval, defaults.HealthCheckInterval)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !defaults.EnableAutoReconnect {
|
|
||||||
t.Error("Expected EnableAutoReconnect to be true by default")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestApplyDefaultsEnablesAutoReconnect(t *testing.T) {
|
func TestApplyDefaultsHealthCheckInterval(t *testing.T) {
|
||||||
// Create a config without setting EnableAutoReconnect
|
cfg := ManagerConfig{}
|
||||||
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
|
|
||||||
cfg.ApplyDefaults()
|
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 {
|
if cfg.HealthCheckInterval != 15*time.Second {
|
||||||
t.Errorf("Expected health check interval to be 15s, got %v", cfg.HealthCheckInterval)
|
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) {
|
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{
|
conn := &healthCheckStubConnection{
|
||||||
healthErr: NewConnectionError("primary", "health check", fmt.Errorf("sql: database is closed")),
|
healthErr: NewConnectionError("primary", "health check", fmt.Errorf("sql: database is closed")),
|
||||||
}
|
}
|
||||||
@@ -284,7 +275,7 @@ func TestPerformHealthCheckReconnectsClosedConnections(t *testing.T) {
|
|||||||
|
|
||||||
mgr.performHealthCheck()
|
mgr.performHealthCheck()
|
||||||
|
|
||||||
if conn.reconnectCalls != 1 {
|
if conn.reconnectCalls != 0 {
|
||||||
t.Fatalf("expected reconnect attempt for closed database handle, got %d", conn.reconnectCalls)
|
t.Fatalf("health check must not close the shared pool via Reconnect, got %d", conn.reconnectCalls)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+39
-15
@@ -1,6 +1,8 @@
|
|||||||
package dbmanager
|
package dbmanager
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"sync"
|
||||||
|
|
||||||
"github.com/prometheus/client_golang/prometheus"
|
"github.com/prometheus/client_golang/prometheus"
|
||||||
"github.com/prometheus/client_golang/prometheus/promauto"
|
"github.com/prometheus/client_golang/prometheus/promauto"
|
||||||
)
|
)
|
||||||
@@ -34,8 +36,8 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// connectionWaitCount tracks how many times connections had to wait for availability
|
// connectionWaitCount tracks how many times connections had to wait for availability
|
||||||
connectionWaitCount = promauto.NewGaugeVec(
|
connectionWaitCount = promauto.NewCounterVec(
|
||||||
prometheus.GaugeOpts{
|
prometheus.CounterOpts{
|
||||||
Name: "dbmanager_connection_wait_count",
|
Name: "dbmanager_connection_wait_count",
|
||||||
Help: "Number of times connections had to wait for availability",
|
Help: "Number of times connections had to wait for availability",
|
||||||
},
|
},
|
||||||
@@ -43,8 +45,8 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// connectionWaitDuration tracks total time connections spent waiting
|
// connectionWaitDuration tracks total time connections spent waiting
|
||||||
connectionWaitDuration = promauto.NewGaugeVec(
|
connectionWaitDuration = promauto.NewCounterVec(
|
||||||
prometheus.GaugeOpts{
|
prometheus.CounterOpts{
|
||||||
Name: "dbmanager_connection_wait_duration_seconds",
|
Name: "dbmanager_connection_wait_duration_seconds",
|
||||||
Help: "Total time connections spent waiting for availability",
|
Help: "Total time connections spent waiting for availability",
|
||||||
},
|
},
|
||||||
@@ -61,8 +63,8 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// connectionLifetimeClosed tracks connections closed due to max lifetime
|
// connectionLifetimeClosed tracks connections closed due to max lifetime
|
||||||
connectionLifetimeClosed = promauto.NewGaugeVec(
|
connectionLifetimeClosed = promauto.NewCounterVec(
|
||||||
prometheus.GaugeOpts{
|
prometheus.CounterOpts{
|
||||||
Name: "dbmanager_connection_lifetime_closed_total",
|
Name: "dbmanager_connection_lifetime_closed_total",
|
||||||
Help: "Total connections closed due to exceeding max lifetime",
|
Help: "Total connections closed due to exceeding max lifetime",
|
||||||
},
|
},
|
||||||
@@ -70,8 +72,8 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// connectionIdleClosed tracks connections closed due to max idle time
|
// connectionIdleClosed tracks connections closed due to max idle time
|
||||||
connectionIdleClosed = promauto.NewGaugeVec(
|
connectionIdleClosed = promauto.NewCounterVec(
|
||||||
prometheus.GaugeOpts{
|
prometheus.CounterOpts{
|
||||||
Name: "dbmanager_connection_idle_closed_total",
|
Name: "dbmanager_connection_idle_closed_total",
|
||||||
Help: "Total connections closed due to exceeding max idle time",
|
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), "idle").Set(float64(connStats.Idle))
|
||||||
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "in_use").Set(float64(connStats.InUse))
|
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "in_use").Set(float64(connStats.InUse))
|
||||||
|
|
||||||
// Wait stats
|
// sql.DBStats values are cumulative, so add only the growth since
|
||||||
connectionWaitCount.With(labels).Set(float64(connStats.WaitCount))
|
// the last publish to keep these true counters.
|
||||||
connectionWaitDuration.With(labels).Set(connStats.WaitDuration.Seconds())
|
prev := lastPublished.swap(name, connStats)
|
||||||
|
connectionWaitCount.With(labels).Add(float64(connStats.WaitCount - prev.WaitCount))
|
||||||
// Lifetime/idle closed
|
connectionWaitDuration.With(labels).Add((connStats.WaitDuration - prev.WaitDuration).Seconds())
|
||||||
connectionLifetimeClosed.With(labels).Set(float64(connStats.MaxLifetimeClosed))
|
connectionLifetimeClosed.With(labels).Add(float64(connStats.MaxLifetimeClosed - prev.MaxLifetimeClosed))
|
||||||
connectionIdleClosed.With(labels).Set(float64(connStats.MaxIdleClosed))
|
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()
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -7,6 +7,8 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"go.mongodb.org/mongo-driver/mongo"
|
"go.mongodb.org/mongo-driver/mongo"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ExistingDBProvider wraps an existing *sql.DB connection
|
// ExistingDBProvider wraps an existing *sql.DB connection
|
||||||
@@ -44,16 +46,27 @@ func (p *ExistingDBProvider) Connect(ctx context.Context, cfg ConnectionConfig)
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close closes the underlying database connection
|
// Refresh verifies the wrapped database is still reachable. The pool belongs to
|
||||||
func (p *ExistingDBProvider) Close() error {
|
// the caller and cannot be re-dialed here, so it is never closed to "reconnect".
|
||||||
p.mu.Lock()
|
func (p *ExistingDBProvider) Refresh(ctx context.Context) error {
|
||||||
defer p.mu.Unlock()
|
p.mu.RLock()
|
||||||
|
defer p.mu.RUnlock()
|
||||||
|
|
||||||
if p.db == nil {
|
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
|
// HealthCheck verifies the connection is alive
|
||||||
|
|||||||
@@ -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:")
|
db, err := sql.Open("sqlite3", ":memory:")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Failed to open database: %v", err)
|
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)
|
t.Errorf("Expected Close to succeed, got error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Verify the database is closed
|
// The caller owns the database, so Close must leave it open
|
||||||
err = db.Ping()
|
defer db.Close()
|
||||||
if err == nil {
|
if err := db.Ping(); err != nil {
|
||||||
t.Error("Expected database to be closed")
|
t.Errorf("Expected caller's database to stay open, got: %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -41,9 +41,10 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
clientOpts.SetMaxPoolSize(maxPoolSize)
|
clientOpts.SetMaxPoolSize(maxPoolSize)
|
||||||
}
|
}
|
||||||
|
|
||||||
if cfg.GetMaxIdleConns() != nil {
|
// MaxIdleConns is a ceiling on idle connections, not a pre-warmed minimum
|
||||||
minPoolSize := uint64(*cfg.GetMaxIdleConns())
|
// (MinPoolSize), so only the idle-time limit maps onto the Mongo pool.
|
||||||
clientOpts.SetMinPoolSize(minPoolSize)
|
if cfg.GetConnMaxIdleTime() != nil {
|
||||||
|
clientOpts.SetMaxConnIdleTime(*cfg.GetConnMaxIdleTime())
|
||||||
}
|
}
|
||||||
|
|
||||||
// Set timeouts
|
// Set timeouts
|
||||||
@@ -65,12 +66,11 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
var client *mongo.Client
|
var client *mongo.Client
|
||||||
var lastErr error
|
var lastErr error
|
||||||
|
|
||||||
retryAttempts := 3
|
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
|
||||||
retryDelay := 1 * time.Second
|
|
||||||
|
|
||||||
for attempt := 0; attempt < retryAttempts; attempt++ {
|
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||||
if attempt > 0 {
|
if attempt > 0 {
|
||||||
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
|
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Info("Retrying MongoDB connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
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 {
|
if err != nil {
|
||||||
lastErr = err
|
lastErr = err
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Warn("Failed to connect to MongoDB", "error", err)
|
logger.Warn("Failed to connect to MongoDB: %v", err)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -101,7 +101,7 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
lastErr = err
|
lastErr = err
|
||||||
_ = client.Disconnect(ctx)
|
_ = client.Disconnect(ctx)
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Warn("Failed to ping MongoDB", "error", err)
|
logger.Warn("Failed to ping MongoDB: %v", err)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -35,12 +35,11 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
var db *sql.DB
|
var db *sql.DB
|
||||||
var lastErr error
|
var lastErr error
|
||||||
|
|
||||||
retryAttempts := 3 // Default retry attempts
|
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
|
||||||
retryDelay := 1 * time.Second
|
|
||||||
|
|
||||||
for attempt := 0; attempt < retryAttempts; attempt++ {
|
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||||
if attempt > 0 {
|
if attempt > 0 {
|
||||||
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
|
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Info("Retrying MSSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
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 {
|
if err != nil {
|
||||||
lastErr = err
|
lastErr = err
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Warn("Failed to open MSSQL connection", "error", err)
|
logger.Warn("Failed to open MSSQL connection: %v", err)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -71,7 +70,7 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
lastErr = err
|
lastErr = err
|
||||||
db.Close()
|
db.Close()
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Warn("Failed to ping MSSQL database", "error", err)
|
logger.Warn("Failed to ping MSSQL database: %v", err)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -3,12 +3,12 @@ package providers
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"math"
|
"math"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
_ "github.com/jackc/pgx/v5/stdlib" // PostgreSQL driver
|
|
||||||
"go.mongodb.org/mongo-driver/mongo"
|
"go.mongodb.org/mongo-driver/mongo"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
@@ -16,10 +16,11 @@ import (
|
|||||||
|
|
||||||
// PostgresProvider implements Provider for PostgreSQL databases
|
// PostgresProvider implements Provider for PostgreSQL databases
|
||||||
type PostgresProvider struct {
|
type PostgresProvider struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
config ConnectionConfig
|
connector *pgConnector
|
||||||
listener *PostgresListener
|
config ConnectionConfig
|
||||||
mu sync.Mutex
|
listener *PostgresListener
|
||||||
|
mu sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewPostgresProvider creates a new PostgreSQL provider
|
// NewPostgresProvider creates a new PostgreSQL provider
|
||||||
@@ -29,22 +30,24 @@ func NewPostgresProvider() *PostgresProvider {
|
|||||||
|
|
||||||
// Connect establishes a PostgreSQL connection
|
// Connect establishes a PostgreSQL connection
|
||||||
func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
||||||
// Build DSN
|
connCfg, err := buildPGXConfig(cfg)
|
||||||
dsn, err := cfg.BuildDSN()
|
|
||||||
if err != nil {
|
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
|
// Connect with retry logic
|
||||||
var db *sql.DB
|
|
||||||
var lastErr error
|
var lastErr error
|
||||||
|
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
|
||||||
|
|
||||||
retryAttempts := 3 // Default retry attempts
|
connected := false
|
||||||
retryDelay := 1 * time.Second
|
|
||||||
|
|
||||||
for attempt := 0; attempt < retryAttempts; attempt++ {
|
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||||
if attempt > 0 {
|
if attempt > 0 {
|
||||||
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
|
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Info("Retrying PostgreSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
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 {
|
select {
|
||||||
case <-time.After(delay):
|
case <-time.After(delay):
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
|
db.Close()
|
||||||
return ctx.Err()
|
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
|
// Test the connection with context timeout
|
||||||
connectCtx, cancel := context.WithTimeout(ctx, cfg.GetConnectTimeout())
|
connectCtx, cancel := context.WithTimeout(ctx, cfg.GetConnectTimeout())
|
||||||
err = db.PingContext(connectCtx)
|
err = db.PingContext(connectCtx)
|
||||||
@@ -73,18 +67,18 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
|
|||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
lastErr = err
|
lastErr = err
|
||||||
db.Close()
|
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Warn("Failed to ping PostgreSQL database", "error", err)
|
logger.Warn("Failed to ping PostgreSQL database: %v", err)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Connection successful
|
connected = true
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if !connected {
|
||||||
|
db.Close()
|
||||||
return fmt.Errorf("failed to connect after %d attempts: %w", retryAttempts, lastErr)
|
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.db = db
|
||||||
|
p.connector = connector
|
||||||
p.config = cfg
|
p.config = cfg
|
||||||
|
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
@@ -112,34 +107,55 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close closes the PostgreSQL connection
|
// Refresh retires every pooled connection and dials fresh ones on demand,
|
||||||
func (p *PostgresProvider) Close() error {
|
// without closing the *sql.DB. Handles already handed out keep working:
|
||||||
// Close listener if it exists
|
// connections in use finish their current query and are then discarded.
|
||||||
p.mu.Lock()
|
func (p *PostgresProvider) Refresh(ctx context.Context) error {
|
||||||
if p.listener != nil {
|
if p.db == nil || p.connector == nil {
|
||||||
if err := p.listener.Close(); err != nil {
|
return fmt.Errorf("database connection is not initialized")
|
||||||
p.mu.Unlock()
|
|
||||||
return fmt.Errorf("failed to close listener: %w", err)
|
|
||||||
}
|
|
||||||
p.listener = nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
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()
|
p.mu.Unlock()
|
||||||
|
|
||||||
if p.db == nil {
|
if listener != nil {
|
||||||
return nil
|
if err := listener.Close(); err != nil {
|
||||||
|
errs = append(errs, fmt.Errorf("failed to close listener: %w", err))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
err := p.db.Close()
|
if p.db != nil {
|
||||||
if err != nil {
|
if err := p.db.Close(); err != nil {
|
||||||
return fmt.Errorf("failed to close PostgreSQL connection: %w", err)
|
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() {
|
return errors.Join(errs...)
|
||||||
logger.Info("PostgreSQL connection closed: name=%s", p.config.GetName())
|
|
||||||
}
|
|
||||||
|
|
||||||
p.db = nil
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// HealthCheck verifies the PostgreSQL connection is alive
|
// HealthCheck verifies the PostgreSQL connection is alive
|
||||||
|
|||||||
@@ -23,6 +23,10 @@ type PostgresListener struct {
|
|||||||
// Channel subscriptions
|
// Channel subscriptions
|
||||||
channels map[string]NotificationHandler
|
channels map[string]NotificationHandler
|
||||||
mu sync.RWMutex
|
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
|
// Lifecycle management
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
@@ -30,6 +34,7 @@ type PostgresListener struct {
|
|||||||
closed bool
|
closed bool
|
||||||
closeMu sync.Mutex
|
closeMu sync.Mutex
|
||||||
reconnectC chan struct{}
|
reconnectC chan struct{}
|
||||||
|
startOnce sync.Once // background goroutines start exactly once
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewPostgresListener creates a new PostgreSQL listener
|
// 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 {
|
func (l *PostgresListener) Connect(ctx context.Context) error {
|
||||||
dsn, err := l.config.BuildDSN()
|
conn, err := l.dial(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to build DSN: %w", err)
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse connection config
|
l.swapConn(conn)
|
||||||
connConfig, err := pgx.ParseConfig(dsn)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to parse connection config: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Connect with retry logic
|
l.startOnce.Do(func() {
|
||||||
var conn *pgx.Conn
|
go l.handleNotifications()
|
||||||
var lastErr error
|
go l.handleReconnection()
|
||||||
|
})
|
||||||
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()
|
|
||||||
|
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Info("PostgreSQL listener connected: name=%s", l.config.GetName())
|
logger.Info("PostgreSQL listener connected: name=%s", l.config.GetName())
|
||||||
@@ -122,30 +71,105 @@ func (l *PostgresListener) Connect(ctx context.Context) error {
|
|||||||
return nil
|
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
|
// Listen subscribes to a PostgreSQL notification channel
|
||||||
func (l *PostgresListener) Listen(channel string, handler NotificationHandler) error {
|
func (l *PostgresListener) Listen(channel string, handler NotificationHandler) error {
|
||||||
l.closeMu.Lock()
|
// Take the connection between notification waits (each wait is short).
|
||||||
if l.closed {
|
l.connMu.Lock()
|
||||||
l.closeMu.Unlock()
|
defer l.connMu.Unlock()
|
||||||
return fmt.Errorf("listener is closed")
|
|
||||||
}
|
|
||||||
l.closeMu.Unlock()
|
|
||||||
|
|
||||||
l.mu.Lock()
|
conn, err := l.currentConn()
|
||||||
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()))
|
|
||||||
if err != nil {
|
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)
|
return fmt.Errorf("failed to listen on channel %s: %w", channel, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Store the handler
|
l.mu.Lock()
|
||||||
l.channels[channel] = handler
|
l.channels[channel] = handler
|
||||||
|
l.mu.Unlock()
|
||||||
|
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Info("Listening on channel: name=%s, channel=%s", l.config.GetName(), channel)
|
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
|
// Unlisten unsubscribes from a PostgreSQL notification channel
|
||||||
func (l *PostgresListener) Unlisten(channel string) error {
|
func (l *PostgresListener) Unlisten(channel string) error {
|
||||||
l.closeMu.Lock()
|
l.connMu.Lock()
|
||||||
if l.closed {
|
defer l.connMu.Unlock()
|
||||||
l.closeMu.Unlock()
|
|
||||||
return fmt.Errorf("listener is closed")
|
|
||||||
}
|
|
||||||
l.closeMu.Unlock()
|
|
||||||
|
|
||||||
l.mu.Lock()
|
conn, err := l.currentConn()
|
||||||
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()))
|
|
||||||
if err != nil {
|
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)
|
return fmt.Errorf("failed to unlisten from channel %s: %w", channel, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Remove the handler
|
l.mu.Lock()
|
||||||
delete(l.channels, channel)
|
delete(l.channels, channel)
|
||||||
|
l.mu.Unlock()
|
||||||
|
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Info("Unlistened from channel: name=%s, channel=%s", l.config.GetName(), channel)
|
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
|
// Notify sends a notification to a PostgreSQL channel
|
||||||
func (l *PostgresListener) Notify(ctx context.Context, channel string, payload string) error {
|
func (l *PostgresListener) Notify(ctx context.Context, channel string, payload string) error {
|
||||||
l.closeMu.Lock()
|
l.connMu.Lock()
|
||||||
if l.closed {
|
defer l.connMu.Unlock()
|
||||||
l.closeMu.Unlock()
|
|
||||||
return fmt.Errorf("listener is closed")
|
|
||||||
}
|
|
||||||
l.closeMu.Unlock()
|
|
||||||
|
|
||||||
l.mu.RLock()
|
conn, err := l.currentConn()
|
||||||
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)
|
|
||||||
if err != nil {
|
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 fmt.Errorf("failed to notify channel %s: %w", channel, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
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 {
|
func (l *PostgresListener) Close() error {
|
||||||
l.closeMu.Lock()
|
l.closeMu.Lock()
|
||||||
if l.closed {
|
if l.closed {
|
||||||
@@ -225,27 +235,26 @@ func (l *PostgresListener) Close() error {
|
|||||||
// Cancel context to stop background goroutines
|
// Cancel context to stop background goroutines
|
||||||
l.cancel()
|
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()
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Unlisten from all channels
|
err := closeConnBounded(conn)
|
||||||
for channel := range l.channels {
|
l.connMu.Unlock()
|
||||||
_, _ = l.conn.Exec(context.Background(), fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize()))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Close connection
|
|
||||||
err := l.conn.Close(context.Background())
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to close listener connection: %w", err)
|
return fmt.Errorf("failed to close listener connection: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
l.conn = nil
|
|
||||||
l.channels = make(map[string]NotificationHandler)
|
|
||||||
|
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Info("PostgreSQL listener closed: name=%s", l.config.GetName())
|
logger.Info("PostgreSQL listener closed: name=%s", l.config.GetName())
|
||||||
}
|
}
|
||||||
@@ -262,20 +271,26 @@ func (l *PostgresListener) handleNotifications() {
|
|||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
|
|
||||||
|
l.connMu.Lock()
|
||||||
l.mu.RLock()
|
l.mu.RLock()
|
||||||
conn := l.conn
|
conn := l.conn
|
||||||
l.mu.RUnlock()
|
l.mu.RUnlock()
|
||||||
|
|
||||||
if conn == nil {
|
if conn == nil {
|
||||||
|
l.connMu.Unlock()
|
||||||
// Connection not available, wait for reconnection
|
// Connection not available, wait for reconnection
|
||||||
time.Sleep(100 * time.Millisecond)
|
if !l.sleep(100 * time.Millisecond) {
|
||||||
|
return
|
||||||
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Wait for notification with timeout
|
// Wait for a notification with a short timeout, so Listen/Unlisten/Notify
|
||||||
ctx, cancel := context.WithTimeout(l.ctx, 5*time.Second)
|
// waiting on connMu are served promptly.
|
||||||
|
ctx, cancel := context.WithTimeout(l.ctx, notificationPollInterval)
|
||||||
notification, err := conn.WaitForNotification(ctx)
|
notification, err := conn.WaitForNotification(ctx)
|
||||||
cancel()
|
cancel()
|
||||||
|
l.connMu.Unlock()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Check if context was cancelled
|
// Check if context was cancelled
|
||||||
@@ -291,13 +306,15 @@ func (l *PostgresListener) handleNotifications() {
|
|||||||
|
|
||||||
// Connection error, trigger reconnection
|
// Connection error, trigger reconnection
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Warn("Notification error, triggering reconnection", "error", err)
|
logger.Warn("Notification error, triggering reconnection: %v", err)
|
||||||
}
|
}
|
||||||
select {
|
select {
|
||||||
case l.reconnectC <- struct{}{}:
|
case l.reconnectC <- struct{}{}:
|
||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
time.Sleep(1 * time.Second)
|
if !l.sleep(1 * time.Second) {
|
||||||
|
return
|
||||||
|
}
|
||||||
continue
|
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() {
|
func (l *PostgresListener) handleReconnection() {
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
@@ -333,31 +365,21 @@ func (l *PostgresListener) handleReconnection() {
|
|||||||
logger.Info("Attempting to reconnect listener: name=%s", l.config.GetName())
|
logger.Info("Attempting to reconnect listener: name=%s", l.config.GetName())
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close existing connection
|
ctx, cancel := context.WithTimeout(l.ctx, 30*time.Second)
|
||||||
l.mu.Lock()
|
err := l.reconnect(ctx)
|
||||||
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)
|
|
||||||
cancel()
|
cancel()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
if l.ctx.Err() != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Error("Failed to reconnect listener: name=%s, error=%v", l.config.GetName(), err)
|
logger.Error("Failed to reconnect listener: name=%s, error=%v", l.config.GetName(), err)
|
||||||
}
|
}
|
||||||
// Retry after delay
|
// Retry after delay
|
||||||
time.Sleep(5 * time.Second)
|
if !l.sleep(5 * time.Second) {
|
||||||
|
return
|
||||||
|
}
|
||||||
select {
|
select {
|
||||||
case l.reconnectC <- struct{}{}:
|
case l.reconnectC <- struct{}{}:
|
||||||
default:
|
default:
|
||||||
@@ -365,15 +387,6 @@ func (l *PostgresListener) handleReconnection() {
|
|||||||
continue
|
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() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Info("Listener reconnected successfully: name=%s", l.config.GetName())
|
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
|
// IsConnected returns true if the listener is connected
|
||||||
func (l *PostgresListener) IsConnected() bool {
|
func (l *PostgresListener) IsConnected() bool {
|
||||||
l.mu.RLock()
|
l.mu.RLock()
|
||||||
|
|||||||
@@ -4,17 +4,11 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"errors"
|
"errors"
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"go.mongodb.org/mongo-driver/mongo"
|
"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
|
// Common errors
|
||||||
var (
|
var (
|
||||||
// ErrNotSQLDatabase is returned when attempting SQL operations on a non-SQL database
|
// ErrNotSQLDatabase is returned when attempting SQL operations on a non-SQL database
|
||||||
@@ -63,6 +57,31 @@ type ConnectionConfig interface {
|
|||||||
GetConnMaxLifetime() *time.Duration
|
GetConnMaxLifetime() *time.Duration
|
||||||
GetConnMaxIdleTime() *time.Duration
|
GetConnMaxIdleTime() *time.Duration
|
||||||
GetReadPreference() string
|
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
|
// Provider creates and manages the underlying database connection
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -15,10 +16,9 @@ import (
|
|||||||
|
|
||||||
// SQLiteProvider implements Provider for SQLite databases
|
// SQLiteProvider implements Provider for SQLite databases
|
||||||
type SQLiteProvider struct {
|
type SQLiteProvider struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
dbMu sync.RWMutex
|
dbMu sync.RWMutex
|
||||||
dbFactory func() (*sql.DB, error)
|
config ConnectionConfig
|
||||||
config ConnectionConfig
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewSQLiteProvider creates a new SQLite provider
|
// NewSQLiteProvider creates a new SQLite provider
|
||||||
@@ -26,6 +26,22 @@ func NewSQLiteProvider() *SQLiteProvider {
|
|||||||
return &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
|
// Connect establishes a SQLite connection
|
||||||
func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
||||||
// Build DSN
|
// 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)
|
return fmt.Errorf("failed to ping SQLite database: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Configure connection pool
|
if isMemoryDSN(dsn) {
|
||||||
// Note: SQLite works best with MaxOpenConns=1 for write operations
|
// A private in-memory database exists per connection and disappears when
|
||||||
// but can handle multiple readers
|
// that connection closes, so pin the pool to one connection that is
|
||||||
if cfg.GetMaxOpenConns() != nil {
|
// never recycled.
|
||||||
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
|
|
||||||
} else {
|
|
||||||
// Default to 1 for SQLite to avoid "database is locked" errors
|
|
||||||
db.SetMaxOpenConns(1)
|
db.SetMaxOpenConns(1)
|
||||||
}
|
db.SetMaxIdleConns(1)
|
||||||
|
db.SetConnMaxLifetime(0)
|
||||||
if cfg.GetMaxIdleConns() != nil {
|
db.SetConnMaxIdleTime(0)
|
||||||
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
|
} else {
|
||||||
}
|
// SQLite works best with few writers; default to 1 unless configured.
|
||||||
if cfg.GetConnMaxLifetime() != nil {
|
if cfg.GetMaxOpenConns() != nil {
|
||||||
db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime())
|
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
|
||||||
}
|
} else {
|
||||||
if cfg.GetConnMaxIdleTime() != nil {
|
db.SetMaxOpenConns(1)
|
||||||
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)
|
|
||||||
}
|
}
|
||||||
// Don't fail connection if WAL mode cannot be enabled
|
if cfg.GetMaxIdleConns() != nil {
|
||||||
}
|
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
|
||||||
|
}
|
||||||
// Set busy timeout to handle locked database (minimum 2 minutes = 120000ms)
|
if cfg.GetConnMaxLifetime() != nil {
|
||||||
busyTimeout := cfg.GetQueryTimeout().Milliseconds()
|
db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime())
|
||||||
if busyTimeout < 120000 {
|
}
|
||||||
busyTimeout = 120000 // Enforce minimum of 2 minutes
|
if cfg.GetConnMaxIdleTime() != nil {
|
||||||
}
|
db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime())
|
||||||
_, 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)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
p.dbMu.Lock()
|
||||||
p.db = db
|
p.db = db
|
||||||
|
p.dbMu.Unlock()
|
||||||
p.config = cfg
|
p.config = cfg
|
||||||
|
|
||||||
if cfg.GetEnableLogging() {
|
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
|
// Execute a simple query to verify the database is accessible
|
||||||
var result int
|
var result int
|
||||||
run := func() error { return p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result) }
|
if err := p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result); err != nil {
|
||||||
err := run()
|
|
||||||
if isDBClosed(err) {
|
|
||||||
if reconnErr := p.reconnectDB(); reconnErr == nil {
|
|
||||||
err = run()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("health check failed: %w", err)
|
return fmt.Errorf("health check failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -150,32 +146,12 @@ func (p *SQLiteProvider) HealthCheck(ctx context.Context) error {
|
|||||||
return nil
|
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 {
|
func (p *SQLiteProvider) getDB() *sql.DB {
|
||||||
p.dbMu.RLock()
|
p.dbMu.RLock()
|
||||||
defer p.dbMu.RUnlock()
|
defer p.dbMu.RUnlock()
|
||||||
return p.db
|
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
|
// GetNative returns the native *sql.DB connection
|
||||||
func (p *SQLiteProvider) GetNative() (*sql.DB, error) {
|
func (p *SQLiteProvider) GetNative() (*sql.DB, error) {
|
||||||
if p.db == nil {
|
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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user