diff --git a/README.md b/README.md index 4354c19..31c9d1e 100644 --- a/README.md +++ b/README.md @@ -576,11 +576,32 @@ Centralized management of multiple database connections with support for Postgre - Multiple named database connections - Multi-ORM access (Bun, GORM, Native SQL) sharing the same connection pool - Automatic SQLite schema translation (`schema.table` → `schema_table`) -- Health checks with auto-reconnect +- Background health checks (report status; they never close the pool) - Prometheus metrics for monitoring - Configuration-driven via YAML - Per-connection statistics and management +**How to use it correctly**: + +```go +mgr, err := dbmanager.NewManager(cfg) // or dbmanager.SetupManager(cfg) + GetInstance() +if err != nil { /* handle */ } +if err := mgr.Connect(ctx); err != nil { /* handle */ } // SetupManager does NOT connect +defer mgr.Close() // once, at shutdown + +conn, _ := mgr.GetDefault() +db, _ := conn.Bun() // or conn.GORM() / conn.Native() / conn.Database() +handler := restheadspec.NewHandlerWithBun(db) +``` + +- **Fetch a handle once and keep it.** `Bun()`, `GORM()`, `Native()` and `Database()` return handles over one long-lived `*sql.DB`. You do not need to re-fetch them per request, and they stay valid for the life of the connection. +- **Never close a handle yourself.** Closing a `*bun.DB`, `*gorm.DB` or the `*sql.DB` closes the shared pool for everyone. Only `mgr.Close()` (at shutdown) should close it. After `Close`, the handles are dead. +- **Don't reconnect to recover from errors.** `database/sql` already discards bad connections and dials new ones. The manager does not close the pool on errors or failed health checks. `conn.Reconnect(ctx)` is for explicit operator use only (for example after rotating credentials): on PostgreSQL it retires pooled connections without closing the pool, so held handles keep working. Other databases close and reopen the pool, which invalidates handles you already hold. +- **Bring your own `*sql.DB`.** `dbmanager.NewConnectionFromDB(name, type, db)` wraps a pool you opened. The manager never closes it (`Close` only logs a warning); you own it and must close it. +- **Set deadlines on request contexts.** `query_timeout` is applied to PostgreSQL as `statement_timeout` (server side) and TCP timeouts detect dead sockets, but pass a context with a deadline to your queries so callers fail fast. +- **Pool tuning.** Keep `conn_max_idle_time` below the shortest idle timeout of any NAT, load balancer or pgbouncer between you and the database (typically 60-240s). SQLite `:memory:` is pinned to a single connection. +- **Health checks** run every `health_check_interval` (default 15s; a negative value disables them) and publish Prometheus metrics. `enable_auto_reconnect` is deprecated and ignored. + For documentation, see [pkg/dbmanager/README.md](pkg/dbmanager/README.md). #### Cache diff --git a/audit/pkg/dbmanager.audit.md b/audit/pkg/dbmanager.audit.md new file mode 100644 index 0000000..bb919b3 --- /dev/null +++ b/audit/pkg/dbmanager.audit.md @@ -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=` 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=` (`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). diff --git a/pkg/dbmanager/README.md b/pkg/dbmanager/README.md index 79a7f84..8fb484c 100644 --- a/pkg/dbmanager/README.md +++ b/pkg/dbmanager/README.md @@ -50,7 +50,6 @@ dbmanager: # Health checks health_check_interval: 30s - enable_auto_reconnect: true connections: # Primary PostgreSQL connection @@ -256,7 +255,7 @@ db, _ := mgr.GetDefaultDatabase() | `retry_delay` | duration | 1s | Initial retry delay | | `retry_max_delay` | duration | 10s | Maximum retry delay | | `health_check_interval` | duration | 30s | Interval between health checks | -| `enable_auto_reconnect` | bool | true | Auto-reconnect on health check failure | +| `enable_auto_reconnect` | bool | - | Deprecated and ignored: the manager never closes the pool to recover from errors | ### Connection Configuration @@ -451,7 +450,6 @@ db.NewSelect().Model(&User{}).Scan(ctx) 3. **Enable Health Checks**: Catch connection issues early ```yaml health_check_interval: 30s - enable_auto_reconnect: true ``` 4. **Use Appropriate ORM**: Choose based on your needs diff --git a/pkg/dbmanager/config.go b/pkg/dbmanager/config.go index e690827..2986708 100644 --- a/pkg/dbmanager/config.go +++ b/pkg/dbmanager/config.go @@ -2,6 +2,10 @@ package dbmanager import ( "fmt" + "net" + "net/url" + "strconv" + "strings" "time" "github.com/bitechdev/ResolveSpec/pkg/config" @@ -57,9 +61,15 @@ type ManagerConfig struct { RetryDelay time.Duration `mapstructure:"retry_delay"` RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"` - // Health checks + // Health checks. A zero HealthCheckInterval selects the default (15s); a + // negative value disables the background health checker. HealthCheckInterval time.Duration `mapstructure:"health_check_interval"` - EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"` + + // Deprecated: ignored. The manager never closes a pool to recover from an + // error because database/sql already replaces broken connections; closing + // it would invalidate every handle handed out. Use Connection.Reconnect for + // an explicit, handle-preserving refresh. + EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"` } // ConnectionConfig defines configuration for a single database connection @@ -103,6 +113,11 @@ type ConnectionConfig struct { ConnectTimeout time.Duration `mapstructure:"connect_timeout"` QueryTimeout time.Duration `mapstructure:"query_timeout"` + // Retry policy for the initial connect (inherited from the manager config) + RetryAttempts int `mapstructure:"retry_attempts"` + RetryDelay time.Duration `mapstructure:"retry_delay"` + RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"` + // Features EnableTracing bool `mapstructure:"enable_tracing"` EnableMetrics bool `mapstructure:"enable_metrics"` @@ -129,7 +144,6 @@ func DefaultManagerConfig() ManagerConfig { RetryDelay: 1 * time.Second, RetryMaxDelay: 10 * time.Second, HealthCheckInterval: 15 * time.Second, - EnableAutoReconnect: true, } } @@ -161,11 +175,6 @@ func (c *ManagerConfig) ApplyDefaults() { if c.HealthCheckInterval == 0 { c.HealthCheckInterval = defaults.HealthCheckInterval } - // EnableAutoReconnect defaults to true - apply if not explicitly set - // Since this is a boolean, we apply the default unconditionally when it's false - if !c.EnableAutoReconnect { - c.EnableAutoReconnect = defaults.EnableAutoReconnect - } } // Validate validates the manager configuration @@ -222,9 +231,18 @@ func (cc *ConnectionConfig) ApplyDefaults(global *ManagerConfig) { } if cc.QueryTimeout == 0 { cc.QueryTimeout = 2 * time.Minute // Default to 2 minutes - } else if cc.QueryTimeout < 2*time.Minute { - // Enforce minimum of 2 minutes - cc.QueryTimeout = 2 * time.Minute + } + + if global != nil { + if cc.RetryAttempts == 0 { + cc.RetryAttempts = global.RetryAttempts + } + if cc.RetryDelay == 0 { + cc.RetryDelay = global.RetryDelay + } + if cc.RetryMaxDelay == 0 { + cc.RetryMaxDelay = global.RetryMaxDelay + } } // Default ORM @@ -314,108 +332,122 @@ func (cc *ConnectionConfig) BuildDSN() (string, error) { } } +// buildPostgresDSN builds a postgres:// URL so credentials and other values are +// escaped rather than spliced into a key=value string. statement_timeout is +// applied by the provider as a runtime parameter. func (cc *ConnectionConfig) buildPostgresDSN() string { - dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s", - cc.Host, cc.Port, cc.User, cc.Password, cc.Database) - + q := url.Values{} if cc.SSLMode != "" { - dsn += fmt.Sprintf(" sslmode=%s", cc.SSLMode) + q.Set("sslmode", cc.SSLMode) } else { - dsn += " sslmode=disable" + // prefer: use TLS when the server offers it, without failing on + // servers that do not. + q.Set("sslmode", "prefer") } - if cc.Schema != "" { - dsn += fmt.Sprintf(" search_path=%s", cc.Schema) + q.Set("search_path", cc.Schema) } - // Add statement_timeout for query execution timeout (in milliseconds) - if cc.QueryTimeout > 0 { - timeoutMs := int(cc.QueryTimeout.Milliseconds()) - dsn += fmt.Sprintf(" statement_timeout=%d", timeoutMs) + u := url.URL{ + Scheme: "postgres", + Host: hostPort(cc.Host, cc.Port), + Path: "/" + cc.Database, + RawQuery: q.Encode(), } - - return dsn + if cc.User != "" || cc.Password != "" { + u.User = url.UserPassword(cc.User, cc.Password) + } + return u.String() } +func hostPort(host string, port int) string { + if port == 0 { + return host + } + // JoinHostPort brackets IPv6 literals. + return net.JoinHostPort(host, strconv.Itoa(port)) +} + +// buildSQLiteDSN puts per-connection settings in the DSN as _pragma parameters +// so every pooled connection gets them, not just the one that ran an Exec. func (cc *ConnectionConfig) buildSQLiteDSN() string { filepath := cc.FilePath if filepath == "" { filepath = ":memory:" } - // Add query parameters for timeouts - // Note: SQLite driver supports _timeout parameter (in milliseconds) + var pragmas []string if cc.QueryTimeout > 0 { - timeoutMs := int(cc.QueryTimeout.Milliseconds()) - filepath += fmt.Sprintf("?_timeout=%d", timeoutMs) + pragmas = append(pragmas, fmt.Sprintf("busy_timeout(%d)", cc.QueryTimeout.Milliseconds())) + } + if filepath != ":memory:" { + pragmas = append(pragmas, "journal_mode(WAL)") + } + if len(pragmas) == 0 { + return filepath } - return filepath + q := url.Values{} + for _, p := range pragmas { + q.Add("_pragma", p) + } + sep := "?" + if strings.Contains(filepath, "?") { + sep = "&" + } + return filepath + sep + q.Encode() } func (cc *ConnectionConfig) buildMSSQLDSN() string { // Format: sqlserver://username:password@host:port?database=dbname - dsn := fmt.Sprintf("sqlserver://%s:%s@%s:%d?database=%s", - cc.User, cc.Password, cc.Host, cc.Port, cc.Database) - + q := url.Values{} + q.Set("database", cc.Database) if cc.Schema != "" { - dsn += fmt.Sprintf("&schema=%s", cc.Schema) + q.Set("schema", cc.Schema) } - - // Add connection timeout (in seconds) if cc.ConnectTimeout > 0 { - timeoutSec := int(cc.ConnectTimeout.Seconds()) - dsn += fmt.Sprintf("&connection timeout=%d", timeoutSec) + sec := strconv.Itoa(int(cc.ConnectTimeout.Seconds())) + q.Set("connection timeout", sec) + q.Set("dial timeout", sec) } - - // Add dial timeout for TCP connection (in seconds) - if cc.ConnectTimeout > 0 { - dialTimeoutSec := int(cc.ConnectTimeout.Seconds()) - dsn += fmt.Sprintf("&dial timeout=%d", dialTimeoutSec) - } - - // Add read timeout (in seconds) - enforces timeout for reading data if cc.QueryTimeout > 0 { - readTimeoutSec := int(cc.QueryTimeout.Seconds()) - dsn += fmt.Sprintf("&read timeout=%d", readTimeoutSec) + q.Set("read timeout", strconv.Itoa(int(cc.QueryTimeout.Seconds()))) } - return dsn + u := url.URL{ + Scheme: "sqlserver", + Host: hostPort(cc.Host, cc.Port), + RawQuery: q.Encode(), + } + if cc.User != "" || cc.Password != "" { + u.User = url.UserPassword(cc.User, cc.Password) + } + return u.String() } func (cc *ConnectionConfig) buildMongoDSN() string { // Format: mongodb://username:password@host:port/database?authSource=admin - var dsn string - - if cc.User != "" && cc.Password != "" { - dsn = fmt.Sprintf("mongodb://%s:%s@%s:%d/%s", - cc.User, cc.Password, cc.Host, cc.Port, cc.Database) - } else { - dsn = fmt.Sprintf("mongodb://%s:%d/%s", cc.Host, cc.Port, cc.Database) - } - - params := "" + q := url.Values{} if cc.AuthSource != "" { - params += fmt.Sprintf("authSource=%s", cc.AuthSource) + q.Set("authSource", cc.AuthSource) } if cc.ReplicaSet != "" { - if params != "" { - params += "&" - } - params += fmt.Sprintf("replicaSet=%s", cc.ReplicaSet) + q.Set("replicaSet", cc.ReplicaSet) } if cc.ReadPreference != "" { - if params != "" { - params += "&" - } - params += fmt.Sprintf("readPreference=%s", cc.ReadPreference) + q.Set("readPreference", cc.ReadPreference) } - if params != "" { - dsn += "?" + params + u := url.URL{ + Scheme: "mongodb", + Host: hostPort(cc.Host, cc.Port), + Path: "/" + cc.Database, + RawQuery: q.Encode(), } - - return dsn + if cc.User != "" && cc.Password != "" { + u.User = url.UserPassword(cc.User, cc.Password) + } + return u.String() } // FromConfig converts config.DBManagerConfig to internal ManagerConfig @@ -487,3 +519,6 @@ func (cc *ConnectionConfig) GetConnMaxIdleTime() *time.Duration { return cc.Conn func (cc *ConnectionConfig) GetQueryTimeout() time.Duration { return cc.QueryTimeout } func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics } func (cc *ConnectionConfig) GetReadPreference() string { return cc.ReadPreference } +func (cc *ConnectionConfig) GetRetryAttempts() int { return cc.RetryAttempts } +func (cc *ConnectionConfig) GetRetryDelay() time.Duration { return cc.RetryDelay } +func (cc *ConnectionConfig) GetRetryMaxDelay() time.Duration { return cc.RetryMaxDelay } diff --git a/pkg/dbmanager/config_dsn_test.go b/pkg/dbmanager/config_dsn_test.go new file mode 100644 index 0000000..f868048 --- /dev/null +++ b/pkg/dbmanager/config_dsn_test.go @@ -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) + } + } +} diff --git a/pkg/dbmanager/connection.go b/pkg/dbmanager/connection.go index 27ef03f..27e61d1 100644 --- a/pkg/dbmanager/connection.go +++ b/pkg/dbmanager/connection.go @@ -3,6 +3,7 @@ package dbmanager import ( "context" "database/sql" + "errors" "fmt" "sync" "time" @@ -14,6 +15,7 @@ import ( "github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" + "github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers" ) // Connection represents a single named database connection @@ -82,6 +84,9 @@ type sqlConnection struct { // State connected bool mu sync.RWMutex + // lifecycleMu serialises Connect/Close/Reconnect against health-check pings. + // Lock order: lifecycleMu before mu. + lifecycleMu sync.RWMutex // Health check lastHealthCheck time.Time @@ -110,9 +115,16 @@ func (c *sqlConnection) Type() DatabaseType { // Connect establishes the database connection func (c *sqlConnection) Connect(ctx context.Context) error { + c.lifecycleMu.Lock() + defer c.lifecycleMu.Unlock() c.mu.Lock() defer c.mu.Unlock() + return c.connectLocked(ctx) +} + +// connectLocked requires lifecycleMu and mu held for writing. +func (c *sqlConnection) connectLocked(ctx context.Context) error { if c.connected { return ErrAlreadyConnected } @@ -127,17 +139,29 @@ func (c *sqlConnection) Connect(ctx context.Context) error { // Close closes the database connection and all ORM instances func (c *sqlConnection) Close() error { + c.lifecycleMu.Lock() + defer c.lifecycleMu.Unlock() c.mu.Lock() defer c.mu.Unlock() + return c.closeLocked() +} + +// closeLocked requires lifecycleMu and mu held for writing. The connection is +// always marked disconnected and its cached handles dropped, even when closing +// fails, so accessors never hand out handles over a half-closed pool. +func (c *sqlConnection) closeLocked() error { if !c.connected { return nil } - // Close Bun if initialized - if c.bunDB != nil { + var errs []error + + // Close Bun if initialized. bun.DB.Close closes the underlying *sql.DB, so + // skip it when the pool belongs to the caller. + if o, ok := c.provider.(interface{ OwnsDB() bool }); c.bunDB != nil && (!ok || o.OwnsDB()) { if err := c.bunDB.Close(); err != nil { - return NewConnectionError(c.name, "close bun", err) + errs = append(errs, NewConnectionError(c.name, "close bun", err)) } } @@ -145,7 +169,7 @@ func (c *sqlConnection) Close() error { // Close the provider (which closes the underlying sql.DB) if err := c.provider.Close(); err != nil { - return NewConnectionError(c.name, "close", err) + errs = append(errs, NewConnectionError(c.name, "close", err)) } c.connected = false @@ -156,39 +180,75 @@ func (c *sqlConnection) Close() error { c.gormAdapter = nil c.nativeAdapter = nil - return nil + return errors.Join(errs...) } -// HealthCheck verifies the connection is alive +// HealthCheck verifies the connection is alive. The network ping runs without +// holding mu, so handle accessors are never blocked behind a slow ping. func (c *sqlConnection) HealthCheck(ctx context.Context) error { if c == nil { return fmt.Errorf("connection is nil") } - c.mu.Lock() - defer c.mu.Unlock() - c.lastHealthCheck = time.Now() + // lifecycleMu (read) keeps Close/Reconnect from tearing the provider down + // mid-ping without blocking the accessors that only need mu. + c.lifecycleMu.RLock() + defer c.lifecycleMu.RUnlock() - if !c.connected { - c.healthCheckStatus = "disconnected" + c.mu.RLock() + connected := c.connected + provider := c.provider + c.mu.RUnlock() + + if !connected { + c.setHealth("disconnected") return ErrConnectionClosed } - if err := c.provider.HealthCheck(ctx); err != nil { - c.healthCheckStatus = "unhealthy: " + err.Error() + if err := provider.HealthCheck(ctx); err != nil { + c.setHealth("unhealthy: " + err.Error()) return NewConnectionError(c.name, "health check", err) } - c.healthCheckStatus = "healthy" + c.setHealth("healthy") return nil } -// Reconnect closes and re-establishes the connection -func (c *sqlConnection) Reconnect(ctx context.Context) error { - if err := c.Close(); err != nil { +func (c *sqlConnection) setHealth(status string) { + c.mu.Lock() + c.lastHealthCheck = time.Now() + c.healthCheckStatus = status + c.mu.Unlock() +} + +// Reconnect refreshes the connection as a single critical section. +// +// Providers that support it (PostgreSQL) retire their pooled connections and +// dial fresh ones without closing the *sql.DB, so handles handed out earlier +// keep working. Other providers fall back to Close+Connect, which invalidates +// earlier handles; that is meant for explicit operator use only, since +// *sql.DB already replaces broken connections by itself. +func (c *sqlConnection) Reconnect(ctx context.Context) (err error) { + c.lifecycleMu.Lock() + defer c.lifecycleMu.Unlock() + c.mu.Lock() + defer c.mu.Unlock() + + defer func() { RecordReconnectAttempt(c.name, c.dbType, err == nil) }() + + if c.connected { + if r, ok := c.provider.(providers.Refresher); ok { + if err := r.Refresh(ctx); err != nil { + return NewConnectionError(c.name, "reconnect", err) + } + return nil + } + } + + if err := c.closeLocked(); err != nil { return err } - return c.Connect(ctx) + return c.connectLocked(ctx) } // Native returns the native *sql.DB connection @@ -250,6 +310,10 @@ func (c *sqlConnection) Bun() (*bun.DB, error) { return c.bunDB, nil } + if !c.connected { + return nil, ErrConnectionClosed + } + // Get native connection first native, err := c.provider.GetNative() if err != nil { @@ -283,6 +347,10 @@ func (c *sqlConnection) GORM() (*gorm.DB, error) { return c.gormDB, nil } + if !c.connected { + return nil, ErrConnectionClosed + } + // Get native connection first native, err := c.provider.GetNative() if err != nil { @@ -359,39 +427,18 @@ func (c *sqlConnection) Stats() *ConnectionStats { return stats } -func (c *sqlConnection) reconnectForAdapter() error { - timeout := c.config.ConnectTimeout - if timeout <= 0 { - timeout = 10 * time.Second - } - - ctx, cancel := context.WithTimeout(context.Background(), timeout) - defer cancel() - - return c.Reconnect(ctx) -} - +// The adapter factories only re-fetch the current handle. They must not close +// the shared pool: *sql.DB discards bad connections on its own, and closing it +// here would break every other holder of the pool. func (c *sqlConnection) reopenNativeForAdapter() (*sql.DB, error) { - if err := c.reconnectForAdapter(); err != nil { - return nil, err - } - return c.Native() } func (c *sqlConnection) reopenBunForAdapter() (*bun.DB, error) { - if err := c.reconnectForAdapter(); err != nil { - return nil, err - } - return c.Bun() } func (c *sqlConnection) reopenGORMForAdapter() (*gorm.DB, error) { - if err := c.reconnectForAdapter(); err != nil { - return nil, err - } - return c.GORM() } @@ -512,15 +559,8 @@ func (c *sqlConnection) getNativeAdapter() (common.Database, error) { // Create a native adapter based on database type switch c.dbType { - case DatabaseTypePostgreSQL: - c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)). - WithDBFactory(c.reopenNativeForAdapter). - SetMetricsEnabled(c.config.EnableMetrics) - case DatabaseTypeSQLite: - c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)). - WithDBFactory(c.reopenNativeForAdapter). - SetMetricsEnabled(c.config.EnableMetrics) - case DatabaseTypeMSSQL: + case DatabaseTypePostgreSQL, DatabaseTypeSQLite, DatabaseTypeMSSQL: + // The adapter takes the driver name so it can adjust its dialect. c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)). WithDBFactory(c.reopenNativeForAdapter). SetMetricsEnabled(c.config.EnableMetrics) @@ -572,8 +612,9 @@ type mongoConnection struct { client *mongo.Client // State - connected bool - mu sync.RWMutex + connected bool + mu sync.RWMutex + lifecycleMu sync.RWMutex // see sqlConnection.lifecycleMu // Health check lastHealthCheck time.Time @@ -601,9 +642,16 @@ func (c *mongoConnection) Type() DatabaseType { // Connect establishes the MongoDB connection func (c *mongoConnection) Connect(ctx context.Context) error { + c.lifecycleMu.Lock() + defer c.lifecycleMu.Unlock() c.mu.Lock() defer c.mu.Unlock() + return c.connectLocked(ctx) +} + +// connectLocked requires lifecycleMu and mu held for writing. +func (c *mongoConnection) connectLocked(ctx context.Context) error { if c.connected { return ErrAlreadyConnected } @@ -615,6 +663,7 @@ func (c *mongoConnection) Connect(ctx context.Context) error { // Get the mongo client client, err := c.provider.GetMongo() if err != nil { + _ = c.provider.Close() return NewConnectionError(c.name, "get mongo client", err) } @@ -625,49 +674,75 @@ func (c *mongoConnection) Connect(ctx context.Context) error { // Close closes the MongoDB connection func (c *mongoConnection) Close() error { + c.lifecycleMu.Lock() + defer c.lifecycleMu.Unlock() c.mu.Lock() defer c.mu.Unlock() + return c.closeLocked() +} + +// closeLocked requires lifecycleMu and mu held for writing. The connection is +// marked disconnected even when the provider fails to close cleanly. +func (c *mongoConnection) closeLocked() error { if !c.connected { return nil } - if err := c.provider.Close(); err != nil { - return NewConnectionError(c.name, "close", err) - } + err := c.provider.Close() c.connected = false c.client = nil + if err != nil { + return NewConnectionError(c.name, "close", err) + } return nil } -// HealthCheck verifies the MongoDB connection is alive +// HealthCheck verifies the MongoDB connection is alive. The ping runs without +// holding mu so handle accessors are never blocked behind it. func (c *mongoConnection) HealthCheck(ctx context.Context) error { - c.mu.Lock() - defer c.mu.Unlock() + c.lifecycleMu.RLock() + defer c.lifecycleMu.RUnlock() - c.lastHealthCheck = time.Now() + c.mu.RLock() + connected := c.connected + c.mu.RUnlock() - if !c.connected { - c.healthCheckStatus = "disconnected" + if !connected { + c.setHealth("disconnected") return ErrConnectionClosed } if err := c.provider.HealthCheck(ctx); err != nil { - c.healthCheckStatus = "unhealthy: " + err.Error() + c.setHealth("unhealthy: " + err.Error()) return NewConnectionError(c.name, "health check", err) } - c.healthCheckStatus = "healthy" + c.setHealth("healthy") return nil } -// Reconnect closes and re-establishes the MongoDB connection -func (c *mongoConnection) Reconnect(ctx context.Context) error { - if err := c.Close(); err != nil { +func (c *mongoConnection) setHealth(status string) { + c.mu.Lock() + c.lastHealthCheck = time.Now() + c.healthCheckStatus = status + c.mu.Unlock() +} + +// Reconnect closes and re-establishes the MongoDB connection atomically. +func (c *mongoConnection) Reconnect(ctx context.Context) (err error) { + c.lifecycleMu.Lock() + defer c.lifecycleMu.Unlock() + c.mu.Lock() + defer c.mu.Unlock() + + defer func() { RecordReconnectAttempt(c.name, DatabaseTypeMongoDB, err == nil) }() + + if err := c.closeLocked(); err != nil { return err } - return c.Connect(ctx) + return c.connectLocked(ctx) } // MongoDB returns the MongoDB client diff --git a/pkg/dbmanager/factory_test.go b/pkg/dbmanager/factory_test.go index 1e71c75..38c0312 100644 --- a/pkg/dbmanager/factory_test.go +++ b/pkg/dbmanager/factory_test.go @@ -4,13 +4,8 @@ import ( "context" "database/sql" "testing" - "time" _ "github.com/mattn/go-sqlite3" - "gorm.io/gorm" - - "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" - "github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers" ) func TestNewConnectionFromDB(t *testing.T) { @@ -213,157 +208,3 @@ func TestNewConnectionFromDB_PostgreSQL(t *testing.T) { t.Errorf("Expected type DatabaseTypePostgreSQL, got '%s'", conn.Type()) } } - -func TestDatabaseNativeAdapterReconnectFactory(t *testing.T) { - conn := newSQLConnection("test-native", DatabaseTypeSQLite, ConnectionConfig{ - Name: "test-native", - Type: DatabaseTypeSQLite, - FilePath: ":memory:", - DefaultORM: string(ORMTypeNative), - ConnectTimeout: 2 * time.Second, - }, providers.NewSQLiteProvider()) - - ctx := context.Background() - if err := conn.Connect(ctx); err != nil { - t.Fatalf("Failed to connect: %v", err) - } - defer conn.Close() - - db, err := conn.Database() - if err != nil { - t.Fatalf("Failed to get database adapter: %v", err) - } - - adapter, ok := db.(*database.PgSQLAdapter) - if !ok { - t.Fatalf("Expected PgSQLAdapter, got %T", db) - } - - underlyingBefore, ok := adapter.GetUnderlyingDB().(*sql.DB) - if !ok { - t.Fatalf("Expected underlying *sql.DB, got %T", adapter.GetUnderlyingDB()) - } - - if err := underlyingBefore.Close(); err != nil { - t.Fatalf("Failed to close underlying database: %v", err) - } - - if _, err := db.Exec(ctx, "SELECT 1"); err != nil { - t.Fatalf("Expected native adapter to reconnect, got error: %v", err) - } - - underlyingAfter, ok := adapter.GetUnderlyingDB().(*sql.DB) - if !ok { - t.Fatalf("Expected reconnected *sql.DB, got %T", adapter.GetUnderlyingDB()) - } - - if underlyingAfter == underlyingBefore { - t.Fatal("Expected adapter to swap to a fresh *sql.DB after reconnect") - } -} - -func TestDatabaseBunAdapterReconnectFactory(t *testing.T) { - conn := newSQLConnection("test-bun", DatabaseTypeSQLite, ConnectionConfig{ - Name: "test-bun", - Type: DatabaseTypeSQLite, - FilePath: ":memory:", - DefaultORM: string(ORMTypeBun), - ConnectTimeout: 2 * time.Second, - }, providers.NewSQLiteProvider()) - - ctx := context.Background() - if err := conn.Connect(ctx); err != nil { - t.Fatalf("Failed to connect: %v", err) - } - defer conn.Close() - - db, err := conn.Database() - if err != nil { - t.Fatalf("Failed to get database adapter: %v", err) - } - - adapter, ok := db.(*database.BunAdapter) - if !ok { - t.Fatalf("Expected BunAdapter, got %T", db) - } - - underlyingBefore, ok := adapter.GetUnderlyingDB().(interface{ Close() error }) - if !ok { - t.Fatalf("Expected underlying Bun DB with Close method, got %T", adapter.GetUnderlyingDB()) - } - - if err := underlyingBefore.Close(); err != nil { - t.Fatalf("Failed to close underlying Bun database: %v", err) - } - - if _, err := db.Exec(ctx, "SELECT 1"); err != nil { - t.Fatalf("Expected Bun adapter to reconnect, got error: %v", err) - } - - underlyingAfter := adapter.GetUnderlyingDB() - if underlyingAfter == underlyingBefore { - t.Fatal("Expected adapter to swap to a fresh Bun DB after reconnect") - } -} - -func TestDatabaseGormAdapterReconnectFactory(t *testing.T) { - conn := newSQLConnection("test-gorm", DatabaseTypeSQLite, ConnectionConfig{ - Name: "test-gorm", - Type: DatabaseTypeSQLite, - FilePath: ":memory:", - DefaultORM: string(ORMTypeGORM), - ConnectTimeout: 2 * time.Second, - }, providers.NewSQLiteProvider()) - - ctx := context.Background() - if err := conn.Connect(ctx); err != nil { - t.Fatalf("Failed to connect: %v", err) - } - defer conn.Close() - - db, err := conn.Database() - if err != nil { - t.Fatalf("Failed to get database adapter: %v", err) - } - - adapter, ok := db.(*database.GormAdapter) - if !ok { - t.Fatalf("Expected GormAdapter, got %T", db) - } - - gormBefore, ok := adapter.GetUnderlyingDB().(*gorm.DB) - if !ok { - t.Fatalf("Expected underlying *gorm.DB, got %T", adapter.GetUnderlyingDB()) - } - - sqlBefore, err := gormBefore.DB() - if err != nil { - t.Fatalf("Failed to get underlying *sql.DB: %v", err) - } - - if err := sqlBefore.Close(); err != nil { - t.Fatalf("Failed to close underlying database: %v", err) - } - - count, err := db.NewSelect().Table("sqlite_master").Count(ctx) - if err != nil { - t.Fatalf("Expected GORM query builder to reconnect, got error: %v", err) - } - if count < 0 { - t.Fatalf("Expected non-negative count, got %d", count) - } - - gormAfter, ok := adapter.GetUnderlyingDB().(*gorm.DB) - if !ok { - t.Fatalf("Expected reconnected *gorm.DB, got %T", adapter.GetUnderlyingDB()) - } - - sqlAfter, err := gormAfter.DB() - if err != nil { - t.Fatalf("Failed to get reconnected *sql.DB: %v", err) - } - - if sqlAfter == sqlBefore { - t.Fatal("Expected GORM adapter to use a fresh *sql.DB after reconnect") - } -} diff --git a/pkg/dbmanager/lifecycle_test.go b/pkg/dbmanager/lifecycle_test.go new file mode 100644 index 0000000..0d5391b --- /dev/null +++ b/pkg/dbmanager/lifecycle_test.go @@ -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) + } +} diff --git a/pkg/dbmanager/manager.go b/pkg/dbmanager/manager.go index 1ae4dfe..46acec7 100644 --- a/pkg/dbmanager/manager.go +++ b/pkg/dbmanager/manager.go @@ -2,9 +2,7 @@ package dbmanager import ( "context" - "errors" "fmt" - "strings" "sync" "time" @@ -49,6 +47,7 @@ type connectionManager struct { // Background health check healthTicker *time.Ticker stopChan chan struct{} + healthMu sync.Mutex // guards healthTicker and stopChan wg sync.WaitGroup } @@ -100,7 +99,9 @@ func ResetInstance() { defer instanceMu.Unlock() if instance != nil { - _ = instance.Close() + if err := instance.Close(); err != nil { + logger.Error("Failed to close manager during reset: %v", err) + } } instance = nil } @@ -116,7 +117,6 @@ func NewManager(cfg ManagerConfig) (Manager, error) { mgr := &connectionManager{ connections: make(map[string]Connection), config: cfg, - stopChan: make(chan struct{}), } return mgr, nil @@ -195,11 +195,26 @@ func (m *connectionManager) SetDefaultDatabase(name string) error { // Connect establishes all configured database connections func (m *connectionManager) Connect(ctx context.Context) error { - m.mu.Lock() - defer m.mu.Unlock() - - // Create connections from configuration + // Dial outside m.mu so a slow connect never blocks Get/Stats/HealthCheck. + m.mu.RLock() + names := make([]string, 0, len(m.config.Connections)) for name := range m.config.Connections { + if _, exists := m.connections[name]; !exists { + names = append(names, name) + } + } + m.mu.RUnlock() + + opened := make(map[string]Connection, len(names)) + closeOpened := func() { + for name, conn := range opened { + if err := conn.Close(); err != nil { + logger.Error("Failed to close connection after failed Connect: name=%s, error=%v", name, err) + } + } + } + + for _, name := range names { // Get a copy of the connection config connCfg := m.config.Connections[name] // Apply global defaults to connection config @@ -209,25 +224,39 @@ func (m *connectionManager) Connect(ctx context.Context) error { // Create connection using factory conn, err := createConnection(connCfg) if err != nil { + closeOpened() return fmt.Errorf("failed to create connection '%s': %w", name, err) } // Connect if err := conn.Connect(ctx); err != nil { + closeOpened() return fmt.Errorf("failed to connect '%s': %w", name, err) } - m.connections[name] = conn + opened[name] = conn logger.Info("Database connection established: name=%s, type=%s", name, connCfg.Type) } + m.mu.Lock() + for name, conn := range opened { + if _, exists := m.connections[name]; exists { + // Lost a race with a concurrent Connect; drop our duplicate. + _ = conn.Close() + continue + } + m.connections[name] = conn + } + total := len(m.connections) + m.mu.Unlock() + // Always start background health checks if m.config.HealthCheckInterval > 0 { m.startHealthChecker() logger.Info("Background health checker started: interval=%v", m.config.HealthCheckInterval) } - logger.Info("Database manager initialized: connections=%d", len(m.connections)) + logger.Info("Database manager initialized: connections=%d", total) return nil } @@ -246,7 +275,7 @@ func (m *connectionManager) Close() error { for name, conn := range m.connections { if err := conn.Close(); err != nil { errors = append(errors, fmt.Errorf("failed to close connection '%s': %w", name, err)) - logger.Error("Failed to close connection", "name", name, "error", err) + logger.Error("Failed to close connection: name=%s, error=%v", name, err) } else { logger.Info("Connection closed: name=%s", name) } @@ -311,11 +340,17 @@ func (m *connectionManager) Stats() *ManagerStats { // startHealthChecker starts background health checking func (m *connectionManager) startHealthChecker() { + m.healthMu.Lock() + defer m.healthMu.Unlock() + if m.healthTicker != nil { return // Already running } - m.healthTicker = time.NewTicker(m.config.HealthCheckInterval) + ticker := time.NewTicker(m.config.HealthCheckInterval) + stop := make(chan struct{}) + m.healthTicker = ticker + m.stopChan = stop m.wg.Add(1) go func() { @@ -324,9 +359,9 @@ func (m *connectionManager) startHealthChecker() { for { select { - case <-m.healthTicker.C: + case <-ticker.C: m.performHealthCheck() - case <-m.stopChan: + case <-stop: logger.Info("Health checker stopped") return } @@ -334,14 +369,19 @@ func (m *connectionManager) startHealthChecker() { }() } -// stopHealthChecker stops background health checking +// stopHealthChecker stops background health checking. Safe to call repeatedly. func (m *connectionManager) stopHealthChecker() { - if m.healthTicker != nil { - m.healthTicker.Stop() - close(m.stopChan) - m.wg.Wait() - m.healthTicker = nil + m.healthMu.Lock() + defer m.healthMu.Unlock() + + if m.healthTicker == nil { + return } + m.healthTicker.Stop() + close(m.stopChan) + m.wg.Wait() + m.healthTicker = nil + m.stopChan = nil } // performHealthCheck performs a health check on all connections @@ -362,40 +402,14 @@ func (m *connectionManager) performHealthCheck() { } m.mu.RUnlock() + defer m.PublishMetrics() + for _, item := range connections { if err := item.conn.HealthCheck(ctx); err != nil { - logger.Warn("Health check failed", - "connection", item.name, - "error", err) - - // Only reconnect when the client handle itself is closed/disconnected. - // For transient database restarts or network blips, *sql.DB can recover - // on its own; forcing Close()+Connect() here invalidates any cached ORM - // wrappers and callers that still hold the old handle. - if m.config.EnableAutoReconnect && shouldReconnectAfterHealthCheck(err) { - logger.Info("Attempting reconnection: connection=%s", item.name) - if err := item.conn.Reconnect(ctx); err != nil { - logger.Error("Reconnection failed", - "connection", item.name, - "error", err) - } else { - logger.Info("Reconnection successful: connection=%s", item.name) - } - } else if m.config.EnableAutoReconnect { - logger.Info("Skipping reconnect for transient health check failure: connection=%s", item.name) - } + // Do not reconnect here: *sql.DB discards bad connections and dials + // new ones by itself, while Reconnect closes the pool and breaks + // every handle already handed out. Reconnect is operator-only. + logger.Warn("Health check failed: connection=%s, error=%v", item.name, err) } } } - -func shouldReconnectAfterHealthCheck(err error) bool { - if err == nil { - return false - } - - if errors.Is(err, ErrConnectionClosed) { - return true - } - - return strings.Contains(err.Error(), "sql: database is closed") -} diff --git a/pkg/dbmanager/manager_test.go b/pkg/dbmanager/manager_test.go index 2f3d27f..096692d 100644 --- a/pkg/dbmanager/manager_test.go +++ b/pkg/dbmanager/manager_test.go @@ -21,19 +21,30 @@ type healthCheckStubConnection struct { reconnectCalls int } -func (c *healthCheckStubConnection) Name() string { return "stub" } -func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL } -func (c *healthCheckStubConnection) Bun() (*bun.DB, error) { return nil, fmt.Errorf("not implemented") } -func (c *healthCheckStubConnection) GORM() (*gorm.DB, error) { return nil, fmt.Errorf("not implemented") } -func (c *healthCheckStubConnection) Native() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") } -func (c *healthCheckStubConnection) DB() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") } -func (c *healthCheckStubConnection) Database() (common.Database, error) { return nil, fmt.Errorf("not implemented") } -func (c *healthCheckStubConnection) MongoDB() (*mongo.Client, error) { return nil, fmt.Errorf("not implemented") } -func (c *healthCheckStubConnection) Connect(ctx context.Context) error { return nil } -func (c *healthCheckStubConnection) Close() error { return nil } -func (c *healthCheckStubConnection) HealthCheck(ctx context.Context) error { return c.healthErr } -func (c *healthCheckStubConnection) Reconnect(ctx context.Context) error { c.reconnectCalls++; return nil } -func (c *healthCheckStubConnection) Stats() *ConnectionStats { return &ConnectionStats{} } +func (c *healthCheckStubConnection) Name() string { return "stub" } +func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL } +func (c *healthCheckStubConnection) Bun() (*bun.DB, error) { return nil, fmt.Errorf("not implemented") } +func (c *healthCheckStubConnection) GORM() (*gorm.DB, error) { + return nil, fmt.Errorf("not implemented") +} +func (c *healthCheckStubConnection) Native() (*sql.DB, error) { + return nil, fmt.Errorf("not implemented") +} +func (c *healthCheckStubConnection) DB() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") } +func (c *healthCheckStubConnection) Database() (common.Database, error) { + return nil, fmt.Errorf("not implemented") +} +func (c *healthCheckStubConnection) MongoDB() (*mongo.Client, error) { + return nil, fmt.Errorf("not implemented") +} +func (c *healthCheckStubConnection) Connect(ctx context.Context) error { return nil } +func (c *healthCheckStubConnection) Close() error { return nil } +func (c *healthCheckStubConnection) HealthCheck(ctx context.Context) error { return c.healthErr } +func (c *healthCheckStubConnection) Reconnect(ctx context.Context) error { + c.reconnectCalls++ + return nil +} +func (c *healthCheckStubConnection) Stats() *ConnectionStats { return &ConnectionStats{} } func TestBackgroundHealthChecker(t *testing.T) { // Create a SQLite in-memory database @@ -117,41 +128,21 @@ func TestDefaultHealthCheckInterval(t *testing.T) { t.Errorf("Expected default health check interval to be %v, got %v", expectedInterval, defaults.HealthCheckInterval) } - - if !defaults.EnableAutoReconnect { - t.Error("Expected EnableAutoReconnect to be true by default") - } } -func TestApplyDefaultsEnablesAutoReconnect(t *testing.T) { - // Create a config without setting EnableAutoReconnect - cfg := ManagerConfig{ - Connections: map[string]ConnectionConfig{ - "test": { - Name: "test", - Type: DatabaseTypeSQLite, - FilePath: ":memory:", - }, - }, - } - - // Verify it's false initially (Go's zero value for bool) - if cfg.EnableAutoReconnect { - t.Error("Expected EnableAutoReconnect to be false before ApplyDefaults") - } - - // Apply defaults +func TestApplyDefaultsHealthCheckInterval(t *testing.T) { + cfg := ManagerConfig{} cfg.ApplyDefaults() - - // Verify it's now true - if !cfg.EnableAutoReconnect { - t.Error("Expected EnableAutoReconnect to be true after ApplyDefaults") - } - - // Verify health check interval is also set if cfg.HealthCheckInterval != 15*time.Second { t.Errorf("Expected health check interval to be 15s, got %v", cfg.HealthCheckInterval) } + + // A negative interval disables the background checker and is preserved. + cfg = ManagerConfig{HealthCheckInterval: -1} + cfg.ApplyDefaults() + if cfg.HealthCheckInterval >= 0 { + t.Errorf("Expected negative interval to be preserved, got %v", cfg.HealthCheckInterval) + } } func TestManagerHealthCheck(t *testing.T) { @@ -270,7 +261,7 @@ func TestPerformHealthCheckSkipsReconnectForTransientFailures(t *testing.T) { } } -func TestPerformHealthCheckReconnectsClosedConnections(t *testing.T) { +func TestPerformHealthCheckNeverReconnects(t *testing.T) { conn := &healthCheckStubConnection{ healthErr: NewConnectionError("primary", "health check", fmt.Errorf("sql: database is closed")), } @@ -284,7 +275,7 @@ func TestPerformHealthCheckReconnectsClosedConnections(t *testing.T) { mgr.performHealthCheck() - if conn.reconnectCalls != 1 { - t.Fatalf("expected reconnect attempt for closed database handle, got %d", conn.reconnectCalls) + if conn.reconnectCalls != 0 { + t.Fatalf("health check must not close the shared pool via Reconnect, got %d", conn.reconnectCalls) } } diff --git a/pkg/dbmanager/metrics.go b/pkg/dbmanager/metrics.go index 3878891..b2960e2 100644 --- a/pkg/dbmanager/metrics.go +++ b/pkg/dbmanager/metrics.go @@ -1,6 +1,8 @@ package dbmanager import ( + "sync" + "github.com/prometheus/client_golang/prometheus" "github.com/prometheus/client_golang/prometheus/promauto" ) @@ -34,8 +36,8 @@ var ( ) // connectionWaitCount tracks how many times connections had to wait for availability - connectionWaitCount = promauto.NewGaugeVec( - prometheus.GaugeOpts{ + connectionWaitCount = promauto.NewCounterVec( + prometheus.CounterOpts{ Name: "dbmanager_connection_wait_count", Help: "Number of times connections had to wait for availability", }, @@ -43,8 +45,8 @@ var ( ) // connectionWaitDuration tracks total time connections spent waiting - connectionWaitDuration = promauto.NewGaugeVec( - prometheus.GaugeOpts{ + connectionWaitDuration = promauto.NewCounterVec( + prometheus.CounterOpts{ Name: "dbmanager_connection_wait_duration_seconds", Help: "Total time connections spent waiting for availability", }, @@ -61,8 +63,8 @@ var ( ) // connectionLifetimeClosed tracks connections closed due to max lifetime - connectionLifetimeClosed = promauto.NewGaugeVec( - prometheus.GaugeOpts{ + connectionLifetimeClosed = promauto.NewCounterVec( + prometheus.CounterOpts{ Name: "dbmanager_connection_lifetime_closed_total", Help: "Total connections closed due to exceeding max lifetime", }, @@ -70,8 +72,8 @@ var ( ) // connectionIdleClosed tracks connections closed due to max idle time - connectionIdleClosed = promauto.NewGaugeVec( - prometheus.GaugeOpts{ + connectionIdleClosed = promauto.NewCounterVec( + prometheus.CounterOpts{ Name: "dbmanager_connection_idle_closed_total", Help: "Total connections closed due to exceeding max idle time", }, @@ -114,13 +116,13 @@ func (m *connectionManager) PublishMetrics() { connectionPoolSize.WithLabelValues(name, string(connStats.Type), "idle").Set(float64(connStats.Idle)) connectionPoolSize.WithLabelValues(name, string(connStats.Type), "in_use").Set(float64(connStats.InUse)) - // Wait stats - connectionWaitCount.With(labels).Set(float64(connStats.WaitCount)) - connectionWaitDuration.With(labels).Set(connStats.WaitDuration.Seconds()) - - // Lifetime/idle closed - connectionLifetimeClosed.With(labels).Set(float64(connStats.MaxLifetimeClosed)) - connectionIdleClosed.With(labels).Set(float64(connStats.MaxIdleClosed)) + // sql.DBStats values are cumulative, so add only the growth since + // the last publish to keep these true counters. + prev := lastPublished.swap(name, connStats) + connectionWaitCount.With(labels).Add(float64(connStats.WaitCount - prev.WaitCount)) + connectionWaitDuration.With(labels).Add((connStats.WaitDuration - prev.WaitDuration).Seconds()) + connectionLifetimeClosed.With(labels).Add(float64(connStats.MaxLifetimeClosed - prev.MaxLifetimeClosed)) + connectionIdleClosed.With(labels).Add(float64(connStats.MaxIdleClosed - prev.MaxIdleClosed)) } } } @@ -134,3 +136,25 @@ func RecordReconnectAttempt(name string, dbType DatabaseType, success bool) { reconnectAttempts.WithLabelValues(name, string(dbType), result).Inc() } + +// publishedStats remembers the cumulative pool stats last exported per +// connection so counters can be advanced by the delta. +type publishedStats struct { + mu sync.Mutex + last map[string]ConnectionStats +} + +var lastPublished = &publishedStats{last: make(map[string]ConnectionStats)} + +// swap stores cur and returns the previous value. A counter reset (a new pool +// after Close+Connect) is treated as starting from zero. +func (p *publishedStats) swap(name string, cur *ConnectionStats) ConnectionStats { + p.mu.Lock() + defer p.mu.Unlock() + prev := p.last[name] + if cur.WaitCount < prev.WaitCount || cur.MaxIdleClosed < prev.MaxIdleClosed || cur.MaxLifetimeClosed < prev.MaxLifetimeClosed { + prev = ConnectionStats{} + } + p.last[name] = *cur + return prev +} diff --git a/pkg/dbmanager/pg_live_test.go b/pkg/dbmanager/pg_live_test.go new file mode 100644 index 0000000..550f1e0 --- /dev/null +++ b/pkg/dbmanager/pg_live_test.go @@ -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") + } +} diff --git a/pkg/dbmanager/providers/existing_db.go b/pkg/dbmanager/providers/existing_db.go index 9b56a75..38407b8 100644 --- a/pkg/dbmanager/providers/existing_db.go +++ b/pkg/dbmanager/providers/existing_db.go @@ -7,6 +7,8 @@ import ( "sync" "go.mongodb.org/mongo-driver/mongo" + + "github.com/bitechdev/ResolveSpec/pkg/logger" ) // ExistingDBProvider wraps an existing *sql.DB connection @@ -44,16 +46,27 @@ func (p *ExistingDBProvider) Connect(ctx context.Context, cfg ConnectionConfig) return nil } -// Close closes the underlying database connection -func (p *ExistingDBProvider) Close() error { - p.mu.Lock() - defer p.mu.Unlock() +// Refresh verifies the wrapped database is still reachable. The pool belongs to +// the caller and cannot be re-dialed here, so it is never closed to "reconnect". +func (p *ExistingDBProvider) Refresh(ctx context.Context) error { + p.mu.RLock() + defer p.mu.RUnlock() if p.db == nil { - return nil + return fmt.Errorf("database connection is nil") } + return p.db.PingContext(ctx) +} - return p.db.Close() +// OwnsDB reports whether Close releases the wrapped database. It never does: +// the *sql.DB was opened by the caller, who is responsible for closing it. +func (p *ExistingDBProvider) OwnsDB() bool { return false } + +// Close is a no-op for the wrapped database. The pool belongs to the caller, so +// closing it here would break the caller's other users of it. +func (p *ExistingDBProvider) Close() error { + logger.Warn("Not closing externally provided database: name=%s; the caller owns this *sql.DB and must close it", p.name) + return nil } // HealthCheck verifies the connection is alive diff --git a/pkg/dbmanager/providers/existing_db_test.go b/pkg/dbmanager/providers/existing_db_test.go index d00e998..b72daf6 100644 --- a/pkg/dbmanager/providers/existing_db_test.go +++ b/pkg/dbmanager/providers/existing_db_test.go @@ -164,7 +164,7 @@ func TestExistingDBProvider_Stats(t *testing.T) { } } -func TestExistingDBProvider_Close(t *testing.T) { +func TestExistingDBProvider_Close_LeavesDBOpen(t *testing.T) { db, err := sql.Open("sqlite3", ":memory:") if err != nil { t.Fatalf("Failed to open database: %v", err) @@ -177,10 +177,10 @@ func TestExistingDBProvider_Close(t *testing.T) { t.Errorf("Expected Close to succeed, got error: %v", err) } - // Verify the database is closed - err = db.Ping() - if err == nil { - t.Error("Expected database to be closed") + // The caller owns the database, so Close must leave it open + defer db.Close() + if err := db.Ping(); err != nil { + t.Errorf("Expected caller's database to stay open, got: %v", err) } } diff --git a/pkg/dbmanager/providers/mongodb.go b/pkg/dbmanager/providers/mongodb.go index e832870..99def9c 100644 --- a/pkg/dbmanager/providers/mongodb.go +++ b/pkg/dbmanager/providers/mongodb.go @@ -41,9 +41,10 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error clientOpts.SetMaxPoolSize(maxPoolSize) } - if cfg.GetMaxIdleConns() != nil { - minPoolSize := uint64(*cfg.GetMaxIdleConns()) - clientOpts.SetMinPoolSize(minPoolSize) + // MaxIdleConns is a ceiling on idle connections, not a pre-warmed minimum + // (MinPoolSize), so only the idle-time limit maps onto the Mongo pool. + if cfg.GetConnMaxIdleTime() != nil { + clientOpts.SetMaxConnIdleTime(*cfg.GetConnMaxIdleTime()) } // Set timeouts @@ -65,12 +66,11 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error var client *mongo.Client var lastErr error - retryAttempts := 3 - retryDelay := 1 * time.Second + retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg) for attempt := 0; attempt < retryAttempts; attempt++ { if attempt > 0 { - delay := calculateBackoff(attempt, retryDelay, 10*time.Second) + delay := calculateBackoff(attempt, retryDelay, retryMaxDelay) if cfg.GetEnableLogging() { logger.Info("Retrying MongoDB connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay) } @@ -87,7 +87,7 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error if err != nil { lastErr = err if cfg.GetEnableLogging() { - logger.Warn("Failed to connect to MongoDB", "error", err) + logger.Warn("Failed to connect to MongoDB: %v", err) } continue } @@ -101,7 +101,7 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error lastErr = err _ = client.Disconnect(ctx) if cfg.GetEnableLogging() { - logger.Warn("Failed to ping MongoDB", "error", err) + logger.Warn("Failed to ping MongoDB: %v", err) } continue } diff --git a/pkg/dbmanager/providers/mssql.go b/pkg/dbmanager/providers/mssql.go index bad58f9..0df8abb 100644 --- a/pkg/dbmanager/providers/mssql.go +++ b/pkg/dbmanager/providers/mssql.go @@ -35,12 +35,11 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error var db *sql.DB var lastErr error - retryAttempts := 3 // Default retry attempts - retryDelay := 1 * time.Second + retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg) for attempt := 0; attempt < retryAttempts; attempt++ { if attempt > 0 { - delay := calculateBackoff(attempt, retryDelay, 10*time.Second) + delay := calculateBackoff(attempt, retryDelay, retryMaxDelay) if cfg.GetEnableLogging() { logger.Info("Retrying MSSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay) } @@ -57,7 +56,7 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error if err != nil { lastErr = err if cfg.GetEnableLogging() { - logger.Warn("Failed to open MSSQL connection", "error", err) + logger.Warn("Failed to open MSSQL connection: %v", err) } continue } @@ -71,7 +70,7 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error lastErr = err db.Close() if cfg.GetEnableLogging() { - logger.Warn("Failed to ping MSSQL database", "error", err) + logger.Warn("Failed to ping MSSQL database: %v", err) } continue } diff --git a/pkg/dbmanager/providers/pgconnector.go b/pkg/dbmanager/providers/pgconnector.go new file mode 100644 index 0000000..057508a --- /dev/null +++ b/pkg/dbmanager/providers/pgconnector.go @@ -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 +} diff --git a/pkg/dbmanager/providers/pgconnector_test.go b/pkg/dbmanager/providers/pgconnector_test.go new file mode 100644 index 0000000..86b1ccb --- /dev/null +++ b/pkg/dbmanager/providers/pgconnector_test.go @@ -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") + } +} diff --git a/pkg/dbmanager/providers/postgres.go b/pkg/dbmanager/providers/postgres.go index 0391c23..f594919 100644 --- a/pkg/dbmanager/providers/postgres.go +++ b/pkg/dbmanager/providers/postgres.go @@ -3,12 +3,12 @@ package providers import ( "context" "database/sql" + "errors" "fmt" "math" "sync" "time" - _ "github.com/jackc/pgx/v5/stdlib" // PostgreSQL driver "go.mongodb.org/mongo-driver/mongo" "github.com/bitechdev/ResolveSpec/pkg/logger" @@ -16,10 +16,11 @@ import ( // PostgresProvider implements Provider for PostgreSQL databases type PostgresProvider struct { - db *sql.DB - config ConnectionConfig - listener *PostgresListener - mu sync.Mutex + db *sql.DB + connector *pgConnector + config ConnectionConfig + listener *PostgresListener + mu sync.Mutex } // NewPostgresProvider creates a new PostgreSQL provider @@ -29,22 +30,24 @@ func NewPostgresProvider() *PostgresProvider { // Connect establishes a PostgreSQL connection func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) error { - // Build DSN - dsn, err := cfg.BuildDSN() + connCfg, err := buildPGXConfig(cfg) if err != nil { - return fmt.Errorf("failed to build DSN: %w", err) + return err } + // The connector and *sql.DB are created once; the pool is never closed to + // recover from errors (see Refresh). + connector := newPGConnector(connCfg) + db := sql.OpenDB(connector) + // Connect with retry logic - var db *sql.DB var lastErr error + retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg) - retryAttempts := 3 // Default retry attempts - retryDelay := 1 * time.Second - + connected := false for attempt := 0; attempt < retryAttempts; attempt++ { if attempt > 0 { - delay := calculateBackoff(attempt, retryDelay, 10*time.Second) + delay := calculateBackoff(attempt, retryDelay, retryMaxDelay) if cfg.GetEnableLogging() { logger.Info("Retrying PostgreSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay) } @@ -52,20 +55,11 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er select { case <-time.After(delay): case <-ctx.Done(): + db.Close() return ctx.Err() } } - // Open database connection - db, err = sql.Open("pgx", dsn) - if err != nil { - lastErr = err - if cfg.GetEnableLogging() { - logger.Warn("Failed to open PostgreSQL connection", "error", err) - } - continue - } - // Test the connection with context timeout connectCtx, cancel := context.WithTimeout(ctx, cfg.GetConnectTimeout()) err = db.PingContext(connectCtx) @@ -73,18 +67,18 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er if err != nil { lastErr = err - db.Close() if cfg.GetEnableLogging() { - logger.Warn("Failed to ping PostgreSQL database", "error", err) + logger.Warn("Failed to ping PostgreSQL database: %v", err) } continue } - // Connection successful + connected = true break } - if err != nil { + if !connected { + db.Close() return fmt.Errorf("failed to connect after %d attempts: %w", retryAttempts, lastErr) } @@ -103,6 +97,7 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er } p.db = db + p.connector = connector p.config = cfg if cfg.GetEnableLogging() { @@ -112,34 +107,55 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er return nil } -// Close closes the PostgreSQL connection -func (p *PostgresProvider) Close() error { - // Close listener if it exists - p.mu.Lock() - if p.listener != nil { - if err := p.listener.Close(); err != nil { - p.mu.Unlock() - return fmt.Errorf("failed to close listener: %w", err) - } - p.listener = nil +// Refresh retires every pooled connection and dials fresh ones on demand, +// without closing the *sql.DB. Handles already handed out keep working: +// connections in use finish their current query and are then discarded. +func (p *PostgresProvider) Refresh(ctx context.Context) error { + if p.db == nil || p.connector == nil { + return fmt.Errorf("database connection is not initialized") } + + connCfg, err := buildPGXConfig(p.config) + if err != nil { + return err + } + p.connector.swap(connCfg) + + pingCtx, cancel := context.WithTimeout(ctx, p.config.GetConnectTimeout()) + defer cancel() + if err := p.db.PingContext(pingCtx); err != nil { + return fmt.Errorf("failed to ping after refresh: %w", err) + } + return nil +} + +// Close closes the PostgreSQL connection. A listener failure does not stop the +// pool from being closed. +func (p *PostgresProvider) Close() error { + var errs []error + + p.mu.Lock() + listener := p.listener + p.listener = nil p.mu.Unlock() - if p.db == nil { - return nil + if listener != nil { + if err := listener.Close(); err != nil { + errs = append(errs, fmt.Errorf("failed to close listener: %w", err)) + } } - err := p.db.Close() - if err != nil { - return fmt.Errorf("failed to close PostgreSQL connection: %w", err) + if p.db != nil { + if err := p.db.Close(); err != nil { + errs = append(errs, fmt.Errorf("failed to close PostgreSQL connection: %w", err)) + } else if p.config.GetEnableLogging() { + logger.Info("PostgreSQL connection closed: name=%s", p.config.GetName()) + } + p.db = nil + p.connector = nil } - if p.config.GetEnableLogging() { - logger.Info("PostgreSQL connection closed: name=%s", p.config.GetName()) - } - - p.db = nil - return nil + return errors.Join(errs...) } // HealthCheck verifies the PostgreSQL connection is alive diff --git a/pkg/dbmanager/providers/postgres_listener.go b/pkg/dbmanager/providers/postgres_listener.go index c417c73..ee9e48f 100644 --- a/pkg/dbmanager/providers/postgres_listener.go +++ b/pkg/dbmanager/providers/postgres_listener.go @@ -23,6 +23,10 @@ type PostgresListener struct { // Channel subscriptions channels map[string]NotificationHandler mu sync.RWMutex + // connMu serialises use of the single pgx.Conn: it is not safe for + // concurrent use, so the notification wait and LISTEN/UNLISTEN/NOTIFY take + // turns. Lock order: connMu before mu. + connMu sync.Mutex // Lifecycle management ctx context.Context @@ -30,6 +34,7 @@ type PostgresListener struct { closed bool closeMu sync.Mutex reconnectC chan struct{} + startOnce sync.Once // background goroutines start exactly once } // NewPostgresListener creates a new PostgreSQL listener @@ -44,76 +49,20 @@ func NewPostgresListener(cfg ConnectionConfig) *PostgresListener { } } -// Connect establishes a dedicated connection for listening +// Connect establishes a dedicated connection for listening and starts the +// background loops (once per listener). func (l *PostgresListener) Connect(ctx context.Context) error { - dsn, err := l.config.BuildDSN() + conn, err := l.dial(ctx) if err != nil { - return fmt.Errorf("failed to build DSN: %w", err) + return err } - // Parse connection config - connConfig, err := pgx.ParseConfig(dsn) - if err != nil { - return fmt.Errorf("failed to parse connection config: %w", err) - } + l.swapConn(conn) - // Connect with retry logic - var conn *pgx.Conn - var lastErr error - - retryAttempts := 3 - retryDelay := 1 * time.Second - - for attempt := 0; attempt < retryAttempts; attempt++ { - if attempt > 0 { - delay := calculateBackoff(attempt, retryDelay, 10*time.Second) - if l.config.GetEnableLogging() { - logger.Info("Retrying PostgreSQL listener connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay) - } - - select { - case <-time.After(delay): - case <-ctx.Done(): - return ctx.Err() - } - } - - conn, err = pgx.ConnectConfig(ctx, connConfig) - if err != nil { - lastErr = err - if l.config.GetEnableLogging() { - logger.Warn("Failed to connect PostgreSQL listener", "error", err) - } - continue - } - - // Test the connection - if err = conn.Ping(ctx); err != nil { - lastErr = err - conn.Close(ctx) - if l.config.GetEnableLogging() { - logger.Warn("Failed to ping PostgreSQL listener", "error", err) - } - continue - } - - // Connection successful - break - } - - if err != nil { - return fmt.Errorf("failed to connect listener after %d attempts: %w", retryAttempts, lastErr) - } - - l.mu.Lock() - l.conn = conn - l.mu.Unlock() - - // Start notification handler - go l.handleNotifications() - - // Start reconnection handler - go l.handleReconnection() + l.startOnce.Do(func() { + go l.handleNotifications() + go l.handleReconnection() + }) if l.config.GetEnableLogging() { logger.Info("PostgreSQL listener connected: name=%s", l.config.GetName()) @@ -122,30 +71,105 @@ func (l *PostgresListener) Connect(ctx context.Context) error { return nil } +// dial opens and verifies a new dedicated connection, with retries. +func (l *PostgresListener) dial(ctx context.Context) (*pgx.Conn, error) { + connConfig, err := buildPGXConfig(l.config) + if err != nil { + return nil, err + } + + var lastErr error + + retryAttempts, retryDelay, retryMaxDelay := retryPolicy(l.config) + + for attempt := 0; attempt < retryAttempts; attempt++ { + if attempt > 0 { + delay := calculateBackoff(attempt, retryDelay, retryMaxDelay) + if l.config.GetEnableLogging() { + logger.Info("Retrying PostgreSQL listener connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay) + } + + select { + case <-time.After(delay): + case <-ctx.Done(): + return nil, ctx.Err() + } + } + + conn, err := pgx.ConnectConfig(ctx, connConfig) + if err != nil { + lastErr = err + if l.config.GetEnableLogging() { + logger.Warn("Failed to connect PostgreSQL listener: %v", err) + } + continue + } + + // Test the connection + if err = conn.Ping(ctx); err != nil { + lastErr = err + closeConnBounded(conn) + if l.config.GetEnableLogging() { + logger.Warn("Failed to ping PostgreSQL listener: %v", err) + } + continue + } + + return conn, nil + } + + return nil, fmt.Errorf("failed to connect listener after %d attempts: %w", retryAttempts, lastErr) +} + +// closeConnBounded closes a pgx connection without ever waiting on a dead socket. +func closeConnBounded(conn *pgx.Conn) error { + ctx, cancel := context.WithTimeout(context.Background(), listenerCloseTimeout) + defer cancel() + return conn.Close(ctx) +} + +const ( + listenerCloseTimeout = 2 * time.Second + notificationPollInterval = 500 * time.Millisecond +) + +// currentConn returns the live connection, or an error if the listener is +// closed or not yet connected. +func (l *PostgresListener) currentConn() (*pgx.Conn, error) { + l.closeMu.Lock() + closed := l.closed + l.closeMu.Unlock() + if closed { + return nil, fmt.Errorf("listener is closed") + } + + l.mu.RLock() + conn := l.conn + l.mu.RUnlock() + if conn == nil { + return nil, fmt.Errorf("listener connection is not initialized") + } + return conn, nil +} + // Listen subscribes to a PostgreSQL notification channel func (l *PostgresListener) Listen(channel string, handler NotificationHandler) error { - l.closeMu.Lock() - if l.closed { - l.closeMu.Unlock() - return fmt.Errorf("listener is closed") - } - l.closeMu.Unlock() + // Take the connection between notification waits (each wait is short). + l.connMu.Lock() + defer l.connMu.Unlock() - l.mu.Lock() - defer l.mu.Unlock() - - if l.conn == nil { - return fmt.Errorf("listener connection is not initialized") - } - - // Execute LISTEN command - _, err := l.conn.Exec(l.ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{channel}.Sanitize())) + conn, err := l.currentConn() if err != nil { + return err + } + + if _, err := conn.Exec(l.ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{channel}.Sanitize())); err != nil { return fmt.Errorf("failed to listen on channel %s: %w", channel, err) } - // Store the handler + l.mu.Lock() l.channels[channel] = handler + l.mu.Unlock() if l.config.GetEnableLogging() { logger.Info("Listening on channel: name=%s, channel=%s", l.config.GetName(), channel) @@ -156,28 +180,21 @@ func (l *PostgresListener) Listen(channel string, handler NotificationHandler) e // Unlisten unsubscribes from a PostgreSQL notification channel func (l *PostgresListener) Unlisten(channel string) error { - l.closeMu.Lock() - if l.closed { - l.closeMu.Unlock() - return fmt.Errorf("listener is closed") - } - l.closeMu.Unlock() + l.connMu.Lock() + defer l.connMu.Unlock() - l.mu.Lock() - defer l.mu.Unlock() - - if l.conn == nil { - return fmt.Errorf("listener connection is not initialized") - } - - // Execute UNLISTEN command - _, err := l.conn.Exec(l.ctx, fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize())) + conn, err := l.currentConn() if err != nil { + return err + } + + if _, err := conn.Exec(l.ctx, fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize())); err != nil { return fmt.Errorf("failed to unlisten from channel %s: %w", channel, err) } - // Remove the handler + l.mu.Lock() delete(l.channels, channel) + l.mu.Unlock() if l.config.GetEnableLogging() { logger.Info("Unlistened from channel: name=%s, channel=%s", l.config.GetName(), channel) @@ -188,31 +205,24 @@ func (l *PostgresListener) Unlisten(channel string) error { // Notify sends a notification to a PostgreSQL channel func (l *PostgresListener) Notify(ctx context.Context, channel string, payload string) error { - l.closeMu.Lock() - if l.closed { - l.closeMu.Unlock() - return fmt.Errorf("listener is closed") - } - l.closeMu.Unlock() + l.connMu.Lock() + defer l.connMu.Unlock() - l.mu.RLock() - conn := l.conn - l.mu.RUnlock() - - if conn == nil { - return fmt.Errorf("listener connection is not initialized") - } - - // Execute NOTIFY command - _, err := conn.Exec(ctx, "SELECT pg_notify($1, $2)", channel, payload) + conn, err := l.currentConn() if err != nil { + return err + } + + if _, err := conn.Exec(ctx, "SELECT pg_notify($1, $2)", channel, payload); err != nil { return fmt.Errorf("failed to notify channel %s: %w", channel, err) } return nil } -// Close closes the listener and all subscriptions +// Close closes the listener and all subscriptions. Closing the connection drops +// every subscription server-side, so no UNLISTEN round trips are needed, and +// the close itself is bounded so a dead socket cannot hang the caller. func (l *PostgresListener) Close() error { l.closeMu.Lock() if l.closed { @@ -225,27 +235,26 @@ func (l *PostgresListener) Close() error { // Cancel context to stop background goroutines l.cancel() + // The cancelled ctx makes the notification wait return promptly, releasing + // connMu; closing the conn while it is being read would race inside pgx. + l.connMu.Lock() l.mu.Lock() - defer l.mu.Unlock() + conn := l.conn + l.conn = nil + l.channels = make(map[string]NotificationHandler) + l.mu.Unlock() - if l.conn == nil { + if conn == nil { + l.connMu.Unlock() return nil } - // Unlisten from all channels - for channel := range l.channels { - _, _ = l.conn.Exec(context.Background(), fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize())) - } - - // Close connection - err := l.conn.Close(context.Background()) + err := closeConnBounded(conn) + l.connMu.Unlock() if err != nil { return fmt.Errorf("failed to close listener connection: %w", err) } - l.conn = nil - l.channels = make(map[string]NotificationHandler) - if l.config.GetEnableLogging() { logger.Info("PostgreSQL listener closed: name=%s", l.config.GetName()) } @@ -262,20 +271,26 @@ func (l *PostgresListener) handleNotifications() { default: } + l.connMu.Lock() l.mu.RLock() conn := l.conn l.mu.RUnlock() if conn == nil { + l.connMu.Unlock() // Connection not available, wait for reconnection - time.Sleep(100 * time.Millisecond) + if !l.sleep(100 * time.Millisecond) { + return + } continue } - // Wait for notification with timeout - ctx, cancel := context.WithTimeout(l.ctx, 5*time.Second) + // Wait for a notification with a short timeout, so Listen/Unlisten/Notify + // waiting on connMu are served promptly. + ctx, cancel := context.WithTimeout(l.ctx, notificationPollInterval) notification, err := conn.WaitForNotification(ctx) cancel() + l.connMu.Unlock() if err != nil { // Check if context was cancelled @@ -291,13 +306,15 @@ func (l *PostgresListener) handleNotifications() { // Connection error, trigger reconnection if l.config.GetEnableLogging() { - logger.Warn("Notification error, triggering reconnection", "error", err) + logger.Warn("Notification error, triggering reconnection: %v", err) } select { case l.reconnectC <- struct{}{}: default: } - time.Sleep(1 * time.Second) + if !l.sleep(1 * time.Second) { + return + } continue } @@ -322,7 +339,22 @@ func (l *PostgresListener) handleNotifications() { } } -// handleReconnection manages automatic reconnection +// sleep waits for d or until the listener is closed; it reports whether the +// listener is still running. +func (l *PostgresListener) sleep(d time.Duration) bool { + t := time.NewTimer(d) + defer t.Stop() + select { + case <-t.C: + return true + case <-l.ctx.Done(): + return false + } +} + +// handleReconnection manages automatic reconnection. It runs as a single +// goroutine and dials replacement connections directly rather than through the +// public Connect, so no extra loops are started. func (l *PostgresListener) handleReconnection() { for { select { @@ -333,31 +365,21 @@ func (l *PostgresListener) handleReconnection() { logger.Info("Attempting to reconnect listener: name=%s", l.config.GetName()) } - // Close existing connection - l.mu.Lock() - if l.conn != nil { - l.conn.Close(context.Background()) - l.conn = nil - } - - // Save current subscriptions - channels := make(map[string]NotificationHandler) - for ch, handler := range l.channels { - channels[ch] = handler - } - l.mu.Unlock() - - // Attempt reconnection - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - err := l.Connect(ctx) + ctx, cancel := context.WithTimeout(l.ctx, 30*time.Second) + err := l.reconnect(ctx) cancel() if err != nil { + if l.ctx.Err() != nil { + return + } if l.config.GetEnableLogging() { logger.Error("Failed to reconnect listener: name=%s, error=%v", l.config.GetName(), err) } // Retry after delay - time.Sleep(5 * time.Second) + if !l.sleep(5 * time.Second) { + return + } select { case l.reconnectC <- struct{}{}: default: @@ -365,15 +387,6 @@ func (l *PostgresListener) handleReconnection() { continue } - // Resubscribe to all channels - for channel, handler := range channels { - if err := l.Listen(channel, handler); err != nil { - if l.config.GetEnableLogging() { - logger.Error("Failed to resubscribe to channel: name=%s, channel=%s, error=%v", l.config.GetName(), channel, err) - } - } - } - if l.config.GetEnableLogging() { logger.Info("Listener reconnected successfully: name=%s", l.config.GetName()) } @@ -381,6 +394,51 @@ func (l *PostgresListener) handleReconnection() { } } +// reconnect replaces the connection and resubscribes every channel on the new +// connection before publishing it, so the notification loop never touches a +// half-initialised conn. +func (l *PostgresListener) reconnect(ctx context.Context) error { + conn, err := l.dial(ctx) + if err != nil { + return err + } + + l.mu.RLock() + channels := make([]string, 0, len(l.channels)) + for ch := range l.channels { + channels = append(channels, ch) + } + l.mu.RUnlock() + + for _, ch := range channels { + if _, err := conn.Exec(ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{ch}.Sanitize())); err != nil { + closeConnBounded(conn) + return fmt.Errorf("failed to resubscribe to channel %s: %w", ch, err) + } + } + + if l.ctx.Err() != nil { + closeConnBounded(conn) + return l.ctx.Err() + } + l.swapConn(conn) + return nil +} + +// swapConn installs conn and closes the previous one. The old connection is +// closed under connMu so it is never closed while another goroutine is using it. +func (l *PostgresListener) swapConn(conn *pgx.Conn) { + l.connMu.Lock() + l.mu.Lock() + old := l.conn + l.conn = conn + l.mu.Unlock() + if old != nil { + closeConnBounded(old) + } + l.connMu.Unlock() +} + // IsConnected returns true if the listener is connected func (l *PostgresListener) IsConnected() bool { l.mu.RLock() diff --git a/pkg/dbmanager/providers/provider.go b/pkg/dbmanager/providers/provider.go index a541f2b..20ea1d4 100644 --- a/pkg/dbmanager/providers/provider.go +++ b/pkg/dbmanager/providers/provider.go @@ -4,17 +4,11 @@ import ( "context" "database/sql" "errors" - "strings" "time" "go.mongodb.org/mongo-driver/mongo" ) -// isDBClosed reports whether err indicates the *sql.DB has been closed. -func isDBClosed(err error) bool { - return err != nil && strings.Contains(err.Error(), "sql: database is closed") -} - // Common errors var ( // ErrNotSQLDatabase is returned when attempting SQL operations on a non-SQL database @@ -63,6 +57,31 @@ type ConnectionConfig interface { GetConnMaxLifetime() *time.Duration GetConnMaxIdleTime() *time.Duration GetReadPreference() string + GetRetryAttempts() int + GetRetryDelay() time.Duration + GetRetryMaxDelay() time.Duration +} + +// retryPolicy returns the configured retry settings, falling back to defaults. +func retryPolicy(cfg ConnectionConfig) (attempts int, delay, maxDelay time.Duration) { + attempts, delay, maxDelay = cfg.GetRetryAttempts(), cfg.GetRetryDelay(), cfg.GetRetryMaxDelay() + if attempts <= 0 { + attempts = 3 + } + if delay <= 0 { + delay = time.Second + } + if maxDelay <= 0 { + maxDelay = 10 * time.Second + } + return +} + +// Refresher is implemented by providers that can retire their pooled +// connections and dial fresh ones without closing the shared *sql.DB, so +// handles already handed out keep working. +type Refresher interface { + Refresh(ctx context.Context) error } // Provider creates and manages the underlying database connection diff --git a/pkg/dbmanager/providers/sqlite.go b/pkg/dbmanager/providers/sqlite.go index 4306b8d..abe3186 100644 --- a/pkg/dbmanager/providers/sqlite.go +++ b/pkg/dbmanager/providers/sqlite.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "fmt" + "strings" "sync" "time" @@ -15,10 +16,9 @@ import ( // SQLiteProvider implements Provider for SQLite databases type SQLiteProvider struct { - db *sql.DB - dbMu sync.RWMutex - dbFactory func() (*sql.DB, error) - config ConnectionConfig + db *sql.DB + dbMu sync.RWMutex + config ConnectionConfig } // NewSQLiteProvider creates a new SQLite provider @@ -26,6 +26,22 @@ func NewSQLiteProvider() *SQLiteProvider { return &SQLiteProvider{} } +// isMemoryDSN reports whether the SQLite DSN refers to a private in-memory +// database (each pooled connection would get its own empty database). +func isMemoryDSN(dsn string) bool { + path := dsn + if i := strings.IndexByte(path, '?'); i >= 0 { + path = path[:i] + } + if path == ":memory:" || path == "" { + return true + } + if strings.Contains(dsn, "mode=memory") && !strings.Contains(dsn, "cache=shared") { + return true + } + return path == "file::memory:" && !strings.Contains(dsn, "cache=shared") +} + // Connect establishes a SQLite connection func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) error { // Build DSN @@ -50,48 +66,35 @@ func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) erro return fmt.Errorf("failed to ping SQLite database: %w", err) } - // Configure connection pool - // Note: SQLite works best with MaxOpenConns=1 for write operations - // but can handle multiple readers - if cfg.GetMaxOpenConns() != nil { - db.SetMaxOpenConns(*cfg.GetMaxOpenConns()) - } else { - // Default to 1 for SQLite to avoid "database is locked" errors + if isMemoryDSN(dsn) { + // A private in-memory database exists per connection and disappears when + // that connection closes, so pin the pool to one connection that is + // never recycled. db.SetMaxOpenConns(1) - } - - if cfg.GetMaxIdleConns() != nil { - db.SetMaxIdleConns(*cfg.GetMaxIdleConns()) - } - if cfg.GetConnMaxLifetime() != nil { - db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime()) - } - if cfg.GetConnMaxIdleTime() != nil { - db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime()) - } - - // Enable WAL mode for better concurrent access - _, err = db.ExecContext(ctx, "PRAGMA journal_mode=WAL") - if err != nil { - if cfg.GetEnableLogging() { - logger.Warn("Failed to enable WAL mode for SQLite", "error", err) + db.SetMaxIdleConns(1) + db.SetConnMaxLifetime(0) + db.SetConnMaxIdleTime(0) + } else { + // SQLite works best with few writers; default to 1 unless configured. + if cfg.GetMaxOpenConns() != nil { + db.SetMaxOpenConns(*cfg.GetMaxOpenConns()) + } else { + db.SetMaxOpenConns(1) } - // Don't fail connection if WAL mode cannot be enabled - } - - // Set busy timeout to handle locked database (minimum 2 minutes = 120000ms) - busyTimeout := cfg.GetQueryTimeout().Milliseconds() - if busyTimeout < 120000 { - busyTimeout = 120000 // Enforce minimum of 2 minutes - } - _, err = db.ExecContext(ctx, fmt.Sprintf("PRAGMA busy_timeout=%d", busyTimeout)) - if err != nil { - if cfg.GetEnableLogging() { - logger.Warn("Failed to set busy timeout for SQLite", "error", err) + if cfg.GetMaxIdleConns() != nil { + db.SetMaxIdleConns(*cfg.GetMaxIdleConns()) + } + if cfg.GetConnMaxLifetime() != nil { + db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime()) + } + if cfg.GetConnMaxIdleTime() != nil { + db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime()) } } + p.dbMu.Lock() p.db = db + p.dbMu.Unlock() p.config = cfg if cfg.GetEnableLogging() { @@ -132,14 +135,7 @@ func (p *SQLiteProvider) HealthCheck(ctx context.Context) error { // Execute a simple query to verify the database is accessible var result int - run := func() error { return p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result) } - err := run() - if isDBClosed(err) { - if reconnErr := p.reconnectDB(); reconnErr == nil { - err = run() - } - } - if err != nil { + if err := p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result); err != nil { return fmt.Errorf("health check failed: %w", err) } @@ -150,32 +146,12 @@ func (p *SQLiteProvider) HealthCheck(ctx context.Context) error { return nil } -// WithDBFactory configures a factory used to reopen the database connection if it is closed. -func (p *SQLiteProvider) WithDBFactory(factory func() (*sql.DB, error)) *SQLiteProvider { - p.dbFactory = factory - return p -} - func (p *SQLiteProvider) getDB() *sql.DB { p.dbMu.RLock() defer p.dbMu.RUnlock() return p.db } -func (p *SQLiteProvider) reconnectDB() error { - if p.dbFactory == nil { - return fmt.Errorf("no db factory configured for reconnect") - } - newDB, err := p.dbFactory() - if err != nil { - return err - } - p.dbMu.Lock() - p.db = newDB - p.dbMu.Unlock() - return nil -} - // GetNative returns the native *sql.DB connection func (p *SQLiteProvider) GetNative() (*sql.DB, error) { if p.db == nil { diff --git a/pkg/dbmanager/providers/tcp_timeout_linux.go b/pkg/dbmanager/providers/tcp_timeout_linux.go new file mode 100644 index 0000000..c838321 --- /dev/null +++ b/pkg/dbmanager/providers/tcp_timeout_linux.go @@ -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 + } +} diff --git a/pkg/dbmanager/providers/tcp_timeout_other.go b/pkg/dbmanager/providers/tcp_timeout_other.go new file mode 100644 index 0000000..fa23a4c --- /dev/null +++ b/pkg/dbmanager/providers/tcp_timeout_other.go @@ -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 +} diff --git a/pkg/dbmanager/restart_live_test.go b/pkg/dbmanager/restart_live_test.go new file mode 100644 index 0000000..4a9f957 --- /dev/null +++ b/pkg/dbmanager/restart_live_test.go @@ -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) + } +}