Compare commits

..
16 Commits
Author SHA1 Message Date
Hein d3a99550d9 fix(security): name GetTemplate results to satisfy gocritic
Tests / Race Detector (push) Failing after 25s
Tests / Unit Tests (push) Failing after 25s
Tests / Integration Tests (push) Failing after 27s
Build , Vet Test, and Lint / Build (push) Successful in 1m3s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m30s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m31s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 1m31s
2026-09-30 13:46:03 +02:00
Hein 8a94d884e7 fix(security): verify passwords, bind row-security args, fail closed on panic
Verify bcrypt passwords in Direct mode and the shipped procedures, hash on
register/reset, ignore client-supplied roles and level at registration and
drop the password from the jwt_login payload. Legacy cleartext upgrade is
opt-in. Row security templates now bind the user as a parameter, validate
identifiers, attach via common.SelectQuery and fail the request if the
filter cannot be attached. ApplyColumnSecurity and GetRowSecurityTemplate
convert panics to errors and the hooks fail closed. Update audit status.
2026-09-30 13:44:59 +02:00
Hein f9c948ca4e fix(tracing): address audit findings
Default to TLS export with Insecure/TLSConfig/Headers options, parent-based
ratio sampling (default 0.1), and no query string or Host in span attributes.
Name spans by route template, record status and panics (re-raised), guard the
tracer with atomic.Pointer, reject double init, add init timeout and attribute
length limit, and move to semconv v1.26.0. Add tracing.insecure and
tracing.sample_rate config keys and tests.
2026-09-30 13:39:50 +02:00
Hein 164ba2b240 fix(security): remove lock contention, cache rules, stop leaked goroutines
Load column/row security without holding locks across provider calls,
honour pOverwrite with a 30s TTL and pruning, cap tokens per
Authorization header, detach session activity updates with a timeout,
stop OAuth2 cleanup goroutines via Close, move nil-map checks inside
locks, and drop O(n^2) string building. Update audit status.
2026-09-30 13:31:03 +02:00
Hein 97fe88b3a6 fix(modelregistry): address audit findings
Replace try-lock/sleep scheme with blocking locks, add ErrModelNotFound/
ErrModelExists/ErrInvalidModel sentinels, make RegisterModelWithRules
atomic, snapshot in IterateModels, guard defaultRegistry access, cap the
pointer-unwrap depth, and recover panics in callbacks and reflection.
Security hooks now allow-by-default only on ErrModelNotFound. Add tests.
2026-09-30 13:26:55 +02:00
Hein a4e1abc1df fix(logger): address audit findings
Redact and rate-limit error tracker fan-out, cap panic stack capture,
sanitise stdlib fallback output, strip contexts in Info/Debug, sync the
replaced logger, add Sync, UpdateLoggerE and CatchPanicRethrow, cache
the PID, and add tests.
2026-09-30 13:19:52 +02:00
Hein d7cb111496 docs(audit): add audit report for pkg/common 2026-09-30 13:16:47 +02:00
Hein f66930c3c9 chore(gosec): enable gosec and address findings 2026-09-30 13:16:12 +02:00
Hein 9533c3a0ed fix(cache): harden providers and default cache handling
Make cache-write failures non-fatal in GetOrSet/Remember, replace the
unsynchronised default cache with an atomic pointer, fix tag-index leaks
and write-lock-on-read in the memory provider, add a janitor, closed
state and byte copies, hash/namespace memcache keys with CAS tag index,
gate Clear behind AllowFlush, return ErrNotFound without the key, and
allowlist Redis stats. Mark audit status.
2026-09-30 13:13:06 +02:00
Hein 652621a70e ci: add race detector job 2026-09-30 13:09:44 +02:00
Hein c7b4530689 fix(race): make make test-race pass across ./pkg/...
Add a test-race target and run test-unit over ./pkg/...

Production races:
- logger: guard Logger/errorTracker with an RWMutex
- security: copy UserContext for the async session-activity goroutine
  and track it with a WaitGroup

Test fixes:
- eventbroker, websocketspec: use atomics for state shared with workers
- security: wait for async activity updates before touching sqlmock
- mqttspec: build full HookContext, set SubscriptionID, pin the
  in-memory SQLite to one connection

Update the cross-cutting audit (X1) with status and findings.
2026-09-30 13:07:47 +02:00
Hein e1cf72834e fix(config): lock Manager, harden defaults handling and path/IP helpers
Guard viper with an RWMutex, make the singleton race-free and stop
NewManager replacing the global (add SetConfigManager), write saved
configs 0600, search CWD last, add ConfigFileUsed and Config.Validate,
nil-safe PathsConfig.Set, confine PathsConfig.Join, bound GetIPs DNS
lookup, and drop dead Unmarshal in SetConfig. Mark audit status.
2026-09-30 13:02:40 +02:00
HeinandClaude Sonnet 5.5 3657aa94cc fix(dbmanager): explicitly ignore best-effort listener close errors
Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
2026-09-30 12:40:22 +02:00
HeinandClaude Sonnet 5.5 da1af1487e fix(dbmanager): keep the pool alive across errors and restarts
Implements the fixes from audit/pkg/dbmanager.audit.md.

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

Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
2026-09-30 12:40:14 +02:00
Hein bc8bff7955 docs(audit): add audit reports for pkg/testmodels and pkg/tracing
Tests / Unit Tests (push) Failing after 24s
Tests / Integration Tests (push) Failing after 26s
Build , Vet Test, and Lint / Build (push) Successful in 1m7s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 1m31s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m33s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m34s
2026-09-29 17:15:00 +02:00
Hein a74eebc7f3 fix(json-columns): skip JSON select columns with no model scan target
Tests / Unit Tests (push) Failing after 28s
Tests / Integration Tests (push) Failing after 29s
Build , Vet Test, and Lint / Build (push) Successful in 1m12s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m43s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 1m46s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m48s
Requesting a JSON sub-field (e.g. jsonvalue->'product'->>'cost') that has
no matching bun scanonly field on the model made the whole read fail with
"bun: ModelX does not have column Y", since bun scans SELECT results
straight into the typed model struct.

Add reflection.HasColumn to check whether the model can actually receive
a given column (including scanonly fields, walking embedded structs), and
gate the JSON select-column expression on it in ApplySelectColumns
(shared by websocketspec/mqttspec) and the resolvespec/restheadspec
handlers. When there's no scan target, drop just that column with a
warning instead of erroring the whole request.
2026-09-28 12:30:56 +02:00
108 changed files with 15682 additions and 1669 deletions
+11
View File
@@ -27,6 +27,17 @@ jobs:
with:
name: coverage-report
path: coverage.html
race-tests:
name: Race Detector
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
- name: Set up Go
uses: actions/setup-go@v6
with:
go-version: "1.24"
- name: Run unit tests with the race detector
run: go test -race -count=1 ./pkg/...
integration-tests:
name: Integration Tests
runs-on: ubuntu-latest
+1
View File
@@ -30,6 +30,7 @@
"linters": {
"enable": [
"gocritic",
"gosec",
"misspell",
"revive"
],
+12 -4
View File
@@ -1,11 +1,18 @@
.PHONY: test test-unit test-integration docker-up docker-down clean
.PHONY: test test-unit test-race test-integration docker-up docker-down clean
GOLANGCI_LINT := $(shell go env GOPATH)/bin/golangci-lint
# Run all unit tests
test-unit:
@echo "Running unit tests..."
@go test ./pkg/resolvespec ./pkg/restheadspec -v -cover
@go test ./pkg/... -v -cover
# Run all unit tests under the race detector (kept separate from coverage:
# race builds are 2-10x slower). Only races on executed paths are reported,
# so this covers every package rather than a subset.
test-race:
@echo "Running unit tests with the race detector..."
@go test -race -count=1 ./pkg/...
# Run all integration tests (requires PostgreSQL)
test-integration:
@@ -13,7 +20,7 @@ test-integration:
@go test -tags=integration ./pkg/resolvespec ./pkg/restheadspec -v
# Run all tests (unit + integration)
test: test-unit test-integration
test: test-unit test-race test-integration
release-version: ## Create and push a release with specific version (use: make release-version VERSION=v1.2.3 or make release-version to auto-increment)
@if [ -z "$(VERSION)" ]; then \
@@ -113,7 +120,8 @@ coverage-integration:
help:
@echo "Available targets:"
@echo " test-unit - Run unit tests"
@echo " test-unit - Run unit tests for all packages (./pkg/...)"
@echo " test-race - Run unit tests for all packages with -race"
@echo " test-integration - Run integration tests (requires PostgreSQL)"
@echo " test - Run all tests"
@echo " docker-up - Start PostgreSQL container"
+22 -1
View File
@@ -576,11 +576,32 @@ Centralized management of multiple database connections with support for Postgre
- Multiple named database connections
- Multi-ORM access (Bun, GORM, Native SQL) sharing the same connection pool
- Automatic SQLite schema translation (`schema.table` → `schema_table`)
- Health checks with auto-reconnect
- Background health checks (report status; they never close the pool)
- Prometheus metrics for monitoring
- Configuration-driven via YAML
- Per-connection statistics and management
**How to use it correctly**:
```go
mgr, err := dbmanager.NewManager(cfg) // or dbmanager.SetupManager(cfg) + GetInstance()
if err != nil { /* handle */ }
if err := mgr.Connect(ctx); err != nil { /* handle */ } // SetupManager does NOT connect
defer mgr.Close() // once, at shutdown
conn, _ := mgr.GetDefault()
db, _ := conn.Bun() // or conn.GORM() / conn.Native() / conn.Database()
handler := restheadspec.NewHandlerWithBun(db)
```
- **Fetch a handle once and keep it.** `Bun()`, `GORM()`, `Native()` and `Database()` return handles over one long-lived `*sql.DB`. You do not need to re-fetch them per request, and they stay valid for the life of the connection.
- **Never close a handle yourself.** Closing a `*bun.DB`, `*gorm.DB` or the `*sql.DB` closes the shared pool for everyone. Only `mgr.Close()` (at shutdown) should close it. After `Close`, the handles are dead.
- **Don't reconnect to recover from errors.** `database/sql` already discards bad connections and dials new ones. The manager does not close the pool on errors or failed health checks. `conn.Reconnect(ctx)` is for explicit operator use only (for example after rotating credentials): on PostgreSQL it retires pooled connections without closing the pool, so held handles keep working. Other databases close and reopen the pool, which invalidates handles you already hold.
- **Bring your own `*sql.DB`.** `dbmanager.NewConnectionFromDB(name, type, db)` wraps a pool you opened. The manager never closes it (`Close` only logs a warning); you own it and must close it.
- **Set deadlines on request contexts.** `query_timeout` is applied to PostgreSQL as `statement_timeout` (server side) and TCP timeouts detect dead sockets, but pass a context with a deadline to your queries so callers fail fast.
- **Pool tuning.** Keep `conn_max_idle_time` below the shortest idle timeout of any NAT, load balancer or pgbouncer between you and the database (typically 60-240s). SQLite `:memory:` is pinned to a single connection.
- **Health checks** run every `health_check_interval` (default 15s; a negative value disables them) and publish Prometheus metrics. `enable_auto_reconnect` is deprecated and ignored.
For documentation, see [pkg/dbmanager/README.md](pkg/dbmanager/README.md).
#### Cache
+677
View File
@@ -0,0 +1,677 @@
# Audit: cross-cutting findings across `pkg/*`
| | |
|---|---|
| **Scope** | all 23 packages under `pkg/` (64 065 non-test lines) |
| **Audit date** | 2026-09-29 |
| **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 |
This file records findings that are **not specific to one package** — they are
properties of the repository or patterns repeated across many packages. The
per-package audits reference this file rather than restating them.
## Findings
| # | Severity | Axis | Finding |
|---|---|---|---|
| X1 | **High** | locking | `-race` is never run anywhere; no package is ever race-checked |
| X2 | **High** | testing | `go test` runs against 2 of 23 packages; the other 21 are only compiled and vetted |
| X3 | **High** | testing | Every integration-test step is `continue-on-error: true` — integration failures cannot fail CI |
| X10 | **High** | security | Whole subsystems are declared, configured, documented and tested but never installed — including every protective middleware and the metrics provider |
| X4 | **Medium** | security | `gosec` is not enabled in `.golangci.json`; no SAST runs on a package set full of dynamic SQL |
| X5 | **Medium** | locking | Unsynchronized package-level mutable globals are the dominant concurrency pattern |
| X6 | **Medium** | security | Insecure-by-default transport across the board: `sslmode: disable`, `WithInsecure()`, no TLS in cache configs |
| X7 | **Medium** | panic handling | Panic handling is inconsistent and, where it exists, tends to fail open |
| X8 | **Medium** | security | `logger.Warn`/`Error` forward every message to Sentry unscrubbed, and error strings routinely embed attacker data *(partly fixed 2026-09-30: redaction and rate limiting added in `pkg/logger`; call sites still embed attacker data)* |
| X9 | **Low** | testing | Test coverage is extremely uneven: 5 packages have no test file at all |
The table is ordered by severity; the sections below are in ID order, since other
audit files reference these findings by number.
---
### X1. High — `-race` is never run
Verified by grep: the string `-race` does not appear in `Makefile`,
`.github/workflows/tests.yml`, `.github/workflows/maint.yml` or
`.github/workflows/make_tag.yml`.
Every test invocation in the repository:
```makefile
# Makefile:8
@go test ./pkg/resolvespec ./pkg/restheadspec -v -cover
# Makefile:13
@go test -tags=integration ./pkg/resolvespec ./pkg/restheadspec -v
# Makefile:97
@go test -tags=integration ./pkg/resolvespec ./pkg/restheadspec -v
# Makefile:103
@go test ./pkg/resolvespec ./pkg/restheadspec -coverprofile=coverage.out
# Makefile:110
@go test -tags=integration ./pkg/resolvespec ./pkg/restheadspec -coverprofile=coverage-integration.out
```
```yaml
# .github/workflows/tests.yml — unit-tests job
- name: Run unit tests
run: go test ./pkg/resolvespec ./pkg/restheadspec -v -cover
```
**Why this matters.** This audit found unsynchronized concurrent access to
mutable state in **six** packages, and the Go race detector would have flagged
every one of them on the first run:
| Package | Racing state | Reference |
|---|---|---|
| `pkg/cache` | `defaultCache` read/written by concurrent request handlers | `cache.audit.md` finding 3 |
| `pkg/config` | `*viper.Viper` has no internal lock; `configInstance` singleton | `config.audit.md` findings 1, 2 |
| `pkg/logger` | `Logger`, `errorTracker` globals | `logger.audit.md` finding 1 |
| `pkg/modelregistry` | `defaultRegistry` read by 6 functions without the lock | `modelregistry.audit.md` findings 2, 8 *(fixed 2026-09-30)* |
| `pkg/tracing` | `tracer` global | `tracing.audit.md` finding 5 *(fixed 2026-09-30)* |
| `pkg/errortracking` | `sentry.Init` mutates process globals | `errortracking.audit.md` finding 2 |
**Failure scenario.** `pkg/config` finding 1 is the sharpest illustration. A
concurrent `Manager.Set`/`Manager.Get` pair reaches viper's internal maps, which
have no mutex. A concurrent map read and write in Go is not a panic — it is
`fatal error: concurrent map read and map write`, which **`recover()` cannot
catch**. The process dies instantly, mid-request, with no graceful shutdown and
no error-tracker report. That is a remotely-triggerable hard crash, and it
cannot be found by inspection at scale — it is precisely what `-race` exists to
find. The detector has been in Go since 1.1 and costs one flag.
**Recommendation.** Add a race job that covers everything, and keep it separate
from the coverage run (race builds are ~2–10× slower):
```makefile
test-race:
@go test -race -count=1 ./pkg/...
```
```yaml
race-tests:
name: Race Detector
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
- uses: actions/setup-go@v6
with: { go-version: "1.24" }
- run: go test -race -count=1 ./pkg/...
```
Expect it to fail on the first run — that is the point. Fix `pkg/logger`,
`pkg/config` and `pkg/cache` first, since they are the shared dependencies. Note
that the race detector only reports races that **actually execute**, so X1 and X2
have to be fixed together: a race detector pointed at packages with no tests
finds nothing.
**Status (2026-09-30) — resolved for the packages with tests.** `make test-race` now exists
(`go test -race -count=1 ./pkg/...`), `test-unit` covers `./pkg/...`, and `test`
depends on both. The CI workflow (`.github/workflows/tests.yml`) now has a `race-tests`
job running the same command. The first full run was not clean:
| Package | Race | Kind |
|---|---|---|
| `pkg/logger` | `Logger` / `errorTracker` reassigned while other goroutines log (hit via `pkg/server` tests) | **production** — now guarded by an `RWMutex` (`getLogger`, `setLogger`, `getErrorTracker`); the exported `Logger` var is kept for compatibility |
| `pkg/security` | `DatabaseAuthenticator.Authenticate` passed `&userCtx` to the async session-activity goroutine while also returning it to the caller | **production** — the goroutine now gets a copy and is tracked by a `WaitGroup` so tests can wait for it |
| `pkg/security` tests | async activity update used sqlmock concurrently with the test adding expectations | test — tests wait via `authenticateSync` |
| `pkg/eventbroker`, `pkg/websocketspec` tests | handler/hook closures mutated a plain `bool`/`int` from worker goroutines | test — now `atomic` |
| `pkg/mqttspec` tests | not a race: hand-built `HookContext` lacked `TableName`/`Model`/`ModelPtr`, the unsubscribe test set `Data` instead of `SubscriptionID`, and `:memory:` SQLite gave each pooled connection its own empty database | test — fixed; these were failing without `-race` too |
`pkg/cache`, `pkg/config`, `pkg/modelregistry`, `pkg/tracing` and
`pkg/errortracking` are listed above but did **not** trip the detector: their
racing paths are not exercised by the current tests, which is the point made in
the paragraph above about X1 and X2 needing to be fixed together. Adding
concurrent tests for those globals is still outstanding.
Known limitation: `pkg/security` tests are not repeatable with `-count>1` (a
package-level capability cache carries over between runs), so the race target
keeps `-count=1`.
---
### X2. High — `go test` runs against 2 of 23 packages
Every `go test` invocation in the repository names exactly
`./pkg/resolvespec ./pkg/restheadspec`. No invocation uses `./...` or
`./pkg/...`.
The test bodies that exist but are never executed by CI:
| Package | Test files | Test lines | Run by CI? |
|---|---|---|---|
| `restheadspec` | 19 | 5 123 | **yes** |
| `resolvespec` | 8 | 2 379 | **yes** |
| `security` | 15 | 6 359 | no |
| `common` | 10 | 3 644 | no |
| `reflection` | 8 | 3 404 | no |
| `websocketspec` | 6 | 3 092 | no |
| `funcspec` | 3 | 2 416 | no |
| `spectypes` | 7 | 2 367 | no |
| `eventbroker` | 4 | 1 527 | no |
| `mqttspec` | 3 | 1 408 | no |
| `middleware` | 5 | 1 127 | no |
| `openapi` | 2 | 1 022 | no |
| `server` | 2 | 694 | no |
| `dbmanager` | 2 | 659 | no |
| `config` | 1 | 608 | no |
| `cache` | 1 | 69 | no |
| `errortracking` | 1 | 67 | no |
| `metrics` | 1 | 64 | no |
| `resolvemcp` | 1 | 34 | no |
| `logger` | 0 | 0 | — |
| `modelregistry` | 1 | ~150 | yes (`-race`) *(added 2026-09-30)* |
| `testmodels` | 0 | 0 | — |
| `tracing` | 1 | ~90 | yes *(added 2026-09-30)* |
**Failure scenario.** `pkg/security` has 6 359 lines of tests — the largest test
body in the repository — and **not one of them runs in CI**. A change that breaks
authentication, column-level security or row-security templates merges green.
The `maint.yml` job named "Run Vet Tests" is misleading: it runs `go mod
download`, `go mod verify` and `go vet ./...` and contains **no `go test` step at
all** (verified by grep). So the only signal on 21 of 23 packages is "it
compiles and vet is happy".
This directly explains the density of findings in this audit. The
`pkg/modelregistry` authorization fail-open (`modelregistry.audit.md` finding 1)
and the `pkg/cache`/`pkg/security` auth-outage-on-cache-failure
(`cache.audit.md` finding 1) are both the kind of defect a single unit test would
have caught, in packages that have never been tested.
**Recommendation.** Change every invocation to `./pkg/...`:
```makefile
test-unit:
@go test ./pkg/... -v -cover
```
```yaml
- name: Run unit tests
run: go test ./pkg/... -v -cover
```
If some currently-unrun package fails immediately, that is a bug report, not a
reason to keep the narrow list. Quarantine individual failing tests with
`t.Skip` and a `TODO` referencing an issue, so the *package* stays in the set.
---
### X3. High — integration failures cannot fail CI
`.github/workflows/tests.yml`, `integration-tests` job — every meaningful step
carries `continue-on-error: true`:
```yaml
- name: Run resolvespec integration tests
continue-on-error: true
env:
TEST_DATABASE_URL: "host=localhost user=postgres password=postgres dbname=resolvespec_test port=5432 sslmode=disable"
run: go test -tags=integration ./pkg/resolvespec -v -coverprofile=coverage-resolvespec-integration.out
- name: Run restheadspec integration tests
continue-on-error: true
...
```
**Failure scenario.** The integration suites are the only tests that exercise
real SQL generation against a real PostgreSQL — i.e. the only automated check on
the identifier-quoting and filter-construction paths that this audit's threat
model cares most about. Because both steps are `continue-on-error`, a SQL
injection regression, a broken join, or a total suite failure (wrong DSN, missing
migration) shows as a green check mark with a collapsed red step that nobody
opens. The job has no step that fails, so the job always passes. This is
strictly worse than not having the tests, because it creates the appearance of
coverage.
Note the integration DSN itself uses `sslmode=disable`, consistent with X6.
**Recommendation.** Remove `continue-on-error` from the two `go test` steps.
Keep it only on the coverage-report generation and artifact-upload steps, which
genuinely should not fail a build. If the suites are currently flaky, fix or
skip the flaky tests individually — `continue-on-error` on the whole step
disables the signal entirely.
---
### X4. Medium — `gosec` is not enabled — **RESOLVED**
> **Status (2026-09-30):** `gosec` is now in `linters.enable` and the repository lints clean
> (0 issues). The initial run produced 115 findings. Real fixes: login-form values in
> `security/oauth_server.go` are now HTML-escaped (G705), and `SqlSparseVector` index
> parsing uses `ParseInt(..., 10, 32)` (G109). The remaining ~110 sites carry
> `//nolint:gosec // Gxxx: <reason>` comments. The G201/G701 reasons (identifiers from
> trusted config or internal/validated names) and the G115 range claims were not
> individually audited and still need review. The text below describes the state before the change.
`.golangci.json` (`version: 2`) enables exactly three linters beyond the v2
standard set:
```json
"linters": {
"enable": [
"gocritic",
"misspell",
"revive"
],
```
golangci-lint v2's standard set (`errcheck`, `govet`, `ineffassign`,
`staticcheck`, `unused`) is on by default, so those do run. **`gosec` does not** —
it appears in the file only inside an exclusion rule for `_test.go`:
```json
{
"linters": [
"dupl",
"errcheck",
"gocritic",
"gosec"
],
"path": "_test\\.go"
},
```
Listing a linter in `exclusions.rules` does not enable it. The `lint` job in
`.github/workflows/maint.yml:40-57` does run golangci-lint over the whole
repository with `version: latest`, so the config is applied — it simply never
asks for the security checks.
**Failure scenario.** This repository builds SQL by string construction from
attacker-controlled schema, table, column and filter names (see
`restheadspec.audit.md` and `common.audit.md`). `gosec`'s `G201`/`G202`
(SQL string formatting/concatenation) are exactly the rules that would flag a
new `fmt.Sprintf` into a query, which is the single most likely way a SQL
injection enters this codebase. Also unenabled and relevant: `G104` (unhandled
errors — this audit found ~20 discarded errors in `pkg/cache` alone), `G304`
(file path from variable — relevant to `PathsConfig.Join`, `config.audit.md`
finding 14), `G402` (bad TLS settings — X6), `G404` (weak random).
**Recommendation.** Add `gosec` to `linters.enable` and triage the initial
findings. Expect noise on the SQL rules given the architecture; suppress
individual verified-safe sites with `//nolint:gosec // G201: identifier is
validated by X` comments that name the invariant, rather than disabling the rule
globally. That converts each suppression into a reviewable claim.
Consider also `bodyclose`, `rowserrcheck` and `sqlclosecheck` for a
database-heavy codebase, and `contextcheck` given how many methods here accept a
`ctx` and ignore it.
---
### X5. Medium — unsynchronized mutable package globals are the dominant pattern
Nine of the twenty-three packages expose mutable process-wide state through
package-level variables, and most guard it with nothing:
| Package | Global | Guarded? |
|---|---|---|
| `pkg/logger` | `Logger *zap.SugaredLogger` (`logger.go:15`), `errorTracker` (`:16`) | **no** — and `Logger` is exported |
| `pkg/cache` | `defaultCache *Cache` (`cache.go:10`) | **no** |
| `pkg/config` | `configInstance *Manager` (`manager.go:15`) | **no** |
| `pkg/tracing` | `tracer` | **yes** *(fixed 2026-09-30)* — `atomic.Pointer` |
| `pkg/modelregistry` | `defaultRegistry` | **yes** *(fixed 2026-09-30)* — guarded by `registriesMutex`; all access via `GetDefaultRegistry()` |
| `pkg/metrics` | `globalProvider` (`interfaces.go:50-51`) | **yes** — `globalProviderMu sync.RWMutex` |
`pkg/metrics` is the model the others should follow:
```go
// pkg/metrics/interfaces.go:50-72
var (
globalProviderMu sync.RWMutex
globalProvider Provider
)
func SetProvider(p Provider) {
globalProviderMu.Lock()
globalProvider = p
globalProviderMu.Unlock()
}
func GetProvider() Provider {
globalProviderMu.RLock()
p := globalProvider
globalProviderMu.RUnlock()
if p == nil {
return &NoOpProvider{}
}
return p
}
```
Note that it also returns a working `NoOpProvider` rather than `nil`, so callers
need no nil check — the pattern `pkg/logger` and `pkg/cache` should copy.
**Failure scenario.** Beyond the data races in X1, the shared failure mode is
**lazy initialization on the request path**. `cache.GetDefaultCache()`
(`cache.go:48`) and `config.GetConfigManager()` (`manager.go:18`) both
`if x == nil { x = construct() }` with no `sync.Once`. Under concurrent first
traffic, several instances are constructed and all but one are silently
discarded, so writes go to an orphaned object — a cache that is permanently 100%
miss, or two `Manager`s disagreeing about configuration. It presents as "the
cache doesn't work" with no error anywhere.
`pkg/logger.Logger` being **exported** and mutable is its own hazard: any
package, or any consumer of this library, can reassign the process logger
mid-flight while other goroutines are calling `Logger.Infow`.
**Recommendation.** For each global: `atomic.Pointer[T]` for
single-pointer swaps, `sync.Once` for lazy defaults, `sync.RWMutex` for
multi-field state. Unexport `logger.Logger` behind accessors. Where a nil global
is possible, return a no-op implementation instead of `nil`, as
`pkg/metrics.GetProvider` does.
---
### X6. Medium — insecure transport is the default everywhere
Every network dependency defaults to cleartext, and in two cases there is no way
to configure otherwise:
| Component | Default | Configurable? | Reference |
|---|---|---|---|
| PostgreSQL | `sslmode: disable` (`config/manager.go:242`) | yes, via config | `config.audit.md` finding 3 |
| OTLP traces | `otlptracegrpc.WithInsecure()` hardcoded (`tracing/tracing.go:41`) *(fixed 2026-09-30: TLS default, `Insecure` opt-in)* | **yes** | `tracing.audit.md` finding 1 |
| Redis (cache) | no `TLSConfig` set | **no** — `RedisConfig` has no TLS field | `cache.audit.md` finding 15 |
| Memcache | no TLS | **no** | `cache.audit.md` finding 15 |
| CORS | `allowed_origins: ["*"]`, `allowed_headers: ["*"]` (`config/manager.go:214-216`) | yes | `config.audit.md` finding 3 |
| DB user | `user: postgres` with blank password (`config/manager.go:239-240`) | yes | `config.audit.md` finding 3 |
The `tracing.go:41` case is the most pointed, because the code knows better:
```go
otlptracegrpc.WithInsecure(), // Use WithTLSCredentials in production
```
The comment names the fix and the config struct provides no way to apply it.
**Failure scenario.** The cache holds `UserContext` — identity and authorization
data — keyed by the raw bearer token (`security/providers.go:398`). With no TLS,
anything on the path between the service and Redis can read session contents and
the `AUTH` password, then **write** a forged `auth:session:<token>` entry.
`GetOrSet` returns a cache hit without consulting the database, so a forged entry
is a complete authentication bypass. Meanwhile the trace exporter ships full
request URLs including query strings (`tracing.audit.md` finding 2) in cleartext
to the collector.
**Recommendation.** Invert every default: TLS on unless explicitly disabled.
Concretely — add `TLS`/`TLSSkipVerify`/`TLSCACertFile` to `cache.RedisConfig`
and `tracing.Config`; change the `sslmode` default to `require`; change
`cors.allowed_origins` to `[]` and require an explicit list; remove the default
`postgres`/blank-password credentials so a misconfigured deployment fails to
start rather than connecting to a local database as a superuser. Add a startup
validation pass that logs a prominent warning for each insecure setting actually
in effect.
---
### X7. Medium — panic handling is inconsistent, and where it exists it fails open
Three different conventions coexist:
1. **`logger.CatchPanic(location)`** (`logger/logger.go:184`) — recovers, logs,
reports, and **swallows**. Both call sites are security enforcement:
`security/provider.go:302` (`ApplyColumnSecurity`) and `:443`
(`GetRowSecurityTemplate`). See `logger.audit.md` finding 4.
2. **`logger.HandlePanic(method, r)`** (`logger/logger.go:197`) — converts the
panic to an `error` the caller must handle. This is the correct shape.
3. **Nothing at all.** `pkg/cache` has zero `recover()` calls in 1 538 lines;
so do several other packages.
**Failure scenario (fail-open).** `ApplyColumnSecurity` panics — a nil map, a
bad type assertion on a rule, a reflection edge case. `CatchPanic` recovers and
the function returns normally, so the caller believes column security was
applied. It was not. The response contains the columns the security layer was
supposed to strip. The panic is logged, but the request succeeds with elevated
data exposure. A security control whose failure mode is "allow" is the wrong
default; it must be "deny".
**Failure scenario (panic under a lock).** `pkg/cache` holds `m.mu` across
`m.items[key] = ...` (`provider_memory.go:111`). After `Close()` sets
`items = nil` that assignment panics. With no recover in the package the panic
propagates to whatever handler exists upstream; if that handler recovers, `m.mu`
is **never unlocked** and every subsequent cache operation blocks forever. The
process stays alive and wedged — worse than a crash, because health checks that
do not touch the cache keep passing.
**Recommendation.** Establish one convention and apply it:
- **Request boundaries** (HTTP handlers, event consumers, goroutines): recover,
log with stack, report to the error tracker, return 500 / nack. A `go`
statement without a deferred recover is a process-kill waiting to happen —
`security/providers.go:447` (`go a.updateSessionActivity(...)`) is one.
- **Security enforcement**: recover, log, and **fail closed** — return an error
that the caller must propagate as a denial. Never `CatchPanic`.
- **Internal helpers**: do not recover. Let the boundary handle it.
- **Anything holding a lock**: prefer `defer mu.Unlock()` (already the pattern in
`pkg/cache`) so a panic cannot leak the lock, and keep panicking code out of
critical sections.
Add a `CatchPanicFailClosed(location string, err *error)` helper so the
fail-closed variant is as easy to reach for as `CatchPanic`.
---
### X8. Medium — attacker data reaches Sentry unscrubbed
Two facts compose badly:
`pkg/logger/logger.go:125-140` — every `Error` (and every `Warn`, `:108-123`)
forwards the fully-formatted message to the error tracker:
```go
func Error(template string, args ...interface{}) {
ctx, remainingArgs := extractContext(args...)
message := fmt.Sprintf(template, remainingArgs...)
...
if errorTracker != nil {
errorTracker.CaptureMessage(ctx, message, errortracking.SeverityError, map[string]interface{}{
"process_id": os.Getpid(),
})
}
}
```
And `pkg/errortracking` installs **no `BeforeSend` scrubber**
(`errortracking.audit.md` finding 1), so the message goes to Sentry verbatim.
Meanwhile error strings across the codebase interpolate attacker-controlled
values, sometimes secrets:
| Site | Interpolated value |
|---|---|
| `cache/cache_manager.go:26`, `:40` | the full cache key — for the session cache, **the raw bearer token** |
| `security/providers.go:391` | the raw `Authorization` header, logged at `Warn` when multiple tokens are present |
| `config/manager.go:164` | config file paths |
| throughout `restheadspec` | schema, table, column and filter values from the request |
**Failure scenario.** `security/providers.go:391` is live today:
```go
logger.Warn("Multiple authentication tokens provided in Authorization header (%d tokens). This is unusual and may indicate a misconfigured client. Header: %s", len(tokens), sessionToken)
```
A client sends two bearer tokens. `logger.Warn` formats the full header value
into the message and forwards it to Sentry, where a **valid session credential**
is now stored by a third party, visible to everyone with Sentry access, retained
per Sentry's policy, and replayable for the token's lifetime. No attacker
sophistication is required — the trigger is a single extra header, and the
codebase invites it by logging the header contents as the diagnostic.
**Recommendation.**
1. Add a `BeforeSend` hook in `pkg/errortracking` that redacts
`Authorization`, `Cookie`, `Set-Cookie`, anything matching
`(?i)(token|password|secret|apikey|api_key|bearer)\s*[:=]\s*\S+`, and
long high-entropy strings. This is the one change that bounds the whole class.
2. Never log a credential, even truncated. Change `providers.go:391` to log
`len(tokens)` only.
3. Replace `fmt.Errorf("key not found: %s", key)` with a sentinel
`cache.ErrNotFound` (`cache.audit.md` finding 7).
4. Key the session cache on `sha256(token)`, as
`security/keystore_database.go:287` already does for API keys.
5. Add sampling / rate limiting to the tracker fan-out
(`logger.audit.md` finding 3) so an error storm is not also a cost and
availability event.
---
### X9. Low — five packages have no tests at all
`pkg/logger`, `pkg/modelregistry`, `pkg/testmodels`, `pkg/tracing` have zero
`*_test.go` files. `pkg/resolvemcp` has 34 lines, `pkg/metrics` 64,
`pkg/errortracking` 67, `pkg/cache` 69.
**Failure scenario.** `pkg/modelregistry` is untested and contains this audit's
only **Critical** authorization finding: `GetModel` returns a "registry locked"
error under write-lock contention, which `security/hooks.go:274-294` converts
into `return nil // model not registered, allow by default`
(`modelregistry.audit.md` finding 1; *fixed 2026-09-30, regression tests added*). A twenty-line test that registers a model
from one goroutine while reading it from another would demonstrate the fail-open
immediately. The package guards a security boundary and has never been tested.
`pkg/logger` being untested matters for a different reason: it is imported by
almost every other package, so a defect there (the format-string sink in `Info`
and `Debug`, `logger.audit.md` finding 6) is repo-wide.
**Recommendation.** Prioritize by blast radius, not by size:
1. `pkg/modelregistry` — concurrent register/read; assert `GetModelRulesByName`
never returns a "locked" error that a caller could read as "not registered".
2. `pkg/logger` — nil-`Logger` fallback paths, format-string handling, and that
`Warn`/`Error` do not forward secrets once a scrubber exists.
3. `pkg/cache` — concurrent `GetDefaultCache`, the expired-item TOCTOU, and that
`tagToKeys` does not grow after eviction.
4. `pkg/tracing`, `pkg/metrics`, `pkg/errortracking` — construction and no-op
paths; these are mostly configuration surfaces.
Combine with X1 and X2: tests that are not run, and tests run without `-race`,
do not close these gaps.
---
### X10. High — configured subsystems that are never installed
Three separate subsystems are fully built — typed config, defaults, tests,
documentation — and then never connected to anything that runs.
**1. Every protective middleware.** `pkg/middleware` provides rate limiting, IP
blacklisting, request-size limiting and input sanitization. Non-test callers:
| Constructor | Non-test callers |
|---|---|
| `middleware.NewRateLimiter` | **0** |
| `middleware.NewIPBlacklist` | **0** |
| `middleware.NewRequestSizeLimiter` | **0** |
| `middleware.DefaultSanitizer` | **0** outside the package |
| `middleware.StrictSanitizer` | **0** |
| `middleware.PanicRecovery` | 1 — `pkg/server/manager.go:466` |
`pkg/server/manager.go` is the only file outside the package that imports it, and
only for `PanicRecovery`. The config that exists to drive the rest —
`MiddlewareConfig.RateLimitRPS`, `.RateLimitBurst`, `.MaxRequestSize`
(`pkg/config/config.go:123-125`), defaulted at `pkg/config/manager.go:209-211` —
has **no reader anywhere in the module**.
**2. The metrics provider.** `metrics.SetProvider` and
`metrics.NewPrometheusProvider` have **0 non-test callers**, so
`metrics.GetProvider()` returns `&NoOpProvider{}`
(`pkg/metrics/interfaces.go:63-72`) for the process lifetime. Every instrumented
call site in the repository — 39 DB-query sites in
`pkg/common/adapters/database`, the HTTP middleware, the event-broker counters,
and the sole `RecordPanic` call at `pkg/middleware/panic.go:19` — writes to a
no-op. `MetricsConfig.Enabled` and `.Provider` are likewise never read, and
`pkg/config` has no `metrics` section at all.
**3. The configured CORS policy.** `config.CORSConfig`
(`pkg/config/config.go:128-134`), defaulted at `pkg/config/manager.go:214-217`,
is never read. The policy that actually applies comes from a **different type of
the same name**, `common.CORSConfig`, built by `common.DefaultCORSConfig()`
(`pkg/common/cors.go:19-48`), which derives allowed origins from the configured
server instances and the host's local IPs and ignores `cors.allowed_origins`
entirely. It is called from ten sites across `pkg/resolvespec` and
`pkg/restheadspec`.
**Failure scenario.** Each of these is a silent, config-shaped lie, and they fail
in the same way: the operator's mental model of the deployment is wrong in the
direction of believing a control exists.
- **Under the hostile-client threat model there is no rate limit and no
request-body limit in the serving path.** `max_request_size: 10485760` is
configured and unenforced, so a single unauthenticated `POST` with a
multi-gigabyte body is read into memory and OOM-kills the process; unlimited
request rate exhausts the 25-connection default pool
(`pkg/config/manager.go:224`) just as cheaply. Both are one-line attacks
against controls the configuration says are active. An operator lowering
`rate_limit_rps` during an incident observes no change and will reasonably
conclude the attack exceeds the limit rather than that no limit exists.
- **There is no telemetry with which to notice any of it.** No request counts, no
latency histograms, no `panics_total`, no DB-query metrics — the one signal
that would show an attack in progress is wired end to end and discarded at the
last step. This is also why the metrics cardinality defects
(`metrics.audit.md` findings 2 and 5) are only latent: they become live the
moment someone installs the provider that the config implies is already there.
- **Tightening `cors.allowed_origins` does nothing.** The value is ignored, so a
hardening change lands, reviews clean, deploys, and changes no behaviour. Two
types named `CORSConfig` in two packages is the mechanism; nothing warns.
The common thread is that none of this fails visibly. It compiles, the tests pass
(`pkg/middleware` has the repo's best test ratio — 1 127 test lines to 799 code
lines — all of it exercising code nothing calls), CI is green, and the config file
documents features that are absent. Under X2 these packages are not even in the
tested set, so the tests that do exist are not run.
**Recommendation.**
1. **Wire the middleware chain** in `pkg/server` from `MiddlewareConfig`,
outermost first: size limiter → rate limiter → blacklist → `PanicRecovery`
(innermost, so it sees handler panics; `trackRequestsMiddleware` at
`manager.go:540` correctly stays outside). Fix the trusted-proxy handling
(`middleware.audit.md` findings 2 and 3) **before** mounting the two IP-based
layers, and do not mount the sanitizer at all until findings 5–7 there are
resolved — as written it corrupts filter values and can synthesize a
`javascript:` URI.
2. **Install a metrics provider** from config, gated on `metrics.enabled`, and
add the missing `metrics` section to `pkg/config`. Bound the label sets first
(`metrics.audit.md` findings 2 and 5) — installing the provider as-is converts
two latent cardinality DoS findings into live ones.
3. **Delete the duplicate `CORSConfig`** or make `common.DefaultCORSConfig()`
read `config.CORSConfig`. Two types with one name, one of them ignored, is a
trap regardless of which way it is resolved.
4. **Make the class of defect detectable.** Log at startup which middleware,
metrics provider and CORS policy are active, so an unwired subsystem is
visible in the first ten lines of a boot log instead of during an incident.
A CI check that every `mapstructure` field in `pkg/config` has at least one
reader would have caught all three of these; so would enabling `unused` in
`.golangci.json` for exported-but-unreferenced constructors.
---
## Recommended order of work
1. **X2 + X1** — point `go test` at `./pkg/...` and add a `-race` job. Everything
else in this audit is easier to verify once these exist, and they will
surface the six data races on their own.
2. **X3** — remove `continue-on-error` from the integration `go test` steps.
3. **X10** — mount the request-size limiter and rate limiter. Until this is
done the service has no volumetric protection at all, and no metrics with
which to see that. Fix `middleware.audit.md` findings 2 and 3 in the same
change, since mounting the IP-based layers without them adds attack surface.
4. **X8 item 1 and 2** — add the Sentry `BeforeSend` scrubber and stop logging
the `Authorization` header. Small, self-contained, stops an active credential
leak.
5. **X7** — decide the panic convention; make the two `CatchPanic` sites in
`pkg/security` fail closed, and stop returning the panic value to the client
(`pkg/middleware/panic.go:28`).
6. **X5** — fix the globals in `pkg/logger`, `pkg/config`, `pkg/cache`
(the shared dependencies) first.
7. **X6** — add TLS fields and invert the defaults.
8. **X4** — enable `gosec` and triage.
9. **X9** — backfill tests, in the order listed above.
## Per-package audits
`cache` · `common` · `config` · `dbmanager` · `errortracking` · `eventbroker` ·
`funcspec` · `logger` · `metrics` · `middleware` · `modelregistry` · `mqttspec` ·
`openapi` · `reflection` · `resolvemcp` · `resolvespec` · `restheadspec` ·
`security` · `server` · `spectypes` · `testmodels` · `tracing` · `websocketspec`
Each is `audit/pkg/<name>.audit.md`.
File diff suppressed because it is too large Load Diff
+407
View File
@@ -0,0 +1,407 @@
# Audit: `pkg/common`
| | |
|---|---|
| **Package** | `github.com/bitechdev/ResolveSpec/pkg/common` (+ `adapters/database`, `adapters/router`) |
| **Files** | `sql_helpers.go` (1060), `recursive_crud.go` (645), `validation.go` (444), `json_column.go` (402), `interfaces.go` (311), `spatial_helpers.go` (317), `handler_utils.go` (309), `json_condition.go` (219), `types.go` (192), `cors.go` (156), `handler_example.go` (97); `adapters/database/bun.go` (1767), `pgsql.go` (1600), `gorm.go` (1018), `query_metrics.go` (335), `pgsql_preload_example.go` (275), `pgsql_example.go` (176), `test_helpers.go` (132), `utils.go` (117); `adapters/router/mux.go` (238), `bunrouter.go` (214) |
| **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 for `sql_helpers.go`, `validation.go`, `cors.go`, `recursive_crud.go`, `json_column.go`/`json_condition.go` and the reconnect/transaction paths of the three DB adapters; medium for the rest; the `*_example.go` files were skimmed. Findings 1–3 were verified with throw-away probe tests, which were deleted afterwards |
## Summary
`pkg/common` is the shared core behind every spec handler. It contains the
`Database` / `SelectQuery` abstraction and its Bun, GORM and raw-`pgx`
adapters, the request-option types, column validation, JSON-column parsing,
nested (recursive) CRUD, CORS, and a set of SQL string helpers. The spec
packages feed **client-supplied raw SQL fragments** through those helpers:
`x-custom-sql-w`, `x-custom-sql-or`, `x-custom-sql-join`, preload `where`,
sort expressions and cursor filters.
The central problem is that **`SanitizeWhereClause` / `validateWhereClauseSecurity`
is a keyword denylist applied to raw SQL**, and the result is concatenated
straight into the query. A denylist can't make arbitrary client SQL safe, and
this one misses subqueries, functions, comments and parenthesis balancing.
With the helpers exactly as the handlers call them, a client can:
- escape the outer parentheses and OR past every filter the server adds
afterwards (**row security, tenant filters, the PK filter**). This is
verified. Row security is inert anyway today (`security.audit.md` finding 2),
but this bug will defeat it as soon as that is fixed;
- read any table the DB role can see through a subquery;
- stall a connection with `pg_sleep`;
- bypass the keyword list with a comment (`delete/**/from`).
Meanwhile, legitimate filters that merely *contain* a word like `update` are
silently dropped, and the query runs **unfiltered** (fail-open).
Other headline findings:
- `SetCORSHeaders` reflects **any** Origin with `Allow-Credentials: true` and
ignores `AllowedOrigins`.
- Nested CRUD updates and deletes child rows by primary key alone. A client can
modify or delete (or re-parent) any row in a related table.
- Sort validation lets arbitrary SQL through whenever a custom join has no
alias.
On the question that started this audit (idle connections becoming unusable),
the relevant part of `pkg/common` is the adapters' reconnect logic (finding 5).
It is only partly wired into the Bun and pgx adapters, and it's what calls
`dbmanager`'s destructive `Reconnect` (`dbmanager.audit.md` findings 1–2).
## Findings
| # | Severity | Axis | Finding |
|---|---|---|---|
| 1 | **Critical** | security | Client raw-SQL WHERE (`x-custom-sql-w`/`-or`, preload where, cursor) is protected only by a keyword denylist: parenthesis escape defeats server-added filters; subqueries, `pg_sleep` and comment bypasses all pass (verified) |
| 2 | **Critical** | security | `SetCORSHeaders` reflects any `Origin` and sends `Access-Control-Allow-Credentials: true`; `AllowedOrigins` is never consulted |
| 3 | **High** | security | Sort validation: an empty join alias makes `strings.Contains(col, "")` accept **any** sort string, and `(…)` sort expressions allow arbitrary subqueries (verified) |
| 4 | **High** | security | Nested CRUD (`recursive_crud.go`) updates and deletes child rows by `WHERE pk = ?` only, with no parent/ownership constraint; a client-supplied `_request` switches the operation per object |
| 5 | **High** | locking / availability | Adapter reconnect is inconsistent (Bun and pgx query builders never reconnect) and, where it exists, calls dbmanager's pool-closing `Reconnect`; `BunAdapter.CommitTx`/`RollbackTx` are silent no-ops |
| 6 | **Medium** | security / correctness | `SanitizeWhereClause` fails open: on a denylist hit it returns `""`, so the client's filter is dropped and the query returns unfiltered rows; false positives on ordinary data (`'awaiting update'`, `last_update`) |
| 7 | **Medium** | security | Any column name starting with `cql` passes `ValidateColumn` unconditionally |
| 8 | **Medium** | slowness | Request bodies are read with unbounded `io.ReadAll` in both router adapters |
| 9 | **Medium** | logging | Failed queries log the fully interpolated SQL, and nested CRUD logs full row data, at `Error`, which is forwarded to Sentry (`_CROSS-CUTTING.audit.md` X8) |
| 10 | **Low** | correctness | `stripEmptyComparisonClauses` regexes rewrite SQL without respecting string literals; quote tracking ignores `''`; `qualifyColumnInCondition` compiles a regex per call |
| 11 | **Low** | locking | Adapter fields read without their mutex (`BunAdapter.NewSelect` `db: b.db`, `DriverName`, `PgSQLAdapter.GetUnderlyingDB`) race with `reconnectDB` |
| 12 | **Info** | — | `json_column.go` / `json_condition.go` are well built: allowlisted casts, path bound as a single `text[]` parameter, identifiers validated and quoted |
---
### 1. Critical — Client raw-SQL WHERE is guarded only by a keyword denylist
`sql_helpers.go:118-161` (`validateWhereClauseSecurity`), `169-305`
(`SanitizeWhereClause`), `375-395` (`EnsureOuterParentheses`).
Call sites that pass **client-controlled** strings:
| Source | Call site |
|---|---|
| `x-custom-sql-w` | `restheadspec/handler.go:692-699` → `query.Where(...)` |
| `x-custom-sql-or` | `restheadspec/handler.go:703-710` → `query.WhereOr(...)` |
| `x-custom-sql-join` | `restheadspec/headers.go:666`, `1330` (sanitized with `tableName ""`) |
| preload `where` | `resolvespec/handler.go:2386, 2458`; `restheadspec/handler.go:618, 1170` |
| cursor filters | `resolvespec/handler.go:458`; `restheadspec/handler.go:916`; `resolvemcp/handler.go:304` |
The pipeline is `AddTablePrefixToColumns` → `SanitizeWhereClause` →
`EnsureOuterParentheses` → `query.Where(s)`, with **no bind arguments**. The
only security check is a substring search for `delete `, `update `, `drop `,
`;delete`, and similar.
A probe reproduced the handler pipeline and then appended a server-side filter
`Where("tenant = ?", 5)`, which is what a row-security or tenant hook does:
| Client `x-custom-sql-w` | Resulting SQL / effect |
|---|---|
| `1=1)) OR ((1=1` | `WHERE ((1=1)) OR ((1=1)) AND (tenant = 5)`: **every tenant's rows**, because `AND` binds tighter than `OR` |
| `id = 1 or (select count(*) from pg_shadow) > 0` | passes unchanged, so boolean-oracle exfiltration from any readable table works |
| `id = 1 and pg_sleep(5) is not null` | passes; each request pins a pool connection for as long as the client likes |
| `id = 1; delete/**/from items` | passes, because the comment defeats `"delete "` (whether it executes depends on the driver's multi-statement handling) |
`EnsureOuterParentheses` only checks whether the string *already* starts and
ends with a matching pair. It never checks that the parentheses inside are
balanced, which is what the escape relies on. `x-custom-sql-or` is worse by
design: `WhereOr` ORs the client clause against **every** condition already on
the query, so it needs no escape at all to widen a server-side filter.
Row security currently has no effect (`security.audit.md` finding 2). Fixing
that type assertion will **not** give tenant isolation while these headers
exist, and the same escape defeats the server's own PK scoping
(`restheadspec/handler.go:759-766`).
**Failure scenario.** An authenticated user of tenant A sends
`X-Custom-SQL-W: 1=1)) OR ((1=1` on a list endpoint and receives tenant B's
rows. Or they send
`X-Custom-SQL-W: (select substr(passwd,1,1) from pg_shadow limit 1) = 'm'` and
extract data one character at a time.
**Recommendation.** Stop accepting raw SQL from clients. Remove the
`x-custom-sql-*` headers from the public surface, or gate them behind an
explicit server-side allowlist per endpoint. Route client filtering through
the structured `FilterOption` path, which validates column names and binds
values. If raw fragments have to stay for trusted internal callers:
- parse them properly (for example with `pg_query_go`) and allow only column
references, literals and comparison operators;
- reject subqueries and function calls;
- verify that parentheses are balanced outside string literals;
- apply server-side security predicates last, as a wrapper
`WHERE (server) AND (client)`, and never let `WhereOr` attach at top level.
---
### 2. Critical — CORS reflects every Origin with credentials
`cors.go:117-155`:
```go
origin := r.Header("Origin")
if origin == "" {
origin = "*"
} else { ... Vary: Origin }
w.SetHeader("Access-Control-Allow-Origin", origin)
...
requestedHeaders := r.Header("Access-Control-Request-Headers")
if requestedHeaders != "" {
w.SetHeader("Access-Control-Allow-Headers", requestedHeaders)
}
...
if origin != "*" {
w.SetHeader("Access-Control-Allow-Credentials", "true")
}
```
`DefaultCORSConfig` (`cors.go:19-48`) carefully builds `AllowedOrigins` from the
server config, and `SetCORSHeaders` **never reads it**. Any site the victim
visits can make credentialed cross-origin requests and read the responses.
Allowed request headers are also reflected, so `Authorization` and every
`X-Custom-SQL-*` header pass preflight. `SetCORSHeaders` is called on every
route in `resolvespec/resolvespec.go` (lines 56-347), and `restheadspec` follows
the same pattern.
**Failure scenario.** A user logged in through cookie auth (`SetSessionCookie`/`GetSessionCookie`,
`pkg/security/middleware.go:512-540`) visits `evil.example`. Its script calls
`fetch("https://api/…/users", {credentials:"include"})` and reads every record
the user can see. Combined with finding 1, it can read other tenants' records
too.
**Recommendation.** Send `Allow-Origin: <origin>` and `Allow-Credentials`
only when `origin` is in `config.AllowedOrigins`, matched exactly. Otherwise
omit the CORS headers. Check requested headers against `AllowedHeaders`
instead of echoing them. Build `exposeHeaders` in a fresh slice: `append` onto
`config.AllowedHeaders` can write into a shared backing array.
---
### 3. High — Sort validation bypasses
`validation.go:271-301`:
```go
foundJoin := false
for _, j := range options.JoinAliases {
if strings.Contains(sort.Column, j) { // j may be ""
```
`restheadspec/headers.go:674-678` deliberately appends `""` to `JoinAliases`
when `extractJoinAlias` can't find an alias (for example
`LEFT JOIN t ON …` with no alias, or a LATERAL join without one).
`strings.Contains(x, "")` is always `true`, so **any** sort string is
accepted. `restheadspec/handler.go:790-793` then passes anything containing a
`.` or wrapped in `(…)` to `OrderExpr` **verbatim**. Even with a real alias,
the check is a substring test, so a sort like `j.id, (select …)` passes for
alias `j`.
Separately, `(…)` sort expressions are checked by `IsSafeSortExpression`
(`validation.go:381-427`), another denylist. It blocks DML keywords, comments
and `;`, but allows subqueries and functions.
Probe results: with `JoinAliases: [""]`, sort `x.id, (select pg_sleep(10))`
was kept. With no joins, sort `(select passwd from pg_shadow limit 1)` was kept.
**Failure scenario.** A client sends a custom join with no alias plus an
arbitrary ORDER BY expression, which gives injection in ORDER BY: time-based
DoS, or data extraction via `ORDER BY (CASE WHEN (subquery) THEN a ELSE b END)`.
**Recommendation.** Skip empty aliases. Match `alias + "."` as a prefix, then
validate the column after the dot against the joined table. Drop client
supplied sort *expressions*, or restrict them to a server-registered set
(the `cql` computed columns already provide this).
---
### 4. High — Nested CRUD modifies arbitrary related rows
`recursive_crud.go:64-67, 144-196, 344-380, 395-520`.
- Children are updated with `UPDATE <related> SET … WHERE pk = ?`, deleted with
`DELETE FROM <related> WHERE pk = ?`, and both use the child's PK **from the
request body**. Nothing checks that the child belongs to the parent being
written or to the caller's tenant.
- For updates, the parent's FK is injected into the child data
(`recursive_crud.go:495-520`), so updating a foreign child **moves it under
the attacker's parent** as well.
- `_request` (`recursive_crud.go:64-67, 205-212`) lets the client choose
`insert`/`update`/`delete` for each nested object, independent of the HTTP
method or the operation the top-level handler authorised.
- These statements go straight to `p.db`, so the spec handlers' Before*/After*
hooks, and any row-security or audit hooks, don't run for nested rows.
**Failure scenario.** A client sends a `PUT /orders/1` whose body includes
`"lines": [{"id": 9999, "_request": "delete"}]`. Row 9999 of `order_lines` is
deleted even if it belongs to another customer's order.
**Recommendation.** For has-many and has-one children, add
`AND <fk> = <parentID>` to update and delete statements, and treat
`RowsAffected() == 0` as a forbidden or not-found error. Run the same hook
chain (including row security) for nested rows. Allow `_request` only for
operations the top-level request is authorised to perform.
---
### 5. High — Reconnect logic is partial, and it triggers pool destruction
`adapters/database/bun.go:131-143, 167-186, 226-273, 1298-1318`;
`pgsql.go:58-79, 81, 134, 160, 220`; `gorm.go:55, 122-134`.
- **Coverage is uneven.** `BunAdapter` only retries after reconnecting in
`Exec`, `Query`, `BeginTx` and `RunInTransaction`. `NewSelect`, `NewInsert`,
`NewUpdate` and `NewDelete` capture `getDB()` once, and `BunSelectQuery.Scan`,
`ScanModel`, `Count` and `Exists` call bun directly with no retry. Those are
the paths every read handler uses. `PgSQLAdapter` query builders have no
reconnect either. Only `GormAdapter` wires `reconnect` into its
select, insert, update and delete builders.
- **Where it exists, it's harmful.** `reconnectDB` calls the dbmanager factory,
which runs `sqlConnection.Reconnect` and closes the pool shared by every
other adapter and handle (`dbmanager.audit.md` findings 1–2). Concurrent
failures each call the factory.
- **Detection is a substring match.** `isDBClosed` (`pgsql.go:72`) matches
`"sql: database is closed"`. That only happens *after* someone closed the
pool, so the reconnect mechanism mainly exists to recover from damage it
causes itself. It does nothing for the real idle-socket failure (a hang, or
`driver.ErrBadConn`, which `database/sql` already retries).
- `BunAdapter.CommitTx` / `RollbackTx` (`bun.go:239-249`) return `nil` without
doing anything. A caller using the `BeginTx`-less path gets
"committed" when nothing happened. `BunTxAdapter` is correct.
**Recommendation.** Remove adapter-level reconnect entirely and rely on
`database/sql`'s pool (see the fix order in `dbmanager.audit.md`). Make
`BunAdapter.CommitTx`/`RollbackTx` return an explicit
"not in a transaction" error. Add a per-query `context.WithTimeout` in the
adapters as the single place where query deadlines are enforced.
---
### 6. Medium — `SanitizeWhereClause` fails open and has false positives
`sql_helpers.go:176-179`:
```go
if err := validateWhereClauseSecurity(where); err != nil {
logger.Debug("Security validation failed for WHERE clause: %v", err)
return ""
}
```
Every caller treats `""` as "no filter" and skips `query.Where`. So a clause
the sanitizer rejects is **removed**, and the request runs unfiltered instead
of failing. The denylist is a substring match on the whole clause, string
literals included, so ordinary filters trip it. The probe showed
`status = 'awaiting update approval'` and `last_update > '2020-01-01'` both
returning `""`, which gives an unfiltered list. For a preload `where` or a
cursor filter, that means returning rows the client asked to exclude, or
breaking pagination.
**Recommendation.** Return an error and make the handler respond `400`. Never
turn a rejected filter into "no filter".
---
### 7. Medium — `cql*` columns bypass column validation
`validation.go:107-110` accepts any column that starts with `cql`
(case-insensitive), with no further checks. The probe showed
`IsValidColumn("cql1); drop")` returning `true`. The computed-column mechanism
only ever generates `cql1…cqlN` (`restheadspec/headers.go:871, 1380`). Whether a
client-supplied `cql…` string reaches SQL unquoted depends on the downstream
handler (`restheadspec/handler.go:490-509`, `cursor.go:189`). The validator
shouldn't be where that decision is made.
**Recommendation.** Accept only `^cql[0-9]+$`, and only when that computed
column was actually registered for the request.
---
### 8. Medium — Unbounded request body reads
`adapters/router/mux.go:101-115` uses `io.ReadAll(h.req.Body)` with no
`http.MaxBytesReader`, and `bunrouter.go:102-115` delegates to it. The
request-size middleware exists but isn't mounted (`_CROSS-CUTTING.audit.md`
X10), so a single request can make the process buffer gigabytes.
**Recommendation.** Wrap the body in `http.MaxBytesReader` inside the adapter,
with a configurable limit (for example 10 MB) and a sensible default.
---
### 9. Medium — Sensitive data in error logs
- `bun.go:1311-1315` (and the equivalent in `ScanModel`/`Count`, and in
`pgsql.go` / `gorm.go`) logs `b.query.String()`, the SQL with **all argument
values interpolated**, at `Error` on every failed query. That includes
filter values, emails, and tokens used as lookup keys.
- `recursive_crud.go` logs `data=%+v` (whole rows, including password or secret
columns) at `Error` on every failed nested write (lines 121, 153, 312, 342,
352, 509, 528, 550).
- `logger.Error` is forwarded to Sentry unscrubbed (`_CROSS-CUTTING.audit.md`
X8), and a hostile client can trigger failing queries at will.
**Recommendation.** Log the query with placeholders, not interpolated. Log
column names, not values. Put full dumps behind a debug flag.
---
### 10. Low — Fragile SQL string rewriting
- `reEmptyCompMid` / `reEmptyCompEnd` (`sql_helpers.go:66-80`) run over the
whole SQL string, including string literals and subqueries, and silently
delete text that matches `col = and`. That can change a query's meaning.
- The quote tracking in `splitByAND` / `findOperatorOutsideParentheses` /
`stripWrappingParens` toggles on every `'`, so an escaped `''` inside a
literal flips the state.
- `qualifyColumnInCondition` (`sql_helpers.go:751-760`) compiles a regex on
every call in the per-request path.
---
### 11. Low — Unsynchronised adapter field reads
`BunAdapter` protects `db` with `dbMu` in `getDB`/`reconnectDB`. But
`NewSelect` stores `db: b.db` (`bun.go:170`, used for count queries), and
`DriverName` reads `b.db` (`bun.go:280`), both without the lock.
`PgSQLAdapter.GetUnderlyingDB` (`pgsql.go:220`) does the same. These are data
races with `reconnectDB`, and `-race` would flag them
(`_CROSS-CUTTING.audit.md` X1). In practice the count query can run against
the old, closed pool.
---
### 12. Info — JSON column parsing is sound
`json_column.go` / `json_condition.go` are a good model for how the rest of
this package should handle client input:
- The base column must match `^[A-Za-z_][A-Za-z0-9_]*$` and is quoted with
`QuoteIdent`.
- Casts go through an allowlist.
- The JSON path is always bound as one `?::text[]` parameter, with depth and
segment-size limits.
- The dotted shorthand counts as JSON only when reflection confirms that the
base is a JSON column.
- The alias is validated and quoted.
No findings.
---
## Panic handling
The adapter methods (`Scan`, `ScanModel`, `Count`, `Exec`, `Query`,
`RunInTransaction`) recover and convert panics with `logger.HandlePanic`.
`PgSQLAdapter.RunInTransaction` rolls back on panic before re-raising or
converting it. `BunAdapter.RunInTransaction` relies on bun's `RunInTx`, which
also rolls back. No panic paths were found in `sql_helpers.go`,
`validation.go` or the JSON parser that are reachable from client input.
Slicing is length-guarded. `recursive_crud.go` recurses over the model's
relation graph. For self-referential models, depth is bounded only by the
JSON decoder's nesting limit, which makes it a slowness issue rather than a
crash.
## Test coverage
`sql_helpers_test.go` and `validation_test.go` test the *intended* behaviour
of the sanitizer and validator. None of them test hostile inputs. Each probe
case in findings 1, 3, 6 and 7 is a one-line table entry and should be added
as a regression test that asserts rejection. `cors.go` and `recursive_crud.go`
have no security-focused tests.
+419
View File
@@ -0,0 +1,419 @@
# Audit — `pkg/config`
- **Date:** 2026-09-29
- **Scope:** `pkg/config/{config,dbmanager,manager,paths,server}.go` (1023 LOC source, 608 LOC tests)
- **Axes:** thread locking/waiting · slowness · security · panic handling & logging
- **Threat model:** hostile internet client. Config itself is operator-controlled, so the security
focus here is **insecure defaults that the internet-facing layers inherit**, plus secret handling.
## Summary
Viper-backed configuration with a singleton `Manager`, a large `setDefaults` table, and per-section
validators. Two serious issues:
1. **`Manager` is a data race by construction.** It wraps a `*viper.Viper`, which has **no internal
locking** (verified: no `sync.Mutex`/`RWMutex` anywhere in `viper@v1.21.0/viper.go`'s `Viper`
struct), and exposes `Get`/`Set` as concurrently-callable methods on an unsynchronised lazy
singleton. A concurrent `Set` + `Get` is a concurrent map write → **`fatal error`, not a
recoverable panic**.
2. **The default configuration is insecure on every axis that matters** — wildcard CORS,
`sslmode=disable`, `user: postgres` with a blank password — and `Load()` silently succeeds when
no config file is found, so a misdeployment lands on exactly those defaults with no warning.
| # | Severity | Axis | Finding |
|---|----------|------|---------|
| 1 | **Critical** | Locking | `Manager.Set`/`Get` over a lock-free `*viper.Viper` → concurrent map write → process-fatal |
| 2 | **High** | Locking | `GetConfigManager()` is an unsynchronised lazy singleton; `NewManager()` also clobbers the global as a side effect |
| 3 | **High** | Security | **OPEN (deferred)** Insecure defaults: `cors.allowed_origins: ["*"]`, `allowed_headers: ["*"]`, `sslmode: disable`, `user: postgres` + blank password |
| 4 | **High** | Security | `SaveConfig` writes all secrets in plaintext at mode `0644` (viper default, never overridden) |
| 5 | Medium | Security | `AddConfigPath(".")` is searched first — CWD config injection |
| 6 | Medium | Observability | `Load()` swallows `ConfigFileNotFoundError` with no log at all |
| 7 | Medium | Correctness | `PathsConfig.Set` on a nil map panics; every sibling method nil-guards |
| 8 | Medium | Locking | `PathsConfig` is a bare `map[string]string` with a mutating `Set` — concurrent access is process-fatal |
| 9 | Medium | Slowness | `GetIPs()` does an uncontexted `net.LookupIP` — blocks on the resolver timeout |
| 10 | Medium | Correctness | `SetConfig` does a pointless `Unmarshal` into a discarded map whose error fails the call |
| 11 | Low | Panic | `GetIPs()` recovers to `fmt.Println`, bypassing the logger, and returns zeroed named results |
| 12 | Low | Security | No validation of `middleware.*` / `event_broker.worker_count` — `0` workers is accepted |
| 13 | Low | Correctness | `ServersConfig.GetDefault()` returns a pointer to a copy of a map value |
| 14 | Low | Security | `PathsConfig.Join` does not confine the result to the base path |
## Resolution status (2026-09-30)
- **#1** — Fixed: `sync.RWMutex` guards every viper access, options included
- **#2** — Fixed: mutex-guarded singleton; `NewManager` no longer touches the global (new `SetConfigManager` publishes explicitly)
- **#4** — Fixed: `SetConfigPermissions(0o600)` plus `chmod 0600` after write (secrets are not stripped)
- **#5** — Fixed: search order is `/etc/resolvespec`, `$HOME/.resolvespec`, `./config`, `.` (CWD last, not dropped)
- **#6** — Partly fixed: `ConfigFileUsed()` added; no log line because `pkg/config` cannot import `logger` (import cycle)
- **#7** — Fixed: `Set` has a pointer receiver and allocates
- **#8** — Not fixed: still a bare map; `Set` documented as not concurrency-safe
- **#9** — Fixed: `LookupIPAddr` with a 2s timeout, fallback normalised to bare IPs and populates the slice
- **#10** — Fixed: dead `Unmarshal` removed, `SetConfig` is atomic
- **#11** — Fixed: recover removed (nothing in the function can panic)
- **#12** — Partly fixed: `Config.Validate()` added, but it is not called from `GetConfig()`. The `*` CORS+credentials check is not implemented
- **#13** — Documented only: `GetDefault` returns a pointer to a copy
- **#14** — Fixed: `Join` errors if the result escapes the base
- **#3** — Open: default flips deferred by decision (breaking change).
- Tests: `pkg/config/hardening_test.go`.
---
## Findings
### 1. `Manager` exposes a lock-free viper as a concurrent API (Critical, Locking)
`manager.go:10-13`, `manager.go:133-158`
```go
type Manager struct {
v *viper.Viper
}
...
func (m *Manager) Get(key string) interface{} { return m.v.Get(key) }
func (m *Manager) GetString(key string) string { return m.v.GetString(key) }
func (m *Manager) Set(key string, value interface{}) { m.v.Set(key, value) }
```
`viper.Viper` carries its configuration in plain maps (`override`, `config`, `defaults`, `aliases`,
…) and has **no mutex**. Verified against the module in use:
```
$ grep -n 'sync\.\|Lock()' $(go env GOMODCACHE)/github.com/spf13/viper@v1.21.0/viper.go
319: initWG := sync.WaitGroup{} # inside WatchConfig only
340: eventsWG := sync.WaitGroup{} # inside WatchConfig only
```
`Set` writes to `v.override`; `Get` reads across those maps. Because `GetConfigManager()` hands the
*same* `*Manager` to every caller, any code path that calls `Manager.Set` at runtime while another
goroutine reads config is a concurrent map read/write. Go's runtime detects this and issues
`fatal error: concurrent map read and map write` — which **`recover()` cannot catch**, so none of
the panic handlers elsewhere in the codebase will save the process.
This is latent-but-loaded: it needs one runtime `Set` to become a crash. `SetConfig`
(`manager.go:107-131`) performs eleven `m.v.Set` calls, so any dynamic reconfiguration triggers it.
**Recommendation:** add a `sync.RWMutex` to `Manager` and take it in every method that touches
`m.v` (including the `Option` functions at `manager.go:60-85`, which also mutate viper). Better:
load once into an immutable `*Config` at startup and pass that value around, keeping `Manager`
confined to startup.
### 2. Unsynchronised lazy singleton (High, Locking)
`manager.go:15-45`
```go
var configInstance *Manager
func GetConfigManager() *Manager {
if configInstance == nil {
configInstance = NewManager()
}
return configInstance
}
```
Classic check-then-act race: two concurrent first calls both see `nil`, both build a `Manager`,
and the two callers get *different* instances — so a `Set` through one is invisible through the
other. The unsynchronised pointer write races with the read.
Worse, `NewManager()` (`manager.go:27-45`) assigns `configInstance = &Manager{v: v}` at line 43 as
a **side effect**. So a caller who deliberately builds an isolated manager silently replaces the
global one, and `NewManagerWithOptions` (`manager.go:48-54`) publishes a half-configured manager to
the global *before* applying its options — another goroutine can observe the instance mid-mutation.
**Recommendation:** `sync.Once` for the singleton; remove the global assignment from `NewManager`.
### 3. Insecure-by-default configuration (High, Security)
`manager.go:203-247`:
```go
v.SetDefault("cors.allowed_origins", []string{"*"})
v.SetDefault("cors.allowed_headers", []string{"*"})
...
v.SetDefault("dbmanager.connections.default.user", "postgres")
v.SetDefault("dbmanager.connections.default.password", "")
v.SetDefault("dbmanager.connections.default.sslmode", "disable")
```
Each of these is inherited by an internet-facing layer:
- **`allowed_origins: ["*"]` + `allowed_headers: ["*"]`** — any origin may make cross-origin calls
with arbitrary headers. Whether this is exploitable depends on whether the CORS middleware also
sets `Access-Control-Allow-Credentials`; see `audit/pkg/middleware.audit.md` for that
determination. Even without credentials, wildcard origin plus wildcard headers defeats any
header-based CSRF defence and lets a malicious page read responses from a
network-position-authenticated deployment (IP allowlisted, mTLS-terminated, VPN).
- **`sslmode: disable`** — DB traffic unencrypted by default. Every row that crosses the wire,
including whatever the internet-facing handlers select, is plaintext on the network.
- **`user: postgres` with an empty password** — the default connection targets the PostgreSQL
superuser. Combined with the identifier-handling concerns in
`audit/pkg/common.audit.md` / `audit/pkg/restheadspec.audit.md`, running as superuser removes the
last line of defence (least-privilege) against a query-construction bug.
Because of finding 6, a deployment with a missing or misnamed config file runs on **all** of these
simultaneously and reports success.
**Recommendation:** default to `sslmode: require`, no default DB user/password (fail loudly if
unset), and `cors.allowed_origins: []` with wildcard requiring an explicit opt-in. Add a
`Config.Validate()` that refuses `allowed_origins: ["*"]` together with credentials.
### 4. `SaveConfig` writes secrets in plaintext at 0644 (High, Security)
`manager.go:160-166`
```go
func (m *Manager) SaveConfig(path string) error {
if err := m.v.WriteConfigAs(path); err != nil { ... }
}
```
`WriteConfigAs` serialises the **entire** merged configuration. That includes
`dbmanager.connections.*.password`, `cache.redis.password`, `event_broker.redis.password` and
`error_tracking.dsn` (a Sentry DSN is a credential).
Viper writes with `v.configPermissions`, which defaults to `0o644`
(`viper@v1.21.0/viper.go:198`). `SetConfigPermissions` is **never called anywhere in this repo**
(verified by grep), so the file is world-readable. Any local user or any other container sharing
the mount can read the DB superuser password.
**Recommendation:** call `v.SetConfigPermissions(0o600)` in `NewManager`; better, strip secret keys
before writing and document that secrets come from env/secret-manager only.
### 5. Current-working-directory config injection (Medium, Security)
`manager.go:32-36`
```go
v.AddConfigPath(".")
v.AddConfigPath("./config")
v.AddConfigPath("/etc/resolvespec")
v.AddConfigPath("$HOME/.resolvespec")
```
Viper searches these **in order** and takes the first hit, so `./config.yaml` wins over
`/etc/resolvespec/config.yaml`. For a daemon this is backwards: the CWD is the least trustworthy of
the four. If the process is ever started with its CWD in a shared or user-writable directory (a
tmp dir, a bind-mounted volume, `/` in some container setups), an attacker with local write
capability redirects the DB connection, disables TLS, or points `error_tracking.dsn` at their own
collector — turning finding 1 of `audit/pkg/errortracking.audit.md` into a full exfiltration path.
**Recommendation:** search `/etc/resolvespec` first, drop `"."` from the default list (keep it
available via `WithConfigPath`), and log the resolved path at startup (`v.ConfigFileUsed()`).
### 6. `Load()` is silent about a missing config file (Medium, Observability)
`manager.go:87-97`
```go
if err := m.v.ReadInConfig(); err != nil {
if _, ok := err.(viper.ConfigFileNotFoundError); !ok {
return fmt.Errorf("error reading config file: %w", err)
}
// Config file not found; will rely on defaults and env vars
}
return nil
```
The comment is the only trace. No log line, no returned indicator, no `ConfigFileUsed()` report.
A typo in the filename, a wrong working directory, or a container that forgot to mount the
ConfigMap is indistinguishable from a deliberate defaults-only run — and the defaults are the ones
in finding 3.
**Recommendation:** log at info level whether a file was used and which one; expose
`ConfigFileUsed()` on `Manager` so startup can print it.
### 7. `PathsConfig.Set` panics on a nil map (Medium, Panic handling)
`paths.go:38-40`
```go
func (pc PathsConfig) Set(name, path string) {
pc[name] = path
}
```
`PathsConfig` is `map[string]string` (`config.go:200`). `Get`, `GetOrDefault`, `Has` and `List` all
begin with `if pc == nil`. `Set` does not — and assignment to a nil map is
`panic: assignment to entry in nil map`.
`Config.Paths` is populated by `mapstructure`, which leaves the map nil when the `paths` key is
absent from the file. `setDefaults` does register `paths.data_dir` etc. (`manager.go:249-253`), so
the map is non-nil on the normal `GetConfig()` path — but a `Config` built in code
(`config.Config{}`) or produced by a partial unmarshal has a nil `Paths`, and `Set` on it panics.
Nothing in `pkg/` currently calls `Set` (verified by grep), so this is a latent API defect.
**Recommendation:** nil-guard consistently, or change the receiver to `*PathsConfig` so `Set` can
allocate.
### 8. `PathsConfig` has no synchronisation (Medium, Locking)
Same type: a bare map with a mutating `Set` and reading `Get`/`Has`/`List`/`EnsureDir`/`AbsPath`/
`Join`. If any consumer calls `Set` at runtime while request handlers resolve paths, that is a
concurrent map write — again the **unrecoverable** `fatal error` class, not a panic.
Currently unused outside the package, so severity is capped at Medium. If the intent is a runtime
path registry, it needs a mutex and an unexported map.
### 9. `GetIPs()` blocks on an uncontexted DNS lookup (Medium, Slowness)
`server.go:113-149`
```go
hostname, _ = os.Hostname()
...
addrs, err := net.LookupIP(hostname)
```
`net.LookupIP` has no context and no timeout override — it blocks for the resolver's own timeout,
which on a misconfigured or slow-resolver host is 5 s per attempt and up to ~15–20 s with retries
across `/etc/resolv.conf` entries. In a container whose hostname is not in DNS (the normal case)
this fails, but only *after* the resolver gives up.
There is no caller in `pkg/` today, so it is not on the request path yet. It is exported and
named like a utility, so the risk is that it lands on one.
Secondary correctness problem in the same function: the fallback branch (`server.go:139-147`)
appends `a.String()` for a `net.Addr` from `net.InterfaceAddrs()`, which renders as CIDR
(`192.168.1.5/24`), into the same comma-joined string that the primary branch fills with bare IPs.
Consumers get two formats from one field. That branch also never appends to `ipaddrlist`, so the
third return value is empty whenever the fallback is taken.
**Recommendation:** `net.DefaultResolver.LookupIPAddr(ctx, host)` with a short deadline; cache the
result; normalise the fallback to bare IPs via `net.Addr.(*net.IPNet).IP`.
### 10. `SetConfig` does dead work that can fail the call (Medium, Correctness)
`manager.go:107-131`
```go
configMap := make(map[string]interface{})
if err := m.v.Unmarshal(&configMap); err != nil {
return fmt.Errorf("failed to prepare config map: %w", err)
}
// configMap is never read again
m.v.Set("servers", cfg.Servers)
...
```
`configMap` is written and then never used. The comment says "Marshal the config to a map structure
that viper can use", but it unmarshals *viper's current state* into a throwaway map — it has
nothing to do with `cfg`. The only effect is that a decode error in the **existing** config makes
`SetConfig` fail for no reason. It also does a full reflective decode of the whole config tree on
every call.
Note also that `SetConfig` stores Go structs into viper via `Set`, and the eleven `Set` calls are
not atomic — a concurrent `GetConfig()` observes a torn config (new `servers`, old `cors`), on top
of finding 1's race.
**Recommendation:** delete the `configMap` block.
### 11. `GetIPs()` panic handling bypasses the logger (Low, Panic handling)
`server.go:114-118`
```go
defer func() {
if err := recover(); err != nil {
fmt.Println("Recovered in GetIPs", err)
}
}()
```
- Writes to stdout with `fmt.Println` rather than `logger.Error`/`logger.HandlePanic`, so the event
never reaches the error tracker and is invisible to structured log collection.
- No stack trace captured.
- The function's results are named (`hostname, ipList string, ipNetList []net.IP`) but the body
builds `iplist`/`ipaddrlist` **locals** and only assigns via the `return` statements. On a panic,
the deferred recover swallows it and the function returns the *zero* named values — `ipNetList`
is nil rather than the empty slice callers might expect. Silent empty success.
`pkg/config` is otherwise the only package outside `pkg/logger` that hand-rolls a recover instead
of using the shared helpers.
**Recommendation:** use `defer logger.CatchPanic("GetIPs")()`, or drop the recover — there is no
panicking operation in this function for it to catch.
### 12. No validation of numeric/limit settings (Low, Security)
`ServerInstanceConfig.Validate` (`server.go:37-68`) and `ServersConfig.Validate`
(`server.go:71-95`) are good — port range, mutually-exclusive TLS modes, cert/key pairing,
AutoTLS domains. But nothing validates:
- `middleware.rate_limit_rps` / `rate_limit_burst` — `0` disables rate limiting silently.
- `middleware.max_request_size` — `0` may mean unlimited depending on the middleware; see
`audit/pkg/middleware.audit.md`.
- `event_broker.worker_count` (default 10) — `0` means no consumers; see
`audit/pkg/eventbroker.audit.md` for whether that deadlocks publishers or drops events.
- `dbmanager.max_open_conns`, retry counts/delays — negative or zero values.
- `cors.allowed_origins: ["*"]` in combination with credentials.
There is also no top-level `Config.Validate()` that calls the section validators, so nothing
guarantees `ServersConfig.Validate` ever runs.
**Recommendation:** add `func (c *Config) Validate() error` that fans out to every section, and
call it from `GetConfig()`.
### 13. `GetDefault()` returns a pointer to a copy (Low, Correctness)
`server.go:98-110`
```go
instance, ok := sc.Instances[sc.DefaultServer]
...
return &instance, nil
```
`instance` is a copy of the map value. A caller that mutates through the returned pointer — which
the `*ServerInstanceConfig` receiver on `ApplyGlobalDefaults` (`server.go:12`) invites — changes
only the copy, and `sc.Instances` is unaffected. This is exactly the shape of bug where timeouts
appear to be applied but aren't.
**Recommendation:** make `Instances` a `map[string]*ServerInstanceConfig`, or return by value.
### 14. `PathsConfig.Join` does not confine to the base (Low, Security)
`paths.go:96-104`
```go
parts := append([]string{base}, elem...)
return filepath.Join(parts...), nil
```
`filepath.Join` calls `Clean`, which *resolves* `..` rather than rejecting it: `Join("data",
"../../etc/passwd")` returns `../etc/passwd`. Any consumer that passes a request-derived segment
gets directory traversal out of the configured base. No consumer does today, hence Low, but the
method's name promises confinement it does not provide.
**Recommendation:** after joining, verify `strings.HasPrefix(filepath.Clean(result), filepath.Clean(base)+string(os.PathSeparator))`, or use `os.Root`/`filepath.Localize` on the elements.
---
## What looks right
- `ServerInstanceConfig.Validate` / `ServersConfig.Validate` (`server.go:37-95`) are thorough:
port bounds, mutual exclusion of the three TLS modes, cert/key co-presence, AutoTLS domain
requirement, and a key-vs-`Name` consistency check on the instances map. This is the strongest
code in the package.
- `ApplyGlobalDefaults` (`server.go:12-32`) uses `*time.Duration` fields so "unset" is
distinguishable from "zero" — the right modelling choice, and it copies into a fresh local
before taking its address rather than aliasing the loop/parameter variable.
- `Load()` correctly distinguishes `ConfigFileNotFoundError` from real read errors instead of
treating every failure as fatal (the *silence* is the problem, not the branch).
- `SetEnvPrefix("RESOLVESPEC")` + `SetEnvKeyReplacer(".", "_")` + `AutomaticEnv`
(`manager.go:38-41`) is the correct trio for env overrides, and because every key has a
registered default, `AutomaticEnv` actually resolves nested keys — so secrets *can* be supplied
via env instead of the file. That's the mitigation for finding 4, and it should be documented as
the only supported way to pass secrets.
- The defaults table is comprehensive and one place — easy to review, which is how findings 3 and
12 were found.
- Test coverage is reasonable for a config package (608 LOC of tests against 1023 of source),
though it does not cover concurrency, `SaveConfig` permissions, or `PathsConfig.Set`.
## Suggested follow-up
1. Lock `Manager` or make config immutable after load (findings 1, 2). Until then, treat
`Manager.Set` as unsafe to call after startup and consider removing it from the public API.
2. Flip the insecure defaults and add `Config.Validate()` (findings 3, 12).
3. `SetConfigPermissions(0o600)` and secret-stripping in `SaveConfig` (finding 4).
4. Reorder the config search path and log the resolved file (findings 5, 6).
5. Delete the dead `Unmarshal` in `SetConfig` (finding 10).
+630
View File
@@ -0,0 +1,630 @@
# Audit: `pkg/dbmanager`
| | |
|---|---|
| **Package** | `github.com/bitechdev/ResolveSpec/pkg/dbmanager` (+ `providers/`) |
| **Files** | `config.go` (489), `connection.go` (722), `manager.go` (401), `metrics.go` (136), `errors.go` (82), `factory.go` (67), `providers/postgres.go` (231), `providers/postgres_listener.go` (401), `providers/sqlite.go` (216), `providers/mongodb.go` (214), `providers/mssql.go` (184), `providers/existing_db.go` (111), `providers/provider.go` (89); tests `factory_test.go` (369), `manager_test.go` (290), `providers/existing_db_test.go` (194), `providers/postgres_listener_example_test.go` (229) |
| **Audit date** | 2026-09-30 |
| **Axes** | thread locking/waiting, slowness, security, panic handling & logging |
| **Threat model** | hostile internet client; request bodies, headers, query params, schema/table/column names all attacker-controlled |
| **Depth** | deep (hot package; every request's DB handle comes from here). Several findings were checked with a throw-away probe test against SQLite, and the probe was deleted afterwards |
## Summary
`pkg/dbmanager` owns every database pool in the process. It wraps a
`*sql.DB` (or a `mongo.Client`) in a `sqlConnection` and hands out lazily-built
`*bun.DB`, `*gorm.DB`, raw `*sql.DB` and `common.Database` adapters over it. A
background health checker pings each connection every 15 s, and it can
**reconnect**, which closes the pool and opens a new one.
This audit was started to answer one question: **"why does a database
connection that has been idle for a while become unusable?"** Several defects
in this package combine to give exactly that symptom. They are findings 1–5,
and the [Idle-connection failure chain](#idle-connection-failure-chain) section
below puts them together.
The root design problem is that **`Reconnect` destroys the shared `*sql.DB`**.
`*sql.DB` is already a self-healing pool: it throws away bad connections and
dials new ones. So "reconnecting" a pool is almost never needed, and here it
has a large blast radius. Every `*bun.DB`, `*gorm.DB` and `*sql.DB` handed out
before the reconnect now points at a closed pool, and it stays closed. Only the
`common.Database` adapters carry a factory that can re-fetch a handle, and even
they only use it on a subset of code paths (see `common.audit.md` finding 5).
Those adapter factories also *trigger* `Reconnect` themselves, so one stale
handle closes the pool for everyone else. `Reconnect` isn't atomic, so
concurrent callers turn this into a storm.
The other major theme is **missing client-side deadlines**. `QueryTimeout` is
only ever sent to the server as `statement_timeout`, which does nothing when
the TCP peer has vanished. No `context.WithTimeout` is applied to request
queries, and pgx's dialer sets no `TCP_USER_TIMEOUT`. So the first query on a
pooled connection whose peer silently disappeared (NAT/firewall idle drop,
failover, a pgbouncer restart) can block for minutes. One Close path does this
while holding the connection's write lock, which stalls every request.
## Findings
| # | Severity | Axis | Finding | Status |
|---|---|---|---|---|
| 1 | **Critical** | locking / availability | `Reconnect` closes the shared `*sql.DB`, so every `*bun.DB` / `*gorm.DB` / `*sql.DB` handed out earlier is permanently dead ("sql: database is closed") | Fixed |
| 2 | **High** | locking | Adapter reconnect factories call `Reconnect` on the *shared* connection, and `Reconnect` is not atomic, so one stale handle starts a reconnect storm that repeatedly closes the pool under in-flight requests | Fixed |
| 3 | **High** | slowness / locking | `sqlConnection.HealthCheck` holds the write lock across a network ping for up to 5 s; every `Bun()`/`GORM()`/`Native()`/`Database()`/`Stats()` call blocks for that time | Fixed |
| 4 | **High** | slowness | No client-side query deadline and no `TCP_USER_TIMEOUT`: a query on a silently-dead idle socket blocks for minutes (up to about 15 min); `QueryTimeout` is server-side only, and is forced to at least 2 min | Fixed |
| 5 | **High** | locking / slowness | `PostgresListener.Close` runs `UNLISTEN` with `context.Background()` while `sqlConnection.mu` (write), `PostgresProvider.mu` and `listener.mu` are all held; on a dead socket this freezes every request for minutes | Fixed |
| 6 | **High** | locking / leak | `PostgresListener.Connect` starts a new goroutine pair on every (re)connect; the old pair keeps running, so two loops call `WaitForNotification` on one `pgx.Conn` concurrently, which triggers more reconnects | Fixed |
| 7 | **High** | panic handling | `Connect → Close → Connect → Close` panics with "close of closed channel"; after the first cycle the health checker also exits immediately and silently | Fixed |
| 8 | **Medium** | availability | SQLite: `:memory:` with a 25-connection pool gives every connection its own empty database, and `ConnMaxIdleTime` then silently discards data; `busy_timeout` / WAL pragmas are applied to only one pooled connection | Fixed |
| 9 | **Medium** | availability | Partial failure in `sqlConnection.Close` leaves `connected=true` over a closed pool; partial failure in `Manager.Connect` leaks the connections already opened | Fixed |
| 10 | **Medium** | security | DSN builders concatenate unescaped credentials (postgres key=value, mssql/mongo URLs); `sslmode` defaults to `disable` | Fixed |
| 11 | **Medium** | config | Several config knobs are ignored or impossible to turn off: `EnableAutoReconnect`, `HealthCheckInterval`, `RetryAttempts`/`RetryDelay`/`RetryMaxDelay`, SQLite `_timeout`, and `statement_timeout` when a DSN is given | Fixed |
| 12 | **Medium** | locking | `Manager.Connect` holds `m.mu` across every network dial (up to 3 retries × `ConnectTimeout` per connection) | Fixed |
| 13 | **Low** | observability | `PublishMetrics` / `RecordReconnectAttempt` are never called, so all dbmanager metrics are permanently zero; `*_total` metrics are gauges | Fixed |
| 14 | **Low** | correctness | `Bun()`/`GORM()` do not check `connected`; `getNativeAdapter` uses `PgSQLAdapter` for SQLite and MSSQL; `ExistingDBProvider` applies no pool settings and closes the caller's DB | Fixed (partly, see notes) |
| 15 | **Low** | logging | `Close` / `performHealthCheck` pass key-value pairs to the printf-style logger, which produces `%!(EXTRA ...)` output; `ResetInstance` discards the close error | Fixed |
## Remediation status
Implemented 2026-09-30. `go build ./...` and `go test -race ./pkg/dbmanager/...`
pass. The Postgres behaviour was also verified against a live server (tests are
skipped unless `PG_LIVE=1` / `PG_RESTART_DIR` is set).
**Design decisions taken**
- No automatic reconnect. Adapter factories and the health checker never close
the pool; they only re-fetch the current handle. `*sql.DB` replaces bad
connections itself. `EnableAutoReconnect` is deprecated and ignored.
- `Reconnect` is atomic (one critical section) and operator-only. On PostgreSQL
it goes through a custom `driver.Connector` (`providers/pgconnector.go`): it
bumps a generation, stale pooled connections are discarded, and the `*sql.DB`
is never closed, so held Bun/GORM/`*sql.DB` handles keep working. Other
providers still close and reopen.
- Client-side deadlines are applied at the driver level rather than in the
adapters (a `context.WithTimeout` around a query is cancelled before the
caller has read the rows).
**Per finding**
1. Fixed. Postgres refresh keeps the pool; explicit `Reconnect` on other
providers still invalidates handles (documented in the README).
2. Fixed. Adapter factories no longer call `Reconnect`; `Reconnect` is a single
critical section under `lifecycleMu` + `mu`.
3. Fixed. The ping runs without `mu`; `lifecycleMu` (read) only keeps
`Close`/`Reconnect` from tearing the provider down mid-ping. Same for Mongo.
4. Fixed. TCP keepalive and `TCP_USER_TIMEOUT` (30 s, Linux) via `DialFunc`;
the reuse-time liveness ping is capped at 5 s; `statement_timeout` is set as
a runtime parameter so it also applies to a supplied DSN; the 2-minute floor
on `QueryTimeout` is removed. `SetConnMaxIdleTime` tuning remains a
configuration matter (documented in the README).
5. Fixed. Listener `Close` sends no `UNLISTEN`, closes with a 2 s bound, and
holds no lock across network I/O.
6. Fixed. Background goroutines start once (`sync.Once`); reconnect dials a
replacement, re-`LISTEN`s, then swaps it in; sleeps honour `ctx.Done()`.
Additionally, all use of the single `pgx.Conn` is serialised (`connMu`, 500 ms
notification poll), fixing "conn busy" from `Listen`/`Unlisten`/`Notify`, and
old connections are closed under `connMu` (a race found by the live test).
7. Fixed. Stop channel is created per start, guarded by `healthMu`; `Close` is
idempotent; `Connect` is idempotent.
8. Fixed. `:memory:` is pinned to one connection with no idle/lifetime limits;
`busy_timeout`/WAL are `_pragma` DSN parameters; `_timeout` and the dead
reconnect code are removed.
9. Fixed. `Close` always marks disconnected and returns joined errors;
`PostgresProvider.Close` closes the pool even if the listener fails;
`Manager.Connect` closes connections it opened when a later one fails.
10. Fixed. Postgres, MSSQL and Mongo DSNs are built as escaped URLs; default
`sslmode` is now `prefer` (was `disable`).
11. Fixed. Retry settings reach every provider; a negative
`HealthCheckInterval` disables the health checker; `EnableAutoReconnect`
deprecated; `statement_timeout` applies with a supplied DSN.
12. Fixed. `Manager.Connect` dials outside `m.mu` and publishes results under it.
13. Fixed. `PublishMetrics` runs on each health-check tick, `Reconnect` records
`RecordReconnectAttempt`, and the wait/closed metrics are true counters
(delta-tracked).
14. Partly fixed. `Bun()`/`GORM()` check `connected`; Mongo no longer maps
`MaxIdleConns` to `MinPoolSize`. `ExistingDBProvider`: `Close` is now a no-op
that logs a warning (the caller owns the `*sql.DB`; the connection's `Close`
also skips `bun.DB.Close`), and `Reconnect` only pings. Pool settings are
still not applied to a caller-owned pool. The `getNativeAdapter` claim was
stale: the adapter already receives the driver name; the three duplicate
cases were merged. Mongo `Stats()` is still empty.
15. Fixed. Printf-style logger calls corrected; `ResetInstance` logs the close
error. Unscrubbed driver errors in Sentry (X8) are not addressed here.
**Behaviour changes**
- Removed tests that closed the pool from outside and expected an adapter to
swap in a new one (three adapter tests, and the health-check reconnect test,
now asserting it never reconnects).
- `sslmode` default `prefer`; `NewConnectionFromDB` connections are no longer
closed by the manager.
**Regression tests added:** `lifecycle_test.go` (double Connect/Close cycle,
idempotent Connect, concurrent Reconnect, adapter factory leaves pool open,
accessors not blocked by health check, Close marks disconnected, existing-DB
Reconnect/Close leave the caller's pool open), `config_dsn_test.go`,
`providers/pgconnector_test.go`, `pg_live_test.go` (refresh keeps handles,
listener Listen/Notify) and `restart_live_test.go` (server crash and restart).
---
## Idle-connection failure chain
This is how findings 1–5 combine into "the connection sat idle and then could
not be used":
1. The app is idle. A NAT, firewall, load balancer or pgbouncer silently drops
the idle TCP flows. No FIN or RST reaches the process.
2. The next request takes a pooled connection. pgx's `ResetSession` pings it
because it has been idle for more than 1 s, and that ping uses the request
ctx, **which has no deadline** (finding 4). The write goes into the kernel
buffer and the read blocks until TCP retransmission gives up, which can
take minutes.
Meanwhile the health checker's 5 s ping times out and holds `c.mu`
**exclusively** for the whole time (finding 3), so every request trying to
get a handle queues behind it.
3. Eventually something returns "sql: database is closed" or
`ErrConnectionClosed`. That can be an adapter that hit a closed pool, or a
partial `Close` (finding 9). An adapter's `dbFactory` or the health checker
then calls `Reconnect` (finding 2).
4. `Reconnect` closes the `*sql.DB` (finding 1). If the Postgres listener has
subscriptions, `Close` first sends `UNLISTEN` on its own dead socket with no
deadline, still holding the write lock (finding 5), which freezes the
process again.
5. When the reconnect completes, every handle captured before it is
permanently broken. That includes the `*gorm.DB` given to
`resolvespec.NewHandlerWithGORM` in `cmd/testserver/main.go:142,56`, any
`*bun.DB` passed to `NewHandlerWithBun`, and every Bun `NewSelect`/`NewInsert`
path. **From this point on, every request that goes through those handles
fails until the process is restarted.** Concurrent failures run their own
`Reconnect`s, and each one closes the pool the previous one just opened
(finding 2).
### Fix order for this symptom
1. **Stop closing the pool to recover from connection errors.** Remove
`WithDBFactory(c.reopen*ForAdapter)` → `Reconnect`, and remove the
health-check → `Reconnect` path for SQL providers. `*sql.DB` already discards
bad connections (`driver.ErrBadConn`, `ResetSession`,
`SetConnMaxIdleTime`/`SetConnMaxLifetime`). Keep `Reconnect` for explicit
operator use only, and make it atomic (finding 2).
2. Give every request a deadline. Wrap the request ctx in
`context.WithTimeout(ctx, QueryTimeout)` in the adapters, or at the handler
boundary.
3. Set `SetConnMaxIdleTime` **below** the shortest idle timeout of any
middlebox (typically 60–240 s for cloud NATs and LBs) so idle connections
are recycled before they can be dropped silently. Also set TCP keepalive and
`TCP_USER_TIMEOUT` through a custom `pgconn.Config.DialFunc`.
4. Ping without the write lock (finding 3), and give the listener's `Close`
bounded ctxs (finding 5).
---
### 1. Critical — `Reconnect` kills every previously issued handle
`connection.go:129-160` (`Close`) and `connection.go:187-192` (`Reconnect`):
```go
func (c *sqlConnection) Close() error {
c.mu.Lock()
...
if c.bunDB != nil {
if err := c.bunDB.Close(); err != nil { // closes the shared *sql.DB
...
if err := c.provider.Close(); err != nil { // closes it again (idempotent)
...
c.nativeDB = nil
c.bunDB = nil
c.gormDB = nil
c.bunAdapter = nil
...
}
func (c *sqlConnection) Reconnect(ctx context.Context) error {
if err := c.Close(); err != nil {
return err
}
return c.Connect(ctx)
}
```
`Bun()`, `GORM()` and `Native()` return the handle itself, and callers keep
it: every spec package has a `NewHandlerWithGORM(*gorm.DB)` /
`NewHandlerWithBun(*bun.DB)` constructor, and `cmd/testserver/main.go:142` does
exactly this. After `Reconnect`, the cached fields are nilled, a new pool is
built, and the handles the callers hold point at a `*sql.DB` whose `closed`
flag is set forever.
Verified with a probe: I obtained `conn.GORM()`, called `conn.Reconnect(ctx)`,
then ran a query through the old handle. It returned
`sql: database is closed`, and a fresh `conn.GORM()` worked.
The comment in `manager.go:371-374` shows the authors already knew about this
("forcing Close()+Connect() here invalidates any cached ORM wrappers and callers
that still hold the old handle"). Their mitigation was to narrow *when* the
health checker reconnects. But the adapters' own `dbFactory` still reconnects
unconditionally (finding 2).
**Failure scenario.** Any event that triggers a reconnect turns every
long-lived handler into a permanent 500 generator: a single adapter query hitting
"database is closed", or a health check returning `ErrConnectionClosed`. The
process does not recover without a restart. The same thing happens after a
normal `Manager.Close()` + `Connect()` in tests or hot-reload code.
**Recommendation.** Treat the `*sql.DB` as immortal for the life of the
`sqlConnection`. Don't close it to "reconnect": `database/sql` already replaces
broken connections. If a real re-dial is ever needed (for example after
changing credentials), build the new pool, atomically swap it in, and close the
old one only after a grace period. Give the handles returned by
`Bun()`/`GORM()`/`Native()` stable identity; one way is a `driver.Connector`
that indirects to the current pool.
---
### 2. High — Adapter-triggered, non-atomic `Reconnect` causes a reconnect storm
`connection.go:362-397` and `connection.go:431/474/517-525`:
```go
func (c *sqlConnection) reconnectForAdapter() error {
...
return c.Reconnect(ctx) // Close() then Connect(): two separate lock scopes
}
...
WithDBFactory(c.reopenBunForAdapter).
```
The adapters (`pkg/common/adapters/database/bun.go:131`, `gorm.go`,
`pgsql.go`) call `dbFactory` whenever an operation returns an error that
matches `"sql: database is closed"`. So:
- **One stale handle closes the pool for everyone.** If an adapter holds a
`*sql.DB` from before a previous reconnect, its first query fails with
"database is closed". Its factory then calls `c.Reconnect`, which closes the
*current, healthy* pool that every other adapter and request is using right
now.
- **`Reconnect` isn't atomic.** `Close` and `Connect` each take `c.mu`
separately. Under N concurrent failures, one goroutine closes and reconnects
while the others either close the brand-new pool again or fail with
`already connected`. The probe used 20 concurrent `Reconnect`s: 9 returned
"already connected", and every successful reconnect closed the pool the
previous winner had just handed to its adapter. Each of those adapters then
sees "database is closed" on its next query, and the cycle continues.
**Failure scenario.** A burst of traffic arrives just after a reconnect. Each
in-flight request whose adapter still holds the old pool triggers another
`Reconnect`, and each of those closes the pool that the previous request
reopened. The service flaps until traffic stops.
**Recommendation.** Remove the adapter → `Reconnect` path (see finding 1). If
it is kept, make `Reconnect` a single critical section, and add a generation
counter: a caller that saw generation N only reconnects if the current
generation is still N; otherwise it just re-fetches the handle.
---
### 3. High — Health check holds the write lock across a network ping
`connection.go:163-185`:
```go
func (c *sqlConnection) HealthCheck(ctx context.Context) error {
c.mu.Lock() // exclusive
defer c.mu.Unlock()
...
if err := c.provider.HealthCheck(ctx); err != nil { // PingContext, 5 s timeout
```
Every handle accessor takes `c.mu.RLock()` first (`connection.go:199, 238, 271,
308, 335, 403, 441, 484`). While the health checker (every 15 s, `manager.go:348`)
is pinging, **every request that needs a DB handle waits**. On a healthy
network this is a few ms. On a dead idle socket it's the full 5 s ping timeout
(`providers/postgres.go:155`, inside a 10 s outer ctx).
Verified with a probe: while `c.mu` was held, `conn.Bun()` blocked for the whole
hold (200 ms in the test).
**Failure scenario.** A network blip or a silently dropped idle connection
makes the ping hang. Every 15 s the whole API pauses for up to 5 s. This fits
reports of "idle, then slow or unusable".
**Recommendation.** Snapshot `provider` under `RLock`, release the lock, ping,
then take the lock only to write `healthCheckStatus` / `lastHealthCheck`. Better
still, keep the status in an `atomic.Value`.
---
### 4. High — No client-side query deadline; `QueryTimeout` is server-side only and floored at 2 min
`config.go:223-228`:
```go
if cc.QueryTimeout == 0 {
cc.QueryTimeout = 2 * time.Minute
} else if cc.QueryTimeout < 2*time.Minute {
cc.QueryTimeout = 2 * time.Minute
}
```
`config.go:331-335` turns this into `statement_timeout=<ms>` in the Postgres DSN,
and it only does that when the DSN is *built*. A user-supplied `DSN` gets no
timeout at all. Nothing anywhere in the request path wraps ctx in a deadline.
`pkg/config`'s `query_timeout: 30s` default is silently raised to 2 min.
`statement_timeout` is enforced by the **server**, so it only helps if the
server is reachable. On a silently dropped connection:
- pgconn's default dialer is `&net.Dialer{}`: Go's default keepalive (15 s idle,
15 s interval, 9 probes) and **no `TCP_USER_TIMEOUT`**.
- Once a query has been written, there is unacknowledged data, so keepalive does
not apply. The socket then waits for TCP retransmission to give up
(`tcp_retries2`), which takes about 15 min on Linux defaults.
- `database/sql` calls pgx's `ResetSession`, which pings a connection that has
been idle for more than 1 s. That ping uses the **request ctx**, so with no
deadline it blocks just as long.
**Failure scenario.** An idle period longer than the NAT or LB idle timeout
causes the next request to hang for minutes rather than failing fast and being
retried on a fresh connection. With `MaxOpenConns` = 25, 25 such requests
exhaust the pool and every later request blocks on `db.conn()`.
**Recommendation.**
- Apply `context.WithTimeout(ctx, QueryTimeout)` in the adapters, or in a
handler middleware.
- Remove the 2-minute floor, and honour the configured value.
- Set `SetConnMaxIdleTime` below the middlebox idle timeout.
- Configure `pgconn.Config.DialFunc` with a `net.Dialer` that has `KeepAlive`
set and a `Control` func setting `TCP_USER_TIMEOUT` (for example 30 s).
- Apply `statement_timeout` through `RuntimeParams` so it also works with a
supplied DSN.
---
### 5. High — Listener `Close` does unbounded network I/O under three locks
`providers/postgres_listener.go:216-244`, reached from
`providers/postgres.go:116-126`, which is reached from `connection.go:147`:
```go
// sqlConnection.Close holds c.mu (write)
// PostgresProvider.Close holds p.mu
// PostgresListener.Close holds l.mu:
for channel := range l.channels {
_, _ = l.conn.Exec(context.Background(), fmt.Sprintf("UNLISTEN %s", ...))
}
err := l.conn.Close(context.Background())
```
If the listener's socket is dead, and it usually is in the situation that
triggers a reconnect, each `UNLISTEN` waits for a reply that never comes. This
is the same unbounded wait as in finding 4, and `c.mu` is held **for writing**
the whole time. Every request blocks. `bunDB` has already been closed at this
point, so there is no fallback either.
Also, if `listener.Close` returns an error, `PostgresProvider.Close` returns
early. `sqlConnection.Close` then returns with `connected=true` over a closed
pool (finding 9).
**Failure scenario.** An app with any `LISTEN` subscription hits a network
partition. The health checker or an adapter calls `Reconnect`, and the process
stops serving database requests for as long as the kernel takes to kill the
socket.
**Recommendation.** Skip `UNLISTEN` entirely, because closing the connection
drops all subscriptions server-side. Close with `context.WithTimeout(…, 2*time.Second)`.
Don't do network I/O while holding `l.mu`, and don't close the listener
inside `sqlConnection.Close`'s write lock.
---
### 6. High — Listener leaks a goroutine pair per reconnect, and they race on one `pgx.Conn`
`providers/postgres_listener.go:48-120` (Connect), `257-324` (handleNotifications),
`326-370` (handleReconnection).
`Connect()` ends by starting `go l.handleNotifications()` and
`go l.handleReconnection()`. `handleReconnection` responds to a reconnect
signal by calling `l.Connect(ctx)`, which starts **another** pair. The old pair
keeps running on the same `l.ctx`. After N reconnects there are N+1
notification loops. Each one snapshots `l.conn` and calls
`conn.WaitForNotification`. `pgx.Conn` is **not** safe for concurrent use, so
the second caller gets a "conn busy" error. That error isn't a timeout, so it
sends another reconnect signal, which adds another pair.
`handleReconnection` also waits with `time.Sleep(5 * time.Second)` instead of
selecting on `l.ctx.Done()`, so `Close` can't interrupt it. And `Listen` runs
`l.conn.Exec(LISTEN …)` while holding `l.mu`, which blocks `handleReconnection`
for as long as that Exec takes.
Once the parent `PostgresProvider` is closed (for example by any `Reconnect`,
finding 1), subscribers holding the old `*PostgresListener` get
"listener is closed" forever. Nothing re-subscribes them on the new provider.
**Failure scenario.** A flaky network causes a few listener reconnects. The
goroutine count grows without bound, notifications are delivered twice or
dropped, and CPU rises because of the busy/reconnect spiral.
**Recommendation.** Start the goroutines once, in the constructor or the first
`Connect`. Have `handleReconnection` dial a new conn without calling the public
`Connect`. Guard `WaitForNotification` so only one loop owns the conn. Replace
`time.Sleep` with `select { case <-time.After(d): case <-l.ctx.Done(): }`.
---
### 7. High — Second `Close` panics; health checker silently dead after first cycle
`manager.go:119, 313-345`:
```go
stopChan: make(chan struct{}), // created once, in the constructor
...
func (m *connectionManager) stopHealthChecker() {
if m.healthTicker != nil {
m.healthTicker.Stop()
close(m.stopChan) // never recreated
m.wg.Wait()
m.healthTicker = nil
}
}
```
After `Connect → Close`, `stopChan` is closed. A second `Connect` calls
`startHealthChecker`, which creates a new ticker and goroutine. That goroutine's
`select` sees the closed `stopChan` right away and **exits**, so health
checking is silently off. A second `Close` finds `healthTicker != nil` and
calls `close(m.stopChan)` again, which **panics**: `close of closed channel`.
`startHealthChecker` and `stopHealthChecker` also read and write `healthTicker`
without `m.mu` held (`Close` calls `stopHealthChecker` before locking), so a
concurrent `Connect`/`Close` pair is a data race.
Calling `Connect` twice without `Close` also leaks: `m.connections[name] = conn`
overwrites the previous connection without closing it.
**Failure scenario.** Anything that cycles the manager can crash the process
during shutdown: graceful restart, config hot-reload, or test suites using
`ResetInstance`.
**Recommendation.** Create `stopChan` in `startHealthChecker`. Guard both
functions with `m.mu`, or a dedicated mutex. Make `Connect` idempotent, or have
it close existing connections first.
---
### 8. Medium — SQLite: in-memory data loss and per-connection pragmas
`providers/sqlite.go:54-90`, `config.go:140-141, 202-204`:
- `ManagerConfig.ApplyDefaults` always gives `MaxOpenConns` a value (25), so the
"SQLite works best with MaxOpenConns=1" branch at `sqlite.go:60` never runs.
The probe reported `MaxOpenConnections=25`.
- With `:memory:` (the documented test setup), each pooled connection opens its
**own** private database. The probe created a table on one connection, and a
second connection reported `no such table: t`. `ConnMaxIdleTime` (default
5 min) then closes idle connections and their data with them.
- `PRAGMA journal_mode=WAL` and `PRAGMA busy_timeout` are `Exec`'d once on
whichever pooled connection runs them. `busy_timeout` is per-connection, so
the other 24 get `database is locked` immediately under write contention.
- `BuildDSN` adds `?_timeout=<ms>` (`config.go:347-351`), but
`glebarez/go-sqlite` only recognises `_pragma`, `_txlock` and `_time_format`,
so this parameter is silently ignored.
- `SQLiteProvider.reconnectDB` (`sqlite.go:165`) needs a `dbFactory` that
nothing ever sets, so it is dead code.
**Recommendation.** For SQLite, force `MaxOpenConns=1` for `:memory:` (or use
`file::memory:?cache=shared`), and never set an idle timeout there. Pass the
pragmas in the DSN (`_pragma=busy_timeout(5000)&_pragma=journal_mode(WAL)`) so
every connection gets them.
---
### 9. Medium — Partial-failure states in `Close` and `Connect`
- `connection.go:137-149`: if `bunDB.Close()` or `provider.Close()` fails, for
example because the listener's Close failed (finding 5), `Close` returns
early with `connected = true` and the pool already closed. Every accessor then
returns a handle to a closed pool until someone calls `Close` again.
- `manager.go:197-231`: if connection *k* of *n* fails to connect, `Connect`
returns an error. Connections 1…k-1 stay open but are never stored in
`m.connections`, so `Close` can't reach them and they leak.
**Recommendation.** In `Close`, mark the connection disconnected and nil the
fields regardless of errors, and return a joined error. In `Connect`, close any
connections opened so far when a later one fails.
---
### 10. Medium — DSN builders don't escape credentials; TLS off by default
`config.go` `buildPostgresDSN` / `buildMSSQLDSN` / `buildMongoDSN` use
`fmt.Sprintf` with raw `User`/`Password`/`Database` values:
- Postgres key=value format: a password containing a space or `'` breaks
parsing. A password like `x sslmode=disable` *overrides earlier parameters*.
- MSSQL and Mongo URLs: `@`, `:`, `/`, `?` or `&` in the password corrupt the
URL. They need `url.QueryEscape` / `url.UserPassword`.
- `sslmode` defaults to `disable` (`config.go:322-325`); see
`_CROSS-CUTTING.audit.md` X6.
These values come from config, not from clients, so this isn't directly
exploitable by the threat model. It is a correctness and hardening problem,
and it becomes a security problem wherever DSN parts come from a tenant or
operator UI.
**Recommendation.** Build the Postgres DSN as a URL with `url.URL{User: url.UserPassword(...)}`,
or quote key=value values properly. Default `sslmode` to `prefer` or `require`.
---
### 11. Medium — Config knobs that are ignored or cannot be disabled
- `config.go:161-168`: `HealthCheckInterval == 0` and
`EnableAutoReconnect == false` are both treated as "unset" and replaced with
the defaults (15 s, `true`). **Auto-reconnect, the trigger for findings 1–2,
cannot be switched off from config.**
- `RetryAttempts`, `RetryDelay` and `RetryMaxDelay` are defaulted and copied,
but no provider reads them. Every provider hardcodes `retryAttempts := 3`
and `retryDelay := 1 * time.Second`.
- `statement_timeout` is only added when the DSN is built (finding 4), and
SQLite `_timeout` is ignored by the driver (finding 8).
**Recommendation.** Use `*bool` / `*time.Duration`, or an explicit
`Disable…` flag, for the values that can legitimately be zero or false. Wire
the retry settings into the providers, or delete them.
---
### 12. Medium — `Manager.Connect` holds the manager lock across network dials
`manager.go:197-231` holds `m.mu` (write) while dialing every configured
connection, each with up to 3 attempts, backoff, and `ConnectTimeout`.
`GetConnection`, `HealthCheck`, `Stats` and the health checker all wait
behind it. That's harmless at startup, but it serialises the whole manager if
`Connect` is ever called at runtime (hot-reload, lazy init).
**Recommendation.** Dial outside the lock, then lock only to publish the
results into `m.connections`.
---
### 13. Low — dbmanager metrics are never published
`metrics.go` defines Prometheus collectors plus `PublishMetrics` and
`RecordReconnectAttempt`. A grep over the repository finds **no callers** of
either. The connection-pool gauges (open, in-use, idle, wait count) are exactly
what would have shown the idle-connection problem, and they are always zero.
The `*_total` names are registered as gauges, not counters.
**Recommendation.** Call `PublishMetrics` from the health-check tick, call
`RecordReconnectAttempt` from `Reconnect`, and make the totals counters.
---
### 14. Low — Assorted correctness issues
- `Native()` checks `c.connected` (`connection.go:214`); `Bun()` and `GORM()`
don't. After a partial `Close` they can build ORM wrappers over a nil or
closed DB.
- `getNativeAdapter` (`connection.go:500-525`) wraps SQLite and MSSQL in
`PgSQLAdapter`, which quotes and builds SQL in Postgres dialect.
- `ExistingDBProvider` (`NewConnectionFromDB`) applies no pool settings and no
idle or lifetime limits, and its `Close` closes the caller's `*sql.DB`.
- `MongoProvider` uses `MaxIdleConns` as `MinPoolSize`, and `Stats()` returns an
empty struct.
---
### 15. Low — Logging defects
- `manager.go:247, 367-369, 378-380` call `logger.Error("…", "name", name, "error", err)`.
`pkg/logger` is printf-style, so these print `%!(EXTRA string=name, …)`, and
the error text is buried in exactly the log lines needed during an outage.
- `ResetInstance` discards the error from `Close`.
- Connection errors wrap driver errors that can include the DSN host and user.
Together with `_CROSS-CUTTING.audit.md` X8, they reach Sentry unscrubbed.
---
## Test coverage
`manager_test.go` and `factory_test.go` cover construction and config defaults.
Nothing tests `Reconnect` while handles are held, concurrent `Reconnect`, a
`Connect`/`Close` cycle run twice, health-check lock hold time, or listener
reconnection. Each of findings 1, 2, 3, 6 and 7 can be reproduced with a short
SQLite-backed test (the probes used for this audit took about 20 lines each).
Add them as regression tests when the fixes land, and run them with `-race`
(`_CROSS-CUTTING.audit.md` X1).
+220
View File
@@ -0,0 +1,220 @@
# Audit — `pkg/errortracking`
- **Date:** 2026-09-29
- **Scope:** `pkg/errortracking/{interfaces,noop,sentry,factory}.go` (260 LOC, 4 source files + 1 test file, 67 LOC)
- **Axes:** thread locking/waiting · slowness · security · panic handling & logging
- **Threat model:** hostile internet client; error messages and `extra` maps may contain attacker-shaped content.
## Summary
Small, clean abstraction: a `Provider` interface, a no-op implementation, a Sentry implementation,
and a config-driven factory. The concurrency story is fine — `sentry.Hub` is internally
mutex-guarded and the provider holds no mutable state of its own. The real exposure is **what
this package sends out of the trust boundary**: it is the egress point for every `Warn`/`Error`
in the codebase (see `audit/pkg/logger.audit.md` findings 2 and 3) and it applies **no scrubbing
whatsoever**.
| # | Severity | Axis | Finding |
|---|----------|------|---------|
| 1 | **High** | Security | No `BeforeSend` scrubber — messages, stack traces and `extra` leave the trust boundary verbatim |
| 2 | Medium | Security | `sentry.Init` mutates process-global state; `NewSentryProvider` can be called repeatedly and silently replaces the global client |
| 3 | Medium | Slowness | `Flush(timeout int)` is second-granularity only; combined with `Close()` gives up to 7 s of shutdown stall |
| 4 | Medium | Slowness | `CapturePanic` stringifies the whole stack trace into an `extra` field on every panic |
| 5 | Low | Security | `AttachStacktrace: true` is hardcoded — source paths and function names of the deployment leak to the SaaS |
| 6 | Low | Correctness | `CaptureError` produces an `Exception` with a nil `Stacktrace` for plain `errors.New` values |
| 7 | Low | Correctness | Config-provided `SampleRate == 0` silently means "send everything", not "send nothing" |
| 8 | Low | Architecture | `factory.go` imports `pkg/config`, coupling the lowest-level package to the config layer |
---
## Findings
### 1. No scrubbing before egress (High, Security)
`sentry.go:29-42`
```go
err := sentry.Init(sentry.ClientOptions{
Dsn: config.DSN,
Environment: config.Environment,
Release: config.Release,
Debug: config.Debug,
AttachStacktrace: true,
SampleRate: config.SampleRate,
TracesSampleRate: config.TracesSampleRate,
})
```
`BeforeSend` is not set. Neither is `BeforeSendTransaction`. Nothing in `CaptureError`
(`sentry.go:46`), `CaptureMessage` (`sentry.go:75`) or `CapturePanic` (`sentry.go:97`) inspects or
redacts its inputs; all three copy straight into `event.Message` / `event.Exception.Value` /
`event.Contexts["extra"]` and hand it to `hub.CaptureEvent`.
Because `pkg/logger.Error`/`Warn` forward every formatted message here unconditionally, the set of
things that can reach Sentry is "every error string produced anywhere in ResolveSpec". In this
codebase that includes driver errors (which embed DSNs and sometimes credentials on connect
failure), SQL fragments with bound values, and identifiers taken from request headers.
Under the hostile-client threat model this is an **attacker-reachable exfiltration channel**: shape
an input that lands in an error message, and its content is written to a third-party system
outside the operator's control.
**Recommendation:** set `BeforeSend` to run a redaction pass over `Message`,
`Exception[].Value` and `Contexts` — at minimum strip `password=`, `://user:pass@`, `Bearer `,
and anything matching the configured DSN patterns. Consider an `extra`-key allowlist rather than
passing the caller's map through (`sentry.go:70`, `92`, `114-121`).
### 2. `sentry.Init` mutates process-global state (Medium, Security/Correctness)
`sentry.go:29` calls the package-level `sentry.Init`, which installs a global client, and
`sentry.go:40` then captures `sentry.CurrentHub()`. Consequences:
- Calling `NewSentryProvider` twice (two `NewProviderFromConfig` calls, or a config reload)
replaces the global client. Any previously-created `SentryProvider` keeps a `hub` pointer whose
client has been swapped underneath it — events start going to the *new* DSN. If the two configs
have different environments or DSNs, events are misrouted with no error.
- Events enqueued on the old client at swap time may be dropped without flush.
- It means this "provider" abstraction is a lie: you cannot actually have two Sentry providers
with different configs in one process.
**Recommendation:** build a dedicated client with `sentry.NewClient(opts)` and bind it to an
owned `sentry.NewHub(client, scope)` rather than touching the global. That also makes `Close()`
able to genuinely release resources.
### 3. Coarse, additive shutdown flush (Medium, Slowness)
`sentry.go:125-128`
```go
func (s *SentryProvider) Flush(timeout int) bool {
return sentry.Flush(time.Duration(timeout) * time.Second)
}
```
`timeout` is an `int` interpreted as whole seconds — the interface (`interfaces.go:30`) cannot
express 500 ms. `Close()` (`sentry.go:131-134`) then runs a *second* `sentry.Flush(2s)`.
`pkg/logger.CloseErrorTracking` (`logger.go:69-75`) calls `Flush(5)` then `Close()`, so a graceful
shutdown blocks for **up to 7 seconds** in this package alone, before the HTTP drain and DB close
budgets in `pkg/server`. If the Sentry endpoint is unreachable (the common case during an
outage — which is when you are restarting) both flushes run to full timeout.
Note `Flush` also flushes the *global* client, not `s.hub`'s, which is the same object today only
because of finding 2.
**Recommendation:** change the interface to `Flush(context.Context) bool` or
`Flush(time.Duration) bool`; have `Close` not re-flush; and pass the server's shutdown deadline
through instead of hardcoding 5.
### 4. Whole stack trace stringified into `extra` on every panic (Medium, Slowness)
`sentry.go:117-119`
```go
if stackTrace != nil {
extraCtx["stack_trace"] = string(stackTrace)
}
```
The caller (`pkg/logger.CatchPanicCallback`, `HandlePanic`) already produced the trace via
`debug.Stack()`. Here it is copied again into a string and shipped as a context field. Per
recovered panic that's two full copies of a multi-kilobyte trace plus a network event. With
panics recovered rather than fatal on the request path, a reliably-panicking input is a cheap
amplification primitive (see `audit/pkg/logger.audit.md` finding 5).
Sentry also truncates large context values server-side, so much of this payload is wasted.
**Recommendation:** put the trace in `Exception[0].Stacktrace` as structured frames (which Sentry
groups and displays properly) rather than a blob in `extra`, and cap the byte length.
### 5. `AttachStacktrace: true` hardcoded (Low, Security)
`sentry.go:35`. Not configurable. Every event carries absolute source paths, package layout and
function names of the build. That's mostly a reconnaissance leak to whoever can read the Sentry
project rather than to the internet attacker, but it should be an operator choice, especially for
on-prem deployments sending to a hosted DSN.
### 6. Nil stack trace for plain errors (Low, Correctness)
`sentry.go:62`
```go
Stacktrace: sentry.ExtractStacktrace(err),
```
`ExtractStacktrace` only finds a trace if the error implements `StackTrace()`/`Callers()`
(`pkg/errors`-style). Nearly all errors in this codebase come from `fmt.Errorf`, so this returns
`nil` and the Sentry event has an exception with no frames — grouping falls back to the message
string, which (because messages embed request-specific values) fragments what should be one issue
into thousands.
**Recommendation:** fall back to `sentry.NewStacktrace()` when extraction yields nil, and set an
explicit `event.Fingerprint` derived from a stable prefix rather than the full message.
### 7. `SampleRate == 0` means "send everything" (Low, Correctness)
`factory.go:20-27` passes `cfg.SampleRate` through untouched, and `pkg/config/manager.go`
registers **no default** for `error_tracking.sample_rate`. So an operator who leaves it out gets
`0.0`, and `sentry-go@v0.46.2` `client.go:339-341` rewrites `0.0` → `1.0`.
Verified in the module cache:
```go
if options.SampleRate == 0.0 {
options.SampleRate = 1.0
}
```
Fail-open rather than fail-closed, which is arguably the right choice for an error tracker — but
it means an operator who *intends* to disable sampling by setting `0` gets the opposite, silently.
**Recommendation:** make `SampleRate` a `*float64` in the config struct, or register an explicit
default in `setDefaults`, and validate/log the effective value at init.
### 8. `factory.go` imports `pkg/config` (Low, Architecture)
`factory.go:6` — `errortracking` is imported by `pkg/logger`, which is imported by essentially
everything. Pulling `pkg/config` (and therefore `viper`) into that dependency chain means the
lowest-level logging path transitively depends on the configuration layer. It works today only
because `pkg/config` imports nothing from ResolveSpec; the first time it wants to log, there is
an import cycle.
**Recommendation:** move `NewProviderFromConfig` into `pkg/config`-adjacent wiring code (or take
a small local options struct instead of `config.ErrorTrackingConfig`) so `errortracking` stays a
leaf.
---
## What looks right
- **Concurrency is genuinely fine.** `SentryProvider` holds only an immutable `*sentry.Hub`;
`sentry.Hub` guards its own state with a mutex, and `CaptureEvent` hands off to a background
worker with a bounded queue, so it does not block the caller and does not need a lock here.
- `GetHubFromContext(ctx)` with fallback to `s.hub` (`sentry.go:53-56`, `81-84`, `103-106`) is the
correct Sentry idiom and preserves per-request scope when middleware installs a hub.
- Nil-input guards on all three capture methods (`sentry.go:47`, `76`, `98`) — a nil error, empty
message or nil recovered value is dropped rather than producing a junk event.
- `event.Contexts` is safe to index: `sentry.NewEvent()` initialises the map, so
`event.Contexts["extra"] = ...` cannot nil-panic.
- `NoOpProvider` means a disabled tracker is always safe to call — no nil checks needed at call
sites beyond the one in `pkg/logger`.
- `factory.go:15-17` correctly refuses to start with `provider: sentry` and an empty DSN rather
than silently no-oping.
## Panic handling
The package neither panics nor recovers, which is correct for its role — it is the *sink* for
panic reports, not a place that should be generating them. The nil-guards in finding "what looks
right" cover the realistic nil-deref paths. One residual: `CapturePanic` ranges over `extra`
(`sentry.go:115`) without a nil check, which is safe in Go (ranging a nil map yields zero
iterations) — noted only to confirm it was checked.
## Suggested follow-up
1. Add `BeforeSend` redaction (finding 1). This is the highest-value single change in the package.
2. Stop using the global Sentry client (finding 2) — unblocks real multi-provider support and a
meaningful `Close()`.
3. Widen `Flush` to a duration/context (finding 3) and wire it to the server shutdown budget.
4. Add tests for the Sentry path. The existing test file covers only `NoOpProvider`, severity
string mapping and interface satisfaction — `SentryProvider`'s capture methods have no
coverage at all. `sentry-go` ships a test transport that makes this straightforward.
+285
View File
@@ -0,0 +1,285 @@
# Audit — `pkg/logger`
- **Date:** 2026-09-29
- **Scope:** `pkg/logger/logger.go` (211 LOC, 1 file, no tests)
- **Axes:** thread locking/waiting · slowness · security · panic handling & logging
- **Threat model:** hostile internet client; request bodies, headers, params and identifiers are attacker-controlled.
## Summary
`pkg/logger` is a thin package-global wrapper over `zap.SugaredLogger` plus a fan-out to
`pkg/errortracking`. It is the single most widely imported package in the repo, so its defects
are systemic. Two classes of problem dominate: **unsynchronised global mutable state** (a real
data race between logger re-initialisation and request-path logging), and **unbounded,
unsampled, unscrubbed egress of formatted messages to a third-party error tracker** on every
`Warn`/`Error` call — which under hostile input is both a data-leak and a cost/latency
amplification channel.
There are **zero tests** in this package.
| # | Severity | Axis | Finding |
|---|----------|------|---------|
| 1 | **High** | Locking | Unsynchronised writes to `Logger` / `errorTracker` globals race with every log call |
| 2 | **High** | Security | Every `Warn`/`Error` message is shipped verbatim to Sentry — no scrubbing, no allowlist |
| 3 | **High** | Slowness | No rate limit, sampling or dedup on error-tracker fan-out; attacker-triggerable |
| 4 | **High** | Panic | `CatchPanic` swallows panics unconditionally — and both call sites are security enforcement functions (fail-open) |
| 5 | Medium | Slowness | `debug.Stack()` + full stack stringification on every recovered panic |
| 6 | Medium | Security | `log.Printf(template, args...)` fallback is a format-string sink for caller-supplied text |
| 7 | Medium | Security | No CRLF/control-char sanitisation on the stdlib fallback path → log injection |
| 8 | Medium | Correctness | `Info`/`Debug` do not strip `context.Context` args; `Warn`/`Error` do |
| 9 | Low | Correctness | `UpdateLogger` leaks the previous zap logger / file descriptor |
| 10 | Low | Correctness | No `Sync()` exported → buffered log lines lost on exit |
| 11 | Low | Slowness | `os.Getpid()` called on every log line |
| 12 | Low | Observability | `UpdateLogger` build failure degrades silently to stdlib `log` |
## Resolution status (2026-09-30)
- **#1** — Fixed (earlier race work): `stateMu` RWMutex with `getLogger`/`swapLogger`/`getErrorTracker`; the exported `Logger` var is kept for compatibility
- **#2** — Fixed: messages are scrubbed before `CaptureMessage` (URL credentials, `password=`/`token=`/`secret=`/`api_key=` values, `Bearer`/`Basic` tokens). Local logs are unchanged. Sentry `BeforeSend` and structured-field allowlisting are not done
- **#3** — Partly fixed: global token bucket (burst 50, 20/s) plus per-severity/template dedup (1s, 1024 keys). Panics are not limited. The `error_tracking.sample_rate` default (Sentry maps 0 to 1.0) is still unset in `config/manager.go`
- **#4** — Partly fixed: `CatchPanicRethrow` added. `pkg/security/provider.go:302` and `:443` still use the swallowing `CatchPanic`; left for the security audit pass
- **#5** — Fixed: stack captured with `runtime.Stack` into a 16 KiB buffer. Per-fingerprint panic rate limiting not done
- **#6** — Fixed: `Info`/`Debug` format first and fall back with `log.Printf("%s", ...)`. `gosec` was enabled separately
- **#7** — Fixed: CR/LF and other control characters are escaped on the stdlib fallback path
- **#8** — Fixed: `Info`/`Debug` strip `context.Context` args
- **#9** — Fixed: the replaced logger is synced on `UpdateLogger`
- **#10** — Fixed: `logger.Sync()` added. Not yet called from the server shutdown path
- **#11** — Fixed: PID cached in a package var
- **#12** — Partly fixed: `UpdateLoggerE` returns the build error and a failed build keeps the previous logger. `Init` still returns nothing
- Tests: `pkg/logger/logger_test.go` (run with `-race`).
---
## Findings
### 1. Unsynchronised global mutable state — data race (High, Locking)
`logger.go:14-15`
```go
var Logger *zap.SugaredLogger
var errorTracker errortracking.Provider
```
`Logger` is written by `Init` → `UpdateLogger` (`logger.go:51`) and by `UpdateLoggerPath`
(`logger.go:29`). `errorTracker` is written by `InitErrorTracking` (`logger.go:57`) and read by
`GetErrorTracker`, `CloseErrorTracking`, `Warn`, `Error`, `CatchPanicCallback`, `HandlePanic`.
Every read site (`logger.go:100`, `108`, `123`, `139`, `156`, `199`) is unguarded. There is no
mutex, no `atomic.Value`, no `sync.Once`.
- **Benign case:** everything is initialised once in `main` before goroutines start. Then it's fine.
- **Real case:** `UpdateLoggerPath` is an exported, runtime-callable API. A config reload, a
log-rotation hook, or a test helper calling it while HTTP handlers log concurrently is an
unsynchronised write to an interface value and a pointer, concurrent with reads. Under the Go
memory model this is undefined behaviour; in practice a torn interface read (type word from the
new value, data word from the old) faults.
- `CloseErrorTracking` (`logger.go:69`) does a read-check-then-use on `errorTracker` with no
guard, so a concurrent `InitErrorTracking(nil)` yields a nil-interface dereference inside
`Flush`.
**Recommendation:** store both behind `atomic.Pointer`/`atomic.Value` (or an `sync.RWMutex`),
and gate first-time init behind `sync.Once`. Run the test suite with `-race` — see finding 12 of
`audit/pkg/config.audit.md` for the same pattern in the config singleton.
### 2. Unscrubbed message egress to third-party error tracker (High, Security)
`logger.go:110-118` and `logger.go:126-134`
```go
message := fmt.Sprintf(template, remainingArgs...)
...
errorTracker.CaptureMessage(ctx, message, errortracking.SeverityError, ...)
```
*Every* `Warn` and `Error` call in the entire codebase has its fully-formatted message sent to
the configured provider (Sentry, in practice). There is no allowlist, no redaction hook, and
`pkg/errortracking/sentry.go` configures no `BeforeSend` scrubber.
Concretely, formatted error strings across `pkg/` embed: SQL fragments and bound values, DB
connection strings, schema/table/column identifiers, filter expressions built from request
input, and raw request bodies in a few handlers. Under the hostile-client threat model this is
two problems at once:
- **Outbound data leak:** secrets that appear in wrapped driver errors (DSNs, credentials from
`pq`/`pgx` connect failures) leave the trust boundary to a SaaS endpoint.
- **Attacker-controlled exfil channel:** an attacker who can shape a value that ends up in an
error message gets that value written to a third-party system — useful for exfiltrating data
read out of the DB via an induced error.
**Recommendation:** add a redaction step before `CaptureMessage`/`CapturePanic` (regex-strip
DSN/`password=`/bearer-token shapes at minimum), and set Sentry's `BeforeSend` as a second
layer. Prefer passing structured fields with an explicit allowlist over shipping the rendered
string.
### 3. No rate limiting or sampling on error-tracker fan-out (High, Slowness)
`logger.go:113`, `logger.go:129`
An unauthenticated request that reliably produces one `Error` log (a malformed filter, an unknown
column, a bad JSON body — all of which the spec handlers log at error level) becomes one Sentry
event. At even modest request rates this means:
- Sentry quota burn → a direct billing-DoS.
- `sentry-go` enqueues onto a bounded worker queue; once saturated events are dropped, so the
*real* errors are the ones lost.
- `pkg/errortracking/sentry.go:34` passes `SampleRate` straight through from config, and
`config/manager.go` sets **no default** for it. `sentry-go@v0.46.2` `client.go:339` maps
`SampleRate == 0.0` → `1.0`, so the out-of-the-box behaviour is *send 100% of events*.
**Recommendation:** default `error_tracking.sample_rate` to something < 1.0 for the message path,
and put a token-bucket or a fingerprint-dedup in front of `CaptureMessage`. Keep panics at 100%.
### 4. `CatchPanic` swallows panics unconditionally, fail-open at both call sites (High, Panic handling)
`logger.go:145-176`
```go
func CatchPanicCallback(location string, cb func(err any), args ...interface{}) func() {
...
if err := recover(); err != nil { ... if cb != nil { cb(err) } }
}
```
The recovered value is logged and then discarded. There is no variant that logs-and-re-panics
and no way for the caller to signal "this panic means state is corrupt, take the process down".
This is the right default for an HTTP handler boundary. The two current call sites are **not**
handler boundaries:
- `pkg/security/provider.go:302` — `defer logger.CatchPanic("ApplyColumnSecurity")()`
- `pkg/security/provider.go:443` — `defer logger.CatchPanic("GetRowSecurityTemplate")()`
Both are *security enforcement* functions. Swallowing a panic there means the column-security
filter or row-security template silently does not get applied, and the caller — which has no way
to learn a panic occurred, since `CatchPanic` returns nothing and sets no error — proceeds as if
security was applied. That is a fail-open security control; see
`audit/pkg/security.audit.md` for the full write-up of those two sites.
Separately: a panic while a mutex is held does not release that mutex unless an intervening
`defer Unlock` exists, so swallowing converts a crash into a permanent deadlock at any
lock-holding call site.
**Recommendation:** add `CatchPanicRethrow(location string)` for internal use and reserve the
swallowing form for the outermost request/goroutine boundary. Document which is which.
### 5. Full stack capture on every recovered panic (Medium, Slowness)
`logger.go:158` and `logger.go:197`
```go
callstack := debug.Stack()
```
`debug.Stack()` stops the world briefly and allocates; `HandlePanic` then formats the whole trace
into a string *and* ships it to Sentry. Because panics on the request path are recovered rather
than fatal (finding 4), an attacker who finds one reliably-panicking input turns each request
into a stack capture + string build + network event. That is a solid amplification factor over a
normal request.
**Recommendation:** cap the captured stack (`runtime.Stack` into a fixed 8–16 KiB buffer rather
than `debug.Stack()`'s grow-until-it-fits loop), and rate-limit identical panic fingerprints.
### 6. Format-string sink in the stdlib fallback (Medium, Security)
`logger.go:100`, `logger.go:142` (and `108`/`123` with `"%s"`, correctly)
```go
func Info(template string, args ...interface{}) {
if Logger == nil {
log.Printf(template, args...) // template is the caller's, args may be empty
```
`Info` and `Debug` pass `template` directly to `log.Printf`. If any caller ever does
`logger.Info(someUserString)` — the idiomatic-looking single-argument call — a `%s` or `%n` in
that string is interpreted as a verb, producing `%!s(MISSING)` garbage and mangled logs. Note
`Warn`/`Error` already avoid this on the fallback path by using `log.Printf("%s", message)`;
`Info`/`Debug` do not.
A grep of `pkg/` found **no** current single-argument call sites, so this is a latent API footgun
rather than a live bug — but it is one that costs one line to close.
**Recommendation:** mirror `Warn`'s shape: format first, then `log.Printf("%s", message)`.
`govet` runs by default under golangci-lint v2's standard set, and its `printf` analyser infers
wrappers like these — so once the fallback is fixed, call sites are checked at build time for
free. (Note `gosec` is *not* in `.golangci.json`'s `linters.enable` list; it appears only in the
exclusion rules. Worth enabling repo-wide.)
### 7. No log-injection sanitisation on the fallback path (Medium, Security)
On the zap path, the JSON encoder escapes newlines and control characters, so injected content
can't forge a log record. On the `Logger == nil` fallback path, `log.Printf` writes raw bytes: a
value containing `\n2026-09-29 ... level=info authorized=true` forges a plausible second log
line. Combined with finding 12 (silent degradation to the fallback path) this is reachable
without the operator noticing the encoder changed.
**Recommendation:** strip/escape `\r`, `\n` and other C0 control characters from formatted
messages before the stdlib write.
### 8. `Info`/`Debug` don't strip `context.Context` arguments (Medium, Correctness)
`extractContext` (`logger.go:79-98`) exists precisely so callers can pass a `ctx` as a trailing
variadic arg. `Warn` (`logger.go:106`) and `Error` (`logger.go:121`) call it. `Info`
(`logger.go:99`) and `Debug` (`logger.go:137`) **do not** — they pass every arg to `Sprintf`.
So `logger.Info("saved %s", name, ctx)` renders as
`saved widget%!(EXTRA *context.valueCtx=context.Background...)`, dumping the context's contents
(which in this codebase carry auth/tenant values) into the log line. That is both noise and a
minor disclosure.
**Recommendation:** call `extractContext` in all four level functions for uniform behaviour.
### 9. `UpdateLogger` leaks the previous logger (Low)
`logger.go:37-53` builds a new zap logger and overwrites `Logger` without calling `Sync()`/close
on the old one. `UpdateLoggerPath` opens a new file sink each call; repeated calls leak a file
descriptor each time and buffered lines in the old logger are lost.
### 10. No `Sync()` on shutdown (Low)
Nothing in the package exposes `Logger.Sync()`, and `CloseErrorTracking` (`logger.go:69`) flushes
only the error tracker. zap buffers writes to file sinks, so the last lines before exit — often
the interesting ones — are dropped. Add `func Sync() error` and call it from the server's
shutdown path alongside `CloseErrorTracking`.
### 11. `os.Getpid()` per log line (Low, Slowness)
`logger.go:102`, `111`, `127`, `140`, `165`, `202`. On Linux `getpid` is cached by the runtime so
this is cheap, but the PID cannot change for the life of the process — cache it in a package var
and drop six calls from the hot path.
### 12. Silent degradation when the logger fails to build (Low, Observability)
`logger.go:45-49`
```go
logger, err := config.Build()
if err != nil { log.Print(err); return }
```
`Logger` stays `nil`, so the whole process silently falls back to unstructured stdlib logging
(and thereby onto the format-string and log-injection paths of findings 6 and 7) with a single
line of warning that itself goes to stderr. A bad `logger.path` in config (unwritable directory)
triggers exactly this.
**Recommendation:** return the error from `Init`/`UpdateLogger` and let the caller decide whether
to fail startup.
---
## What looks right
- `extractContext` correctly ignores second and subsequent contexts rather than fighting over them.
- `Warn`/`Error` use `log.Printf("%s", message)` on the fallback path — the safe form.
- `HandlePanic` returns an `error` rather than swallowing, which lets callers convert a panic into
a normal error return. This is the better of the two panic idioms in the package.
- The `errortracking.Provider` indirection means a nil/noop provider is always safe to call.
## Suggested follow-up
1. Guard the two globals (finding 1) — prerequisite for running the suite under `-race`.
2. Add redaction + sampling in front of the error-tracker fan-out (findings 2, 3).
3. Split `CatchPanic` into swallow/rethrow variants and re-audit the ~60 `recover()` sites
listed in the other package audits against the split (finding 4).
4. Add a test file. Minimum: concurrent `UpdateLogger` + `Error` under `-race`, `Info` with a
`%`-bearing message, and nil-provider paths.
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+502
View File
@@ -0,0 +1,502 @@
# Audit — `pkg/modelregistry`
- **Date:** 2026-09-29
- **Scope:** `pkg/modelregistry/model_registry.go` (381 LOC, 1 file, **no tests**)
- **Axes:** thread locking/waiting · slowness · security · panic handling & logging
- **Threat model:** hostile internet client. This package holds the `ModelRules` that
`pkg/security/hooks.go` consults to authorise read/update/create/delete, so it is **on the
authorisation path**.
## Summary
This package is the highest-risk find in the audit. It has been deliberately reworked to "never
hang" by replacing blocking `Lock`/`RLock` with **bounded `TryLock` retry loops that give up and
return a wrong answer** — and because those wrong answers are consumed by
`pkg/security/hooks.go` as authorisation decisions, the result is an **authorisation control that
fails open under lock contention**.
The comments in the file are explicit about the trade-off ("falls back to the last known value
without synchronization", "the call is a no-op") — so the hazard was known at the time of writing.
What appears not to have been traced is where those degraded results end up. They end up in
`checkModelUpdateAllowed` / `checkModelDeleteAllowed`, which treat any error as *permit*.
There are **no tests** in this package and no `-race` coverage of it anywhere.
| # | Severity | Axis | Finding |
|---|----------|------|---------|
| 1 | **Critical** | Security + Locking | `GetModel`'s "registry locked" error is consumed by `pkg/security/hooks.go` as *allow by default* → authorisation fails open under write-lock contention |
| 2 | **High** | Locking | `GetDefaultRegistry` documents and performs an unsynchronised read of `defaultRegistry` on lock-acquire failure — a data race by design |
| 3 | **High** | Locking | `SetDefaultRegistry` silently no-ops after ~20 ms of contention; caller gets no error |
| 4 | **High** | Security | `RegisterModelWithRules` is non-atomic: the model is visible with permissive `DefaultModelRules` before its real rules are applied (TOCTOU) |
| 5 | Medium | Correctness | `GetAllModels` returns an empty map, and `GetModels` silently skips whole registries, on lock-acquire failure |
| 6 | Medium | Locking | `IterateModels` invokes the caller's callback while holding `RLock` → guaranteed self-deadlock if the callback touches the registry |
| 7 | Medium | Slowness | `time.Sleep(1ms)` spin loops add up to 20 ms of latency per call and defeat mutex fairness/hand-off |
| 8 | Medium | Locking | `defaultRegistry` is read unsynchronised by six package-level functions while `SetDefaultRegistry` writes it under lock |
| 9 | Medium | Locking | Inconsistent discipline: `SetModelRules`/`GetModelRules`/`AddRegistry`/`IterateModels` use blocking locks; the rest use try-locks |
| 10 | Low | Slowness/Locking | Reflection (`TypeOf`, unwrap loop, `reflect.New`) runs while holding the registry **write** lock |
| 11 | Low | Availability | Unbounded unwrap loop: a recursive pointer type (`type T *T`) spins forever holding the write lock (**verified**) |
| 12 | Low | Panic | Package has no `recover` anywhere, and calls a caller-supplied callback under a lock (see 6) |
| 13 | Low | Security | `DefaultModelRules()` grants `CanRead/Update/Create/Delete: true` — registration without explicit rules is fully mutable |
## Resolution (2026-09-30)
Fixed in `pkg/modelregistry/model_registry.go`, `pkg/security/hooks.go`, and new
`pkg/modelregistry/model_registry_test.go` (passes under `-race`).
| # | Status | What changed |
|---|--------|--------------|
| 1 | **Fixed** | Added sentinels `ErrModelNotFound`, `ErrModelExists`, `ErrInvalidModel` (wrapped, `errors.Is`-friendly). `checkModelUpdateAllowed`/`checkModelDeleteAllowed` now allow-by-default **only** on `ErrModelNotFound`; any other error denies. Lookups can no longer return a "locked" error at all. |
| 2 | **Fixed** | `GetDefaultRegistry` uses a plain `RLock`; no unsynchronised fallback. |
| 3 | **Fixed** | `SetDefaultRegistry` uses a blocking `Lock` (cannot silently no-op); a nil registry is ignored. |
| 4 | **Fixed** | `RegisterModelWithRules` and `RegisterModel` share `registerLocked`, which writes model + rules under one lock acquisition. |
| 5 | **Fixed** | `GetAllModels`/`GetModels` use blocking locks and can no longer return empty/partial results due to contention. Signatures unchanged (`GetAllModels` is used through interfaces by resolvespec/restheadspec/openapi). |
| 6 | **Fixed** | `IterateModels` iterates a snapshot; the callback runs with no lock held (regression test re-enters the registry). |
| 7 | **Fixed** | Try-lock/sleep helpers and `lockRetry*` constants removed. |
| 8 | **Fixed** | All package-level functions go through `GetDefaultRegistry()` / `registriesSnapshot()`; `defaultRegistry` is only touched under `registriesMutex`. |
| 9 | **Fixed** | One discipline: blocking locks, snapshot-and-release, documented lock order (`registriesMutex` before a registry's mutex). |
| 10 | **Fixed** | Reflection/validation (`validateModel`) runs before the write lock is taken. |
| 11 | **Fixed** | Unwrap loop capped at 16 levels; `type T *T` now returns `ErrInvalidModel` (tested). |
| 12 | **Fixed** | Sentinel errors added. `IterateModels` recovers a callback panic per model, logs it via `logger.HandlePanic` with the model name, and continues; `validateModel` recovers reflection panics and returns `ErrInvalidModel` so registration fails closed. No lock is held during either, so the registry cannot be wedged. |
| 13 | **Accepted (decision)** | Allow-by-default retained deliberately: `DefaultModelRules()` still grants read/update/create/delete. Callers wanting restrictions must use `RegisterModelWithRules`/`SetModelRules`. |
Tests added: sentinel errors, recursive pointer type, pointer normalisation, atomic
`RegisterModelWithRules` (concurrent reader never sees permissive rules), re-entrant `IterateModels`,
cross-registry `GetModelRulesByName`, and a concurrent `-race` stress test.
Not changed: the `pkg/security` middleware-wiring question (context fast-path) remains tracked in
`audit/pkg/security.audit.md`.
---
## Findings
### 1. Authorisation fails open under lock contention (Critical, Security + Locking)
The mechanism spans two packages.
**Here**, `GetModel` conflates "not found" with "could not lock" into a single `error` return
(`model_registry.go:198-210`):
```go
func (r *DefaultModelRegistry) GetModel(name string) (interface{}, error) {
if !r.tryRLock() {
return nil, fmt.Errorf("failed to get model %s: registry locked", name)
}
defer r.mutex.RUnlock()
model, exists := r.models[name]
if !exists {
return nil, fmt.Errorf("model %s not found", name)
}
return model, nil
}
```
`GetModelRulesByName` (`model_registry.go:364-376`) uses `GetModel` as its existence probe:
```go
for _, registry := range registries {
if _, err := registry.GetModel(name); err == nil {
return registry.GetModelRules(name)
}
}
return ModelRules{}, fmt.Errorf("model %s not found in any registry", name)
```
So a `tryRLock` failure makes the registry look like it does not contain the model.
**In `pkg/security/hooks.go`**, that outcome is interpreted as *permit*
(`pkg/security/hooks.go:274-294`, and identically at `:298-318`):
```go
func checkModelUpdateAllowed(secCtx SecurityContext) error {
rules, ok := GetModelRulesFromContext(secCtx.GetContext())
if !ok {
schema := secCtx.GetSchema()
entity := secCtx.GetEntity()
var err error
if schema != "" {
rules, err = modelregistry.GetModelRulesByName(fmt.Sprintf("%s.%s", schema, entity))
}
if err != nil || schema == "" {
rules, err = modelregistry.GetModelRulesByName(entity)
}
if err != nil {
return nil // model not registered, allow by default
}
}
if !rules.CanUpdate {
return fmt.Errorf("update not allowed for %s", secCtx.GetEntity())
}
return nil
}
```
Note the context fast-path at `hooks.go:275`: if `NewModelAuthMiddleware` already put rules in the
context, the registry is not consulted and this bug does not fire. The registry fallback runs
whenever that middleware is absent or did not resolve rules — so the blast radius depends on
deployment wiring. `audit/pkg/security.audit.md` covers whether that middleware is mandatory.
`return nil` from `checkModelUpdateAllowed` means **the update is authorised**. Same for
`checkModelDeleteAllowed`. `GetModelRules(name)` for the "found" path also uses a blocking
`RLock` (`model_registry.go:253`) — so the two calls in `GetModelRulesByName` don't even use the
same locking discipline.
**Failure scenario.** A model `public.employees` is registered with `CanDelete: false`. A
concurrent `RegisterModel` (or `SetModelRules`, or `RegisterModelWithRules`) holds the write lock
for longer than `lockRetryAttempts * lockRetryDelay` = 20 ms — which is entirely achievable given
finding 10 (reflection under the write lock) and finding 7 (each waiter sleeps in 1 ms
increments, so N waiters serialise). During that window every `DELETE` request against
`public.employees` has `GetModelRulesByName` return an error, `checkModelDeleteAllowed` return
`nil`, and the delete proceeds. The model's `CanDelete: false` is not enforced.
This is remotely triggerable if any request path can cause a model registration or a rules
update; even without that, it is a straightforward race that will fire under load.
**Recommendation, in order of value:**
1. Make the security layer **fail closed**: distinguish a sentinel `ErrModelNotFound` from any
other error, and only allow-by-default on `ErrModelNotFound`. Any other error must deny.
2. Delete the try-lock scheme here entirely and use plain `RLock`/`Lock` (see finding 2 for why
the scheme does not achieve its stated goal anyway).
3. Separate the existence probe from the rules fetch so `GetModelRulesByName` takes each registry's
lock once and returns a typed "found / not found / unavailable" result.
### 2. `GetDefaultRegistry` races by design (High, Locking)
`model_registry.go:71-84`
```go
// GetDefaultRegistry returns the current default registry. It uses a
// bounded TryRLock instead of a blocking RLock so it can never hang;
// if the lock can't be acquired in time it falls back to the last known
// value without synchronization.
func GetDefaultRegistry() *DefaultModelRegistry {
for i := 0; i < lockRetryAttempts; i++ {
if registriesMutex.TryRLock() {
defer registriesMutex.RUnlock()
return defaultRegistry
}
time.Sleep(lockRetryDelay)
}
return defaultRegistry
}
```
The `return defaultRegistry` on line 83 reads a pointer that `SetDefaultRegistry`
(`model_registry.go:89-116`) writes under the write lock. The only time this path is taken is
precisely when a writer holds or is contending for the lock — i.e. the fallback executes
*exactly* in the window where the race is live. The trade is not "hang vs. slightly stale value";
it is "block for 20 ms vs. data race", and a torn/`nil` pointer read here means a nil-pointer
dereference in the caller.
The premise is also wrong: a `sync.RWMutex.RLock` that is only ever held for a map lookup cannot
"hang". The hang this was written to avoid must have had a different root cause — most likely
finding 6 (self-deadlock through `IterateModels`) or a lock-ordering inversion — and the try-lock
scheme papers over it rather than fixing it.
**Recommendation:** revert to `RLock`/`RUnlock`. If a real hang was observed, reproduce it under
`-race` and `GODEBUG=gctrace`/`SIGQUIT` stack dump; the fix belongs at the deadlock, not here.
### 3. `SetDefaultRegistry` silently no-ops (High, Locking)
`model_registry.go:90-100`
```go
acquired := false
for i := 0; i < lockRetryAttempts; i++ {
if registriesMutex.TryLock() { acquired = true; break }
time.Sleep(lockRetryDelay)
}
if !acquired {
return
}
```
The function returns no error. A caller that swaps in a registry — plausibly one with *restrictive*
`ModelRules* — has no way to learn the swap did not happen, and continues believing the new
registry is in effect. Every subsequent authorisation check consults the old registry's rules.
`GetModels` (`model_registry.go:319-329`) has the same shape and returns `nil`.
**Recommendation:** return `error` from `SetDefaultRegistry`; or (better) use a blocking `Lock`,
since this is a startup-time operation where blocking is correct.
### 4. `RegisterModelWithRules` is non-atomic (High, Security)
`model_registry.go:270-282`
```go
func (r *DefaultModelRegistry) RegisterModelWithRules(name string, model interface{}, rules ModelRules) error {
// First register the model
if err := r.RegisterModel(name, model); err != nil {
return err
}
// Then set the rules (we need to lock again for rules)
r.mutex.Lock()
defer r.mutex.Unlock()
r.rules[name] = rules
return nil
}
```
`RegisterModel` releases the write lock before returning, and it initialises the model's rules to
`DefaultModelRules()` (`model_registry.go:191-194`) — which is **permissive**:
`CanRead/CanUpdate/CanCreate/CanDelete` all `true`.
Between the two lock acquisitions, any concurrent `GetModelRulesByName` sees the model registered
with full read/update/create/delete permission, regardless of the restrictive `rules` the caller
passed. The comment "we need to lock again for rules" acknowledges the re-lock without noticing
the gap it opens.
**Failure scenario.** `RegisterModelWithRules("public.audit_log", AuditLog{}, ModelRules{CanRead:
true})` — intended read-only. A `DELETE /public.audit_log/...` that lands in the window is
authorised because `rules.CanDelete` is `true` from the default.
**Recommendation:** add an unexported `registerLocked(name, model, rules)` that writes both maps
under one lock acquisition, and build both public constructors on it. Also change the default
initialisation in `RegisterModel` to deny-by-default, or require rules at registration.
### 5. Degraded results indistinguishable from real results (Medium, Correctness)
Three functions return a plausible-looking answer when they cannot lock:
- `GetAllModels` (`model_registry.go:212-215`) — `return make(map[string]interface{})`, i.e. "the
registry is empty".
- `GetModels` (`model_registry.go:327-329`) — `return nil` on `registriesMutex` failure, and
`model_registry.go:336-338` `continue`s past any individual registry it cannot read, returning a
**partial** list with no indication of truncation.
- `GetDefaultRegistry` — finding 2.
Consumers of `GetModels`/`GetAllModels` (schema introspection, OpenAPI generation, migration
helpers) will emit a document that is missing models, and there is no error to log. Note these two
have no callers in `pkg/` today, which is the only reason this is Medium.
**Recommendation:** return `(T, error)`; never manufacture an empty-but-valid result.
### 6. `IterateModels` calls a user callback under a read lock (Medium, Locking)
`model_registry.go:307-314`
```go
func IterateModels(fn func(name string, model interface{})) {
defaultRegistry.mutex.RLock()
defer defaultRegistry.mutex.RUnlock()
for name, model := range defaultRegistry.models {
fn(name, model)
}
}
```
`fn` is arbitrary caller code running with `defaultRegistry.mutex` read-held. `sync.RWMutex` is not
reentrant, and once a writer is blocked on `Lock` it also blocks *new* readers. So:
- `fn` calling `modelregistry.RegisterModel` / `SetModelRules` → `Lock` waits for the reader, which
is the same goroutine. **Permanent self-deadlock.**
- `fn` calling `GetModel` → `tryRLock` fails for 20 ms and returns "registry locked" for every
model, which is silent nonsense rather than a deadlock (and feeds finding 1).
- `fn` doing anything slow (I/O, a DB call) holds the registry read lock for that whole duration,
blocking all registration and — via the blocked-writer rule — all other readers too.
This is the most likely original cause of the "hang" the try-lock scheme was introduced to work
around.
**Recommendation:** snapshot under the lock, release, then iterate:
```go
func IterateModels(fn func(name string, model interface{})) {
reg := GetDefaultRegistry()
snapshot := reg.GetAllModels() // takes and releases the lock
for name, model := range snapshot {
fn(name, model)
}
}
```
### 7. `time.Sleep` spin loops (Medium, Slowness)
`tryLock` (`model_registry.go:128-136`), `tryRLock` (`:140-148`), and the inline loops in
`GetDefaultRegistry`, `SetDefaultRegistry`, `GetModels`.
```go
for i := 0; i < lockRetryAttempts; i++ {
if r.mutex.TryLock() { return true }
time.Sleep(lockRetryDelay) // 1ms
}
```
Problems:
- **Latency floor.** A contended call costs a multiple of 1 ms even if the lock frees after 10 µs,
because the waiter is asleep. A blocking `Lock` would be handed the mutex in microseconds. So the
"no-hang" scheme is *slower* in the common contended case, not faster.
- **No fairness.** `sync.Mutex` has a starvation-avoidance mode that hands the lock to a waiter
queued > 1 ms. `TryLock` participates in none of it, so a try-lock waiter can be starved
indefinitely by a stream of blocking `Lock` callers (`SetModelRules`, `AddRegistry`,
`IterateModels` all still block) — see finding 9.
- **Timer churn.** 20 timer allocations per contended call.
- Sleeping in a loop scales badly: 50 concurrent callers each sleep and wake 20 times, producing
1000 needless scheduler round-trips for what a mutex does with one park/unpark.
**Recommendation:** delete the try-lock helpers. If a bounded wait is genuinely required for an
SLO, express it as `context`-aware acquisition (a buffered-channel semaphore with a `select` on
`ctx.Done()`), which gives a real deadline *and* a real error — not a silent wrong answer.
### 8. `defaultRegistry` read without the guarding mutex (Medium, Locking)
`SetDefaultRegistry` writes `defaultRegistry` (`model_registry.go:110`) under `registriesMutex`.
These read it **without** taking that mutex:
- `RegisterModel` (`model_registry.go:288`)
- `IterateModels` (`model_registry.go:308`, `311`)
- `SetModelRules` (`model_registry.go:354`)
- `GetModelRules` (`model_registry.go:359`)
- `GetDefaultRegistry`'s fallback (`model_registry.go:83`, finding 2)
A data race on the pointer, and semantically these functions may operate on the *previous* default
registry after a swap — so rules set through `SetModelRules` can land on a registry nobody consults
any more.
**Recommendation:** route every access through one accessor that takes the lock (and make
`defaultRegistry` an `atomic.Pointer[DefaultModelRegistry]` if lock-free reads are wanted — that is
the correct way to get the "never blocks" property finding 2 was reaching for).
### 9. Inconsistent locking discipline (Medium, Locking)
Within one 381-line file:
| Function | `registriesMutex` | `r.mutex` |
|---|---|---|
| `GetDefaultRegistry` | `TryRLock` + fallback | — |
| `SetDefaultRegistry` | `TryLock`, no-op on fail | — |
| `AddRegistry` (`:120`) | blocking `Lock` | — |
| `GetModelByName` (`:293`) | blocking `RLock` | via `GetModel` → `TryRLock` |
| `GetModelRulesByName` (`:365`) | blocking `RLock` | `TryRLock` then blocking `RLock` |
| `GetModels` (`:318`) | `TryRLock`, nil on fail | `tryRLock`, skip on fail |
| `RegisterModel` (`:150`) | — | `tryLock`, error on fail |
| `GetModel` (`:198`) | — | `tryRLock`, error on fail |
| `GetAllModels` (`:212`) | — | `tryRLock`, empty on fail |
| `SetModelRules` (`:237`) | — | blocking `Lock` |
| `GetModelRules` (`:252`) | — | blocking `RLock` |
| `IterateModels` (`:307`) | — | blocking `RLock` |
Four different failure behaviours for the same class of event. The mix also means the try-lock
callers can be starved by the blocking ones (finding 7), so the functions that "can never hang" are
the ones most likely to return garbage.
**Recommendation:** pick one discipline — blocking locks with snapshot-and-release — and apply it
uniformly.
### 10. Reflection under the write lock (Low, Slowness + Locking)
`RegisterModel` holds `r.mutex` (write) from `model_registry.go:151` through `:195`, and inside
that window does `reflect.TypeOf` (`:161`), the unwrap loop (`:169-171`), `reflect.New(...).Elem().Interface()`
(`:181`), and another `reflect.TypeOf` (`:185`). None of that touches `r.models`/`r.rules` and none
of it needs the lock.
This directly lengthens the window that makes finding 1 exploitable. Validate first, then take the
lock only for the two map writes.
### 11. Unbounded unwrap loop on a recursive pointer type (Low, Availability)
`model_registry.go:169-171`
```go
for modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array {
modelType = modelType.Elem()
}
```
`type T *T` is legal Go, and `reflect.Type.Elem()` on it returns itself — so the loop never
terminates. **Verified experimentally:**
```go
type T *T
var x T
tt := reflect.TypeOf(x) // main.T
for tt.Kind() == reflect.Pointer { tt = tt.Elem() } // spins on main.T forever
// → "INFINITE LOOP CONFIRMED after 101 iterations, still main.T"
```
Because the loop runs with the write lock held (finding 10), this doesn't just hang one goroutine —
it wedges the registry permanently, at which point every try-lock caller starts returning
"registry locked", which via finding 1 means **authorisation fails open for the rest of the process
lifetime**.
Requires a pathological model type, so exploitability is near zero; the fix is a one-line depth cap
and it converts a permanent fail-open into an error return.
**Recommendation:** bound the loop (`for depth := 0; depth < 16 && ...; depth++`) and return an
error if the cap is hit.
### 12. No panic handling at all (Low, Panic handling)
The package contains **zero** `recover()` calls and never logs — it does not import `pkg/logger`.
For a pure data structure that is a defensible choice, with two caveats:
- `IterateModels` runs a caller callback under a read lock (finding 6). If `fn` panics, the
`defer RUnlock` does release the lock, so the registry is not wedged — that part is fine — but
the panic propagates to whatever boundary handler exists, and nothing here records which model
was being processed. A `logger`-free package can still name the model in a re-panic.
- Every failure mode in the package is reported as a `fmt.Errorf` string with no wrapping and no
sentinel values, so callers cannot distinguish them (finding 1). That is the panic/error-handling
defect that actually matters here.
**Recommendation:** define `ErrModelNotFound`, `ErrModelExists`, `ErrRegistryUnavailable` as
sentinels and wrap them, so `errors.Is` works at the security layer.
### 13. Permissive default rules (Low, Security)
`DefaultModelRules()` (`model_registry.go:24-36`) returns `CanRead`, `CanUpdate`, `CanCreate`,
`CanDelete` all `true`. `RegisterModel` applies it to any model registered without explicit rules
(`model_registry.go:191-194`), and `GetModelRules` falls back to it as well (`model_registry.go:266`).
The `CanPublic*` flags default to `false` and `SecurityDisabled` to `false`, which is right. But the
authenticated-path flags default open, so `RegisterModel(name, m)` — the form used by
`pkg/testmodels/business.go` `RegisterTestModels` and the `modelregistry.RegisterModel` convenience wrapper —
yields a fully mutable model. Combined with `pkg/security/hooks.go`'s allow-on-error, the system's
default posture at every layer is permit.
**Recommendation:** default to deny and make permissions opt-in, or at minimum log at registration
time when a model is registered without explicit rules.
---
## What looks right
- The struct-vs-pointer validation in `RegisterModel` (`model_registry.go:160-194`) is careful and
well-reasoned: it rejects `nil`, unwraps pointer/slice/array to find the base type, rejects
non-struct kinds with a message naming the original type, normalises a pointer/slice input to a
zero struct value, and re-checks the final type. The error message even tells the caller to use
`MyModel{}` instead of `&MyModel{}`. Good API ergonomics.
- Duplicate registration is rejected (`model_registry.go:156-158`) rather than silently overwriting
— important, since silent overwrite would be a rules-replacement primitive.
- `GetAllModels` returns a **copy** of the map (`model_registry.go:218-222`) rather than the
internal one, so callers cannot mutate registry state or race on it after the lock is dropped.
This is the pattern the rest of the package should follow.
- `GetModelByEntity` (`model_registry.go:225-234`) tries `schema.entity` before bare `entity`,
which is the right precedence and matches what `pkg/security/hooks.go` does.
- `GetModels` de-duplicates by name across registries (`model_registry.go:335-347`), so
registry-order precedence is consistent with `GetModelByName`'s first-match rule.
- Every `defer` for an acquired lock is correctly paired; there is no missing-`Unlock` path. The
problems here are about *which* lock discipline was chosen, not about leaking locks.
## Suggested follow-up
Ordered by risk:
1. **Make `pkg/security/hooks.go` fail closed** (finding 1). This is the single change that
converts a Critical authorisation bypass into a Medium availability issue. It does not require
touching this package.
2. **Remove the try-lock scheme** (findings 2, 3, 5, 7, 9) and fix the underlying hang by
snapshotting in `IterateModels` (finding 6).
3. **Make `RegisterModelWithRules` atomic** (finding 4).
4. Route `defaultRegistry` access through a single locked accessor or `atomic.Pointer` (finding 8).
5. Move reflection out of the write-locked region and cap the unwrap loop (findings 10, 11).
6. **Add tests.** This package has none. Priority cases: `-race` test with concurrent
`RegisterModel` + `GetModelRulesByName` asserting that rules are *never* observed as permissive
for a restrictively-registered model; a test that `GetModelRulesByName` under contention does
not return a "not found"-shaped error; `IterateModels` with a callback that calls back into the
registry (should not deadlock); sentinel-error assertions.
File diff suppressed because it is too large Load Diff
+203
View File
@@ -0,0 +1,203 @@
# Audit — `pkg/testmodels`
- **Date:** 2026-09-29
- **Scope:** `pkg/testmodels/business.go` (161 LOC, 1 file, **no tests**)
- **Axes:** thread locking/waiting · slowness · security · panic handling & logging
- **Threat model:** hostile internet client. These models are registered into
`pkg/modelregistry`, which means any model here becomes a reachable entity for the spec handlers.
## Summary
Six GORM struct definitions (`Department`, `Employee`, `Project`, `ProjectTask`, `Document`,
`Comment`) used as fixtures, plus two registration helpers. No concurrency, no I/O, no panics, no
logging — so three of the four audit axes are trivially clean.
Two real issues: **all six registration errors are discarded**, and this fixture package ships in
`pkg/` (not `_test.go`, not `internal/`) where a consuming application can register test tables
into a production registry.
| # | Severity | Axis | Finding |
|---|----------|------|---------|
| 1 | Medium | Correctness | `RegisterTestModels` discards all six `RegisterModel` error returns |
| 2 | Medium | Security | Fixtures live in exported `pkg/`, registerable into a production model registry |
| 3 | Low | Security | Models are registered via `RegisterModel`, which applies permissive `DefaultModelRules` |
| 4 | Low | Correctness | `GetTestModels()` return order is unrelated to FK dependency order |
| 5 | Low | Correctness | `Document.Path` is an unconstrained filesystem path exposed as a writable API field |
---
## Findings
### 1. All registration errors discarded (Medium, Correctness)
`business.go:142-149`
```go
func RegisterTestModels(registry *modelregistry.DefaultModelRegistry) {
registry.RegisterModel("departments", Department{})
registry.RegisterModel("employees", Employee{})
registry.RegisterModel("projects", Project{})
registry.RegisterModel("project_tasks", ProjectTask{})
registry.RegisterModel("documents", Document{})
registry.RegisterModel("comments", Comment{})
}
```
`RegisterModel` returns `error` and every return value is dropped. The function itself returns
nothing, so a caller cannot detect failure either.
This matters more than usual because of how `pkg/modelregistry.RegisterModel` fails. It has two
error paths (`pkg/modelregistry/model_registry.go:151-158`):
```go
if !r.tryLock() {
return fmt.Errorf("failed to register model %s: registry locked", name)
}
...
if _, exists := r.models[name]; exists {
return fmt.Errorf("model %s already registered", name)
}
```
The first is a **transient lock-contention failure** — see `audit/pkg/modelregistry.audit.md`
finding 7, where a contended `tryLock` gives up after ~20 ms. So under concurrent registration, some
subset of these six models silently fails to register, with no error, no log, and no panic. The
process then runs with, say, `documents` and `comments` missing from the registry.
That is not merely a missing-fixture annoyance. Per `audit/pkg/modelregistry.audit.md` finding 1, an
unregistered model causes `pkg/security/hooks.go:274-294` to take the
`return nil // model not registered, allow by default` branch — so a silently-failed registration
turns into **authorisation fail-open** for that entity.
Note `errcheck` is enabled (golangci-lint v2 standard set) but `.golangci.json` excludes
`"tests?"` paths — `pkg/testmodels` does not match that pattern, so this *should* be flagged
today. Worth checking whether the linter is actually run in CI.
**Recommendation:** return `error`, and use `errors.Join` so a partial failure is reported in full:
```go
func RegisterTestModels(registry *modelregistry.DefaultModelRegistry) error {
return errors.Join(
registry.RegisterModel("departments", Department{}),
registry.RegisterModel("employees", Employee{}),
...
)
}
```
### 2. Fixtures are exported from `pkg/` (Medium, Security)
The package path is `github.com/bitechdev/ResolveSpec/pkg/testmodels`, not a `_test.go` file and not
under `internal/`. Consequences:
- The six structs and both helpers are part of ResolveSpec's **public API surface**. They are
compiled into every binary that imports anything which transitively imports this package.
- A consuming application (or a copy-pasted quickstart) that calls
`testmodels.RegisterTestModels(registry)` against its production registry makes
`departments`, `employees`, `projects`, `project_tasks`, `documents` and `comments` live entities
on the spec handlers, addressable by name. If the production database happens to have tables with
those names — `documents` and `comments` are very common names — the handlers will happily
read and write them under the permissive default rules of finding 3.
- It also means any future model added here for test convenience automatically becomes reachable.
Nothing in `pkg/` currently calls `RegisterTestModels` (only the test tree does), so this is a
packaging hazard rather than a live exposure.
**Recommendation:** move to `internal/testmodels` (blocks external import outright) or to a
`testmodels_test` package / `testdata` helper. If it must stay importable for downstream tests,
document loudly and consider a build tag.
### 3. Registered with permissive default rules (Low, Security)
`RegisterTestModels` uses `RegisterModel`, not `RegisterModelWithRules`. Per
`pkg/modelregistry/model_registry.go:191-194`, that initialises each model with
`DefaultModelRules()`, which grants `CanRead`, `CanUpdate`, `CanCreate` and `CanDelete` — see
`audit/pkg/modelregistry.audit.md` finding 13. `CanPublic*` are `false`, which is the saving grace.
If finding 2 is acted on this becomes moot; if these models are intended to stay registerable, they
should be registered read-only.
### 4. `GetTestModels()` order is not dependency order (Low, Correctness)
`business.go:152-160` returns the models in declaration order:
`Department, Employee, Project, ProjectTask, Document, Comment`.
The FK graph is not satisfied by that order. `Employee.DepartmentID → Department.ID` happens to work,
but `Document.OwnerID → Employee.ID` and `Document.ProjectID → Project.ID` mean `Document` must
follow both, and `ProjectTask.AssigneeID → Employee.ID` and `ProjectTask.ProjectID → Project.ID`
likewise. Coincidentally the declaration order does satisfy these — but nothing enforces it, and
`Employee.ManagerID → Employee.ID` is self-referential, which several migration/auto-migrate paths
handle only if the self-FK is deferred.
Also the two `many2many` joins (`department_projects`, `employee_projects`, declared at
`business.go:20`, `:45`, `:67-68`) are not in the returned list at all, so a caller using
`GetTestModels()` to drive `AutoMigrate` gets the join tables only because GORM infers them from the
tags — a Bun-based migration path (`pkg/common/adapters/database/bun.go`) would not.
**Recommendation:** document that the order is migration-safe and add a comment stating the
constraint, or return an explicitly ordered list with a test that asserts it.
### 5. `Document.Path` is an unconstrained path field (Low, Security)
`business.go:107`
```go
Path string `json:"path"`
```
No validation, no length limit, no `gorm` constraint. As a plain string column it is inert — the
risk only materialises if some handler or downstream consumer uses it to open a file, at which point
an attacker who can `POST`/`PATCH` a `Document` controls a filesystem path (`../../etc/passwd`,
`/proc/self/environ`). The same applies to `ContentType` (`business.go:105`) if it is ever echoed
into a response header unvalidated, and `Size` (`business.go:106`) which is a client-settable
`int64` that can disagree with reality.
Nothing in `pkg/` reads these fields, so this is a note about the fixture's shape rather than a
present vulnerability — but it is a bad example to ship, since fixtures get copied.
**Recommendation:** if these stay, mark `Path` as server-set (a `gorm:"->"` read-only tag, or
exclude it from the writable column set) so the fixture demonstrates the safe pattern.
---
## Axis-by-axis
- **Thread locking / waiting:** nothing to report. The package declares no goroutines, channels,
mutexes or atomics. Its only concurrency exposure is *through* `pkg/modelregistry`, covered in
finding 1 and in that package's audit.
- **Slowness:** nothing to report. `RegisterTestModels` and `GetTestModels` are O(1) with six
elements and are startup-only. The `TableName()` methods (`business.go:23`, `:49`, `:73`, `:96`,
`:119`, `:137`) return constants — no allocation, no reflection.
- **Security:** findings 2, 3, 5 — all about packaging and field shape, none about code behaviour.
- **Panic handling and logging:** the package contains no `panic`, no `recover`, and does not import
`pkg/logger`. For plain struct definitions that is correct. The one place where logging *would*
belong is the discarded errors of finding 1 — silently dropping six error returns is the
panic/error-handling defect in this package, even though no panic is involved.
## What looks right
- Struct tags are consistent and complete: `json` on every field, `gorm:"primaryKey"` on every ID,
`gorm:"uniqueIndex"` on the natural keys (`Department.Code`, `Employee.Email`, `Project.Code`),
and explicit `foreignKey`/`references` on every relation rather than relying on GORM's inference.
That makes these fixtures genuinely useful for exercising the relation-expansion paths in
`pkg/restheadspec` and `pkg/resolvespec`.
- `omitempty` on every relation field prevents empty relation arrays from bloating responses — which
matters, because these fixtures are what the handler tests measure payloads against.
- Nullable FKs are correctly modelled as `*string` (`Employee.ManagerID` `business.go:35`,
`Document.ProjectID` `business.go:109`) rather than empty-string sentinels.
- The self-referential manager/reports pair (`business.go:43-44`) and the two `many2many` relations
give reasonable coverage of the harder relation shapes — a genuinely well-chosen fixture set for
the recursive-preload logic audited in `audit/pkg/restheadspec.audit.md`.
- `TableName()` is defined on the value receiver for all six, so it works whether a value or a
pointer is passed — which matters given `pkg/modelregistry.RegisterModel` normalises pointers to
values.
## Suggested follow-up
1. Return and check errors from `RegisterTestModels` (finding 1). One-line-per-call change, and it
closes a silent path to authorisation fail-open.
2. Decide whether this package belongs in `pkg/` at all (finding 2). `internal/testmodels` is the
low-effort fix.
3. Confirm `golangci-lint` runs in CI and that `errcheck` flags `business.go:143-148` — if it does
not, the exclusion patterns in `.golangci.json` need review, since this is exactly the class of
bug it exists to catch.
+286
View File
@@ -0,0 +1,286 @@
# Audit — `pkg/tracing`
- **Date:** 2026-09-29
- **Scope:** `pkg/tracing/tracing.go` (146 LOC, 1 file, **no tests**)
- **Axes:** thread locking/waiting · slowness · security · panic handling & logging
- **Threat model:** hostile internet client. Span names and attributes here are built directly from
request-controlled data (method, path, full URL, Host header).
## Summary
A thin OpenTelemetry wrapper: `InitTracer`, an HTTP middleware, and helpers. The abstraction is
fine; the **hardcoded choices** are the problem. Three of them are not configurable at all and each
is wrong for a production, internet-facing deployment:
- `otlptracegrpc.WithInsecure()` — trace export is **plaintext**, with a source comment admitting it.
- `sdktrace.AlwaysSample()` — **100% of requests** are traced, with no sampling knob in config.
- `semconv.HTTPURLKey.String(r.URL.String())` — the **full URL including query string** is exported.
Combined: every request's full URL is shipped unencrypted to a collector, and an attacker sets the
export volume. Span names are also built from raw paths, giving unbounded cardinality.
`config.TracingConfig` (`pkg/config/config.go:86-91`) exposes only `Enabled`, `ServiceName`,
`ServiceVersion` and `Endpoint` — there is no field for TLS or sample rate, so these cannot be fixed
by configuration alone.
| # | Severity | Axis | Finding |
|---|----------|------|---------|
| 1 | **High** | Security | `WithInsecure()` hardcoded — traces exported in plaintext, not configurable |
| 2 | **High** | Security | Full URL **including query string** exported as a span attribute |
| 3 | **High** | Slowness | `AlwaysSample()` hardcoded — 100% trace volume, attacker-controlled, no sampling config |
| 4 | Medium | Slowness | Span name is `method + " " + r.URL.Path` — unbounded cardinality from raw path IDs |
| 5 | Medium | Locking | `tracer` global written by `InitTracer`, read unsynchronised by `Middleware`/`StartSpan` |
| 6 | Medium | Observability | `Middleware` records no HTTP status and no error status — spans never show failures |
| 7 | Medium | Panic | `Middleware` does not recover; a downstream panic leaves the span unmarked (`Unset` status) |
| 8 | Low | Slowness | `InitTracer` has no timeout/deadline on exporter or resource creation |
| 9 | Low | Maintenance | `semconv/v1.4.0` (2021) — deprecated attribute names modern collectors no longer index |
| 10 | Low | Security | `SetAttributes`/`AddEvent` pass caller data through with no size or cardinality limit |
## Resolution (2026-09-30)
Fixed in `pkg/tracing/tracing.go`, `pkg/config` (`TracingConfig`, defaults), the package README, and new
`pkg/tracing/tracing_test.go` (passes).
| # | Status | What changed |
|---|--------|--------------|
| 1 | **Fixed** | TLS is the default. `Config` gains `Insecure`, `TLSConfig` and `Headers` (OTLP auth). `tracing.insecure` added to `pkg/config`. **Breaking:** plaintext collectors now need `Insecure: true`. |
| 2 | **Fixed** | Query string and `Host` are no longer exported; attributes are method, `url.path`, scheme, `http.route`, status. `TLSConfig` has no config-file key (code only). |
| 3 | **Fixed** | `ParentBased(TraceIDRatioBased(rate))`; `SampleRate` defaults to 0.1, validated to [0,1]; `tracing.sample_rate` added to `pkg/config`. |
| 4 | **Fixed** | Span name is `METHOD <route template>` from `Request.Pattern`, `<unmatched>` otherwise. `MiddlewareWithRoute(fn)` supports other routers. |
| 5 | **Fixed** | `tracer` is an `atomic.Pointer`; a second `InitTracer` returns an error; the shutdown func resets state. |
| 6 | **Fixed** | Response writer wrapped; `http.response.status_code` recorded, 5xx sets Error status. Preserves `Flush`/`Unwrap`. |
| 7 | **Fixed** | Panics are recorded (`RecordError`, Error status) and re-raised so the panic middleware still responds. Must be installed inside the panic middleware; actual order in `pkg/server` not verified. |
| 8 | **Fixed** | `InitTracerContext(ctx, cfg)` with `InitTimeout` (default 10s); `InitTracer` retained as a wrapper. Exporter is shut down if resource creation fails. |
| 9 | **Fixed** | Moved to `semconv/v1.26.0`. |
| 10 | **Fixed** | `AttributeValueLengthLimit` set via `WithRawSpanLimits` (`AttributeValueLimit`, default 1024). |
Tests added: query redaction and route naming, unmatched route, 5xx status, panic recorded and re-raised,
double-init and invalid sample rate.
---
## Findings
### 1. `WithInsecure()` hardcoded (High, Security)
`tracing.go:38-42`
```go
client := otlptracegrpc.NewClient(
otlptracegrpc.WithEndpoint(config.Endpoint),
otlptracegrpc.WithInsecure(), // Use WithTLSCredentials in production
)
```
The comment names the fix and the code does not implement it, and — critically — `Config`
(`tracing.go:21-27`) has no field to express it:
```go
type Config struct {
ServiceName string
ServiceVersion string
Endpoint string
Enabled bool
}
```
So there is **no supported way** to enable TLS on trace export short of editing this file. Every
span — carrying the full request URL per finding 2 — crosses the network in cleartext, and the
collector endpoint is unauthenticated (no OTLP headers/bearer token option either), so anything that
can reach it can also *inject* fabricated spans.
**Recommendation:** add `Insecure bool`, `TLSConfig *tls.Config` and `Headers map[string]string` to
`Config` (and the matching `tracing.*` keys to `pkg/config`), default to TLS on, and require an
explicit opt-in for insecure. Wire `otlptracegrpc.WithTLSCredentials` / `WithHeaders`.
### 2. Full URL with query string exported (High, Security)
`tracing.go:95-103`
```go
ctx, span := tracer.Start(ctx, r.Method+" "+r.URL.Path,
trace.WithSpanKind(trace.SpanKindServer),
trace.WithAttributes(
semconv.HTTPMethodKey.String(r.Method),
semconv.HTTPURLKey.String(r.URL.String()), // <- full URL, query string included
semconv.HTTPTargetKey.String(r.URL.Path),
semconv.HTTPSchemeKey.String(r.URL.Scheme),
semconv.NetHostNameKey.String(r.Host),
),
)
```
`r.URL.String()` includes `RawQuery`. For this API the query string is where the interesting data
lives: filter expressions, column lists, and — for any client that passes credentials as a query
parameter (`?api_key=`, `?token=`, signed-URL style parameters) — secrets. All of it lands in the
tracing backend, and per finding 1 it gets there in plaintext.
Note `HTTPTargetKey` is also set to `r.URL.Path`, so the *useful* part is already captured
separately; `HTTPURLKey` adds only the sensitive part.
Secondary: `r.Host` comes from the `Host` header, which is client-controlled and unvalidated here —
so an attacker can pollute the `net.host.name` dimension with arbitrary values (cardinality blowup,
and log/dashboard spoofing).
**Recommendation:** export a redacted URL (scheme + host + path, query keys only or dropped
entirely). OTel's own guidance is to strip or redact query parameters for exactly this reason.
### 3. `AlwaysSample()` hardcoded (High, Slowness)
`tracing.go:61-65`
```go
tp := sdktrace.NewTracerProvider(
sdktrace.WithBatcher(exporter),
sdktrace.WithResource(res),
sdktrace.WithSampler(sdktrace.AlwaysSample()),
)
```
Every request produces a recorded, exported span. There is no `SampleRate` in `Config` and no
`tracing.sample_rate` key in `pkg/config/manager.go`'s defaults, so this is not tunable.
Under the hostile-client threat model the request rate — and therefore the span rate, the batch
queue pressure, the serialisation cost and the outbound bandwidth — is set by the attacker. Each
request pays span allocation, attribute encoding (including the full URL string), and a share of
batch export. When the batch queue fills, the SDK drops spans, so a flood also destroys the
observability you need to see the flood.
**Recommendation:** default to `sdktrace.ParentBased(sdktrace.TraceIDRatioBased(rate))` with a
configurable rate (e.g. 0.01–0.1), keeping `AlwaysSample` available for development. `ParentBased`
also means an upstream sampling decision is respected, which `AlwaysSample` currently overrides.
### 4. Unbounded span-name cardinality (Medium, Slowness)
`tracing.go:95` — the span name is `r.Method + " " + r.URL.Path`.
This API's paths embed identifiers (`/api/public/employees/7f3c…`, `/api/<schema>/<entity>/<id>`), so
each distinct ID becomes a distinct span name. Consequences:
- Tracing backends index on span name; unbounded distinct names is the classic cardinality-explosion
cost bomb (and in some backends, a hard limit that starts rejecting data).
- It violates the OTel HTTP convention, which requires a **low-cardinality route template**
(`GET /api/{schema}/{entity}/{id}`), with the concrete value in `http.route`/attributes.
- It is attacker-driven: requests to random paths — including 404s — each mint a new span name.
**Recommendation:** derive the name from the matched route pattern. `pkg/server`'s router
(chi/mux/gin, see `audit/pkg/server.audit.md`) exposes the route template after matching; use it, and
place this middleware after the router so the pattern is available. Fall back to
`r.Method + " " + "<unmatched>"` rather than the raw path.
### 5. Unsynchronised `tracer` global (Medium, Locking)
`tracing.go:19`
```go
var tracer trace.Tracer
```
Written at `tracing.go:77` (`tracer = tp.Tracer(config.ServiceName)`), read at `tracing.go:86`,
`:95` (`Middleware`) and `:116`, `:119` (`StartSpan`). No mutex, no `atomic.Value`.
Same pattern as `pkg/logger`'s `Logger` global (see `audit/pkg/logger.audit.md` finding 1). Benign if
`InitTracer` runs once before any request is served; a race the moment tracing is re-initialised at
runtime. `InitTracer` is exported and callable at any time, and calling it twice also leaks the
first `TracerProvider` (nothing shuts it down) — its batch processor goroutine and gRPC connection
stay alive for the life of the process.
Note the nil checks at `:86` and `:116` are the read-half of the race: a goroutine can observe a
non-nil-but-torn interface value.
**Recommendation:** `atomic.Pointer` or a `sync.Once`-guarded init; return an error from a second
`InitTracer` call, or shut down the previous provider first.
### 6. No HTTP status or error status on spans (Medium, Observability)
`Middleware` (`tracing.go:84-112`) never wraps `w`, so it cannot observe the status code. It sets no
`semconv.HTTPStatusCodeKey` and never calls `span.SetStatus`. Every span therefore has status
`Unset`, which tracing backends render as "OK".
The practical effect: you cannot find failing requests in the traces. A 500-storm and a healthy
period look identical in the span data, which defeats the main reason to run tracing on an
internet-facing service.
**Recommendation:** wrap the `ResponseWriter` to capture the status, set
`semconv.HTTPStatusCodeKey.Int(status)`, and `span.SetStatus(codes.Error, ...)` for 5xx.
### 7. No panic handling in the middleware (Medium, Panic handling)
`Middleware` has `defer span.End()` (`tracing.go:106`) but no `recover()`. If `next.ServeHTTP`
panics:
- The `defer span.End()` **does** run, so no span is leaked — that part is correct.
- But the span is ended with status `Unset` and no exception event, so the panic is invisible in the
trace. The one place a trace would be most valuable records nothing.
- The panic propagates up to whichever handler is outermost. Whether that is
`pkg/middleware/panic.go` depends on middleware ordering — if `tracing.Middleware` is installed
*outside* the panic middleware, the panic escapes to `net/http`'s per-connection recovery, which
kills the connection and logs to the default logger, bypassing `pkg/logger` and the error tracker
entirely. See `audit/pkg/middleware.audit.md` and `audit/pkg/server.audit.md` for the actual order.
This package does not import `pkg/logger` at all, so nothing here can be logged.
**Recommendation:** recover, record `span.RecordError` + `span.SetStatus(codes.Error, …)`, then
re-panic so the dedicated panic middleware still handles the response. Document the required
middleware order.
### 8. No deadline on initialisation (Low, Slowness)
`tracing.go:36` uses `ctx := context.Background()` for both `otlptrace.New` (`:44`) and
`resource.New` (`:51`). `otlptracegrpc` does not block on connect by default, so this is unlikely to
hang today — but `resource.New` with detectors can perform network calls (cloud metadata endpoints),
and an unreachable metadata service is a classic multi-second startup stall. `InitTracer` should
accept a `context.Context` from the caller so startup has a deadline.
### 9. `semconv/v1.4.0` (Low, Maintenance)
`tracing.go:15` pins the 2021 semantic conventions. `http.method`, `http.url`, `http.target`,
`http.scheme`, `net.host.name` were all renamed in v1.20+ (`http.request.method`, `url.full`,
`url.path`, `url.scheme`, `server.address`). Current collectors, dashboards and backend
auto-instrumentation views key off the new names, so these spans will not populate standard HTTP
dashboards.
### 10. No limits on caller-supplied span data (Low, Security)
`StartSpan`, `AddEvent`, `SetAttributes` (`tracing.go:115-145`) forward caller attributes verbatim.
If any caller passes request-derived values (a filter expression, a row payload), span size is
attacker-influenced. The SDK's default limits (128 attributes, 128 events) cap the count but not the
*value* length — a 1 MB string attribute is accepted.
**Recommendation:** set explicit `sdktrace.WithSpanLimits` including `AttributeValueLengthLimit`.
---
## What looks right
- **Disabled path is genuinely free.** `InitTracer` with `Enabled: false` (`tracing.go:31-34`)
returns a no-op shutdown func and never builds an exporter, so a disabled deployment pays nothing
and cannot leak.
- **Nil-tracer guards everywhere.** `Middleware` (`tracing.go:86-89`) passes through untouched and
`StartSpan` (`tracing.go:116-118`) returns the incoming context plus the context's (no-op) span.
So a partially-initialised process degrades safely rather than nil-panicking — a pattern
`pkg/logger` gets right too.
- **Context propagation is correct.** `Extract` from `propagation.HeaderCarrier(r.Header)`
(`tracing.go:92`), a composite `TraceContext` + `Baggage` propagator (`tracing.go:72-75`), and
`r = r.WithContext(ctx)` (`tracing.go:109`) before calling `next` — the span context actually
reaches downstream handlers, which is the part most hand-rolled middlewares get wrong.
- `SpanKindServer` is set correctly (`tracing.go:96`).
- `WithBatcher` rather than a simple/sync span processor (`tracing.go:62`) — export does not block
the request path.
- `InitTracer` returns `tp.Shutdown` (`tracing.go:80`), giving the caller a real flush-on-shutdown
hook with a caller-supplied context, which is better than the fixed-timeout pattern in
`pkg/errortracking` (see that audit, finding 3).
- `RecordError` nil-guards (`tracing.go:140-143`) so `RecordError(ctx, nil)` is a no-op.
## Suggested follow-up
1. Extend `Config` with `Insecure`, TLS credentials, OTLP headers and `SampleRate`; add the matching
keys to `pkg/config` (findings 1, 3). These cannot be fixed without an API change, so they should
go together.
2. Redact the query string from exported attributes (finding 2).
3. Move to route-template span names, which requires positioning the middleware after routing
(finding 4).
4. Capture status code and panics in the middleware (findings 6, 7).
5. Guard the `tracer` global (finding 5) and upgrade `semconv` (finding 9).
6. Add tests: this package has none. A tracetest/in-memory exporter makes assertions on span name,
attributes and status straightforward, and would have caught findings 2, 4 and 6.
+1
View File
@@ -24,6 +24,7 @@ import (
func main() {
// Load configuration
cfgMgr := config.NewManager()
config.SetConfigManager(cfgMgr)
if err := cfgMgr.Load(); err != nil {
log.Fatalf("Failed to load configuration: %v", err)
}
+2 -2
View File
@@ -41,7 +41,9 @@ require (
go.uber.org/zap v1.28.0
golang.org/x/crypto v0.55.0
golang.org/x/oauth2 v0.36.0
golang.org/x/sys v0.47.0
golang.org/x/time v0.15.0
google.golang.org/grpc v1.83.2
gorm.io/driver/postgres v1.6.0
gorm.io/driver/sqlite v1.6.0
gorm.io/driver/sqlserver v1.6.3
@@ -147,11 +149,9 @@ require (
golang.org/x/mod v0.38.0 // indirect
golang.org/x/net v0.58.0 // indirect
golang.org/x/sync v0.22.0 // indirect
golang.org/x/sys v0.47.0 // indirect
golang.org/x/text v0.41.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
google.golang.org/grpc v1.83.2 // indirect
google.golang.org/protobuf v1.36.11 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect
modernc.org/libc v1.72.3 // indirect
+30 -17
View File
@@ -3,23 +3,29 @@ package cache
import (
"context"
"fmt"
"sync/atomic"
"time"
)
var (
defaultCache *Cache
)
var defaultCache atomic.Pointer[Cache]
// swapOwned installs c as the default and closes the displaced cache, which this
// package created and therefore owns.
func swapOwned(c *Cache) {
if old := defaultCache.Swap(c); old != nil && old != c {
_ = old.Close() // best-effort: the displaced provider is being discarded
}
}
// Initialize initializes the cache with a provider.
// If not called, the package will use an in-memory provider by default.
func Initialize(provider Provider) {
defaultCache = NewCache(provider)
swapOwned(NewCache(provider))
}
// UseMemory configures the cache to use in-memory storage.
func UseMemory(opts *Options) error {
provider := NewMemoryProvider(opts)
defaultCache = NewCache(provider)
swapOwned(NewCache(NewMemoryProvider(opts)))
return nil
}
@@ -29,7 +35,7 @@ func UseRedis(config *RedisConfig) error {
if err != nil {
return fmt.Errorf("failed to initialize Redis provider: %w", err)
}
defaultCache = NewCache(provider)
swapOwned(NewCache(provider))
return nil
}
@@ -39,26 +45,33 @@ func UseMemcache(config *MemcacheConfig) error {
if err != nil {
return fmt.Errorf("failed to initialize Memcache provider: %w", err)
}
defaultCache = NewCache(provider)
swapOwned(NewCache(provider))
return nil
}
// GetDefaultCache returns the default cache instance.
// Initializes with in-memory provider if not already initialized.
// Safe for concurrent use.
func GetDefaultCache() *Cache {
if defaultCache == nil {
_ = UseMemory(&Options{
DefaultTTL: 5 * time.Minute,
MaxSize: 10000,
})
if c := defaultCache.Load(); c != nil {
return c
}
return defaultCache
fresh := NewCache(NewMemoryProvider(&Options{
DefaultTTL: 5 * time.Minute,
MaxSize: 10000,
}))
if defaultCache.CompareAndSwap(nil, fresh) {
return fresh
}
_ = fresh.Close() // lost the race; discard our provider
return defaultCache.Load()
}
// SetDefaultCache sets a custom cache instance as the default cache.
// This is useful for testing or when you want to use a pre-configured cache instance.
// The caller keeps ownership of both the new and the displaced cache; neither is closed.
func SetDefaultCache(cache *Cache) {
defaultCache = cache
defaultCache.Store(cache)
}
// GetStats returns cache statistics.
@@ -69,8 +82,8 @@ func GetStats(ctx context.Context) (*CacheStats, error) {
// Close closes the cache and releases resources.
func Close() error {
if defaultCache != nil {
return defaultCache.Close()
if c := defaultCache.Load(); c != nil {
return c.Close()
}
return nil
}
+19 -6
View File
@@ -3,10 +3,17 @@ package cache
import (
"context"
"encoding/json"
"errors"
"fmt"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// ErrNotFound is returned when a key is not in the cache. The key is deliberately
// not included in the error: keys may embed credentials (e.g. session tokens).
var ErrNotFound = errors.New("cache: key not found")
// Cache is the main cache manager that wraps a Provider.
type Cache struct {
provider Provider
@@ -23,7 +30,7 @@ func NewCache(provider Provider) *Cache {
func (c *Cache) Get(ctx context.Context, key string, dest interface{}) error {
data, exists := c.provider.Get(ctx, key)
if !exists {
return fmt.Errorf("key not found: %s", key)
return ErrNotFound
}
if err := json.Unmarshal(data, dest); err != nil {
@@ -37,7 +44,7 @@ func (c *Cache) Get(ctx context.Context, key string, dest interface{}) error {
func (c *Cache) GetBytes(ctx context.Context, key string) ([]byte, error) {
data, exists := c.provider.Get(ctx, key)
if !exists {
return nil, fmt.Errorf("key not found: %s", key)
return nil, ErrNotFound
}
return data, nil
}
@@ -122,9 +129,10 @@ func (c *Cache) GetOrSet(ctx context.Context, key string, dest interface{}, ttl
return fmt.Errorf("loader failed: %w", err)
}
// Store in cache
// Store in cache. A cache-write failure must not fail the call: the authoritative
// value has already been loaded, and the cache is only an optimisation.
if err := c.Set(ctx, key, value, ttl); err != nil {
return fmt.Errorf("failed to cache value: %w", err)
logger.Warn("cache: failed to store loaded value, continuing uncached: %v", err)
}
// Populate dest with the loaded value
@@ -142,6 +150,11 @@ func (c *Cache) GetOrSet(ctx context.Context, key string, dest interface{}, ttl
// Remember is a convenience function that caches the result of a function call.
// It's similar to GetOrSet but returns the value directly.
//
// WARNING: the returned type differs between a hit and a miss. On a hit the value is
// generic decoded JSON (map[string]interface{}, []interface{}, float64, string, ...);
// on a miss it is exactly what loader returned. Do not type-assert the result to a
// concrete type; prefer GetOrSet, which decodes into a typed destination.
func (c *Cache) Remember(ctx context.Context, key string, ttl time.Duration, loader func() (interface{}, error)) (interface{}, error) {
// Try to get from cache first as bytes
data, err := c.GetBytes(ctx, key)
@@ -158,9 +171,9 @@ func (c *Cache) Remember(ctx context.Context, key string, ttl time.Duration, loa
return nil, fmt.Errorf("loader failed: %w", err)
}
// Store in cache
// Cache-write failures are non-fatal (see GetOrSet)
if err := c.Set(ctx, key, value, ttl); err != nil {
return nil, fmt.Errorf("failed to cache value: %w", err)
logger.Warn("cache: failed to store loaded value, continuing uncached: %v", err)
}
return value, nil
+4
View File
@@ -1,3 +1,7 @@
//go:build ignore
// Examples are excluded from the build: they call log.Fatal and are not part of the API.
package cache
import (
+154
View File
@@ -0,0 +1,154 @@
package cache
import (
"context"
"errors"
"fmt"
"strings"
"sync"
"testing"
"time"
)
type failSetProvider struct{ *MemoryProvider }
func (f failSetProvider) Set(context.Context, string, []byte, time.Duration) error {
return errors.New("backend down")
}
func TestGetOrSetSurvivesCacheWriteFailure(t *testing.T) {
c := NewCache(failSetProvider{NewMemoryProvider(nil)})
var out string
err := c.GetOrSet(context.Background(), "k", &out, time.Minute, func() (interface{}, error) { return "v", nil })
if err != nil || out != "v" {
t.Fatalf("got %q, %v", out, err)
}
}
func TestNotFoundDoesNotLeakKey(t *testing.T) {
c := NewCache(NewMemoryProvider(nil))
err := c.Get(context.Background(), "auth:session:SECRET", new(string))
if !errors.Is(err, ErrNotFound) || strings.Contains(err.Error(), "SECRET") {
t.Fatalf("unexpected error %v", err)
}
}
func TestMemoryTagIndexCleanedOnAllRemovals(t *testing.T) {
ctx := context.Background()
m := NewMemoryProvider(&Options{MaxSize: 2})
defer m.Close()
for i := 0; i < 50; i++ {
_ = m.SetWithTags(ctx, fmt.Sprintf("k%d", i), []byte("x"), time.Minute, []string{"t"})
}
m.mu.RLock()
n := len(m.tagToKeys["t"])
m.mu.RUnlock()
if n > 2 {
t.Fatalf("tag index leaked: %d members with MaxSize 2", n)
}
_ = m.Clear(ctx)
if len(m.tagToKeys) != 0 {
t.Fatal("Clear did not reset tag index")
}
}
func TestMemoryClosedNoPanic(t *testing.T) {
m := NewMemoryProvider(nil)
_ = m.Close()
if err := m.Set(context.Background(), "k", []byte("v"), 0); !errors.Is(err, ErrClosed) {
t.Fatalf("got %v", err)
}
if _, ok := m.Get(context.Background(), "k"); ok {
t.Fatal("hit after close")
}
}
func TestMemoryCopiesAndDefaults(t *testing.T) {
ctx := context.Background()
opts := &Options{}
m := NewMemoryProvider(opts)
defer m.Close()
if opts.MaxSize != 0 || m.options.MaxSize != defaultMemoryMaxSize {
t.Fatal("options not copied/defaulted")
}
buf := []byte("abc")
_ = m.Set(ctx, "k", buf, time.Minute)
buf[0] = 'X'
got, _ := m.Get(ctx, "k")
if string(got) != "abc" {
t.Fatalf("stored slice aliased caller: %q", got)
}
got[0] = 'Y'
if again, _ := m.Get(ctx, "k"); string(again) != "abc" {
t.Fatal("returned slice aliases stored value")
}
}
func TestMemoryJanitorRemovesExpired(t *testing.T) {
m := NewMemoryProvider(&Options{CleanupInterval: 10 * time.Millisecond})
defer m.Close()
_ = m.Set(context.Background(), "k", []byte("v"), 5*time.Millisecond)
time.Sleep(100 * time.Millisecond)
m.mu.RLock()
n := len(m.items)
m.mu.RUnlock()
if n != 0 {
t.Fatalf("expired item still stored: %d", n)
}
}
func TestMemoryConcurrent(t *testing.T) {
ctx := context.Background()
m := NewMemoryProvider(&Options{MaxSize: 50})
defer m.Close()
var wg sync.WaitGroup
for g := 0; g < 8; g++ {
wg.Add(1)
go func(g int) {
defer wg.Done()
for i := 0; i < 300; i++ {
k := fmt.Sprintf("k%d", i%80)
_ = m.SetWithTags(ctx, k, []byte("v"), time.Millisecond, []string{"t"})
m.Get(ctx, k)
if i%50 == 0 {
_ = m.DeleteByTag(ctx, "t")
}
}
}(g)
}
wg.Wait()
}
func TestGetDefaultCacheConcurrent(t *testing.T) {
SetDefaultCache(nil)
var wg sync.WaitGroup
res := make([]*Cache, 16)
for i := range res {
wg.Add(1)
go func(i int) { defer wg.Done(); res[i] = GetDefaultCache() }(i)
}
wg.Wait()
for _, c := range res {
if c != res[0] {
t.Fatal("different default caches returned")
}
}
}
func TestMemcacheKeyAndExpiry(t *testing.T) {
if k := memcacheKey(strings.Repeat("a", 300)); len(k) > 250 || !legalMemcacheKey(k) {
t.Fatalf("bad key %q", k)
}
if k := memcacheKey("has space"); !legalMemcacheKey(k) {
t.Fatal("illegal key not normalised")
}
if memcacheKey("a") != "k:a" {
t.Fatal("unexpected prefix")
}
if got := memcacheExpiry(30*24*time.Hour + time.Hour); got < int32(time.Now().Unix()) {
t.Fatalf("expected absolute timestamp, got %d", got)
}
if memcacheExpiry(-time.Second) != 0 || memcacheExpiry(time.Minute) != 60 {
t.Fatal("bad relative expiry")
}
}
+10
View File
@@ -2,9 +2,14 @@ package cache
import (
"context"
"errors"
"time"
)
// ErrFlushNotAllowed is returned by Clear on shared-server providers (Redis, Memcache)
// unless AllowFlush is set in their config, because Clear flushes the whole server/DB.
var ErrFlushNotAllowed = errors.New("cache: Clear flushes the entire server; set AllowFlush in the provider config to permit it")
// Provider defines the interface that all cache providers must implement.
type Provider interface {
// Get retrieves a value from the cache by key.
@@ -58,8 +63,13 @@ type Options struct {
DefaultTTL time.Duration
// MaxSize is the maximum number of items (for in-memory provider).
// 0 selects the default (10000); a negative value means unbounded.
MaxSize int
// CleanupInterval is how often the in-memory provider removes expired items
// (default: 1 minute).
CleanupInterval time.Duration
// EvictionPolicy determines how items are evicted (LRU, LFU, etc).
EvictionPolicy string
}
+238 -119
View File
@@ -2,17 +2,34 @@ package cache
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"time"
"github.com/bradfitz/gomemcache/memcache"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
const (
// memcacheMaxRelativeTTL is the largest expiry memcached treats as relative seconds;
// anything larger is interpreted as an absolute Unix timestamp.
memcacheMaxRelativeTTL = 30 * 24 * 60 * 60
// memcacheMaxTagKeys bounds the per-tag key list so it stays under memcached's 1MB item limit.
memcacheMaxTagKeys = 5000
memcacheCASRetries = 5
)
// MemcacheProvider is a Memcache implementation of the Provider interface.
type MemcacheProvider struct {
client *memcache.Client
options *Options
client *memcache.Client
options *Options
allowFlush bool
}
// MemcacheConfig contains Memcache-specific configuration.
@@ -28,37 +45,46 @@ type MemcacheConfig struct {
// Options contains general cache options
Options *Options
// AllowFlush permits Clear() to run flush_all, which wipes every key on every
// configured server (including data not owned by this cache). Off by default.
AllowFlush bool
}
// NewMemcacheProvider creates a new Memcache cache provider.
func NewMemcacheProvider(config *MemcacheConfig) (*MemcacheProvider, error) {
if config == nil {
config = &MemcacheConfig{
Servers: []string{"localhost:11211"},
}
// Work on a copy so the caller's struct is not mutated
var cfg MemcacheConfig
if config != nil {
cfg = *config
cfg.Servers = append([]string(nil), config.Servers...)
}
if cfg.Options != nil {
o := *cfg.Options
cfg.Options = &o
}
if len(config.Servers) == 0 {
config.Servers = []string{"localhost:11211"}
if len(cfg.Servers) == 0 {
cfg.Servers = []string{"localhost:11211"}
}
if config.MaxIdleConns == 0 {
config.MaxIdleConns = 2
if cfg.MaxIdleConns == 0 {
cfg.MaxIdleConns = 2
}
if config.Timeout == 0 {
config.Timeout = 1 * time.Second
if cfg.Timeout == 0 {
cfg.Timeout = 1 * time.Second
}
if config.Options == nil {
config.Options = &Options{
if cfg.Options == nil {
cfg.Options = &Options{
DefaultTTL: 5 * time.Minute,
}
}
client := memcache.New(config.Servers...)
client.MaxIdleConns = config.MaxIdleConns
client.Timeout = config.Timeout
client := memcache.New(cfg.Servers...)
client.MaxIdleConns = cfg.MaxIdleConns
client.Timeout = cfg.Timeout
// Test connection
if err := client.Ping(); err != nil {
@@ -66,18 +92,72 @@ func NewMemcacheProvider(config *MemcacheConfig) (*MemcacheProvider, error) {
}
return &MemcacheProvider{
client: client,
options: config.Options,
client: client,
options: cfg.Options,
allowFlush: cfg.AllowFlush,
}, nil
}
// memcacheKey maps a caller key to a legal memcached key. User keys live under the "k:"
// prefix, so they can never collide with the "cache:tag:" / "cache:tags:" index keys.
// Keys that are too long or contain illegal bytes (whitespace/control characters)
// are replaced with their SHA-256.
func memcacheKey(key string) string {
k := "k:" + key
if len(k) > 200 || !legalMemcacheKey(k) {
sum := sha256.Sum256([]byte(key))
return "k:h:" + hex.EncodeToString(sum[:])
}
return k
}
func legalMemcacheKey(key string) bool {
for i := 0; i < len(key); i++ {
if key[i] <= ' ' || key[i] == 0x7f {
return false
}
}
return true
}
// memcacheTagKey maps a tag to its index key, hashing if it is not a legal key.
func memcacheTagKey(prefix, name string) string {
k := prefix + name
if len(k) > 200 || !legalMemcacheKey(k) {
sum := sha256.Sum256([]byte(name))
return prefix + "h:" + hex.EncodeToString(sum[:])
}
return k
}
// memcacheExpiry converts a TTL into a memcached expiry value, switching to an
// absolute Unix timestamp above 30 days as the protocol requires.
func memcacheExpiry(ttl time.Duration) int32 {
if ttl <= 0 {
return 0 // never expires
}
secs := int64(ttl.Seconds())
if secs > memcacheMaxRelativeTTL {
return int32(time.Now().Add(ttl).Unix()) //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
}
if secs == 0 {
secs = 1
}
return int32(secs)
}
// Get retrieves a value from the cache by key.
func (m *MemcacheProvider) Get(ctx context.Context, key string) ([]byte, bool) {
item, err := m.client.Get(key)
if err == memcache.ErrCacheMiss {
if ctx.Err() != nil {
return nil, false
}
item, err := m.client.Get(memcacheKey(key))
if errors.Is(err, memcache.ErrCacheMiss) {
return nil, false
}
if err != nil {
// Reported as a miss (the Provider interface cannot express errors), but not silently
logger.Warn("cache: memcache GET failed: %v", err)
return nil, false
}
return item.Value, true
@@ -85,130 +165,158 @@ func (m *MemcacheProvider) Get(ctx context.Context, key string) ([]byte, bool) {
// Set stores a value in the cache with the specified TTL.
func (m *MemcacheProvider) Set(ctx context.Context, key string, value []byte, ttl time.Duration) error {
if err := ctx.Err(); err != nil {
return err
}
if ttl == 0 {
ttl = m.options.DefaultTTL
}
item := &memcache.Item{
Key: key,
return m.client.Set(&memcache.Item{
Key: memcacheKey(key),
Value: value,
Expiration: int32(ttl.Seconds()),
}
return m.client.Set(item)
Expiration: memcacheExpiry(ttl),
})
}
// SetWithTags stores a value in the cache with the specified TTL and tags.
// Note: Tag support in Memcache is limited and less efficient than Redis.
// Note: Tag support in Memcache is limited and less efficient than Redis. The
// tag index is updated with compare-and-swap; if it cannot be updated the value is
// removed again and an error is returned, so an untracked entry is never left behind.
func (m *MemcacheProvider) SetWithTags(ctx context.Context, key string, value []byte, ttl time.Duration, tags []string) error {
if err := ctx.Err(); err != nil {
return err
}
if ttl == 0 {
ttl = m.options.DefaultTTL
}
expiration := int32(ttl.Seconds())
expiration := memcacheExpiry(ttl)
mkey := memcacheKey(key)
// Set the main value
item := &memcache.Item{
Key: key,
Value: value,
Expiration: expiration,
if err := m.client.Set(&memcache.Item{Key: mkey, Value: value, Expiration: expiration}); err != nil {
return err
}
if err := m.client.Set(item); err != nil {
if len(tags) == 0 {
return nil
}
fail := func(err error) error {
_ = m.client.Delete(mkey) // best-effort rollback; the original error is what matters
return err
}
// Store tags for this key
if len(tags) > 0 {
tagsData, err := json.Marshal(tags)
if err != nil {
return fmt.Errorf("failed to marshal tags: %w", err)
}
tagsData, err := json.Marshal(tags)
if err != nil {
return fail(fmt.Errorf("failed to marshal tags: %w", err))
}
if err := m.client.Set(&memcache.Item{
Key: memcacheTagKey("cache:tags:", key),
Value: tagsData,
Expiration: expiration,
}); err != nil {
return fail(err)
}
tagsItem := &memcache.Item{
Key: fmt.Sprintf("cache:tags:%s", key),
Value: tagsData,
Expiration: expiration,
}
if err := m.client.Set(tagsItem); err != nil {
return err
}
// Add key to each tag's key list
for _, tag := range tags {
tagKey := fmt.Sprintf("cache:tag:%s", tag)
// Get existing keys for this tag
var keys []string
if item, err := m.client.Get(tagKey); err == nil {
_ = json.Unmarshal(item.Value, &keys)
}
// Add current key if not already present
found := false
// Tag lists live longer than the entries they index
tagExpiry := memcacheExpiry(ttl + time.Hour)
for _, tag := range tags {
if err := m.updateTagKeys(memcacheTagKey("cache:tag:", tag), tagExpiry, func(keys []string) ([]string, error) {
for _, k := range keys {
if k == key {
found = true
break
return keys, nil
}
}
if !found {
keys = append(keys, key)
if len(keys) >= memcacheMaxTagKeys {
return nil, fmt.Errorf("tag index for %q is full (%d keys)", tag, memcacheMaxTagKeys)
}
// Store updated key list
keysData, err := json.Marshal(keys)
if err != nil {
continue
}
tagItem := &memcache.Item{
Key: tagKey,
Value: keysData,
Expiration: expiration + 3600, // Give tag lists longer TTL
}
_ = m.client.Set(tagItem)
return append(keys, key), nil
}); err != nil {
return fail(err)
}
}
return nil
}
// updateTagKeys applies fn to a tag's key list using compare-and-swap.
func (m *MemcacheProvider) updateTagKeys(tagKey string, expiry int32, fn func([]string) ([]string, error)) error {
for attempt := 0; attempt < memcacheCASRetries; attempt++ {
item, err := m.client.Get(tagKey)
var keys []string
switch {
case errors.Is(err, memcache.ErrCacheMiss):
item = nil
case err != nil:
return err
default:
if err := json.Unmarshal(item.Value, &keys); err != nil {
return fmt.Errorf("failed to unmarshal tag keys: %w", err)
}
}
keys, err = fn(keys)
if err != nil {
return err
}
data, err := json.Marshal(keys)
if err != nil {
return err
}
if item == nil {
err = m.client.Add(&memcache.Item{Key: tagKey, Value: data, Expiration: expiry})
if errors.Is(err, memcache.ErrNotStored) {
continue // someone created it first; retry
}
return err
}
item.Value = data
item.Expiration = expiry
err = m.client.CompareAndSwap(item)
if errors.Is(err, memcache.ErrCASConflict) || errors.Is(err, memcache.ErrNotStored) || errors.Is(err, memcache.ErrCacheMiss) {
continue
}
return err
}
return fmt.Errorf("tag index %q: too much contention, giving up", tagKey)
}
// Delete removes a key from the cache.
func (m *MemcacheProvider) Delete(ctx context.Context, key string) error {
if err := ctx.Err(); err != nil {
return err
}
mkey := memcacheKey(key)
// Get tags for this key
tagsKey := fmt.Sprintf("cache:tags:%s", key)
tagsKey := memcacheTagKey("cache:tags:", key)
if item, err := m.client.Get(tagsKey); err == nil {
var tags []string
if err := json.Unmarshal(item.Value, &tags); err == nil {
// Remove key from each tag's key list
for _, tag := range tags {
tagKey := fmt.Sprintf("cache:tag:%s", tag)
if tagItem, err := m.client.Get(tagKey); err == nil {
var keys []string
if err := json.Unmarshal(tagItem.Value, &keys); err == nil {
// Remove current key from the list
newKeys := make([]string, 0, len(keys))
for _, k := range keys {
if k != key {
newKeys = append(newKeys, k)
}
}
// Update the tag's key list
if keysData, err := json.Marshal(newKeys); err == nil {
tagItem.Value = keysData
_ = m.client.Set(tagItem)
err := m.updateTagKeys(memcacheTagKey("cache:tag:", tag), memcacheExpiry(m.options.DefaultTTL+time.Hour), func(keys []string) ([]string, error) {
out := make([]string, 0, len(keys))
for _, k := range keys {
if k != key {
out = append(out, k)
}
}
return out, nil
})
if err != nil {
logger.Warn("cache: failed to update memcache tag index on delete: %v", err)
}
}
}
// Delete the tags key
_ = m.client.Delete(tagsKey)
if err := m.client.Delete(tagsKey); err != nil && !errors.Is(err, memcache.ErrCacheMiss) {
logger.Warn("cache: failed to delete memcache tags key: %v", err)
}
}
// Delete the actual key
err := m.client.Delete(key)
if err == memcache.ErrCacheMiss {
err := m.client.Delete(mkey)
if errors.Is(err, memcache.ErrCacheMiss) {
return nil
}
return err
@@ -216,11 +324,13 @@ func (m *MemcacheProvider) Delete(ctx context.Context, key string) error {
// DeleteByTag removes all keys associated with the given tag.
func (m *MemcacheProvider) DeleteByTag(ctx context.Context, tag string) error {
tagKey := fmt.Sprintf("cache:tag:%s", tag)
if err := ctx.Err(); err != nil {
return err
}
tagKey := memcacheTagKey("cache:tag:", tag)
// Get all keys associated with this tag
item, err := m.client.Get(tagKey)
if err == memcache.ErrCacheMiss {
if errors.Is(err, memcache.ErrCacheMiss) {
return nil
}
if err != nil {
@@ -232,42 +342,51 @@ func (m *MemcacheProvider) DeleteByTag(ctx context.Context, tag string) error {
return fmt.Errorf("failed to unmarshal tag keys: %w", err)
}
// Delete all keys
var firstErr error
note := func(err error) {
if err != nil && !errors.Is(err, memcache.ErrCacheMiss) && firstErr == nil {
firstErr = err
}
}
for _, key := range keys {
_ = m.client.Delete(key)
// Also delete the tags key for this cache key
tagsKey := fmt.Sprintf("cache:tags:%s", key)
_ = m.client.Delete(tagsKey)
note(m.client.Delete(memcacheKey(key)))
note(m.client.Delete(memcacheTagKey("cache:tags:", key)))
}
// Delete the tag key itself
_ = m.client.Delete(tagKey)
return nil
if firstErr != nil {
return firstErr // keep the tag index so the invalidation can be retried
}
note(m.client.Delete(tagKey))
return firstErr
}
// DeleteByPattern removes all keys matching the pattern.
// Note: Memcache does not support pattern-based deletion natively.
// This is a no-op for memcache and returns an error.
// DeleteByPattern is not supported by Memcache; it always returns an error.
// Use tags (SetWithTags / DeleteByTag) for group invalidation instead.
func (m *MemcacheProvider) DeleteByPattern(ctx context.Context, pattern string) error {
return fmt.Errorf("pattern-based deletion is not supported by Memcache")
}
// Clear removes all items from the cache.
// It runs flush_all on every configured server and therefore requires MemcacheConfig.AllowFlush.
func (m *MemcacheProvider) Clear(ctx context.Context) error {
if !m.allowFlush {
return ErrFlushNotAllowed
}
return m.client.FlushAll()
}
// Exists checks if a key exists in the cache.
func (m *MemcacheProvider) Exists(ctx context.Context, key string) bool {
_, err := m.client.Get(key)
if ctx.Err() != nil {
return false
}
_, err := m.client.Get(memcacheKey(key))
return err == nil
}
// Close closes the provider and releases any resources.
// Close closes the provider and releases idle connections.
func (m *MemcacheProvider) Close() error {
// Memcache client doesn't have a close method
return nil
return m.client.Close()
}
// Stats returns statistics about the cache provider.
+114 -105
View File
@@ -2,20 +2,39 @@ package cache
import (
"context"
"errors"
"fmt"
"regexp"
"sync"
"sync/atomic"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
const (
defaultMemoryMaxSize = 10000
defaultMemoryCleanupInterval = time.Minute
)
// ErrClosed is returned by a provider that has been closed.
var ErrClosed = errors.New("cache: provider closed")
// memoryItem represents a cached item in memory.
type memoryItem struct {
Value []byte
Expiration time.Time
LastAccess time.Time
HitCount int64
Tags []string
lastAccess atomic.Int64 // unix nanos
hitCount atomic.Int64
}
func newMemoryItem(value []byte, expiration time.Time, tags []string) *memoryItem {
buf := make([]byte, len(value))
copy(buf, value)
item := &memoryItem{Value: buf, Expiration: expiration, Tags: tags}
item.lastAccess.Store(time.Now().UnixNano())
return item
}
// isExpired checks if the item has expired.
@@ -34,30 +53,72 @@ type MemoryProvider struct {
options *Options
hits atomic.Int64
misses atomic.Int64
closed bool
done chan struct{}
closeOnce sync.Once
}
// NewMemoryProvider creates a new in-memory cache provider.
// A MaxSize <= 0 selects the default (10000); use MaxSize -1 for an unbounded cache.
// A background goroutine removes expired items until Close is called.
func NewMemoryProvider(opts *Options) *MemoryProvider {
if opts == nil {
opts = &Options{
DefaultTTL: 5 * time.Minute,
MaxSize: 10000,
}
var o Options
if opts != nil {
o = *opts // do not mutate the caller's struct
} else {
o = Options{DefaultTTL: 5 * time.Minute}
}
if o.MaxSize == 0 {
o.MaxSize = defaultMemoryMaxSize
}
if o.CleanupInterval <= 0 {
o.CleanupInterval = defaultMemoryCleanupInterval
}
return &MemoryProvider{
m := &MemoryProvider{
items: make(map[string]*memoryItem),
tagToKeys: make(map[string]map[string]struct{}),
options: opts,
options: &o,
done: make(chan struct{}),
}
go m.janitor(o.CleanupInterval)
return m
}
func (m *MemoryProvider) janitor(interval time.Duration) {
defer logger.CatchPanic("cache.MemoryProvider.janitor")()
t := time.NewTicker(interval)
defer t.Stop()
for {
select {
case <-m.done:
return
case <-t.C:
m.CleanExpired(context.Background())
}
}
}
// removeLocked deletes a key and its tag associations. Caller must hold m.mu for writing.
func (m *MemoryProvider) removeLocked(key string) {
if item, ok := m.items[key]; ok {
for _, tag := range item.Tags {
if ks := m.tagToKeys[tag]; ks != nil {
delete(ks, key)
if len(ks) == 0 {
delete(m.tagToKeys, tag)
}
}
}
}
delete(m.items, key)
}
// Get retrieves a value from the cache by key.
func (m *MemoryProvider) Get(ctx context.Context, key string) ([]byte, bool) {
// First try with read lock for fast path
m.mu.RLock()
item, exists := m.items[key]
if !exists {
if !exists || m.closed {
m.mu.RUnlock()
m.misses.Add(1)
return nil, false
@@ -65,56 +126,29 @@ func (m *MemoryProvider) Get(ctx context.Context, key string) ([]byte, bool) {
if item.isExpired() {
m.mu.RUnlock()
// Upgrade to write lock to delete expired item
// Delete only if the entry is still the same expired one
m.mu.Lock()
delete(m.items, key)
if cur, ok := m.items[key]; ok && cur == item {
m.removeLocked(key)
}
m.mu.Unlock()
m.misses.Add(1)
return nil, false
}
// Update stats and access time with write lock
value := item.Value
item.lastAccess.Store(time.Now().UnixNano())
item.hitCount.Add(1)
out := make([]byte, len(item.Value))
copy(out, item.Value)
m.mu.RUnlock()
// Update access tracking with write lock
m.mu.Lock()
item.LastAccess = time.Now()
item.HitCount++
m.mu.Unlock()
m.hits.Add(1)
return value, true
return out, true
}
// Set stores a value in the cache with the specified TTL.
func (m *MemoryProvider) Set(ctx context.Context, key string, value []byte, ttl time.Duration) error {
m.mu.Lock()
defer m.mu.Unlock()
if ttl == 0 {
ttl = m.options.DefaultTTL
}
var expiration time.Time
if ttl > 0 {
expiration = time.Now().Add(ttl)
}
// Check max size and evict if necessary
if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
if _, exists := m.items[key]; !exists {
m.evictOne()
}
}
m.items[key] = &memoryItem{
Value: value,
Expiration: expiration,
LastAccess: time.Now(),
}
return nil
return m.SetWithTags(ctx, key, value, ttl, nil)
}
// SetWithTags stores a value in the cache with the specified TTL and tags.
@@ -122,6 +156,10 @@ func (m *MemoryProvider) SetWithTags(ctx context.Context, key string, value []by
m.mu.Lock()
defer m.mu.Unlock()
if m.closed {
return ErrClosed
}
if ttl == 0 {
ttl = m.options.DefaultTTL
}
@@ -131,34 +169,14 @@ func (m *MemoryProvider) SetWithTags(ctx context.Context, key string, value []by
expiration = time.Now().Add(ttl)
}
// Check max size and evict if necessary
if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
if _, exists := m.items[key]; !exists {
m.evictOne()
}
if _, exists := m.items[key]; exists {
m.removeLocked(key) // drops old tag associations
} else if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
m.evictOne()
}
// Remove old tag associations if key exists
if oldItem, exists := m.items[key]; exists {
for _, tag := range oldItem.Tags {
if keySet, ok := m.tagToKeys[tag]; ok {
delete(keySet, key)
if len(keySet) == 0 {
delete(m.tagToKeys, tag)
}
}
}
}
m.items[key] = newMemoryItem(value, expiration, tags)
// Store the item
m.items[key] = &memoryItem{
Value: value,
Expiration: expiration,
LastAccess: time.Now(),
Tags: tags,
}
// Add new tag associations
for _, tag := range tags {
if m.tagToKeys[tag] == nil {
m.tagToKeys[tag] = make(map[string]struct{})
@@ -174,19 +192,7 @@ func (m *MemoryProvider) Delete(ctx context.Context, key string) error {
m.mu.Lock()
defer m.mu.Unlock()
// Remove tag associations
if item, exists := m.items[key]; exists {
for _, tag := range item.Tags {
if keySet, ok := m.tagToKeys[tag]; ok {
delete(keySet, key)
if len(keySet) == 0 {
delete(m.tagToKeys, tag)
}
}
}
}
delete(m.items, key)
m.removeLocked(key)
return nil
}
@@ -195,16 +201,13 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
m.mu.Lock()
defer m.mu.Unlock()
// Get all keys associated with this tag
keySet, exists := m.tagToKeys[tag]
if !exists {
return nil // No keys with this tag
}
// Delete all items with this tag
for key := range keySet {
if item, ok := m.items[key]; ok {
// Remove this tag from the item's tag list
newTags := make([]string, 0, len(item.Tags))
for _, t := range item.Tags {
if t != tag {
@@ -212,8 +215,7 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
}
}
// If item has no more tags, delete it
// Otherwise update its tags
// If item has no more tags, delete it; otherwise update its tags
if len(newTags) == 0 {
delete(m.items, key)
} else {
@@ -222,24 +224,24 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
}
}
// Remove the tag mapping
delete(m.tagToKeys, tag)
return nil
}
// DeleteByPattern removes all keys matching the pattern.
// The pattern is a Go regular expression (unanchored); it is compiled before the lock is taken.
func (m *MemoryProvider) DeleteByPattern(ctx context.Context, pattern string) error {
m.mu.Lock()
defer m.mu.Unlock()
re, err := regexp.Compile(pattern)
if err != nil {
return fmt.Errorf("invalid pattern: %w", err)
}
m.mu.Lock()
defer m.mu.Unlock()
for key := range m.items {
if re.MatchString(key) {
delete(m.items, key)
m.removeLocked(key)
}
}
@@ -252,6 +254,7 @@ func (m *MemoryProvider) Clear(ctx context.Context) error {
defer m.mu.Unlock()
m.items = make(map[string]*memoryItem)
m.tagToKeys = make(map[string]map[string]struct{})
m.hits.Store(0)
m.misses.Store(0)
return nil
@@ -270,12 +273,17 @@ func (m *MemoryProvider) Exists(ctx context.Context, key string) bool {
return !item.isExpired()
}
// Close closes the provider and releases any resources.
// Close closes the provider, stops the janitor and releases stored items.
// Later writes return ErrClosed and reads report a miss.
func (m *MemoryProvider) Close() error {
m.closeOnce.Do(func() { close(m.done) })
m.mu.Lock()
defer m.mu.Unlock()
m.items = nil
m.closed = true
m.items = make(map[string]*memoryItem)
m.tagToKeys = make(map[string]map[string]struct{})
return nil
}
@@ -284,7 +292,7 @@ func (m *MemoryProvider) Stats(ctx context.Context) (*CacheStats, error) {
m.mu.RLock()
defer m.mu.RUnlock()
// Clean expired items first
// Count non-expired items (read-only)
validKeys := 0
for _, item := range m.items {
if !item.isExpired() {
@@ -304,24 +312,25 @@ func (m *MemoryProvider) Stats(ctx context.Context) (*CacheStats, error) {
}
// evictOne removes one item from the cache using LRU strategy.
// Note: this is an O(n) scan. Caller must hold m.mu for writing.
func (m *MemoryProvider) evictOne() {
var oldestKey string
var oldestTime time.Time
var oldest int64
for key, item := range m.items {
if item.isExpired() {
delete(m.items, key)
m.removeLocked(key)
return
}
if oldestKey == "" || item.LastAccess.Before(oldestTime) {
if la := item.lastAccess.Load(); oldestKey == "" || la < oldest {
oldestKey = key
oldestTime = item.LastAccess
oldest = la
}
}
if oldestKey != "" {
delete(m.items, oldestKey)
m.removeLocked(oldestKey)
}
}
@@ -333,7 +342,7 @@ func (m *MemoryProvider) CleanExpired(ctx context.Context) int {
count := 0
for key, item := range m.items {
if item.isExpired() {
delete(m.items, key)
m.removeLocked(key)
count++
}
}
+52 -15
View File
@@ -3,15 +3,20 @@ package cache
import (
"context"
"fmt"
"strconv"
"strings"
"time"
"github.com/redis/go-redis/v9"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// RedisProvider is a Redis implementation of the Provider interface.
type RedisProvider struct {
client *redis.Client
options *Options
client *redis.Client
options *Options
allowFlush bool
}
// RedisConfig contains Redis-specific configuration.
@@ -33,16 +38,25 @@ type RedisConfig struct {
// Options contains general cache options
Options *Options
// AllowFlush permits Clear() to run FLUSHDB, which wipes the entire logical Redis DB
// (including data that is not owned by this cache). Off by default.
AllowFlush bool
}
// NewRedisProvider creates a new Redis cache provider.
func NewRedisProvider(config *RedisConfig) (*RedisProvider, error) {
if config == nil {
config = &RedisConfig{
Host: "localhost",
Port: 6379,
DB: 0,
}
// Work on a copy so the caller's struct is not mutated
var cfg RedisConfig
if config != nil {
cfg = *config
} else {
cfg = RedisConfig{Host: "localhost", Port: 6379, DB: 0}
}
config = &cfg
if config.Options != nil {
o := *config.Options
config.Options = &o
}
if config.Host == "" {
@@ -77,8 +91,9 @@ func NewRedisProvider(config *RedisConfig) (*RedisProvider, error) {
}
return &RedisProvider{
client: client,
options: config.Options,
client: client,
options: config.Options,
allowFlush: config.AllowFlush,
}, nil
}
@@ -89,6 +104,8 @@ func (r *RedisProvider) Get(ctx context.Context, key string) ([]byte, bool) {
return nil, false
}
if err != nil {
// Reported as a miss (the Provider interface cannot express errors), but not silently
logger.Warn("cache: redis GET failed: %v", err)
return nil, false
}
return val, true
@@ -194,7 +211,7 @@ func (r *RedisProvider) DeleteByTag(ctx context.Context, tag string) error {
// DeleteByPattern removes all keys matching the pattern.
func (r *RedisProvider) DeleteByPattern(ctx context.Context, pattern string) error {
iter := r.client.Scan(ctx, 0, pattern, 0).Iterator()
iter := r.client.Scan(ctx, 0, pattern, 500).Iterator()
pipe := r.client.Pipeline()
count := 0
@@ -225,7 +242,11 @@ func (r *RedisProvider) DeleteByPattern(ctx context.Context, pattern string) err
}
// Clear removes all items from the cache.
// It runs FLUSHDB and therefore requires RedisConfig.AllowFlush.
func (r *RedisProvider) Clear(ctx context.Context) error {
if !r.allowFlush {
return ErrFlushNotAllowed
}
return r.client.FlushDB(ctx).Err()
}
@@ -244,8 +265,9 @@ func (r *RedisProvider) Close() error {
}
// Stats returns statistics about the cache provider.
// Only an allowlist of numeric counters from INFO is exposed, not the raw output.
func (r *RedisProvider) Stats(ctx context.Context) (*CacheStats, error) {
info, err := r.client.Info(ctx, "stats", "keyspace").Result()
info, err := r.client.Info(ctx, "stats").Result()
if err != nil {
return nil, fmt.Errorf("failed to get Redis stats: %w", err)
}
@@ -255,13 +277,28 @@ func (r *RedisProvider) Stats(ctx context.Context) (*CacheStats, error) {
return nil, fmt.Errorf("failed to get DB size: %w", err)
}
// Parse stats from INFO command
// This is a simplified version - you may want to parse more detailed stats
counters := map[string]int64{}
for _, line := range strings.Split(info, "\n") {
k, v, ok := strings.Cut(strings.TrimSpace(line), ":")
if !ok {
continue
}
switch k {
case "keyspace_hits", "keyspace_misses", "evicted_keys", "expired_keys":
if n, err := strconv.ParseInt(v, 10, 64); err == nil {
counters[k] = n
}
}
}
stats := &CacheStats{
Hits: counters["keyspace_hits"],
Misses: counters["keyspace_misses"],
Keys: dbSize,
ProviderType: "redis",
ProviderStats: map[string]any{
"info": info,
"evicted_keys": counters["evicted_keys"],
"expired_keys": counters["expired_keys"],
},
}
+4 -4
View File
@@ -686,7 +686,7 @@ func (p *PgSQLInsertQuery) Exec(ctx context.Context) (res common.Result, err err
i++
}
query := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)",
query := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)", //nolint:gosec // G201: table identifier is internal/validated; values use placeholders
p.tableName,
strings.Join(columns, ", "),
strings.Join(placeholders, ", "))
@@ -736,7 +736,7 @@ func (p *PgSQLInsertQuery) Scan(ctx context.Context, dest interface{}) (err erro
i++
}
query := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)",
query := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)", //nolint:gosec // G201: table identifier is internal/validated; values use placeholders
p.tableName,
strings.Join(columns, ", "),
strings.Join(placeholders, ", "))
@@ -886,7 +886,7 @@ func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err err
i++
}
query := fmt.Sprintf("UPDATE %s SET %s",
query := fmt.Sprintf("UPDATE %s SET %s", //nolint:gosec // G201: table identifier is internal/validated; values use placeholders
p.tableName,
strings.Join(setClauses, ", "))
@@ -997,7 +997,7 @@ func (p *PgSQLDeleteQuery) Exec(ctx context.Context) (res common.Result, err err
recordQueryMetrics(p.metricsEnabled, "DELETE", p.schema, p.entity, p.tableName, startedAt, err)
}()
query := fmt.Sprintf("DELETE FROM %s", p.tableName)
query := fmt.Sprintf("DELETE FROM %s", p.tableName) //nolint:gosec // G201: table identifier is internal/validated; values use placeholders
if len(p.whereClauses) > 0 {
query += " WHERE " + strings.Join(p.whereClauses, " AND ")
@@ -13,7 +13,7 @@ import (
// Example demonstrates how to use the PgSQL adapter
func ExamplePgSQLAdapter() error {
// Connect to PostgreSQL database
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
db, err := sql.Open("pgx", dsn)
if err != nil {
return fmt.Errorf("failed to open database: %w", err)
@@ -155,7 +155,7 @@ func (u User) TableName() string {
// ExampleWithModel demonstrates using models with the PgSQL adapter
func ExampleWithModel() error {
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
db, err := sql.Open("pgx", dsn)
if err != nil {
return err
@@ -51,7 +51,7 @@ func (c Comment) TableName() string {
// ExamplePreload demonstrates the Preload functionality
func ExamplePreload() error {
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
db, err := sql.Open("pgx", dsn)
if err != nil {
return err
@@ -79,7 +79,7 @@ func ExamplePreload() error {
// ExamplePreloadRelation demonstrates smart PreloadRelation with auto-detection
func ExamplePreloadRelation() error {
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
db, err := sql.Open("pgx", dsn)
if err != nil {
return err
@@ -148,7 +148,7 @@ func ExamplePreloadRelation() error {
// ExampleJoinRelation demonstrates explicit JOIN loading
func ExampleJoinRelation() error {
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
db, err := sql.Open("pgx", dsn)
if err != nil {
return err
@@ -185,7 +185,7 @@ func ExampleJoinRelation() error {
// ExampleScanModel demonstrates ScanModel with struct destinations
func ExampleScanModel() error {
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
db, err := sql.Open("pgx", dsn)
if err != nil {
return err
@@ -221,7 +221,7 @@ func ExampleScanModel() error {
// ExampleCompleteWorkflow demonstrates a complete workflow with preloading
func ExampleCompleteWorkflow() error {
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
db, err := sql.Open("pgx", dsn)
if err != nil {
return err
+10
View File
@@ -4,6 +4,7 @@ import (
"fmt"
"strings"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
)
@@ -72,6 +73,15 @@ func ResolveJSONColumnExpr(model interface{}, tableAlias, token string) (expr st
func ApplySelectColumns(query SelectQuery, model interface{}, tableAlias string, columns []string) SelectQuery {
for _, col := range columns {
if expr, args, alias, ok := ResolveJSONColumnExpr(model, tableAlias, col); ok {
if !reflection.HasColumn(model, alias) {
// No matching scan target on the model (e.g. no
// `bun:"<alias>,scanonly"` field declared for this JSON
// path) - bun would fail to scan the row with "does not
// have column X". Drop the expression rather than erroring;
// the rest of the requested columns still get selected.
logger.Warn("Skipping JSON select column %q: model has no scan target for alias %q", col, alias)
continue
}
query = query.ColumnExpr(expr+" AS "+QuoteIdent(alias), args...)
continue
}
+41
View File
@@ -2,6 +2,7 @@ package common
import (
"reflect"
"strings"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
@@ -162,3 +163,43 @@ func TestBuildJSONFilterCondition_QualifiedAndInjectionSafe(t *testing.T) {
t.Errorf("args = %#v", args)
}
}
// selectCapQuery is a minimal SelectQuery that records Column/ColumnExpr calls
// so ApplySelectColumns' behaviour can be asserted without a real DB.
type selectCapQuery struct {
SelectQuery
columns []string
columnExprs []string
}
func (m *selectCapQuery) Column(cols ...string) SelectQuery {
m.columns = append(m.columns, cols...)
return m
}
func (m *selectCapQuery) ColumnExpr(q string, args ...interface{}) SelectQuery {
m.columnExprs = append(m.columnExprs, q)
return m
}
// jsonSelectModel has a real JSON column (Data) but only ONE pre-declared
// scanonly field for a computed JSON path ("data_city"); "data_age" has no
// matching scan target.
type jsonSelectModel struct {
ID int64 `json:"id" bun:"id,pk"`
Data spectypes.SqlJSONB `json:"data" bun:"data"`
DataCity string `json:"-" bun:"data_city,scanonly"`
}
func TestApplySelectColumns_SkipsJSONColumnWithoutScanTarget(t *testing.T) {
m := jsonSelectModel{}
q := &selectCapQuery{}
ApplySelectColumns(q, m, "", []string{"id", "data.city", "data.age"})
if !reflect.DeepEqual(q.columns, []string{"id"}) {
t.Errorf("columns = %#v, want [id]", q.columns)
}
if len(q.columnExprs) != 1 || !strings.Contains(q.columnExprs[0], `AS "data_city"`) {
t.Errorf("columnExprs = %#v, want exactly one expr aliased data_city", q.columnExprs)
}
}
+30 -1
View File
@@ -1,6 +1,9 @@
package config
import "time"
import (
"fmt"
"time"
)
// Config represents the complete application configuration
type Config struct {
@@ -88,6 +91,12 @@ type TracingConfig struct {
ServiceName string `mapstructure:"service_name"`
ServiceVersion string `mapstructure:"service_version"`
Endpoint string `mapstructure:"endpoint"`
// Insecure exports traces over plaintext gRPC (default false: TLS).
Insecure bool `mapstructure:"insecure"`
// SampleRate is the fraction of root traces sampled; 0 selects the default (0.1).
SampleRate float64 `mapstructure:"sample_rate"`
// Headers are sent with every OTLP export request (e.g. auth tokens).
Headers map[string]string `mapstructure:"headers"`
}
// CacheConfig holds cache provider configuration
@@ -198,3 +207,23 @@ type EventBrokerRetryPolicyConfig struct {
// This is a map of path name to file system path
// Example: "data_dir": "/var/lib/myapp/data"
type PathsConfig map[string]string
// Validate checks every configuration section for invalid or unsafe values.
func (c *Config) Validate() error {
if err := c.Servers.Validate(); err != nil {
return fmt.Errorf("servers: %w", err)
}
if c.Middleware.RateLimitRPS < 0 || c.Middleware.RateLimitBurst < 0 {
return fmt.Errorf("middleware: rate_limit_rps and rate_limit_burst must not be negative")
}
if c.Middleware.MaxRequestSize <= 0 {
return fmt.Errorf("middleware: max_request_size must be greater than 0")
}
if c.EventBroker.Enabled && c.EventBroker.WorkerCount <= 0 {
return fmt.Errorf("event_broker: worker_count must be greater than 0")
}
if c.DBManager.MaxOpenConns < 0 || c.DBManager.MaxIdleConns < 0 || c.DBManager.RetryAttempts < 0 {
return fmt.Errorf("dbmanager: max_open_conns, max_idle_conns and retry_attempts must not be negative")
}
return nil
}
+79
View File
@@ -0,0 +1,79 @@
package config
import (
"os"
"path/filepath"
"sync"
"testing"
)
func TestManagerConcurrentSetGet(t *testing.T) {
m := NewManager()
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(2)
go func() {
defer wg.Done()
for j := 0; j < 200; j++ {
m.Set("x.y", j)
}
}()
go func() {
defer wg.Done()
for j := 0; j < 200; j++ {
_ = m.Get("x.y")
_, _ = m.GetConfig()
}
}()
}
wg.Wait()
}
func TestNewManagerDoesNotReplaceGlobal(t *testing.T) {
g := GetConfigManager()
_ = NewManager()
if GetConfigManager() != g {
t.Fatal("NewManager replaced the global manager")
}
}
func TestSaveConfigPermissions(t *testing.T) {
path := filepath.Join(t.TempDir(), "out.yaml")
if err := os.WriteFile(path, nil, 0o644); err != nil {
t.Fatal(err)
}
if err := NewManager().SaveConfig(path); err != nil {
t.Fatal(err)
}
if fi, _ := os.Stat(path); fi.Mode().Perm() != 0o600 {
t.Fatalf("mode = %v, want 0600", fi.Mode().Perm())
}
}
func TestPathsSetNilAndJoinConfined(t *testing.T) {
var pc PathsConfig
pc.Set("data", "data")
if got, _ := pc.Get("data"); got != "data" {
t.Fatalf("got %q", got)
}
if _, err := pc.Join("data", "../../etc/passwd"); err == nil {
t.Fatal("expected traversal error")
}
if p, err := pc.Join("data", "a", "b"); err != nil || p != filepath.Join("data", "a", "b") {
t.Fatalf("got %q, %v", p, err)
}
}
func TestConfigValidate(t *testing.T) {
cfg, err := NewManager().GetConfig()
if err != nil {
t.Fatal(err)
}
if err := cfg.Validate(); err != nil {
t.Fatalf("defaults should validate: %v", err)
}
cfg.EventBroker.Enabled, cfg.EventBroker.WorkerCount = true, 0
if cfg.Validate() == nil {
t.Fatal("expected worker_count error")
}
}
+81 -19
View File
@@ -2,37 +2,60 @@ package config
import (
"fmt"
"os"
"strings"
"sync"
"github.com/spf13/viper"
)
// Manager handles configuration loading from multiple sources
// Manager handles configuration loading from multiple sources.
// viper.Viper is not safe for concurrent use, so every access to it is guarded by mu.
type Manager struct {
v *viper.Viper
mu sync.RWMutex
v *viper.Viper
}
var configInstance *Manager
var (
configInstance *Manager
configMu sync.Mutex
)
// GetConfigManager returns a singleton configuration manager instance
func GetConfigManager() *Manager {
configMu.Lock()
defer configMu.Unlock()
if configInstance == nil {
configInstance = NewManager()
}
return configInstance
}
// NewManager creates a new configuration manager with defaults
// SetConfigManager publishes m as the global manager returned by GetConfigManager.
// NewManager no longer does this implicitly.
func SetConfigManager(m *Manager) {
configMu.Lock()
defer configMu.Unlock()
configInstance = m
}
// NewManager creates a new, isolated configuration manager with defaults.
// It does not replace the global manager; use SetConfigManager for that.
func NewManager() *Manager {
v := viper.New()
// Set configuration file settings
v.SetConfigName("config")
v.SetConfigType("yaml")
v.AddConfigPath(".")
v.AddConfigPath("./config")
// Most trusted location first; the working directory is the least trustworthy
// and is searched last (viper takes the first match).
v.AddConfigPath("/etc/resolvespec")
v.AddConfigPath("$HOME/.resolvespec")
v.AddConfigPath("./config")
v.AddConfigPath(".")
// Saved configs may contain secrets; never write them world-readable
v.SetConfigPermissions(0o600)
// Enable environment variable support
v.SetEnvPrefix("RESOLVESPEC")
@@ -42,8 +65,7 @@ func NewManager() *Manager {
// Set default values
setDefaults(v)
configInstance = &Manager{v: v}
return configInstance
return &Manager{v: v}
}
// NewManagerWithOptions creates a new configuration manager with custom options
@@ -61,6 +83,8 @@ type Option func(*Manager)
// WithConfigFile sets a specific config file path
func WithConfigFile(path string) Option {
return func(m *Manager) {
m.mu.Lock()
defer m.mu.Unlock()
m.v.SetConfigFile(path)
}
}
@@ -68,6 +92,8 @@ func WithConfigFile(path string) Option {
// WithConfigName sets the config file name (without extension)
func WithConfigName(name string) Option {
return func(m *Manager) {
m.mu.Lock()
defer m.mu.Unlock()
m.v.SetConfigName(name)
}
}
@@ -75,6 +101,8 @@ func WithConfigName(name string) Option {
// WithConfigPath adds a path to search for config files
func WithConfigPath(path string) Option {
return func(m *Manager) {
m.mu.Lock()
defer m.mu.Unlock()
m.v.AddConfigPath(path)
}
}
@@ -82,13 +110,19 @@ func WithConfigPath(path string) Option {
// WithEnvPrefix sets the environment variable prefix
func WithEnvPrefix(prefix string) Option {
return func(m *Manager) {
m.mu.Lock()
defer m.mu.Unlock()
m.v.SetEnvPrefix(prefix)
}
}
// Load attempts to load configuration from file and environment
// Load attempts to load configuration from file and environment.
// A missing config file is not an error (defaults and env vars are used); check
// ConfigFileUsed after Load to see whether a file was actually read.
func (m *Manager) Load() error {
// Try to read config file (not an error if it doesn't exist)
m.mu.Lock()
defer m.mu.Unlock()
if err := m.v.ReadInConfig(); err != nil {
if _, ok := err.(viper.ConfigFileNotFoundError); !ok {
return fmt.Errorf("error reading config file: %w", err)
@@ -99,8 +133,19 @@ func (m *Manager) Load() error {
return nil
}
// ConfigFileUsed returns the config file that was read by Load, or "" if none was
// found (i.e. the manager is running on defaults and environment variables only).
func (m *Manager) ConfigFileUsed() string {
m.mu.RLock()
defer m.mu.RUnlock()
return m.v.ConfigFileUsed()
}
// GetConfig returns the complete configuration
func (m *Manager) GetConfig() (*Config, error) {
m.mu.RLock()
defer m.mu.RUnlock()
var cfg Config
if err := m.v.Unmarshal(&cfg); err != nil {
return nil, fmt.Errorf("failed to unmarshal config: %w", err)
@@ -108,16 +153,11 @@ func (m *Manager) GetConfig() (*Config, error) {
return &cfg, nil
}
// SetConfig sets the complete configuration
// SetConfig sets the complete configuration atomically
func (m *Manager) SetConfig(cfg *Config) error {
configMap := make(map[string]interface{})
m.mu.Lock()
defer m.mu.Unlock()
// Marshal the config to a map structure that viper can use
if err := m.v.Unmarshal(&configMap); err != nil {
return fmt.Errorf("failed to prepare config map: %w", err)
}
// Use viper's merge to apply the config
m.v.Set("servers", cfg.Servers)
m.v.Set("tracing", cfg.Tracing)
m.v.Set("cache", cfg.Cache)
@@ -135,34 +175,54 @@ func (m *Manager) SetConfig(cfg *Config) error {
// Get returns a configuration value by key
func (m *Manager) Get(key string) interface{} {
m.mu.RLock()
defer m.mu.RUnlock()
return m.v.Get(key)
}
// GetString returns a string configuration value
func (m *Manager) GetString(key string) string {
m.mu.RLock()
defer m.mu.RUnlock()
return m.v.GetString(key)
}
// GetInt returns an int configuration value
func (m *Manager) GetInt(key string) int {
m.mu.RLock()
defer m.mu.RUnlock()
return m.v.GetInt(key)
}
// GetBool returns a bool configuration value
func (m *Manager) GetBool(key string) bool {
m.mu.RLock()
defer m.mu.RUnlock()
return m.v.GetBool(key)
}
// Set sets a configuration value
func (m *Manager) Set(key string, value interface{}) {
m.mu.Lock()
defer m.mu.Unlock()
m.v.Set(key, value)
}
// SaveConfig writes the current configuration to the specified path
// SaveConfig writes the current configuration to the specified path.
// The file contains the entire merged configuration, including secrets
// (database/redis passwords, error-tracking DSN), so it is written with mode 0600.
// Prefer supplying secrets via RESOLVESPEC_* environment variables.
func (m *Manager) SaveConfig(path string) error {
m.mu.RLock()
defer m.mu.RUnlock()
if err := m.v.WriteConfigAs(path); err != nil {
return fmt.Errorf("failed to save config to %s: %w", path, err)
}
// viper only applies its permissions when creating the file; tighten a pre-existing one too
if err := os.Chmod(path, 0o600); err != nil {
return fmt.Errorf("failed to restrict permissions on %s: %w", path, err)
}
return nil
}
@@ -190,6 +250,8 @@ func setDefaults(v *viper.Viper) {
v.SetDefault("tracing.service_name", "resolvespec")
v.SetDefault("tracing.service_version", "1.0.0")
v.SetDefault("tracing.endpoint", "")
v.SetDefault("tracing.insecure", false)
v.SetDefault("tracing.sample_rate", 0.1)
// Cache defaults
v.SetDefault("cache.provider", "memory")
+18 -5
View File
@@ -4,6 +4,7 @@ import (
"fmt"
"os"
"path/filepath"
"strings"
)
// Get retrieves a path by name
@@ -34,9 +35,13 @@ func (pc PathsConfig) GetOrDefault(name, defaultPath string) string {
return path
}
// Set sets a path by name
func (pc PathsConfig) Set(name, path string) {
pc[name] = path
// Set sets a path by name. It takes a pointer so a nil map can be allocated.
// PathsConfig is not safe for concurrent mutation; populate it before sharing.
func (pc *PathsConfig) Set(name, path string) {
if *pc == nil {
*pc = make(PathsConfig)
}
(*pc)[name] = path
}
// Has checks if a path exists by name
@@ -92,7 +97,8 @@ func (pc PathsConfig) AbsPath(name string) (string, error) {
return absPath, nil
}
// Join joins path segments with a named base path
// Join joins path segments with a named base path.
// It returns an error if the result would escape the base path.
func (pc PathsConfig) Join(name string, elem ...string) (string, error) {
base, err := pc.Get(name)
if err != nil {
@@ -100,7 +106,14 @@ func (pc PathsConfig) Join(name string, elem ...string) (string, error) {
}
parts := append([]string{base}, elem...)
return filepath.Join(parts...), nil
joined := filepath.Join(parts...)
// filepath.Join resolves ".." rather than rejecting it; ensure the result stays under base
cleanBase := filepath.Clean(base)
if joined != cleanBase && !strings.HasPrefix(joined, cleanBase+string(os.PathSeparator)) {
return "", fmt.Errorf("path %q escapes base path '%s'", joined, name)
}
return joined, nil
}
// List returns all configured path names
+30 -31
View File
@@ -1,10 +1,12 @@
package config
import (
"context"
"fmt"
"net"
"os"
"strings"
"time"
)
// ApplyGlobalDefaults applies global server defaults to this instance
@@ -95,7 +97,8 @@ func (sc *ServersConfig) Validate() error {
return nil
}
// GetDefault returns the default server instance configuration
// GetDefault returns the default server instance configuration.
// The returned pointer refers to a copy: mutating it does not modify sc.Instances.
func (sc *ServersConfig) GetDefault() (*ServerInstanceConfig, error) {
if sc.DefaultServer == "" {
return nil, fmt.Errorf("no default server configured")
@@ -109,41 +112,37 @@ func (sc *ServersConfig) GetDefault() (*ServerInstanceConfig, error) {
return &instance, nil
}
// GetIPs - GetIP for pc
// GetIPs returns the hostname, a comma-separated list of non-loopback IPs and the
// same IPs as []net.IP. The lookup is bounded by a short timeout.
func GetIPs() (hostname string, ipList string, ipNetList []net.IP) {
defer func() {
if err := recover(); err != nil {
fmt.Println("Recovered in GetIPs", err)
}
}()
hostname, _ = os.Hostname()
ipaddrlist := make([]net.IP, 0)
iplist := ""
addrs, err := net.LookupIP(hostname)
if err != nil {
return hostname, iplist, ipaddrlist
}
ipNetList = make([]net.IP, 0)
for _, a := range addrs {
// cfg.LogInfo("\nFound IP Host Address: %s", a)
if strings.Contains(a.String(), "127.0.0.1") {
continue
}
iplist = fmt.Sprintf("%s,%s", iplist, a)
ipaddrlist = append(ipaddrlist, a)
}
if iplist == "" {
iff, _ := net.InterfaceAddrs()
for _, a := range iff {
// cfg.LogInfo("\nFound IP Address: %s", a)
if strings.Contains(a.String(), "127.0.0.1") {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
var ips []string
if addrs, err := net.DefaultResolver.LookupIPAddr(ctx, hostname); err == nil {
for _, a := range addrs {
if a.IP.IsLoopback() {
continue
}
iplist = fmt.Sprintf("%s,%s", iplist, a)
ips = append(ips, a.IP.String())
ipNetList = append(ipNetList, a.IP)
}
}
iplist = strings.TrimLeft(iplist, ",")
return hostname, iplist, ipaddrlist
if len(ips) == 0 {
ifaceAddrs, _ := net.InterfaceAddrs()
for _, a := range ifaceAddrs {
ipn, ok := a.(*net.IPNet)
if !ok || ipn.IP.IsLoopback() {
continue
}
ips = append(ips, ipn.IP.String())
ipNetList = append(ipNetList, ipn.IP)
}
}
return hostname, strings.Join(ips, ","), ipNetList
}
+1 -3
View File
@@ -50,7 +50,6 @@ dbmanager:
# Health checks
health_check_interval: 30s
enable_auto_reconnect: true
connections:
# Primary PostgreSQL connection
@@ -256,7 +255,7 @@ db, _ := mgr.GetDefaultDatabase()
| `retry_delay` | duration | 1s | Initial retry delay |
| `retry_max_delay` | duration | 10s | Maximum retry delay |
| `health_check_interval` | duration | 30s | Interval between health checks |
| `enable_auto_reconnect` | bool | true | Auto-reconnect on health check failure |
| `enable_auto_reconnect` | bool | - | Deprecated and ignored: the manager never closes the pool to recover from errors |
### Connection Configuration
@@ -451,7 +450,6 @@ db.NewSelect().Model(&User{}).Scan(ctx)
3. **Enable Health Checks**: Catch connection issues early
```yaml
health_check_interval: 30s
enable_auto_reconnect: true
```
4. **Use Appropriate ORM**: Choose based on your needs
+106 -71
View File
@@ -2,6 +2,10 @@ package dbmanager
import (
"fmt"
"net"
"net/url"
"strconv"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/config"
@@ -57,9 +61,15 @@ type ManagerConfig struct {
RetryDelay time.Duration `mapstructure:"retry_delay"`
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
// Health checks
// Health checks. A zero HealthCheckInterval selects the default (15s); a
// negative value disables the background health checker.
HealthCheckInterval time.Duration `mapstructure:"health_check_interval"`
EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"`
// Deprecated: ignored. The manager never closes a pool to recover from an
// error because database/sql already replaces broken connections; closing
// it would invalidate every handle handed out. Use Connection.Reconnect for
// an explicit, handle-preserving refresh.
EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"`
}
// ConnectionConfig defines configuration for a single database connection
@@ -103,6 +113,11 @@ type ConnectionConfig struct {
ConnectTimeout time.Duration `mapstructure:"connect_timeout"`
QueryTimeout time.Duration `mapstructure:"query_timeout"`
// Retry policy for the initial connect (inherited from the manager config)
RetryAttempts int `mapstructure:"retry_attempts"`
RetryDelay time.Duration `mapstructure:"retry_delay"`
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
// Features
EnableTracing bool `mapstructure:"enable_tracing"`
EnableMetrics bool `mapstructure:"enable_metrics"`
@@ -129,7 +144,6 @@ func DefaultManagerConfig() ManagerConfig {
RetryDelay: 1 * time.Second,
RetryMaxDelay: 10 * time.Second,
HealthCheckInterval: 15 * time.Second,
EnableAutoReconnect: true,
}
}
@@ -161,11 +175,6 @@ func (c *ManagerConfig) ApplyDefaults() {
if c.HealthCheckInterval == 0 {
c.HealthCheckInterval = defaults.HealthCheckInterval
}
// EnableAutoReconnect defaults to true - apply if not explicitly set
// Since this is a boolean, we apply the default unconditionally when it's false
if !c.EnableAutoReconnect {
c.EnableAutoReconnect = defaults.EnableAutoReconnect
}
}
// Validate validates the manager configuration
@@ -222,9 +231,18 @@ func (cc *ConnectionConfig) ApplyDefaults(global *ManagerConfig) {
}
if cc.QueryTimeout == 0 {
cc.QueryTimeout = 2 * time.Minute // Default to 2 minutes
} else if cc.QueryTimeout < 2*time.Minute {
// Enforce minimum of 2 minutes
cc.QueryTimeout = 2 * time.Minute
}
if global != nil {
if cc.RetryAttempts == 0 {
cc.RetryAttempts = global.RetryAttempts
}
if cc.RetryDelay == 0 {
cc.RetryDelay = global.RetryDelay
}
if cc.RetryMaxDelay == 0 {
cc.RetryMaxDelay = global.RetryMaxDelay
}
}
// Default ORM
@@ -314,108 +332,122 @@ func (cc *ConnectionConfig) BuildDSN() (string, error) {
}
}
// buildPostgresDSN builds a postgres:// URL so credentials and other values are
// escaped rather than spliced into a key=value string. statement_timeout is
// applied by the provider as a runtime parameter.
func (cc *ConnectionConfig) buildPostgresDSN() string {
dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s",
cc.Host, cc.Port, cc.User, cc.Password, cc.Database)
q := url.Values{}
if cc.SSLMode != "" {
dsn += fmt.Sprintf(" sslmode=%s", cc.SSLMode)
q.Set("sslmode", cc.SSLMode)
} else {
dsn += " sslmode=disable"
// prefer: use TLS when the server offers it, without failing on
// servers that do not.
q.Set("sslmode", "prefer")
}
if cc.Schema != "" {
dsn += fmt.Sprintf(" search_path=%s", cc.Schema)
q.Set("search_path", cc.Schema)
}
// Add statement_timeout for query execution timeout (in milliseconds)
if cc.QueryTimeout > 0 {
timeoutMs := int(cc.QueryTimeout.Milliseconds())
dsn += fmt.Sprintf(" statement_timeout=%d", timeoutMs)
u := url.URL{
Scheme: "postgres",
Host: hostPort(cc.Host, cc.Port),
Path: "/" + cc.Database,
RawQuery: q.Encode(),
}
return dsn
if cc.User != "" || cc.Password != "" {
u.User = url.UserPassword(cc.User, cc.Password)
}
return u.String()
}
func hostPort(host string, port int) string {
if port == 0 {
return host
}
// JoinHostPort brackets IPv6 literals.
return net.JoinHostPort(host, strconv.Itoa(port))
}
// buildSQLiteDSN puts per-connection settings in the DSN as _pragma parameters
// so every pooled connection gets them, not just the one that ran an Exec.
func (cc *ConnectionConfig) buildSQLiteDSN() string {
filepath := cc.FilePath
if filepath == "" {
filepath = ":memory:"
}
// Add query parameters for timeouts
// Note: SQLite driver supports _timeout parameter (in milliseconds)
var pragmas []string
if cc.QueryTimeout > 0 {
timeoutMs := int(cc.QueryTimeout.Milliseconds())
filepath += fmt.Sprintf("?_timeout=%d", timeoutMs)
pragmas = append(pragmas, fmt.Sprintf("busy_timeout(%d)", cc.QueryTimeout.Milliseconds()))
}
if filepath != ":memory:" {
pragmas = append(pragmas, "journal_mode(WAL)")
}
if len(pragmas) == 0 {
return filepath
}
return filepath
q := url.Values{}
for _, p := range pragmas {
q.Add("_pragma", p)
}
sep := "?"
if strings.Contains(filepath, "?") {
sep = "&"
}
return filepath + sep + q.Encode()
}
func (cc *ConnectionConfig) buildMSSQLDSN() string {
// Format: sqlserver://username:password@host:port?database=dbname
dsn := fmt.Sprintf("sqlserver://%s:%s@%s:%d?database=%s",
cc.User, cc.Password, cc.Host, cc.Port, cc.Database)
q := url.Values{}
q.Set("database", cc.Database)
if cc.Schema != "" {
dsn += fmt.Sprintf("&schema=%s", cc.Schema)
q.Set("schema", cc.Schema)
}
// Add connection timeout (in seconds)
if cc.ConnectTimeout > 0 {
timeoutSec := int(cc.ConnectTimeout.Seconds())
dsn += fmt.Sprintf("&connection timeout=%d", timeoutSec)
sec := strconv.Itoa(int(cc.ConnectTimeout.Seconds()))
q.Set("connection timeout", sec)
q.Set("dial timeout", sec)
}
// Add dial timeout for TCP connection (in seconds)
if cc.ConnectTimeout > 0 {
dialTimeoutSec := int(cc.ConnectTimeout.Seconds())
dsn += fmt.Sprintf("&dial timeout=%d", dialTimeoutSec)
}
// Add read timeout (in seconds) - enforces timeout for reading data
if cc.QueryTimeout > 0 {
readTimeoutSec := int(cc.QueryTimeout.Seconds())
dsn += fmt.Sprintf("&read timeout=%d", readTimeoutSec)
q.Set("read timeout", strconv.Itoa(int(cc.QueryTimeout.Seconds())))
}
return dsn
u := url.URL{
Scheme: "sqlserver",
Host: hostPort(cc.Host, cc.Port),
RawQuery: q.Encode(),
}
if cc.User != "" || cc.Password != "" {
u.User = url.UserPassword(cc.User, cc.Password)
}
return u.String()
}
func (cc *ConnectionConfig) buildMongoDSN() string {
// Format: mongodb://username:password@host:port/database?authSource=admin
var dsn string
if cc.User != "" && cc.Password != "" {
dsn = fmt.Sprintf("mongodb://%s:%s@%s:%d/%s",
cc.User, cc.Password, cc.Host, cc.Port, cc.Database)
} else {
dsn = fmt.Sprintf("mongodb://%s:%d/%s", cc.Host, cc.Port, cc.Database)
}
params := ""
q := url.Values{}
if cc.AuthSource != "" {
params += fmt.Sprintf("authSource=%s", cc.AuthSource)
q.Set("authSource", cc.AuthSource)
}
if cc.ReplicaSet != "" {
if params != "" {
params += "&"
}
params += fmt.Sprintf("replicaSet=%s", cc.ReplicaSet)
q.Set("replicaSet", cc.ReplicaSet)
}
if cc.ReadPreference != "" {
if params != "" {
params += "&"
}
params += fmt.Sprintf("readPreference=%s", cc.ReadPreference)
q.Set("readPreference", cc.ReadPreference)
}
if params != "" {
dsn += "?" + params
u := url.URL{
Scheme: "mongodb",
Host: hostPort(cc.Host, cc.Port),
Path: "/" + cc.Database,
RawQuery: q.Encode(),
}
return dsn
if cc.User != "" && cc.Password != "" {
u.User = url.UserPassword(cc.User, cc.Password)
}
return u.String()
}
// FromConfig converts config.DBManagerConfig to internal ManagerConfig
@@ -487,3 +519,6 @@ func (cc *ConnectionConfig) GetConnMaxIdleTime() *time.Duration { return cc.Conn
func (cc *ConnectionConfig) GetQueryTimeout() time.Duration { return cc.QueryTimeout }
func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics }
func (cc *ConnectionConfig) GetReadPreference() string { return cc.ReadPreference }
func (cc *ConnectionConfig) GetRetryAttempts() int { return cc.RetryAttempts }
func (cc *ConnectionConfig) GetRetryDelay() time.Duration { return cc.RetryDelay }
func (cc *ConnectionConfig) GetRetryMaxDelay() time.Duration { return cc.RetryMaxDelay }
+95
View File
@@ -0,0 +1,95 @@
package dbmanager
import (
"net/url"
"strings"
"testing"
"time"
)
func TestPostgresDSNEscapesCredentials(t *testing.T) {
cc := ConnectionConfig{
Type: DatabaseTypePostgreSQL, Host: "db", Port: 5432, Database: "app",
User: "u@x", Password: "p w'd sslmode=disable&x=y/?#",
}
dsn := cc.buildPostgresDSN()
u, err := url.Parse(dsn)
if err != nil {
t.Fatalf("DSN is not a valid URL: %v", err)
}
if pw, _ := u.User.Password(); pw != cc.Password {
t.Errorf("password did not round-trip: %q", pw)
}
if u.User.Username() != cc.User {
t.Errorf("user did not round-trip: %q", u.User.Username())
}
if got := u.Query().Get("sslmode"); got != "prefer" {
t.Errorf("sslmode = %q, want prefer (password must not inject parameters)", got)
}
}
func TestMSSQLAndMongoDSNEscapeCredentials(t *testing.T) {
cc := ConnectionConfig{Host: "h", Port: 1, Database: "d", User: "u", Password: "a@b:c/d?e&f"}
for name, dsn := range map[string]string{"mssql": cc.buildMSSQLDSN(), "mongo": cc.buildMongoDSN()} {
u, err := url.Parse(dsn)
if err != nil {
t.Fatalf("%s: %v", name, err)
}
if pw, _ := u.User.Password(); pw != cc.Password {
t.Errorf("%s: password did not round-trip: %q", name, pw)
}
if u.Host != "h:1" {
t.Errorf("%s: host = %q", name, u.Host)
}
}
}
func TestSQLiteDSNUsesPragmas(t *testing.T) {
cc := ConnectionConfig{FilePath: "/tmp/x.db", QueryTimeout: 3 * time.Second}
dsn := cc.buildSQLiteDSN()
if strings.Contains(dsn, "?_timeout=") {
t.Errorf("unsupported _timeout parameter present: %s", dsn)
}
if !strings.Contains(dsn, "busy_timeout%283000%29") || !strings.Contains(dsn, "journal_mode%28WAL%29") {
t.Errorf("expected busy_timeout and WAL pragmas in DSN: %s", dsn)
}
}
func TestQueryTimeoutHonoredWithoutFloor(t *testing.T) {
cc := ConnectionConfig{QueryTimeout: 30 * time.Second}
cc.ApplyDefaults(&ManagerConfig{})
if cc.QueryTimeout != 30*time.Second {
t.Errorf("QueryTimeout = %v, want 30s", cc.QueryTimeout)
}
}
func TestRetryPolicyInherited(t *testing.T) {
g := ManagerConfig{RetryAttempts: 5, RetryDelay: time.Second, RetryMaxDelay: time.Minute}
cc := ConnectionConfig{}
cc.ApplyDefaults(&g)
if cc.GetRetryAttempts() != 5 || cc.GetRetryMaxDelay() != time.Minute {
t.Errorf("retry policy not inherited: %+v", cc)
}
}
func TestSQLiteMemoryPoolPinned(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
if err := mgr.Connect(t.Context()); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
db, _ := conn.Native()
if got := db.Stats().MaxOpenConnections; got != 1 {
t.Fatalf("MaxOpenConnections = %d, want 1 for :memory:", got)
}
if _, err := db.Exec("CREATE TABLE t(a int)"); err != nil {
t.Fatal(err)
}
for i := 0; i < 5; i++ {
var n int
if err := db.QueryRow("SELECT count(*) FROM t").Scan(&n); err != nil {
t.Fatalf("table missing on later use: %v", err)
}
}
}
+143 -68
View File
@@ -3,6 +3,7 @@ package dbmanager
import (
"context"
"database/sql"
"errors"
"fmt"
"sync"
"time"
@@ -14,6 +15,7 @@ import (
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
)
// Connection represents a single named database connection
@@ -82,6 +84,9 @@ type sqlConnection struct {
// State
connected bool
mu sync.RWMutex
// lifecycleMu serialises Connect/Close/Reconnect against health-check pings.
// Lock order: lifecycleMu before mu.
lifecycleMu sync.RWMutex
// Health check
lastHealthCheck time.Time
@@ -110,9 +115,16 @@ func (c *sqlConnection) Type() DatabaseType {
// Connect establishes the database connection
func (c *sqlConnection) Connect(ctx context.Context) error {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
return c.connectLocked(ctx)
}
// connectLocked requires lifecycleMu and mu held for writing.
func (c *sqlConnection) connectLocked(ctx context.Context) error {
if c.connected {
return ErrAlreadyConnected
}
@@ -127,17 +139,29 @@ func (c *sqlConnection) Connect(ctx context.Context) error {
// Close closes the database connection and all ORM instances
func (c *sqlConnection) Close() error {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
return c.closeLocked()
}
// closeLocked requires lifecycleMu and mu held for writing. The connection is
// always marked disconnected and its cached handles dropped, even when closing
// fails, so accessors never hand out handles over a half-closed pool.
func (c *sqlConnection) closeLocked() error {
if !c.connected {
return nil
}
// Close Bun if initialized
if c.bunDB != nil {
var errs []error
// Close Bun if initialized. bun.DB.Close closes the underlying *sql.DB, so
// skip it when the pool belongs to the caller.
if o, ok := c.provider.(interface{ OwnsDB() bool }); c.bunDB != nil && (!ok || o.OwnsDB()) {
if err := c.bunDB.Close(); err != nil {
return NewConnectionError(c.name, "close bun", err)
errs = append(errs, NewConnectionError(c.name, "close bun", err))
}
}
@@ -145,7 +169,7 @@ func (c *sqlConnection) Close() error {
// Close the provider (which closes the underlying sql.DB)
if err := c.provider.Close(); err != nil {
return NewConnectionError(c.name, "close", err)
errs = append(errs, NewConnectionError(c.name, "close", err))
}
c.connected = false
@@ -156,39 +180,75 @@ func (c *sqlConnection) Close() error {
c.gormAdapter = nil
c.nativeAdapter = nil
return nil
return errors.Join(errs...)
}
// HealthCheck verifies the connection is alive
// HealthCheck verifies the connection is alive. The network ping runs without
// holding mu, so handle accessors are never blocked behind a slow ping.
func (c *sqlConnection) HealthCheck(ctx context.Context) error {
if c == nil {
return fmt.Errorf("connection is nil")
}
c.mu.Lock()
defer c.mu.Unlock()
c.lastHealthCheck = time.Now()
// lifecycleMu (read) keeps Close/Reconnect from tearing the provider down
// mid-ping without blocking the accessors that only need mu.
c.lifecycleMu.RLock()
defer c.lifecycleMu.RUnlock()
if !c.connected {
c.healthCheckStatus = "disconnected"
c.mu.RLock()
connected := c.connected
provider := c.provider
c.mu.RUnlock()
if !connected {
c.setHealth("disconnected")
return ErrConnectionClosed
}
if err := c.provider.HealthCheck(ctx); err != nil {
c.healthCheckStatus = "unhealthy: " + err.Error()
if err := provider.HealthCheck(ctx); err != nil {
c.setHealth("unhealthy: " + err.Error())
return NewConnectionError(c.name, "health check", err)
}
c.healthCheckStatus = "healthy"
c.setHealth("healthy")
return nil
}
// Reconnect closes and re-establishes the connection
func (c *sqlConnection) Reconnect(ctx context.Context) error {
if err := c.Close(); err != nil {
func (c *sqlConnection) setHealth(status string) {
c.mu.Lock()
c.lastHealthCheck = time.Now()
c.healthCheckStatus = status
c.mu.Unlock()
}
// Reconnect refreshes the connection as a single critical section.
//
// Providers that support it (PostgreSQL) retire their pooled connections and
// dial fresh ones without closing the *sql.DB, so handles handed out earlier
// keep working. Other providers fall back to Close+Connect, which invalidates
// earlier handles; that is meant for explicit operator use only, since
// *sql.DB already replaces broken connections by itself.
func (c *sqlConnection) Reconnect(ctx context.Context) (err error) {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
defer func() { RecordReconnectAttempt(c.name, c.dbType, err == nil) }()
if c.connected {
if r, ok := c.provider.(providers.Refresher); ok {
if err := r.Refresh(ctx); err != nil {
return NewConnectionError(c.name, "reconnect", err)
}
return nil
}
}
if err := c.closeLocked(); err != nil {
return err
}
return c.Connect(ctx)
return c.connectLocked(ctx)
}
// Native returns the native *sql.DB connection
@@ -250,6 +310,10 @@ func (c *sqlConnection) Bun() (*bun.DB, error) {
return c.bunDB, nil
}
if !c.connected {
return nil, ErrConnectionClosed
}
// Get native connection first
native, err := c.provider.GetNative()
if err != nil {
@@ -283,6 +347,10 @@ func (c *sqlConnection) GORM() (*gorm.DB, error) {
return c.gormDB, nil
}
if !c.connected {
return nil, ErrConnectionClosed
}
// Get native connection first
native, err := c.provider.GetNative()
if err != nil {
@@ -359,39 +427,18 @@ func (c *sqlConnection) Stats() *ConnectionStats {
return stats
}
func (c *sqlConnection) reconnectForAdapter() error {
timeout := c.config.ConnectTimeout
if timeout <= 0 {
timeout = 10 * time.Second
}
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
return c.Reconnect(ctx)
}
// The adapter factories only re-fetch the current handle. They must not close
// the shared pool: *sql.DB discards bad connections on its own, and closing it
// here would break every other holder of the pool.
func (c *sqlConnection) reopenNativeForAdapter() (*sql.DB, error) {
if err := c.reconnectForAdapter(); err != nil {
return nil, err
}
return c.Native()
}
func (c *sqlConnection) reopenBunForAdapter() (*bun.DB, error) {
if err := c.reconnectForAdapter(); err != nil {
return nil, err
}
return c.Bun()
}
func (c *sqlConnection) reopenGORMForAdapter() (*gorm.DB, error) {
if err := c.reconnectForAdapter(); err != nil {
return nil, err
}
return c.GORM()
}
@@ -512,15 +559,8 @@ func (c *sqlConnection) getNativeAdapter() (common.Database, error) {
// Create a native adapter based on database type
switch c.dbType {
case DatabaseTypePostgreSQL:
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
WithDBFactory(c.reopenNativeForAdapter).
SetMetricsEnabled(c.config.EnableMetrics)
case DatabaseTypeSQLite:
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
WithDBFactory(c.reopenNativeForAdapter).
SetMetricsEnabled(c.config.EnableMetrics)
case DatabaseTypeMSSQL:
case DatabaseTypePostgreSQL, DatabaseTypeSQLite, DatabaseTypeMSSQL:
// The adapter takes the driver name so it can adjust its dialect.
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
WithDBFactory(c.reopenNativeForAdapter).
SetMetricsEnabled(c.config.EnableMetrics)
@@ -572,8 +612,9 @@ type mongoConnection struct {
client *mongo.Client
// State
connected bool
mu sync.RWMutex
connected bool
mu sync.RWMutex
lifecycleMu sync.RWMutex // see sqlConnection.lifecycleMu
// Health check
lastHealthCheck time.Time
@@ -601,9 +642,16 @@ func (c *mongoConnection) Type() DatabaseType {
// Connect establishes the MongoDB connection
func (c *mongoConnection) Connect(ctx context.Context) error {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
return c.connectLocked(ctx)
}
// connectLocked requires lifecycleMu and mu held for writing.
func (c *mongoConnection) connectLocked(ctx context.Context) error {
if c.connected {
return ErrAlreadyConnected
}
@@ -615,6 +663,7 @@ func (c *mongoConnection) Connect(ctx context.Context) error {
// Get the mongo client
client, err := c.provider.GetMongo()
if err != nil {
_ = c.provider.Close()
return NewConnectionError(c.name, "get mongo client", err)
}
@@ -625,49 +674,75 @@ func (c *mongoConnection) Connect(ctx context.Context) error {
// Close closes the MongoDB connection
func (c *mongoConnection) Close() error {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
return c.closeLocked()
}
// closeLocked requires lifecycleMu and mu held for writing. The connection is
// marked disconnected even when the provider fails to close cleanly.
func (c *mongoConnection) closeLocked() error {
if !c.connected {
return nil
}
if err := c.provider.Close(); err != nil {
return NewConnectionError(c.name, "close", err)
}
err := c.provider.Close()
c.connected = false
c.client = nil
if err != nil {
return NewConnectionError(c.name, "close", err)
}
return nil
}
// HealthCheck verifies the MongoDB connection is alive
// HealthCheck verifies the MongoDB connection is alive. The ping runs without
// holding mu so handle accessors are never blocked behind it.
func (c *mongoConnection) HealthCheck(ctx context.Context) error {
c.mu.Lock()
defer c.mu.Unlock()
c.lifecycleMu.RLock()
defer c.lifecycleMu.RUnlock()
c.lastHealthCheck = time.Now()
c.mu.RLock()
connected := c.connected
c.mu.RUnlock()
if !c.connected {
c.healthCheckStatus = "disconnected"
if !connected {
c.setHealth("disconnected")
return ErrConnectionClosed
}
if err := c.provider.HealthCheck(ctx); err != nil {
c.healthCheckStatus = "unhealthy: " + err.Error()
c.setHealth("unhealthy: " + err.Error())
return NewConnectionError(c.name, "health check", err)
}
c.healthCheckStatus = "healthy"
c.setHealth("healthy")
return nil
}
// Reconnect closes and re-establishes the MongoDB connection
func (c *mongoConnection) Reconnect(ctx context.Context) error {
if err := c.Close(); err != nil {
func (c *mongoConnection) setHealth(status string) {
c.mu.Lock()
c.lastHealthCheck = time.Now()
c.healthCheckStatus = status
c.mu.Unlock()
}
// Reconnect closes and re-establishes the MongoDB connection atomically.
func (c *mongoConnection) Reconnect(ctx context.Context) (err error) {
c.lifecycleMu.Lock()
defer c.lifecycleMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
defer func() { RecordReconnectAttempt(c.name, DatabaseTypeMongoDB, err == nil) }()
if err := c.closeLocked(); err != nil {
return err
}
return c.Connect(ctx)
return c.connectLocked(ctx)
}
// MongoDB returns the MongoDB client
-159
View File
@@ -4,13 +4,8 @@ import (
"context"
"database/sql"
"testing"
"time"
_ "github.com/mattn/go-sqlite3"
"gorm.io/gorm"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
)
func TestNewConnectionFromDB(t *testing.T) {
@@ -213,157 +208,3 @@ func TestNewConnectionFromDB_PostgreSQL(t *testing.T) {
t.Errorf("Expected type DatabaseTypePostgreSQL, got '%s'", conn.Type())
}
}
func TestDatabaseNativeAdapterReconnectFactory(t *testing.T) {
conn := newSQLConnection("test-native", DatabaseTypeSQLite, ConnectionConfig{
Name: "test-native",
Type: DatabaseTypeSQLite,
FilePath: ":memory:",
DefaultORM: string(ORMTypeNative),
ConnectTimeout: 2 * time.Second,
}, providers.NewSQLiteProvider())
ctx := context.Background()
if err := conn.Connect(ctx); err != nil {
t.Fatalf("Failed to connect: %v", err)
}
defer conn.Close()
db, err := conn.Database()
if err != nil {
t.Fatalf("Failed to get database adapter: %v", err)
}
adapter, ok := db.(*database.PgSQLAdapter)
if !ok {
t.Fatalf("Expected PgSQLAdapter, got %T", db)
}
underlyingBefore, ok := adapter.GetUnderlyingDB().(*sql.DB)
if !ok {
t.Fatalf("Expected underlying *sql.DB, got %T", adapter.GetUnderlyingDB())
}
if err := underlyingBefore.Close(); err != nil {
t.Fatalf("Failed to close underlying database: %v", err)
}
if _, err := db.Exec(ctx, "SELECT 1"); err != nil {
t.Fatalf("Expected native adapter to reconnect, got error: %v", err)
}
underlyingAfter, ok := adapter.GetUnderlyingDB().(*sql.DB)
if !ok {
t.Fatalf("Expected reconnected *sql.DB, got %T", adapter.GetUnderlyingDB())
}
if underlyingAfter == underlyingBefore {
t.Fatal("Expected adapter to swap to a fresh *sql.DB after reconnect")
}
}
func TestDatabaseBunAdapterReconnectFactory(t *testing.T) {
conn := newSQLConnection("test-bun", DatabaseTypeSQLite, ConnectionConfig{
Name: "test-bun",
Type: DatabaseTypeSQLite,
FilePath: ":memory:",
DefaultORM: string(ORMTypeBun),
ConnectTimeout: 2 * time.Second,
}, providers.NewSQLiteProvider())
ctx := context.Background()
if err := conn.Connect(ctx); err != nil {
t.Fatalf("Failed to connect: %v", err)
}
defer conn.Close()
db, err := conn.Database()
if err != nil {
t.Fatalf("Failed to get database adapter: %v", err)
}
adapter, ok := db.(*database.BunAdapter)
if !ok {
t.Fatalf("Expected BunAdapter, got %T", db)
}
underlyingBefore, ok := adapter.GetUnderlyingDB().(interface{ Close() error })
if !ok {
t.Fatalf("Expected underlying Bun DB with Close method, got %T", adapter.GetUnderlyingDB())
}
if err := underlyingBefore.Close(); err != nil {
t.Fatalf("Failed to close underlying Bun database: %v", err)
}
if _, err := db.Exec(ctx, "SELECT 1"); err != nil {
t.Fatalf("Expected Bun adapter to reconnect, got error: %v", err)
}
underlyingAfter := adapter.GetUnderlyingDB()
if underlyingAfter == underlyingBefore {
t.Fatal("Expected adapter to swap to a fresh Bun DB after reconnect")
}
}
func TestDatabaseGormAdapterReconnectFactory(t *testing.T) {
conn := newSQLConnection("test-gorm", DatabaseTypeSQLite, ConnectionConfig{
Name: "test-gorm",
Type: DatabaseTypeSQLite,
FilePath: ":memory:",
DefaultORM: string(ORMTypeGORM),
ConnectTimeout: 2 * time.Second,
}, providers.NewSQLiteProvider())
ctx := context.Background()
if err := conn.Connect(ctx); err != nil {
t.Fatalf("Failed to connect: %v", err)
}
defer conn.Close()
db, err := conn.Database()
if err != nil {
t.Fatalf("Failed to get database adapter: %v", err)
}
adapter, ok := db.(*database.GormAdapter)
if !ok {
t.Fatalf("Expected GormAdapter, got %T", db)
}
gormBefore, ok := adapter.GetUnderlyingDB().(*gorm.DB)
if !ok {
t.Fatalf("Expected underlying *gorm.DB, got %T", adapter.GetUnderlyingDB())
}
sqlBefore, err := gormBefore.DB()
if err != nil {
t.Fatalf("Failed to get underlying *sql.DB: %v", err)
}
if err := sqlBefore.Close(); err != nil {
t.Fatalf("Failed to close underlying database: %v", err)
}
count, err := db.NewSelect().Table("sqlite_master").Count(ctx)
if err != nil {
t.Fatalf("Expected GORM query builder to reconnect, got error: %v", err)
}
if count < 0 {
t.Fatalf("Expected non-negative count, got %d", count)
}
gormAfter, ok := adapter.GetUnderlyingDB().(*gorm.DB)
if !ok {
t.Fatalf("Expected reconnected *gorm.DB, got %T", adapter.GetUnderlyingDB())
}
sqlAfter, err := gormAfter.DB()
if err != nil {
t.Fatalf("Failed to get reconnected *sql.DB: %v", err)
}
if sqlAfter == sqlBefore {
t.Fatal("Expected GORM adapter to use a fresh *sql.DB after reconnect")
}
}
+199
View File
@@ -0,0 +1,199 @@
package dbmanager
import (
"context"
"database/sql"
"sync"
"testing"
"time"
)
func sqliteManagerConfig() ManagerConfig {
return ManagerConfig{
DefaultConnection: "test",
Connections: map[string]ConnectionConfig{
"test": {Name: "test", Type: DatabaseTypeSQLite, FilePath: ":memory:"},
},
HealthCheckInterval: time.Hour,
}
}
func TestManagerConnectCloseCycleTwice(t *testing.T) {
mgr, err := NewManager(sqliteManagerConfig())
if err != nil {
t.Fatal(err)
}
ctx := context.Background()
for i := 0; i < 2; i++ {
if err := mgr.Connect(ctx); err != nil {
t.Fatalf("cycle %d connect: %v", i, err)
}
cm := mgr.(*connectionManager)
cm.healthMu.Lock()
running := cm.healthTicker != nil
cm.healthMu.Unlock()
if !running {
t.Fatalf("cycle %d: health checker not running", i)
}
if err := mgr.Close(); err != nil {
t.Fatalf("cycle %d close: %v", i, err)
}
}
// A further Close must not panic.
if err := mgr.Close(); err != nil {
t.Fatal(err)
}
}
func TestManagerConnectIsIdempotent(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
if got := mgr.Stats().TotalConnections; got != 1 {
t.Fatalf("expected 1 connection, got %d", got)
}
}
func TestConcurrentReconnectIsAtomic(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
var wg sync.WaitGroup
for i := 0; i < 20; i++ {
wg.Add(1)
go func() {
defer wg.Done()
if err := conn.Reconnect(ctx); err != nil {
t.Errorf("reconnect: %v", err)
}
}()
}
wg.Wait()
db, err := conn.Native()
if err != nil {
t.Fatal(err)
}
if err := db.PingContext(ctx); err != nil {
t.Fatalf("pool unusable after concurrent reconnects: %v", err)
}
}
func TestAdapterFactoryDoesNotClosePool(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
sc := conn.(*sqlConnection)
held, _ := sc.Native()
if _, err := sc.reopenNativeForAdapter(); err != nil {
t.Fatal(err)
}
if err := held.PingContext(ctx); err != nil {
t.Fatalf("existing handle broken by adapter factory: %v", err)
}
}
func TestHealthCheckDoesNotBlockAccessors(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
sc := conn.(*sqlConnection)
// Simulate a health check in flight: it holds lifecycleMu (read) only.
sc.lifecycleMu.RLock()
defer sc.lifecycleMu.RUnlock()
done := make(chan struct{})
go func() {
_, _ = sc.Bun()
_, _ = sc.GORM()
close(done)
}()
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("accessors blocked while health check in flight")
}
}
func TestCloseAlwaysMarksDisconnected(t *testing.T) {
mgr, _ := NewManager(sqliteManagerConfig())
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
conn, _ := mgr.GetDefault()
sc := conn.(*sqlConnection)
_, _ = sc.Bun()
_ = mgr.Close()
if _, err := sc.Bun(); err == nil {
t.Fatal("Bun() should fail after Close")
}
if _, err := sc.GORM(); err == nil {
t.Fatal("GORM() should fail after Close")
}
}
func TestReconnectOnExistingDBKeepsCallersPool(t *testing.T) {
db, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatal(err)
}
defer db.Close()
conn := NewConnectionFromDB("existing", DatabaseTypeSQLite, db)
ctx := context.Background()
if err := conn.Connect(ctx); err != nil {
t.Fatal(err)
}
if err := conn.Reconnect(ctx); err != nil {
t.Fatalf("reconnect: %v", err)
}
if err := db.PingContext(ctx); err != nil {
t.Fatalf("caller's pool was closed by Reconnect: %v", err)
}
}
func TestCloseOnExistingDBLeavesCallersPoolOpen(t *testing.T) {
db, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatal(err)
}
defer db.Close()
conn := NewConnectionFromDB("existing", DatabaseTypeSQLite, db)
ctx := context.Background()
if err := conn.Connect(ctx); err != nil {
t.Fatal(err)
}
if _, err := conn.Bun(); err != nil { // bun.DB.Close would close the pool
t.Fatal(err)
}
if err := conn.Close(); err != nil {
t.Fatal(err)
}
if err := db.PingContext(ctx); err != nil {
t.Fatalf("caller's pool was closed: %v", err)
}
}
+66 -52
View File
@@ -2,9 +2,7 @@ package dbmanager
import (
"context"
"errors"
"fmt"
"strings"
"sync"
"time"
@@ -49,6 +47,7 @@ type connectionManager struct {
// Background health check
healthTicker *time.Ticker
stopChan chan struct{}
healthMu sync.Mutex // guards healthTicker and stopChan
wg sync.WaitGroup
}
@@ -100,7 +99,9 @@ func ResetInstance() {
defer instanceMu.Unlock()
if instance != nil {
_ = instance.Close()
if err := instance.Close(); err != nil {
logger.Error("Failed to close manager during reset: %v", err)
}
}
instance = nil
}
@@ -116,7 +117,6 @@ func NewManager(cfg ManagerConfig) (Manager, error) {
mgr := &connectionManager{
connections: make(map[string]Connection),
config: cfg,
stopChan: make(chan struct{}),
}
return mgr, nil
@@ -195,11 +195,26 @@ func (m *connectionManager) SetDefaultDatabase(name string) error {
// Connect establishes all configured database connections
func (m *connectionManager) Connect(ctx context.Context) error {
m.mu.Lock()
defer m.mu.Unlock()
// Create connections from configuration
// Dial outside m.mu so a slow connect never blocks Get/Stats/HealthCheck.
m.mu.RLock()
names := make([]string, 0, len(m.config.Connections))
for name := range m.config.Connections {
if _, exists := m.connections[name]; !exists {
names = append(names, name)
}
}
m.mu.RUnlock()
opened := make(map[string]Connection, len(names))
closeOpened := func() {
for name, conn := range opened {
if err := conn.Close(); err != nil {
logger.Error("Failed to close connection after failed Connect: name=%s, error=%v", name, err)
}
}
}
for _, name := range names {
// Get a copy of the connection config
connCfg := m.config.Connections[name]
// Apply global defaults to connection config
@@ -209,25 +224,39 @@ func (m *connectionManager) Connect(ctx context.Context) error {
// Create connection using factory
conn, err := createConnection(connCfg)
if err != nil {
closeOpened()
return fmt.Errorf("failed to create connection '%s': %w", name, err)
}
// Connect
if err := conn.Connect(ctx); err != nil {
closeOpened()
return fmt.Errorf("failed to connect '%s': %w", name, err)
}
m.connections[name] = conn
opened[name] = conn
logger.Info("Database connection established: name=%s, type=%s", name, connCfg.Type)
}
m.mu.Lock()
for name, conn := range opened {
if _, exists := m.connections[name]; exists {
// Lost a race with a concurrent Connect; drop our duplicate.
_ = conn.Close()
continue
}
m.connections[name] = conn
}
total := len(m.connections)
m.mu.Unlock()
// Always start background health checks
if m.config.HealthCheckInterval > 0 {
m.startHealthChecker()
logger.Info("Background health checker started: interval=%v", m.config.HealthCheckInterval)
}
logger.Info("Database manager initialized: connections=%d", len(m.connections))
logger.Info("Database manager initialized: connections=%d", total)
return nil
}
@@ -246,7 +275,7 @@ func (m *connectionManager) Close() error {
for name, conn := range m.connections {
if err := conn.Close(); err != nil {
errors = append(errors, fmt.Errorf("failed to close connection '%s': %w", name, err))
logger.Error("Failed to close connection", "name", name, "error", err)
logger.Error("Failed to close connection: name=%s, error=%v", name, err)
} else {
logger.Info("Connection closed: name=%s", name)
}
@@ -311,11 +340,17 @@ func (m *connectionManager) Stats() *ManagerStats {
// startHealthChecker starts background health checking
func (m *connectionManager) startHealthChecker() {
m.healthMu.Lock()
defer m.healthMu.Unlock()
if m.healthTicker != nil {
return // Already running
}
m.healthTicker = time.NewTicker(m.config.HealthCheckInterval)
ticker := time.NewTicker(m.config.HealthCheckInterval)
stop := make(chan struct{})
m.healthTicker = ticker
m.stopChan = stop
m.wg.Add(1)
go func() {
@@ -324,9 +359,9 @@ func (m *connectionManager) startHealthChecker() {
for {
select {
case <-m.healthTicker.C:
case <-ticker.C:
m.performHealthCheck()
case <-m.stopChan:
case <-stop:
logger.Info("Health checker stopped")
return
}
@@ -334,14 +369,19 @@ func (m *connectionManager) startHealthChecker() {
}()
}
// stopHealthChecker stops background health checking
// stopHealthChecker stops background health checking. Safe to call repeatedly.
func (m *connectionManager) stopHealthChecker() {
if m.healthTicker != nil {
m.healthTicker.Stop()
close(m.stopChan)
m.wg.Wait()
m.healthTicker = nil
m.healthMu.Lock()
defer m.healthMu.Unlock()
if m.healthTicker == nil {
return
}
m.healthTicker.Stop()
close(m.stopChan)
m.wg.Wait()
m.healthTicker = nil
m.stopChan = nil
}
// performHealthCheck performs a health check on all connections
@@ -362,40 +402,14 @@ func (m *connectionManager) performHealthCheck() {
}
m.mu.RUnlock()
defer m.PublishMetrics()
for _, item := range connections {
if err := item.conn.HealthCheck(ctx); err != nil {
logger.Warn("Health check failed",
"connection", item.name,
"error", err)
// Only reconnect when the client handle itself is closed/disconnected.
// For transient database restarts or network blips, *sql.DB can recover
// on its own; forcing Close()+Connect() here invalidates any cached ORM
// wrappers and callers that still hold the old handle.
if m.config.EnableAutoReconnect && shouldReconnectAfterHealthCheck(err) {
logger.Info("Attempting reconnection: connection=%s", item.name)
if err := item.conn.Reconnect(ctx); err != nil {
logger.Error("Reconnection failed",
"connection", item.name,
"error", err)
} else {
logger.Info("Reconnection successful: connection=%s", item.name)
}
} else if m.config.EnableAutoReconnect {
logger.Info("Skipping reconnect for transient health check failure: connection=%s", item.name)
}
// Do not reconnect here: *sql.DB discards bad connections and dials
// new ones by itself, while Reconnect closes the pool and breaks
// every handle already handed out. Reconnect is operator-only.
logger.Warn("Health check failed: connection=%s, error=%v", item.name, err)
}
}
}
func shouldReconnectAfterHealthCheck(err error) bool {
if err == nil {
return false
}
if errors.Is(err, ErrConnectionClosed) {
return true
}
return strings.Contains(err.Error(), "sql: database is closed")
}
+36 -45
View File
@@ -21,19 +21,30 @@ type healthCheckStubConnection struct {
reconnectCalls int
}
func (c *healthCheckStubConnection) Name() string { return "stub" }
func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL }
func (c *healthCheckStubConnection) Bun() (*bun.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) GORM() (*gorm.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) Native() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) DB() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) Database() (common.Database, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) MongoDB() (*mongo.Client, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) Connect(ctx context.Context) error { return nil }
func (c *healthCheckStubConnection) Close() error { return nil }
func (c *healthCheckStubConnection) HealthCheck(ctx context.Context) error { return c.healthErr }
func (c *healthCheckStubConnection) Reconnect(ctx context.Context) error { c.reconnectCalls++; return nil }
func (c *healthCheckStubConnection) Stats() *ConnectionStats { return &ConnectionStats{} }
func (c *healthCheckStubConnection) Name() string { return "stub" }
func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL }
func (c *healthCheckStubConnection) Bun() (*bun.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) GORM() (*gorm.DB, error) {
return nil, fmt.Errorf("not implemented")
}
func (c *healthCheckStubConnection) Native() (*sql.DB, error) {
return nil, fmt.Errorf("not implemented")
}
func (c *healthCheckStubConnection) DB() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") }
func (c *healthCheckStubConnection) Database() (common.Database, error) {
return nil, fmt.Errorf("not implemented")
}
func (c *healthCheckStubConnection) MongoDB() (*mongo.Client, error) {
return nil, fmt.Errorf("not implemented")
}
func (c *healthCheckStubConnection) Connect(ctx context.Context) error { return nil }
func (c *healthCheckStubConnection) Close() error { return nil }
func (c *healthCheckStubConnection) HealthCheck(ctx context.Context) error { return c.healthErr }
func (c *healthCheckStubConnection) Reconnect(ctx context.Context) error {
c.reconnectCalls++
return nil
}
func (c *healthCheckStubConnection) Stats() *ConnectionStats { return &ConnectionStats{} }
func TestBackgroundHealthChecker(t *testing.T) {
// Create a SQLite in-memory database
@@ -117,41 +128,21 @@ func TestDefaultHealthCheckInterval(t *testing.T) {
t.Errorf("Expected default health check interval to be %v, got %v",
expectedInterval, defaults.HealthCheckInterval)
}
if !defaults.EnableAutoReconnect {
t.Error("Expected EnableAutoReconnect to be true by default")
}
}
func TestApplyDefaultsEnablesAutoReconnect(t *testing.T) {
// Create a config without setting EnableAutoReconnect
cfg := ManagerConfig{
Connections: map[string]ConnectionConfig{
"test": {
Name: "test",
Type: DatabaseTypeSQLite,
FilePath: ":memory:",
},
},
}
// Verify it's false initially (Go's zero value for bool)
if cfg.EnableAutoReconnect {
t.Error("Expected EnableAutoReconnect to be false before ApplyDefaults")
}
// Apply defaults
func TestApplyDefaultsHealthCheckInterval(t *testing.T) {
cfg := ManagerConfig{}
cfg.ApplyDefaults()
// Verify it's now true
if !cfg.EnableAutoReconnect {
t.Error("Expected EnableAutoReconnect to be true after ApplyDefaults")
}
// Verify health check interval is also set
if cfg.HealthCheckInterval != 15*time.Second {
t.Errorf("Expected health check interval to be 15s, got %v", cfg.HealthCheckInterval)
}
// A negative interval disables the background checker and is preserved.
cfg = ManagerConfig{HealthCheckInterval: -1}
cfg.ApplyDefaults()
if cfg.HealthCheckInterval >= 0 {
t.Errorf("Expected negative interval to be preserved, got %v", cfg.HealthCheckInterval)
}
}
func TestManagerHealthCheck(t *testing.T) {
@@ -270,7 +261,7 @@ func TestPerformHealthCheckSkipsReconnectForTransientFailures(t *testing.T) {
}
}
func TestPerformHealthCheckReconnectsClosedConnections(t *testing.T) {
func TestPerformHealthCheckNeverReconnects(t *testing.T) {
conn := &healthCheckStubConnection{
healthErr: NewConnectionError("primary", "health check", fmt.Errorf("sql: database is closed")),
}
@@ -284,7 +275,7 @@ func TestPerformHealthCheckReconnectsClosedConnections(t *testing.T) {
mgr.performHealthCheck()
if conn.reconnectCalls != 1 {
t.Fatalf("expected reconnect attempt for closed database handle, got %d", conn.reconnectCalls)
if conn.reconnectCalls != 0 {
t.Fatalf("health check must not close the shared pool via Reconnect, got %d", conn.reconnectCalls)
}
}
+39 -15
View File
@@ -1,6 +1,8 @@
package dbmanager
import (
"sync"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/promauto"
)
@@ -34,8 +36,8 @@ var (
)
// connectionWaitCount tracks how many times connections had to wait for availability
connectionWaitCount = promauto.NewGaugeVec(
prometheus.GaugeOpts{
connectionWaitCount = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "dbmanager_connection_wait_count",
Help: "Number of times connections had to wait for availability",
},
@@ -43,8 +45,8 @@ var (
)
// connectionWaitDuration tracks total time connections spent waiting
connectionWaitDuration = promauto.NewGaugeVec(
prometheus.GaugeOpts{
connectionWaitDuration = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "dbmanager_connection_wait_duration_seconds",
Help: "Total time connections spent waiting for availability",
},
@@ -61,8 +63,8 @@ var (
)
// connectionLifetimeClosed tracks connections closed due to max lifetime
connectionLifetimeClosed = promauto.NewGaugeVec(
prometheus.GaugeOpts{
connectionLifetimeClosed = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "dbmanager_connection_lifetime_closed_total",
Help: "Total connections closed due to exceeding max lifetime",
},
@@ -70,8 +72,8 @@ var (
)
// connectionIdleClosed tracks connections closed due to max idle time
connectionIdleClosed = promauto.NewGaugeVec(
prometheus.GaugeOpts{
connectionIdleClosed = promauto.NewCounterVec(
prometheus.CounterOpts{
Name: "dbmanager_connection_idle_closed_total",
Help: "Total connections closed due to exceeding max idle time",
},
@@ -114,13 +116,13 @@ func (m *connectionManager) PublishMetrics() {
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "idle").Set(float64(connStats.Idle))
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "in_use").Set(float64(connStats.InUse))
// Wait stats
connectionWaitCount.With(labels).Set(float64(connStats.WaitCount))
connectionWaitDuration.With(labels).Set(connStats.WaitDuration.Seconds())
// Lifetime/idle closed
connectionLifetimeClosed.With(labels).Set(float64(connStats.MaxLifetimeClosed))
connectionIdleClosed.With(labels).Set(float64(connStats.MaxIdleClosed))
// sql.DBStats values are cumulative, so add only the growth since
// the last publish to keep these true counters.
prev := lastPublished.swap(name, connStats)
connectionWaitCount.With(labels).Add(float64(connStats.WaitCount - prev.WaitCount))
connectionWaitDuration.With(labels).Add((connStats.WaitDuration - prev.WaitDuration).Seconds())
connectionLifetimeClosed.With(labels).Add(float64(connStats.MaxLifetimeClosed - prev.MaxLifetimeClosed))
connectionIdleClosed.With(labels).Add(float64(connStats.MaxIdleClosed - prev.MaxIdleClosed))
}
}
}
@@ -134,3 +136,25 @@ func RecordReconnectAttempt(name string, dbType DatabaseType, success bool) {
reconnectAttempts.WithLabelValues(name, string(dbType), result).Inc()
}
// publishedStats remembers the cumulative pool stats last exported per
// connection so counters can be advanced by the delta.
type publishedStats struct {
mu sync.Mutex
last map[string]ConnectionStats
}
var lastPublished = &publishedStats{last: make(map[string]ConnectionStats)}
// swap stores cur and returns the previous value. A counter reset (a new pool
// after Close+Connect) is treated as starting from zero.
func (p *publishedStats) swap(name string, cur *ConnectionStats) ConnectionStats {
p.mu.Lock()
defer p.mu.Unlock()
prev := p.last[name]
if cur.WaitCount < prev.WaitCount || cur.MaxIdleClosed < prev.MaxIdleClosed || cur.MaxLifetimeClosed < prev.MaxLifetimeClosed {
prev = ConnectionStats{}
}
p.last[name] = *cur
return prev
}
+93
View File
@@ -0,0 +1,93 @@
package dbmanager
import (
"context"
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
"os"
"testing"
"time"
)
func TestLivePostgresRefreshKeepsHandles(t *testing.T) {
if os.Getenv("PG_LIVE") == "" {
t.Skip("PG_LIVE not set")
}
mgr, err := NewManager(ManagerConfig{
DefaultConnection: "pg",
Connections: map[string]ConnectionConfig{"pg": {
Name: "pg", Type: DatabaseTypePostgreSQL, Host: "127.0.0.1", Port: 54329,
User: "postgres", Database: "postgres", QueryTimeout: 30 * time.Second,
}},
HealthCheckInterval: -1,
})
if err != nil {
t.Fatal(err)
}
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
held, _ := conn.Bun()
gormDB, _ := conn.GORM()
var pid1, pid2 int
var st string
if err := held.DB.QueryRow("select pg_backend_pid(), current_setting('statement_timeout')").Scan(&pid1, &st); err != nil {
t.Fatal(err)
}
if st != "30s" {
t.Errorf("statement_timeout = %q", st)
}
if err := conn.Reconnect(ctx); err != nil {
t.Fatal(err)
}
if err := held.DB.QueryRow("select pg_backend_pid()").Scan(&pid2); err != nil {
t.Fatalf("held bun handle broken after reconnect: %v", err)
}
if pid1 == pid2 {
t.Error("expected a new backend after reconnect")
}
var n int
if err := gormDB.Raw("select 1").Scan(&n).Error; err != nil || n != 1 {
t.Fatalf("held gorm handle broken: %v", err)
}
}
func TestLiveListenerListenNotify(t *testing.T) {
if os.Getenv("PG_LIVE") == "" {
t.Skip("PG_LIVE not set")
}
cc := ConnectionConfig{Name: "pg", Type: DatabaseTypePostgreSQL, Host: "127.0.0.1", Port: 54329,
User: "postgres", Database: "postgres", ConnectTimeout: 5 * time.Second}
p := providers.NewPostgresProvider()
ctx := context.Background()
if err := p.Connect(ctx, &cc); err != nil {
t.Fatal(err)
}
defer p.Close()
l, err := p.GetListener(ctx)
if err != nil {
t.Fatal(err)
}
got := make(chan string, 4)
for _, ch := range []string{"a", "b"} {
if err := l.Listen(ch, func(c, payload string) { got <- c + ":" + payload }); err != nil {
t.Fatal(err)
}
}
for i := 0; i < 3; i++ {
if err := l.Notify(ctx, "a", "x"); err != nil {
t.Fatalf("notify: %v", err)
}
}
select {
case v := <-got:
if v != "a:x" {
t.Fatalf("got %q", v)
}
case <-time.After(3 * time.Second):
t.Fatal("no notification")
}
}
+19 -6
View File
@@ -7,6 +7,8 @@ import (
"sync"
"go.mongodb.org/mongo-driver/mongo"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// ExistingDBProvider wraps an existing *sql.DB connection
@@ -44,16 +46,27 @@ func (p *ExistingDBProvider) Connect(ctx context.Context, cfg ConnectionConfig)
return nil
}
// Close closes the underlying database connection
func (p *ExistingDBProvider) Close() error {
p.mu.Lock()
defer p.mu.Unlock()
// Refresh verifies the wrapped database is still reachable. The pool belongs to
// the caller and cannot be re-dialed here, so it is never closed to "reconnect".
func (p *ExistingDBProvider) Refresh(ctx context.Context) error {
p.mu.RLock()
defer p.mu.RUnlock()
if p.db == nil {
return nil
return fmt.Errorf("database connection is nil")
}
return p.db.PingContext(ctx)
}
return p.db.Close()
// OwnsDB reports whether Close releases the wrapped database. It never does:
// the *sql.DB was opened by the caller, who is responsible for closing it.
func (p *ExistingDBProvider) OwnsDB() bool { return false }
// Close is a no-op for the wrapped database. The pool belongs to the caller, so
// closing it here would break the caller's other users of it.
func (p *ExistingDBProvider) Close() error {
logger.Warn("Not closing externally provided database: name=%s; the caller owns this *sql.DB and must close it", p.name)
return nil
}
// HealthCheck verifies the connection is alive
+5 -5
View File
@@ -164,7 +164,7 @@ func TestExistingDBProvider_Stats(t *testing.T) {
}
}
func TestExistingDBProvider_Close(t *testing.T) {
func TestExistingDBProvider_Close_LeavesDBOpen(t *testing.T) {
db, err := sql.Open("sqlite3", ":memory:")
if err != nil {
t.Fatalf("Failed to open database: %v", err)
@@ -177,10 +177,10 @@ func TestExistingDBProvider_Close(t *testing.T) {
t.Errorf("Expected Close to succeed, got error: %v", err)
}
// Verify the database is closed
err = db.Ping()
if err == nil {
t.Error("Expected database to be closed")
// The caller owns the database, so Close must leave it open
defer db.Close()
if err := db.Ping(); err != nil {
t.Errorf("Expected caller's database to stay open, got: %v", err)
}
}
+9 -9
View File
@@ -37,13 +37,14 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
// Set connection pool size
if cfg.GetMaxOpenConns() != nil {
maxPoolSize := uint64(*cfg.GetMaxOpenConns())
maxPoolSize := uint64(*cfg.GetMaxOpenConns()) //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
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
}
+5 -6
View File
@@ -35,12 +35,11 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
var db *sql.DB
var lastErr error
retryAttempts := 3 // Default retry attempts
retryDelay := 1 * time.Second
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
for attempt := 0; attempt < retryAttempts; attempt++ {
if attempt > 0 {
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
if cfg.GetEnableLogging() {
logger.Info("Retrying MSSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
}
@@ -57,7 +56,7 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
if err != nil {
lastErr = err
if cfg.GetEnableLogging() {
logger.Warn("Failed to open MSSQL connection", "error", err)
logger.Warn("Failed to open MSSQL connection: %v", err)
}
continue
}
@@ -69,9 +68,9 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
if err != nil {
lastErr = err
db.Close()
db.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
if cfg.GetEnableLogging() {
logger.Warn("Failed to ping MSSQL database", "error", err)
logger.Warn("Failed to ping MSSQL database: %v", err)
}
continue
}
+132
View File
@@ -0,0 +1,132 @@
package providers
import (
"context"
"database/sql/driver"
"fmt"
"net"
"sync/atomic"
"time"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/stdlib"
)
const (
// tcpKeepAlive is how often keepalive probes are sent on idle connections.
tcpKeepAlive = 30 * time.Second
// tcpUserTimeout bounds how long written data may stay unacknowledged before
// the kernel drops the socket. Without it a query on a silently dead peer
// waits for tcp_retries2 (about 15 minutes).
tcpUserTimeout = 30 * time.Second
// resetSessionTimeout bounds the liveness ping database/sql triggers when a
// pooled connection is reused, which otherwise runs on the request context.
resetSessionTimeout = 5 * time.Second
)
// pgConnector is a driver.Connector whose connections can be retired without
// closing the *sql.DB. Reconnecting bumps a generation; connections created
// under an older generation report themselves invalid and database/sql
// discards them and dials new ones. Every handle wrapping the *sql.DB keeps
// working across a reconnect.
type pgConnector struct {
inner atomic.Pointer[connectorState]
generation atomic.Uint64
}
type connectorState struct {
connector driver.Connector
gen uint64
}
func newPGConnector(cfg *pgx.ConnConfig) *pgConnector {
c := &pgConnector{}
c.swap(cfg)
return c
}
// swap installs a new connection config under a fresh generation.
func (c *pgConnector) swap(cfg *pgx.ConnConfig) {
gen := c.generation.Add(1)
c.inner.Store(&connectorState{connector: stdlib.GetConnector(*cfg), gen: gen})
}
func (c *pgConnector) Connect(ctx context.Context) (driver.Conn, error) {
st := c.inner.Load()
conn, err := st.connector.Connect(ctx)
if err != nil {
return nil, err
}
sc, ok := conn.(*stdlib.Conn)
if !ok {
conn.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
return nil, fmt.Errorf("unexpected pgx driver connection type %T", conn)
}
return &pgConn{Conn: sc, owner: c, gen: st.gen}, nil
}
func (c *pgConnector) Driver() driver.Driver {
return stdlib.GetDefaultDriver()
}
// pgConn embeds *stdlib.Conn, so every optional driver interface (context
// queries, Pinger, NamedValueChecker, ...) is promoted unchanged.
type pgConn struct {
*stdlib.Conn
owner *pgConnector
gen uint64
}
func (c *pgConn) stale() bool { return c.gen != c.owner.generation.Load() }
// IsValid implements driver.Validator: stale or closed connections are dropped
// when returned to the pool.
func (c *pgConn) IsValid() bool {
return !c.stale() && !c.Conn.Conn().IsClosed()
}
// ResetSession runs when a pooled connection is reused. It discards stale
// connections and bounds pgx's liveness ping so a dead socket fails in seconds
// rather than blocking on the caller's context.
func (c *pgConn) ResetSession(ctx context.Context) error {
if c.stale() {
return driver.ErrBadConn
}
ctx, cancel := context.WithTimeout(ctx, resetSessionTimeout)
defer cancel()
return c.Conn.ResetSession(ctx)
}
// newDialFunc returns a pgconn dial function with TCP keepalive and, where the
// platform supports it, TCP_USER_TIMEOUT.
func newDialFunc(connectTimeout time.Duration) func(ctx context.Context, network, addr string) (net.Conn, error) {
d := &net.Dialer{
Timeout: connectTimeout,
KeepAlive: tcpKeepAlive,
Control: setTCPUserTimeout(tcpUserTimeout),
}
return d.DialContext
}
// buildPGXConfig parses the DSN and applies client-side hardening: bounded
// dialing, TCP timeouts, and statement_timeout, which is set as a runtime
// parameter so it also applies to caller-supplied DSNs.
func buildPGXConfig(cfg ConnectionConfig) (*pgx.ConnConfig, error) {
dsn, err := cfg.BuildDSN()
if err != nil {
return nil, fmt.Errorf("failed to build DSN: %w", err)
}
cc, err := pgx.ParseConfig(dsn)
if err != nil {
return nil, fmt.Errorf("failed to parse connection config: %w", err)
}
cc.DialFunc = newDialFunc(cfg.GetConnectTimeout())
if cfg.GetQueryTimeout() > 0 {
if _, set := cc.RuntimeParams["statement_timeout"]; !set {
cc.RuntimeParams["statement_timeout"] = fmt.Sprintf("%d", cfg.GetQueryTimeout().Milliseconds())
}
}
return cc, nil
}
@@ -0,0 +1,34 @@
package providers
import (
"database/sql/driver"
"testing"
"time"
"github.com/jackc/pgx/v5"
)
func TestConnectorGenerationInvalidatesConns(t *testing.T) {
cfg, err := pgx.ParseConfig("postgres://u:p@127.0.0.1:1/db")
if err != nil {
t.Fatal(err)
}
c := newPGConnector(cfg)
conn := &pgConn{owner: c, gen: c.generation.Load()}
if conn.stale() {
t.Fatal("fresh connection reported stale")
}
c.swap(cfg)
if !conn.stale() {
t.Fatal("connection from an older generation must be stale")
}
if err := conn.ResetSession(t.Context()); err != driver.ErrBadConn {
t.Fatalf("ResetSession on stale conn = %v, want ErrBadConn", err)
}
}
func TestDialFuncHasTimeout(t *testing.T) {
if newDialFunc(2*time.Second) == nil {
t.Fatal("nil dial func")
}
}
+64 -48
View File
@@ -3,12 +3,12 @@ package providers
import (
"context"
"database/sql"
"errors"
"fmt"
"math"
"sync"
"time"
_ "github.com/jackc/pgx/v5/stdlib" // PostgreSQL driver
"go.mongodb.org/mongo-driver/mongo"
"github.com/bitechdev/ResolveSpec/pkg/logger"
@@ -16,10 +16,11 @@ import (
// PostgresProvider implements Provider for PostgreSQL databases
type PostgresProvider struct {
db *sql.DB
config ConnectionConfig
listener *PostgresListener
mu sync.Mutex
db *sql.DB
connector *pgConnector
config ConnectionConfig
listener *PostgresListener
mu sync.Mutex
}
// NewPostgresProvider creates a new PostgreSQL provider
@@ -29,22 +30,24 @@ func NewPostgresProvider() *PostgresProvider {
// Connect establishes a PostgreSQL connection
func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
// Build DSN
dsn, err := cfg.BuildDSN()
connCfg, err := buildPGXConfig(cfg)
if err != nil {
return fmt.Errorf("failed to build DSN: %w", err)
return err
}
// The connector and *sql.DB are created once; the pool is never closed to
// recover from errors (see Refresh).
connector := newPGConnector(connCfg)
db := sql.OpenDB(connector)
// Connect with retry logic
var db *sql.DB
var lastErr error
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
retryAttempts := 3 // Default retry attempts
retryDelay := 1 * time.Second
connected := false
for attempt := 0; attempt < retryAttempts; attempt++ {
if attempt > 0 {
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
if cfg.GetEnableLogging() {
logger.Info("Retrying PostgreSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
}
@@ -52,20 +55,11 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
select {
case <-time.After(delay):
case <-ctx.Done():
db.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
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() //nolint:gosec // G104: best-effort call, error intentionally ignored
return fmt.Errorf("failed to connect after %d attempts: %w", retryAttempts, lastErr)
}
@@ -103,6 +97,7 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
}
p.db = db
p.connector = connector
p.config = cfg
if cfg.GetEnableLogging() {
@@ -112,34 +107,55 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
return nil
}
// Close closes the PostgreSQL connection
func (p *PostgresProvider) Close() error {
// Close listener if it exists
p.mu.Lock()
if p.listener != nil {
if err := p.listener.Close(); err != nil {
p.mu.Unlock()
return fmt.Errorf("failed to close listener: %w", err)
}
p.listener = nil
// Refresh retires every pooled connection and dials fresh ones on demand,
// without closing the *sql.DB. Handles already handed out keep working:
// connections in use finish their current query and are then discarded.
func (p *PostgresProvider) Refresh(ctx context.Context) error {
if p.db == nil || p.connector == nil {
return fmt.Errorf("database connection is not initialized")
}
connCfg, err := buildPGXConfig(p.config)
if err != nil {
return err
}
p.connector.swap(connCfg)
pingCtx, cancel := context.WithTimeout(ctx, p.config.GetConnectTimeout())
defer cancel()
if err := p.db.PingContext(pingCtx); err != nil {
return fmt.Errorf("failed to ping after refresh: %w", err)
}
return nil
}
// Close closes the PostgreSQL connection. A listener failure does not stop the
// pool from being closed.
func (p *PostgresProvider) Close() error {
var errs []error
p.mu.Lock()
listener := p.listener
p.listener = nil
p.mu.Unlock()
if p.db == nil {
return nil
if listener != nil {
if err := listener.Close(); err != nil {
errs = append(errs, fmt.Errorf("failed to close listener: %w", err))
}
}
err := p.db.Close()
if err != nil {
return fmt.Errorf("failed to close PostgreSQL connection: %w", err)
if p.db != nil {
if err := p.db.Close(); err != nil {
errs = append(errs, fmt.Errorf("failed to close PostgreSQL connection: %w", err))
} else if p.config.GetEnableLogging() {
logger.Info("PostgreSQL connection closed: name=%s", p.config.GetName())
}
p.db = nil
p.connector = nil
}
if p.config.GetEnableLogging() {
logger.Info("PostgreSQL connection closed: name=%s", p.config.GetName())
}
p.db = nil
return nil
return errors.Join(errs...)
}
// HealthCheck verifies the PostgreSQL connection is alive
+217 -159
View File
@@ -23,6 +23,10 @@ type PostgresListener struct {
// Channel subscriptions
channels map[string]NotificationHandler
mu sync.RWMutex
// connMu serialises use of the single pgx.Conn: it is not safe for
// concurrent use, so the notification wait and LISTEN/UNLISTEN/NOTIFY take
// turns. Lock order: connMu before mu.
connMu sync.Mutex
// Lifecycle management
ctx context.Context
@@ -30,6 +34,7 @@ type PostgresListener struct {
closed bool
closeMu sync.Mutex
reconnectC chan struct{}
startOnce sync.Once // background goroutines start exactly once
}
// NewPostgresListener creates a new PostgreSQL listener
@@ -44,76 +49,20 @@ func NewPostgresListener(cfg ConnectionConfig) *PostgresListener {
}
}
// Connect establishes a dedicated connection for listening
// Connect establishes a dedicated connection for listening and starts the
// background loops (once per listener).
func (l *PostgresListener) Connect(ctx context.Context) error {
dsn, err := l.config.BuildDSN()
conn, err := l.dial(ctx)
if err != nil {
return fmt.Errorf("failed to build DSN: %w", err)
return err
}
// Parse connection config
connConfig, err := pgx.ParseConfig(dsn)
if err != nil {
return fmt.Errorf("failed to parse connection config: %w", err)
}
l.swapConn(conn)
// Connect with retry logic
var conn *pgx.Conn
var lastErr error
retryAttempts := 3
retryDelay := 1 * time.Second
for attempt := 0; attempt < retryAttempts; attempt++ {
if attempt > 0 {
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
if l.config.GetEnableLogging() {
logger.Info("Retrying PostgreSQL listener connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
}
select {
case <-time.After(delay):
case <-ctx.Done():
return ctx.Err()
}
}
conn, err = pgx.ConnectConfig(ctx, connConfig)
if err != nil {
lastErr = err
if l.config.GetEnableLogging() {
logger.Warn("Failed to connect PostgreSQL listener", "error", err)
}
continue
}
// Test the connection
if err = conn.Ping(ctx); err != nil {
lastErr = err
conn.Close(ctx)
if l.config.GetEnableLogging() {
logger.Warn("Failed to ping PostgreSQL listener", "error", err)
}
continue
}
// Connection successful
break
}
if err != nil {
return fmt.Errorf("failed to connect listener after %d attempts: %w", retryAttempts, lastErr)
}
l.mu.Lock()
l.conn = conn
l.mu.Unlock()
// Start notification handler
go l.handleNotifications()
// Start reconnection handler
go l.handleReconnection()
l.startOnce.Do(func() {
go l.handleNotifications()
go l.handleReconnection()
})
if l.config.GetEnableLogging() {
logger.Info("PostgreSQL listener connected: name=%s", l.config.GetName())
@@ -122,30 +71,105 @@ func (l *PostgresListener) Connect(ctx context.Context) error {
return nil
}
// dial opens and verifies a new dedicated connection, with retries.
func (l *PostgresListener) dial(ctx context.Context) (*pgx.Conn, error) {
connConfig, err := buildPGXConfig(l.config)
if err != nil {
return nil, err
}
var lastErr error
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(l.config)
for attempt := 0; attempt < retryAttempts; attempt++ {
if attempt > 0 {
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
if l.config.GetEnableLogging() {
logger.Info("Retrying PostgreSQL listener connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
}
select {
case <-time.After(delay):
case <-ctx.Done():
return nil, ctx.Err()
}
}
conn, err := pgx.ConnectConfig(ctx, connConfig)
if err != nil {
lastErr = err
if l.config.GetEnableLogging() {
logger.Warn("Failed to connect PostgreSQL listener: %v", err)
}
continue
}
// Test the connection
if err = conn.Ping(ctx); err != nil {
lastErr = err
_ = closeConnBounded(conn)
if l.config.GetEnableLogging() {
logger.Warn("Failed to ping PostgreSQL listener: %v", err)
}
continue
}
return conn, nil
}
return nil, fmt.Errorf("failed to connect listener after %d attempts: %w", retryAttempts, lastErr)
}
// closeConnBounded closes a pgx connection without ever waiting on a dead socket.
func closeConnBounded(conn *pgx.Conn) error {
ctx, cancel := context.WithTimeout(context.Background(), listenerCloseTimeout)
defer cancel()
return conn.Close(ctx)
}
const (
listenerCloseTimeout = 2 * time.Second
notificationPollInterval = 500 * time.Millisecond
)
// currentConn returns the live connection, or an error if the listener is
// closed or not yet connected.
func (l *PostgresListener) currentConn() (*pgx.Conn, error) {
l.closeMu.Lock()
closed := l.closed
l.closeMu.Unlock()
if closed {
return nil, fmt.Errorf("listener is closed")
}
l.mu.RLock()
conn := l.conn
l.mu.RUnlock()
if conn == nil {
return nil, fmt.Errorf("listener connection is not initialized")
}
return conn, nil
}
// Listen subscribes to a PostgreSQL notification channel
func (l *PostgresListener) Listen(channel string, handler NotificationHandler) error {
l.closeMu.Lock()
if l.closed {
l.closeMu.Unlock()
return fmt.Errorf("listener is closed")
}
l.closeMu.Unlock()
// Take the connection between notification waits (each wait is short).
l.connMu.Lock()
defer l.connMu.Unlock()
l.mu.Lock()
defer l.mu.Unlock()
if l.conn == nil {
return fmt.Errorf("listener connection is not initialized")
}
// Execute LISTEN command
_, err := l.conn.Exec(l.ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{channel}.Sanitize()))
conn, err := l.currentConn()
if err != nil {
return err
}
if _, err := conn.Exec(l.ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{channel}.Sanitize())); err != nil {
return fmt.Errorf("failed to listen on channel %s: %w", channel, err)
}
// Store the handler
l.mu.Lock()
l.channels[channel] = handler
l.mu.Unlock()
if l.config.GetEnableLogging() {
logger.Info("Listening on channel: name=%s, channel=%s", l.config.GetName(), channel)
@@ -156,28 +180,21 @@ func (l *PostgresListener) Listen(channel string, handler NotificationHandler) e
// Unlisten unsubscribes from a PostgreSQL notification channel
func (l *PostgresListener) Unlisten(channel string) error {
l.closeMu.Lock()
if l.closed {
l.closeMu.Unlock()
return fmt.Errorf("listener is closed")
}
l.closeMu.Unlock()
l.connMu.Lock()
defer l.connMu.Unlock()
l.mu.Lock()
defer l.mu.Unlock()
if l.conn == nil {
return fmt.Errorf("listener connection is not initialized")
}
// Execute UNLISTEN command
_, err := l.conn.Exec(l.ctx, fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize()))
conn, err := l.currentConn()
if err != nil {
return err
}
if _, err := conn.Exec(l.ctx, fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize())); err != nil {
return fmt.Errorf("failed to unlisten from channel %s: %w", channel, err)
}
// Remove the handler
l.mu.Lock()
delete(l.channels, channel)
l.mu.Unlock()
if l.config.GetEnableLogging() {
logger.Info("Unlistened from channel: name=%s, channel=%s", l.config.GetName(), channel)
@@ -188,31 +205,24 @@ func (l *PostgresListener) Unlisten(channel string) error {
// Notify sends a notification to a PostgreSQL channel
func (l *PostgresListener) Notify(ctx context.Context, channel string, payload string) error {
l.closeMu.Lock()
if l.closed {
l.closeMu.Unlock()
return fmt.Errorf("listener is closed")
}
l.closeMu.Unlock()
l.connMu.Lock()
defer l.connMu.Unlock()
l.mu.RLock()
conn := l.conn
l.mu.RUnlock()
if conn == nil {
return fmt.Errorf("listener connection is not initialized")
}
// Execute NOTIFY command
_, err := conn.Exec(ctx, "SELECT pg_notify($1, $2)", channel, payload)
conn, err := l.currentConn()
if err != nil {
return err
}
if _, err := conn.Exec(ctx, "SELECT pg_notify($1, $2)", channel, payload); err != nil {
return fmt.Errorf("failed to notify channel %s: %w", channel, err)
}
return nil
}
// Close closes the listener and all subscriptions
// Close closes the listener and all subscriptions. Closing the connection drops
// every subscription server-side, so no UNLISTEN round trips are needed, and
// the close itself is bounded so a dead socket cannot hang the caller.
func (l *PostgresListener) Close() error {
l.closeMu.Lock()
if l.closed {
@@ -225,27 +235,26 @@ func (l *PostgresListener) Close() error {
// Cancel context to stop background goroutines
l.cancel()
// The cancelled ctx makes the notification wait return promptly, releasing
// connMu; closing the conn while it is being read would race inside pgx.
l.connMu.Lock()
l.mu.Lock()
defer l.mu.Unlock()
conn := l.conn
l.conn = nil
l.channels = make(map[string]NotificationHandler)
l.mu.Unlock()
if l.conn == nil {
if conn == nil {
l.connMu.Unlock()
return nil
}
// Unlisten from all channels
for channel := range l.channels {
_, _ = l.conn.Exec(context.Background(), fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize()))
}
// Close connection
err := l.conn.Close(context.Background())
err := closeConnBounded(conn)
l.connMu.Unlock()
if err != nil {
return fmt.Errorf("failed to close listener connection: %w", err)
}
l.conn = nil
l.channels = make(map[string]NotificationHandler)
if l.config.GetEnableLogging() {
logger.Info("PostgreSQL listener closed: name=%s", l.config.GetName())
}
@@ -262,20 +271,26 @@ func (l *PostgresListener) handleNotifications() {
default:
}
l.connMu.Lock()
l.mu.RLock()
conn := l.conn
l.mu.RUnlock()
if conn == nil {
l.connMu.Unlock()
// Connection not available, wait for reconnection
time.Sleep(100 * time.Millisecond)
if !l.sleep(100 * time.Millisecond) {
return
}
continue
}
// Wait for notification with timeout
ctx, cancel := context.WithTimeout(l.ctx, 5*time.Second)
// Wait for a notification with a short timeout, so Listen/Unlisten/Notify
// waiting on connMu are served promptly.
ctx, cancel := context.WithTimeout(l.ctx, notificationPollInterval)
notification, err := conn.WaitForNotification(ctx)
cancel()
l.connMu.Unlock()
if err != nil {
// Check if context was cancelled
@@ -291,13 +306,15 @@ func (l *PostgresListener) handleNotifications() {
// Connection error, trigger reconnection
if l.config.GetEnableLogging() {
logger.Warn("Notification error, triggering reconnection", "error", err)
logger.Warn("Notification error, triggering reconnection: %v", err)
}
select {
case l.reconnectC <- struct{}{}:
default:
}
time.Sleep(1 * time.Second)
if !l.sleep(1 * time.Second) {
return
}
continue
}
@@ -322,7 +339,22 @@ func (l *PostgresListener) handleNotifications() {
}
}
// handleReconnection manages automatic reconnection
// sleep waits for d or until the listener is closed; it reports whether the
// listener is still running.
func (l *PostgresListener) sleep(d time.Duration) bool {
t := time.NewTimer(d)
defer t.Stop()
select {
case <-t.C:
return true
case <-l.ctx.Done():
return false
}
}
// handleReconnection manages automatic reconnection. It runs as a single
// goroutine and dials replacement connections directly rather than through the
// public Connect, so no extra loops are started.
func (l *PostgresListener) handleReconnection() {
for {
select {
@@ -333,31 +365,21 @@ func (l *PostgresListener) handleReconnection() {
logger.Info("Attempting to reconnect listener: name=%s", l.config.GetName())
}
// Close existing connection
l.mu.Lock()
if l.conn != nil {
l.conn.Close(context.Background())
l.conn = nil
}
// Save current subscriptions
channels := make(map[string]NotificationHandler)
for ch, handler := range l.channels {
channels[ch] = handler
}
l.mu.Unlock()
// Attempt reconnection
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
err := l.Connect(ctx)
ctx, cancel := context.WithTimeout(l.ctx, 30*time.Second)
err := l.reconnect(ctx)
cancel()
if err != nil {
if l.ctx.Err() != nil {
return
}
if l.config.GetEnableLogging() {
logger.Error("Failed to reconnect listener: name=%s, error=%v", l.config.GetName(), err)
}
// Retry after delay
time.Sleep(5 * time.Second)
if !l.sleep(5 * time.Second) {
return
}
select {
case l.reconnectC <- struct{}{}:
default:
@@ -365,15 +387,6 @@ func (l *PostgresListener) handleReconnection() {
continue
}
// Resubscribe to all channels
for channel, handler := range channels {
if err := l.Listen(channel, handler); err != nil {
if l.config.GetEnableLogging() {
logger.Error("Failed to resubscribe to channel: name=%s, channel=%s, error=%v", l.config.GetName(), channel, err)
}
}
}
if l.config.GetEnableLogging() {
logger.Info("Listener reconnected successfully: name=%s", l.config.GetName())
}
@@ -381,6 +394,51 @@ func (l *PostgresListener) handleReconnection() {
}
}
// reconnect replaces the connection and resubscribes every channel on the new
// connection before publishing it, so the notification loop never touches a
// half-initialised conn.
func (l *PostgresListener) reconnect(ctx context.Context) error {
conn, err := l.dial(ctx)
if err != nil {
return err
}
l.mu.RLock()
channels := make([]string, 0, len(l.channels))
for ch := range l.channels {
channels = append(channels, ch)
}
l.mu.RUnlock()
for _, ch := range channels {
if _, err := conn.Exec(ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{ch}.Sanitize())); err != nil {
_ = closeConnBounded(conn)
return fmt.Errorf("failed to resubscribe to channel %s: %w", ch, err)
}
}
if l.ctx.Err() != nil {
_ = closeConnBounded(conn)
return l.ctx.Err()
}
l.swapConn(conn)
return nil
}
// swapConn installs conn and closes the previous one. The old connection is
// closed under connMu so it is never closed while another goroutine is using it.
func (l *PostgresListener) swapConn(conn *pgx.Conn) {
l.connMu.Lock()
l.mu.Lock()
old := l.conn
l.conn = conn
l.mu.Unlock()
if old != nil {
_ = closeConnBounded(old)
}
l.connMu.Unlock()
}
// IsConnected returns true if the listener is connected
func (l *PostgresListener) IsConnected() bool {
l.mu.RLock()
+25 -6
View File
@@ -4,17 +4,11 @@ import (
"context"
"database/sql"
"errors"
"strings"
"time"
"go.mongodb.org/mongo-driver/mongo"
)
// isDBClosed reports whether err indicates the *sql.DB has been closed.
func isDBClosed(err error) bool {
return err != nil && strings.Contains(err.Error(), "sql: database is closed")
}
// Common errors
var (
// ErrNotSQLDatabase is returned when attempting SQL operations on a non-SQL database
@@ -63,6 +57,31 @@ type ConnectionConfig interface {
GetConnMaxLifetime() *time.Duration
GetConnMaxIdleTime() *time.Duration
GetReadPreference() string
GetRetryAttempts() int
GetRetryDelay() time.Duration
GetRetryMaxDelay() time.Duration
}
// retryPolicy returns the configured retry settings, falling back to defaults.
func retryPolicy(cfg ConnectionConfig) (attempts int, delay, maxDelay time.Duration) {
attempts, delay, maxDelay = cfg.GetRetryAttempts(), cfg.GetRetryDelay(), cfg.GetRetryMaxDelay()
if attempts <= 0 {
attempts = 3
}
if delay <= 0 {
delay = time.Second
}
if maxDelay <= 0 {
maxDelay = 10 * time.Second
}
return
}
// Refresher is implemented by providers that can retire their pooled
// connections and dial fresh ones without closing the shared *sql.DB, so
// handles already handed out keep working.
type Refresher interface {
Refresh(ctx context.Context) error
}
// Provider creates and manages the underlying database connection
+45 -69
View File
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"fmt"
"strings"
"sync"
"time"
@@ -15,10 +16,9 @@ import (
// SQLiteProvider implements Provider for SQLite databases
type SQLiteProvider struct {
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
config ConnectionConfig
db *sql.DB
dbMu sync.RWMutex
config ConnectionConfig
}
// NewSQLiteProvider creates a new SQLite provider
@@ -26,6 +26,22 @@ func NewSQLiteProvider() *SQLiteProvider {
return &SQLiteProvider{}
}
// isMemoryDSN reports whether the SQLite DSN refers to a private in-memory
// database (each pooled connection would get its own empty database).
func isMemoryDSN(dsn string) bool {
path := dsn
if i := strings.IndexByte(path, '?'); i >= 0 {
path = path[:i]
}
if path == ":memory:" || path == "" {
return true
}
if strings.Contains(dsn, "mode=memory") && !strings.Contains(dsn, "cache=shared") {
return true
}
return path == "file::memory:" && !strings.Contains(dsn, "cache=shared")
}
// Connect establishes a SQLite connection
func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
// Build DSN
@@ -46,52 +62,39 @@ func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) erro
cancel()
if err != nil {
db.Close()
db.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
return fmt.Errorf("failed to ping SQLite database: %w", err)
}
// Configure connection pool
// Note: SQLite works best with MaxOpenConns=1 for write operations
// but can handle multiple readers
if cfg.GetMaxOpenConns() != nil {
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
} else {
// Default to 1 for SQLite to avoid "database is locked" errors
if isMemoryDSN(dsn) {
// A private in-memory database exists per connection and disappears when
// that connection closes, so pin the pool to one connection that is
// never recycled.
db.SetMaxOpenConns(1)
}
if cfg.GetMaxIdleConns() != nil {
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
}
if cfg.GetConnMaxLifetime() != nil {
db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime())
}
if cfg.GetConnMaxIdleTime() != nil {
db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime())
}
// Enable WAL mode for better concurrent access
_, err = db.ExecContext(ctx, "PRAGMA journal_mode=WAL")
if err != nil {
if cfg.GetEnableLogging() {
logger.Warn("Failed to enable WAL mode for SQLite", "error", err)
db.SetMaxIdleConns(1)
db.SetConnMaxLifetime(0)
db.SetConnMaxIdleTime(0)
} else {
// SQLite works best with few writers; default to 1 unless configured.
if cfg.GetMaxOpenConns() != nil {
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
} else {
db.SetMaxOpenConns(1)
}
// Don't fail connection if WAL mode cannot be enabled
}
// Set busy timeout to handle locked database (minimum 2 minutes = 120000ms)
busyTimeout := cfg.GetQueryTimeout().Milliseconds()
if busyTimeout < 120000 {
busyTimeout = 120000 // Enforce minimum of 2 minutes
}
_, err = db.ExecContext(ctx, fmt.Sprintf("PRAGMA busy_timeout=%d", busyTimeout))
if err != nil {
if cfg.GetEnableLogging() {
logger.Warn("Failed to set busy timeout for SQLite", "error", err)
if cfg.GetMaxIdleConns() != nil {
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
}
if cfg.GetConnMaxLifetime() != nil {
db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime())
}
if cfg.GetConnMaxIdleTime() != nil {
db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime())
}
}
p.dbMu.Lock()
p.db = db
p.dbMu.Unlock()
p.config = cfg
if cfg.GetEnableLogging() {
@@ -132,14 +135,7 @@ func (p *SQLiteProvider) HealthCheck(ctx context.Context) error {
// Execute a simple query to verify the database is accessible
var result int
run := func() error { return p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result) }
err := run()
if isDBClosed(err) {
if reconnErr := p.reconnectDB(); reconnErr == nil {
err = run()
}
}
if err != nil {
if err := p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result); err != nil {
return fmt.Errorf("health check failed: %w", err)
}
@@ -150,32 +146,12 @@ func (p *SQLiteProvider) HealthCheck(ctx context.Context) error {
return nil
}
// WithDBFactory configures a factory used to reopen the database connection if it is closed.
func (p *SQLiteProvider) WithDBFactory(factory func() (*sql.DB, error)) *SQLiteProvider {
p.dbFactory = factory
return p
}
func (p *SQLiteProvider) getDB() *sql.DB {
p.dbMu.RLock()
defer p.dbMu.RUnlock()
return p.db
}
func (p *SQLiteProvider) reconnectDB() error {
if p.dbFactory == nil {
return fmt.Errorf("no db factory configured for reconnect")
}
newDB, err := p.dbFactory()
if err != nil {
return err
}
p.dbMu.Lock()
p.db = newDB
p.dbMu.Unlock()
return nil
}
// GetNative returns the native *sql.DB connection
func (p *SQLiteProvider) GetNative() (*sql.DB, error) {
if p.db == nil {
@@ -0,0 +1,24 @@
//go:build linux
package providers
import (
"syscall"
"time"
"golang.org/x/sys/unix"
)
// setTCPUserTimeout returns a net.Dialer Control func setting TCP_USER_TIMEOUT.
func setTCPUserTimeout(d time.Duration) func(network, address string, c syscall.RawConn) error {
return func(network, address string, c syscall.RawConn) error {
var sockErr error
err := c.Control(func(fd uintptr) {
sockErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_TCP, unix.TCP_USER_TIMEOUT, int(d.Milliseconds()))
})
if err != nil {
return err
}
return sockErr
}
}
@@ -0,0 +1,14 @@
//go:build !linux
package providers
import (
"syscall"
"time"
)
// setTCPUserTimeout is a no-op where TCP_USER_TIMEOUT is unavailable; TCP
// keepalive still applies.
func setTCPUserTimeout(time.Duration) func(network, address string, c syscall.RawConn) error {
return nil
}
+69
View File
@@ -0,0 +1,69 @@
package dbmanager
import (
"context"
"os"
"os/exec"
"testing"
"time"
)
func TestLiveServerRestart(t *testing.T) {
dir := os.Getenv("PG_RESTART_DIR")
if dir == "" {
t.Skip("PG_RESTART_DIR not set")
}
mgr, _ := NewManager(ManagerConfig{
DefaultConnection: "pg",
Connections: map[string]ConnectionConfig{"pg": {
Name: "pg", Type: DatabaseTypePostgreSQL, Host: "127.0.0.1", Port: 54329,
User: "postgres", Database: "postgres", ConnectTimeout: 2 * time.Second,
}},
HealthCheckInterval: 500 * time.Millisecond,
})
ctx := context.Background()
if err := mgr.Connect(ctx); err != nil {
t.Fatal(err)
}
defer mgr.Close()
conn, _ := mgr.GetDefault()
db, _ := conn.Bun()
gdb, _ := conn.GORM()
query := func() error { var n int; return db.DB.QueryRow("select 1").Scan(&n) }
if err := query(); err != nil {
t.Fatal(err)
}
run := func(args ...string) {
// Output must not be piped: the daemonised server would hold the pipe open.
if err := exec.Command("pg_ctl", append([]string{"-D", dir}, args...)...).Run(); err != nil {
t.Fatalf("pg_ctl %v: %v", args, err)
}
}
run("-m", "immediate", "-w", "stop") // crash-style shutdown
time.Sleep(1500 * time.Millisecond) // health checks fail meanwhile
if err := query(); err == nil {
t.Fatal("expected failure while server is down")
}
run("-l", dir+"/restart.log", "-o", "-p 54329 -k "+dir+" -c listen_addresses=127.0.0.1", "-w", "start")
var last error
for i := 0; i < 20; i++ {
if last = query(); last == nil {
break
}
t.Logf("attempt %d after restart: %v", i, last)
time.Sleep(200 * time.Millisecond)
}
if last != nil {
t.Fatalf("held bun handle never recovered: %v", last)
}
var n int
if err := gdb.Raw("select 1").Scan(&n).Error; err != nil {
t.Fatalf("held gorm handle: %v", err)
}
time.Sleep(time.Second)
if err := conn.HealthCheck(ctx); err != nil {
t.Fatalf("health check after restart: %v", err)
}
}
+8 -8
View File
@@ -171,7 +171,7 @@ func TestBrokerPublishAsync(t *testing.T) {
// Publish multiple events
for i := 0; i < 5; i++ {
event := NewEvent(EventSourceSystem, "test.event")
event.InstanceID = "test-instance"
event.InstanceID = "test-instance"
if err := broker.PublishAsync(context.Background(), event); err != nil {
t.Fatalf("PublishAsync failed: %v", err)
}
@@ -346,7 +346,7 @@ func TestBrokerStats(t *testing.T) {
// Publish events
for i := 0; i < 3; i++ {
event := NewEvent(EventSourceSystem, "test.event")
event.InstanceID = "test-instance"
event.InstanceID = "test-instance"
broker.PublishSync(context.Background(), event)
}
@@ -413,7 +413,7 @@ func TestBrokerConcurrentPublish(t *testing.T) {
go func() {
defer wg.Done()
event := NewEvent(EventSourceSystem, "test.event")
event.InstanceID = "test-instance"
event.InstanceID = "test-instance"
broker.PublishAsync(context.Background(), event)
}()
}
@@ -450,7 +450,7 @@ func TestBrokerGracefulShutdown(t *testing.T) {
// Publish events
for i := 0; i < 5; i++ {
event := NewEvent(EventSourceSystem, "test.event")
event.InstanceID = "test-instance"
event.InstanceID = "test-instance"
broker.PublishAsync(context.Background(), event)
}
@@ -502,21 +502,21 @@ func TestBrokerProcessingModes(t *testing.T) {
broker.Start(context.Background())
defer broker.Stop(context.Background())
called := false
var called atomic.Bool
broker.Subscribe("test.*", EventHandlerFunc(func(ctx context.Context, event *Event) error {
called = true
called.Store(true)
return nil
}))
event := NewEvent(EventSourceSystem, "test.event")
event.InstanceID = "test-instance"
event.InstanceID = "test-instance"
broker.Publish(context.Background(), event)
if tt.mode == ProcessingModeAsync {
time.Sleep(50 * time.Millisecond)
}
if !called {
if !called.Load() {
t.Error("Expected handler to be called")
}
})
+2 -2
View File
@@ -584,7 +584,7 @@ func (dp *DatabaseProvider) pollEvents() {
dp.stats.EventsConsumed.Add(1)
sub.lastSeenID = event.ID
case <-sub.ctx.Done():
rows.Close()
rows.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
return
default:
// Channel full, skip
@@ -595,7 +595,7 @@ func (dp *DatabaseProvider) pollEvents() {
sub.lastSeenID = event.ID
}
rows.Close()
rows.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
}
}
+221 -57
View File
@@ -5,16 +5,144 @@ import (
"fmt"
"log"
"os"
"regexp"
"runtime"
"runtime/debug"
"strings"
"sync"
"time"
"go.uber.org/zap"
errortracking "github.com/bitechdev/ResolveSpec/pkg/errortracking"
)
// Logger is the active logger. It is kept exported for compatibility, but
// inside this package it must only be accessed through getLogger/setLogger.
var Logger *zap.SugaredLogger
var errorTracker errortracking.Provider
// stateMu guards Logger and errorTracker, which may be replaced while other
// goroutines are logging.
var stateMu sync.RWMutex
func getLogger() *zap.SugaredLogger {
stateMu.RLock()
defer stateMu.RUnlock()
return Logger
}
func swapLogger(l *zap.SugaredLogger) *zap.SugaredLogger {
stateMu.Lock()
defer stateMu.Unlock()
old := Logger
Logger = l
return old
}
func getErrorTracker() errortracking.Provider {
stateMu.RLock()
defer stateMu.RUnlock()
return errorTracker
}
// pid is constant for the life of the process; cache it off the hot path.
var pid = os.Getpid()
// maxStackBytes caps the stack trace captured for a recovered panic.
const maxStackBytes = 16 << 10
func captureStack() []byte {
buf := make([]byte, maxStackBytes)
return buf[:runtime.Stack(buf, false)]
}
// Patterns scrubbed from messages before they leave the process for the
// error tracker. The local log is left untouched.
var redactPatterns = []struct {
re *regexp.Regexp
repl string
}{
{regexp.MustCompile(`([a-zA-Z][a-zA-Z0-9+.\-]*://)[^\s/@:]+:[^\s/@]+@`), "${1}[REDACTED]@"},
{regexp.MustCompile(`(?i)\b(password|passwd|pwd|secret|token|api[_-]?key|access[_-]?key)(\s*[=:]\s*)('[^']*'|"[^"]*"|[^\s,;&]+)`), "${1}${2}[REDACTED]"},
{regexp.MustCompile(`(?i)\b(bearer|basic)\s+[A-Za-z0-9._~+/=\-]+`), "${1} [REDACTED]"},
}
func redact(s string) string {
for _, p := range redactPatterns {
s = p.re.ReplaceAllString(s, p.repl)
}
return s
}
// sanitizeForStdlog escapes CR/LF and other control characters so untrusted
// values cannot forge additional log lines on the stdlib fallback path.
func sanitizeForStdlog(s string) string {
if !strings.ContainsFunc(s, func(r rune) bool { return r < 0x20 || r == 0x7f }) {
return s
}
var b strings.Builder
b.Grow(len(s) + 8)
for _, r := range s {
switch {
case r == '\n':
b.WriteString(`\n`)
case r == '\r':
b.WriteString(`\r`)
case r == '\t':
b.WriteString(`\t`)
case r < 0x20 || r == 0x7f:
fmt.Fprintf(&b, `\x%02x`, r)
default:
b.WriteRune(r)
}
}
return b.String()
}
// Error tracker fan-out limiting: a global token bucket plus per-template
// dedup, so attacker-triggerable errors cannot burn quota or flood the queue.
const (
trackerBurst = 50
trackerRefillPerSec = 20.0
trackerDedupWindow = time.Second
trackerMaxKeys = 1024
)
var limiter = struct {
sync.Mutex
tokens float64
last time.Time
seen map[string]time.Time
}{tokens: trackerBurst, seen: map[string]time.Time{}}
func allowTracker(key string) bool {
now := time.Now()
limiter.Lock()
defer limiter.Unlock()
if !limiter.last.IsZero() {
limiter.tokens += now.Sub(limiter.last).Seconds() * trackerRefillPerSec
if limiter.tokens > trackerBurst {
limiter.tokens = trackerBurst
}
}
limiter.last = now
if t, ok := limiter.seen[key]; ok && now.Sub(t) < trackerDedupWindow {
return false
}
if limiter.tokens < 1 {
return false
}
if len(limiter.seen) >= trackerMaxKeys {
limiter.seen = map[string]time.Time{}
}
limiter.seen[key] = now
limiter.tokens--
return true
}
func Init(dev bool) {
if dev {
@@ -36,7 +164,16 @@ func UpdateLoggerPath(path string, dev bool) {
UpdateLogger(&defaultConfig)
}
// UpdateLogger rebuilds the logger from config. On failure the previous logger
// stays in place; use UpdateLoggerE to get the error.
func UpdateLogger(config *zap.Config) {
if err := UpdateLoggerE(config); err != nil {
log.Printf("logger: failed to build logger, keeping previous: %s", sanitizeForStdlog(err.Error()))
}
}
// UpdateLoggerE is UpdateLogger but returns the build error.
func UpdateLoggerE(config *zap.Config) error {
defaultConfig := zap.NewProductionConfig()
defaultConfig.OutputPaths = []string{"resolvespec.log"}
if config == nil {
@@ -45,32 +182,45 @@ func UpdateLogger(config *zap.Config) {
logger, err := config.Build()
if err != nil {
log.Print(err)
return
return err
}
Logger = logger.Sugar()
old := swapLogger(logger.Sugar())
if old != nil {
_ = old.Sync()
}
Info("ResolveSpec Logger initialized")
return nil
}
// Sync flushes buffered log entries. Call it on shutdown.
func Sync() error {
if lg := getLogger(); lg != nil {
return lg.Sync()
}
return nil
}
// InitErrorTracking initializes the error tracking provider
func InitErrorTracking(provider errortracking.Provider) {
stateMu.Lock()
errorTracker = provider
if errorTracker != nil {
stateMu.Unlock()
if provider != nil {
Info("Error tracking initialized")
}
}
// GetErrorTracker returns the current error tracking provider
func GetErrorTracker() errortracking.Provider {
return errorTracker
return getErrorTracker()
}
// CloseErrorTracking flushes and closes the error tracking provider
func CloseErrorTracking() error {
if errorTracker != nil {
errorTracker.Flush(5)
return errorTracker.Close()
if tracker := getErrorTracker(); tracker != nil {
tracker.Flush(5)
return tracker.Close()
}
return nil
}
@@ -98,53 +248,51 @@ func extractContext(args ...interface{}) (ctx context.Context, filteredArgs []in
}
func Info(template string, args ...interface{}) {
if Logger == nil {
log.Printf(template, args...)
_, args = extractContext(args...)
message := fmt.Sprintf(template, args...)
if lg := getLogger(); lg != nil {
lg.Infow(message, "process_id", pid)
return
}
Logger.Infow(fmt.Sprintf(template, args...), "process_id", os.Getpid())
}
func Warn(template string, args ...interface{}) {
ctx, remainingArgs := extractContext(args...)
message := fmt.Sprintf(template, remainingArgs...)
if Logger == nil {
log.Printf("%s", message)
} else {
Logger.Warnw(message, "process_id", os.Getpid())
}
// Send to error tracker
if errorTracker != nil {
errorTracker.CaptureMessage(ctx, message, errortracking.SeverityWarning, map[string]interface{}{
"process_id": os.Getpid(),
})
}
}
func Error(template string, args ...interface{}) {
ctx, remainingArgs := extractContext(args...)
message := fmt.Sprintf(template, remainingArgs...)
if Logger == nil {
log.Printf("%s", message)
} else {
Logger.Errorw(message, "process_id", os.Getpid())
}
// Send to error tracker
if errorTracker != nil {
errorTracker.CaptureMessage(ctx, message, errortracking.SeverityError, map[string]interface{}{
"process_id": os.Getpid(),
})
}
log.Printf("%s", sanitizeForStdlog(message))
}
func Debug(template string, args ...interface{}) {
if Logger == nil {
log.Printf(template, args...)
_, args = extractContext(args...)
message := fmt.Sprintf(template, args...)
if lg := getLogger(); lg != nil {
lg.Debugw(message, "process_id", pid)
return
}
Logger.Debugw(fmt.Sprintf(template, args...), "process_id", os.Getpid())
log.Printf("%s", sanitizeForStdlog(message))
}
func Warn(template string, args ...interface{}) {
logAndTrack(errortracking.SeverityWarning, template, args)
}
func Error(template string, args ...interface{}) {
logAndTrack(errortracking.SeverityError, template, args)
}
func logAndTrack(sev errortracking.Severity, template string, args []interface{}) {
ctx, remainingArgs := extractContext(args...)
message := fmt.Sprintf(template, remainingArgs...)
if lg := getLogger(); lg == nil {
log.Printf("%s", sanitizeForStdlog(message))
} else if sev == errortracking.SeverityWarning {
lg.Warnw(message, "process_id", pid)
} else {
lg.Errorw(message, "process_id", pid)
}
tracker := getErrorTracker()
if tracker == nil || !allowTracker(string(sev)+"|"+template) {
return
}
tracker.CaptureMessage(ctx, redact(message), sev, map[string]interface{}{
"process_id": pid,
})
}
// CatchPanic - Handle panic
@@ -154,9 +302,11 @@ func CatchPanicCallback(location string, cb func(err any), args ...interface{})
ctx, _ := extractContext(args...)
return func() {
if err := recover(); err != nil {
callstack := debug.Stack()
callstack := captureStack()
lg := getLogger()
tracker := getErrorTracker()
if Logger != nil {
if lg != nil {
Error("Panic in %s : %v", location, err, ctx) // Pass context implicitly
} else {
fmt.Printf("%s:PANIC->%+v", location, err)
@@ -164,10 +314,10 @@ func CatchPanicCallback(location string, cb func(err any), args ...interface{})
}
// Send to error tracker
if errorTracker != nil {
errorTracker.CapturePanic(ctx, err, callstack, map[string]interface{}{
if tracker != nil {
tracker.CapturePanic(ctx, err, callstack, map[string]interface{}{
"location": location,
"process_id": os.Getpid(),
"process_id": pid,
})
}
@@ -185,6 +335,19 @@ func CatchPanic(location string, args ...interface{}) func() {
return CatchPanicCallback(location, nil, args...)
}
// CatchPanicRethrow returns a function to defer that logs and reports a panic
// and then re-panics. Use it inside enforcement/internal code where swallowing
// the panic would let the caller proceed as if the work had succeeded. The
// swallowing CatchPanic is for outermost request/goroutine boundaries only.
func CatchPanicRethrow(location string, args ...interface{}) func() {
return func() {
if r := recover(); r != nil {
_ = HandlePanic(location, r, args...)
panic(r)
}
}
}
// HandlePanic logs a panic and returns it as an error
// This should be called with the result of recover() from a deferred function
// Example usage:
@@ -195,15 +358,16 @@ func CatchPanic(location string, args ...interface{}) func() {
// }
// }()
func HandlePanic(methodName string, r any, args ...interface{}) error {
tracker := getErrorTracker()
ctx, _ := extractContext(args...)
stack := debug.Stack()
stack := captureStack()
Error("Panic in %s: %v\nStack trace:\n%s", methodName, r, string(stack), ctx) // Pass context implicitly
// Send to error tracker
if errorTracker != nil {
errorTracker.CapturePanic(ctx, r, stack, map[string]interface{}{
if tracker != nil {
tracker.CapturePanic(ctx, r, stack, map[string]interface{}{
"method": methodName,
"process_id": os.Getpid(),
"process_id": pid,
})
}
+129
View File
@@ -0,0 +1,129 @@
package logger
import (
"bytes"
"context"
"log"
"strings"
"sync"
"testing"
"go.uber.org/zap"
errortracking "github.com/bitechdev/ResolveSpec/pkg/errortracking"
)
type fakeTracker struct {
mu sync.Mutex
msgs []string
}
func (f *fakeTracker) CaptureError(context.Context, error, errortracking.Severity, map[string]interface{}) {
}
func (f *fakeTracker) CaptureMessage(_ context.Context, m string, _ errortracking.Severity, _ map[string]interface{}) {
f.mu.Lock()
f.msgs = append(f.msgs, m)
f.mu.Unlock()
}
func (f *fakeTracker) CapturePanic(context.Context, interface{}, []byte, map[string]interface{}) {}
func (f *fakeTracker) Flush(int) bool { return true }
func (f *fakeTracker) Close() error { return nil }
func TestStdlibFallbackSafe(t *testing.T) {
swapLogger(nil)
var buf bytes.Buffer
log.SetOutput(&buf)
defer log.SetOutput(nil)
Info("100%s done\nFAKE line", "")
Info("%d%% %s", 5, context.Background())
out := buf.String()
if strings.Count(out, "\n") != 2 {
t.Fatalf("log injection: %q", out)
}
}
func TestContextStrippedFromInfo(t *testing.T) {
swapLogger(nil)
var buf bytes.Buffer
log.SetOutput(&buf)
defer log.SetOutput(nil)
Info("saved %s", "widget", context.WithValue(context.Background(), "k", "secret"))
if strings.Contains(buf.String(), "EXTRA") || strings.Contains(buf.String(), "secret") {
t.Fatalf("context leaked: %q", buf.String())
}
}
func TestRedact(t *testing.T) {
in := "connect postgres://admin:hunter2@db:5432/x failed password=abc123 Authorization: Bearer eyJ.abc-d"
out := redact(in)
for _, leak := range []string{"hunter2", "abc123", "eyJ.abc-d"} {
if strings.Contains(out, leak) {
t.Errorf("%q leaked in %q", leak, out)
}
}
}
func TestTrackerRateLimitAndRedaction(t *testing.T) {
swapLogger(zap.NewNop().Sugar())
ft := &fakeTracker{}
InitErrorTracking(ft)
defer InitErrorTracking(nil)
for i := 0; i < 100; i++ {
Error("same template %d password=x", i)
}
if len(ft.msgs) != 1 {
t.Fatalf("dedup failed: %d events", len(ft.msgs))
}
if strings.Contains(ft.msgs[0], "password=x") {
t.Fatalf("not redacted: %q", ft.msgs[0])
}
}
func TestConcurrentUpdateAndLog(t *testing.T) {
InitErrorTracking(&fakeTracker{})
defer InitErrorTracking(nil)
var wg sync.WaitGroup
for i := 0; i < 4; i++ {
wg.Add(2)
go func() {
defer wg.Done()
for j := 0; j < 50; j++ {
cfg := zap.NewProductionConfig()
cfg.OutputPaths = []string{"stderr"}
_ = UpdateLoggerE(&cfg)
}
}()
go func() {
defer wg.Done()
for j := 0; j < 200; j++ {
Error("e %d", j)
_ = CloseErrorTracking()
}
}()
}
wg.Wait()
_ = Sync()
}
func TestCatchPanicRethrow(t *testing.T) {
swapLogger(zap.NewNop().Sugar())
defer func() {
if recover() == nil {
t.Fatal("expected re-panic")
}
}()
func() {
defer CatchPanicRethrow("x")()
panic("boom")
}()
}
func TestCatchPanicSwallows(t *testing.T) {
swapLogger(zap.NewNop().Sugar())
func() {
defer CatchPanic("x")()
panic("boom")
}()
}
+115 -147
View File
@@ -1,10 +1,12 @@
package modelregistry
import (
"errors"
"fmt"
"reflect"
"sync"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// ModelRules defines the permissions and security settings for a model
@@ -52,6 +54,21 @@ var defaultRegistry = &DefaultModelRegistry{
var registries = []*DefaultModelRegistry{defaultRegistry}
var registriesMutex sync.RWMutex
// Sentinel errors so callers (notably the security layer) can distinguish
// "not registered" from every other failure with errors.Is.
var (
ErrModelNotFound = errors.New("model not found")
ErrModelExists = errors.New("model already registered")
ErrInvalidModel = errors.New("invalid model")
)
// maxUnwrapDepth bounds pointer/slice/array unwrapping so a recursive type
// (type T *T) cannot spin forever.
const maxUnwrapDepth = 16
// Lock ordering: registriesMutex is always taken before a registry's mutex,
// never the reverse. No caller-supplied code ever runs while a lock is held.
// NewModelRegistry creates a new model registry
func NewModelRegistry() *DefaultModelRegistry {
return &DefaultModelRegistry{
@@ -60,44 +77,19 @@ func NewModelRegistry() *DefaultModelRegistry {
}
}
// lockRetryAttempts/lockRetryDelay bound how long the try-lock helpers below
// will spin before giving up, so a contended registriesMutex can never hang
// a caller of GetDefaultRegistry/SetDefaultRegistry.
const (
lockRetryAttempts = 20
lockRetryDelay = 1 * time.Millisecond
)
// GetDefaultRegistry returns the current default registry. It uses a
// bounded TryRLock instead of a blocking RLock so it can never hang;
// if the lock can't be acquired in time it falls back to the last known
// value without synchronization.
// GetDefaultRegistry returns the current default registry.
func GetDefaultRegistry() *DefaultModelRegistry {
for i := 0; i < lockRetryAttempts; i++ {
if registriesMutex.TryRLock() {
defer registriesMutex.RUnlock()
return defaultRegistry
}
time.Sleep(lockRetryDelay)
}
registriesMutex.RLock()
defer registriesMutex.RUnlock()
return defaultRegistry
}
// SetDefaultRegistry replaces the default registry. It uses a bounded
// TryLock instead of a blocking Lock so it can never hang; if the lock
// can't be acquired in time the call is a no-op.
// SetDefaultRegistry replaces the default registry. A nil registry is ignored.
func SetDefaultRegistry(registry *DefaultModelRegistry) {
acquired := false
for i := 0; i < lockRetryAttempts; i++ {
if registriesMutex.TryLock() {
acquired = true
break
}
time.Sleep(lockRetryDelay)
}
if !acquired {
if registry == nil {
return
}
registriesMutex.Lock()
defer registriesMutex.Unlock()
foundAt := -1
@@ -123,99 +115,96 @@ func AddRegistry(registry *DefaultModelRegistry) {
registries = append(registries, registry)
}
// tryLock attempts to acquire the registry's write lock, retrying briefly.
// Returns false if it could not be acquired within the bound.
func (r *DefaultModelRegistry) tryLock() bool {
for i := 0; i < lockRetryAttempts; i++ {
if r.mutex.TryLock() {
return true
}
time.Sleep(lockRetryDelay)
}
return false
// registriesSnapshot returns a copy of the registry list so callers can
// iterate without holding registriesMutex.
func registriesSnapshot() []*DefaultModelRegistry {
registriesMutex.RLock()
defer registriesMutex.RUnlock()
return append([]*DefaultModelRegistry(nil), registries...)
}
// tryRLock attempts to acquire the registry's read lock, retrying briefly.
// Returns false if it could not be acquired within the bound.
func (r *DefaultModelRegistry) tryRLock() bool {
for i := 0; i < lockRetryAttempts; i++ {
if r.mutex.TryRLock() {
return true
// validateModel checks the model is a struct (or pointer/slice/array of one)
// and returns the normalised non-pointer struct value. It takes no locks.
func validateModel(model interface{}) (result interface{}, err error) {
// Reflection on a pathological type must fail the registration, not crash the process.
defer func() {
if r := recover(); r != nil {
result = nil
err = fmt.Errorf("%w: %v", ErrInvalidModel, logger.HandlePanic("modelregistry.validateModel", r))
}
time.Sleep(lockRetryDelay)
}
return false
}
}()
func (r *DefaultModelRegistry) RegisterModel(name string, model interface{}) error {
if !r.tryLock() {
return fmt.Errorf("failed to register model %s: registry locked", name)
}
defer r.mutex.Unlock()
if _, exists := r.models[name]; exists {
return fmt.Errorf("model %s already registered", name)
}
// Validate that model is a non-pointer struct
modelType := reflect.TypeOf(model)
if modelType == nil {
return fmt.Errorf("model cannot be nil")
return nil, fmt.Errorf("%w: model cannot be nil", ErrInvalidModel)
}
originalType := modelType
// Unwrap pointers, slices, and arrays to check the underlying type
for modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array {
for depth := 0; modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array; depth++ {
if depth >= maxUnwrapDepth {
return nil, fmt.Errorf("%w: type %s nests deeper than %d levels", ErrInvalidModel, originalType.String(), maxUnwrapDepth)
}
modelType = modelType.Elem()
}
// Validate that the underlying type is a struct
if modelType.Kind() != reflect.Struct {
return fmt.Errorf("model must be a struct or pointer to struct, got %s", originalType.String())
return nil, fmt.Errorf("%w: model must be a struct or pointer to struct, got %s", ErrInvalidModel, originalType.String())
}
// If a pointer/slice/array was passed, unwrap to the base struct
if originalType != modelType {
// Create a zero value of the struct type
model = reflect.New(modelType).Elem().Interface()
}
// Additional check: ensure model is not a pointer
finalType := reflect.TypeOf(model)
if finalType.Kind() == reflect.Pointer {
return fmt.Errorf("model must be a non-pointer struct, got pointer to %s. Use MyModel{} instead of &MyModel{}", finalType.Elem().Name())
if finalType := reflect.TypeOf(model); finalType.Kind() == reflect.Pointer {
return nil, fmt.Errorf("%w: model must be a non-pointer struct, got pointer to %s. Use MyModel{} instead of &MyModel{}", ErrInvalidModel, finalType.Elem().Name())
}
return model, nil
}
// registerLocked validates the model outside the lock, then writes the model
// and its rules under a single lock acquisition so no reader can observe the
// model without its final rules.
func (r *DefaultModelRegistry) registerLocked(name string, model interface{}, rules ModelRules) error {
model, err := validateModel(model)
if err != nil {
return err
}
r.models[name] = model
// Initialize with default rules if not already set
if _, exists := r.rules[name]; !exists {
r.rules[name] = DefaultModelRules()
r.mutex.Lock()
defer r.mutex.Unlock()
if _, exists := r.models[name]; exists {
return fmt.Errorf("%w: %s", ErrModelExists, name)
}
r.models[name] = model
r.rules[name] = rules
return nil
}
func (r *DefaultModelRegistry) RegisterModel(name string, model interface{}) error {
return r.registerLocked(name, model, DefaultModelRules())
}
func (r *DefaultModelRegistry) GetModel(name string) (interface{}, error) {
if !r.tryRLock() {
return nil, fmt.Errorf("failed to get model %s: registry locked", name)
}
r.mutex.RLock()
defer r.mutex.RUnlock()
model, exists := r.models[name]
if !exists {
return nil, fmt.Errorf("model %s not found", name)
return nil, fmt.Errorf("%w: %s", ErrModelNotFound, name)
}
return model, nil
}
func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
if !r.tryRLock() {
return make(map[string]interface{})
}
r.mutex.RLock()
defer r.mutex.RUnlock()
result := make(map[string]interface{})
result := make(map[string]interface{}, len(r.models))
for k, v := range r.models {
result[k] = v
}
@@ -225,9 +214,13 @@ func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
func (r *DefaultModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) {
// Try full name first
fullName := fmt.Sprintf("%s.%s", schema, entity)
if model, err := r.GetModel(fullName); err == nil {
model, err := r.GetModel(fullName)
if err == nil {
return model, nil
}
if !errors.Is(err, ErrModelNotFound) {
return nil, err
}
// Fallback to entity name only
return r.GetModel(entity)
@@ -238,9 +231,8 @@ func (r *DefaultModelRegistry) SetModelRules(name string, rules ModelRules) erro
r.mutex.Lock()
defer r.mutex.Unlock()
// Check if model exists
if _, exists := r.models[name]; !exists {
return fmt.Errorf("model %s not found", name)
return fmt.Errorf("%w: %s", ErrModelNotFound, name)
}
r.rules[name] = rules
@@ -253,12 +245,10 @@ func (r *DefaultModelRegistry) GetModelRules(name string) (ModelRules, error) {
r.mutex.RLock()
defer r.mutex.RUnlock()
// Check if model exists
if _, exists := r.models[name]; !exists {
return ModelRules{}, fmt.Errorf("model %s not found", name)
return ModelRules{}, fmt.Errorf("%w: %s", ErrModelNotFound, name)
}
// Return rules if set, otherwise return default rules
if rules, exists := r.rules[name]; exists {
return rules, nil
}
@@ -266,84 +256,62 @@ func (r *DefaultModelRegistry) GetModelRules(name string) (ModelRules, error) {
return DefaultModelRules(), nil
}
// RegisterModelWithRules registers a model with specific rules
// RegisterModelWithRules registers a model with specific rules atomically
func (r *DefaultModelRegistry) RegisterModelWithRules(name string, model interface{}, rules ModelRules) error {
// First register the model
if err := r.RegisterModel(name, model); err != nil {
return err
}
// Then set the rules (we need to lock again for rules)
r.mutex.Lock()
defer r.mutex.Unlock()
r.rules[name] = rules
return nil
return r.registerLocked(name, model, rules)
}
// Global convenience functions using the default registry
// RegisterModel registers a model with the default global registry
func RegisterModel(model interface{}, name string) error {
return defaultRegistry.RegisterModel(name, model)
return GetDefaultRegistry().RegisterModel(name, model)
}
// GetModelByName retrieves a model by searching through all registries in order
// Returns the first match found
func GetModelByName(name string) (interface{}, error) {
registriesMutex.RLock()
defer registriesMutex.RUnlock()
for _, registry := range registries {
for _, registry := range registriesSnapshot() {
if model, err := registry.GetModel(name); err == nil {
return model, nil
}
}
return nil, fmt.Errorf("model %s not found in any registry", name)
return nil, fmt.Errorf("%w: %s (in any registry)", ErrModelNotFound, name)
}
// IterateModels iterates over all models in the default global registry
// IterateModels iterates over all models in the default global registry.
// It iterates over a snapshot, so fn may safely call back into the registry.
// A panic in fn is recovered and logged with the model name, and iteration
// continues with the remaining models.
func IterateModels(fn func(name string, model interface{})) {
defaultRegistry.mutex.RLock()
defer defaultRegistry.mutex.RUnlock()
for name, model := range defaultRegistry.models {
fn(name, model)
for name, model := range GetDefaultRegistry().GetAllModels() {
callIsolated(name, model, fn)
}
}
// GetModels returns a list of all models from all registries
// Models are collected in registry order, with duplicates included
func GetModels() []interface{} {
acquired := false
for i := 0; i < lockRetryAttempts; i++ {
if registriesMutex.TryRLock() {
acquired = true
break
func callIsolated(name string, model interface{}, fn func(name string, model interface{})) {
defer func() {
if r := recover(); r != nil {
_ = logger.HandlePanic("modelregistry.IterateModels", r, "model", name)
}
time.Sleep(lockRetryDelay)
}
if !acquired {
return nil
}
defer registriesMutex.RUnlock()
}()
fn(name, model)
}
// GetModels returns a list of all models from all registries.
// Only the first occurrence of each model name is included.
func GetModels() []interface{} {
var models []interface{}
seen := make(map[string]bool)
for _, registry := range registries {
if !registry.tryRLock() {
continue
}
for name, model := range registry.models {
// Only add the first occurrence of each model name
for _, registry := range registriesSnapshot() {
for name, model := range registry.GetAllModels() {
if !seen[name] {
models = append(models, model)
seen[name] = true
}
}
registry.mutex.RUnlock()
}
return models
@@ -351,31 +319,31 @@ func GetModels() []interface{} {
// SetModelRules sets the rules for a specific model in the default registry
func SetModelRules(name string, rules ModelRules) error {
return defaultRegistry.SetModelRules(name, rules)
return GetDefaultRegistry().SetModelRules(name, rules)
}
// GetModelRules retrieves the rules for a specific model from the default registry
func GetModelRules(name string) (ModelRules, error) {
return defaultRegistry.GetModelRules(name)
return GetDefaultRegistry().GetModelRules(name)
}
// GetModelRulesByName retrieves the rules for a model by searching through all registries in order
// Returns the first match found
// Returns the first match found. The error wraps ErrModelNotFound when no registry has the model.
func GetModelRulesByName(name string) (ModelRules, error) {
registriesMutex.RLock()
defer registriesMutex.RUnlock()
for _, registry := range registries {
if _, err := registry.GetModel(name); err == nil {
// Model found in this registry, get its rules
return registry.GetModelRules(name)
for _, registry := range registriesSnapshot() {
rules, err := registry.GetModelRules(name)
if err == nil {
return rules, nil
}
if !errors.Is(err, ErrModelNotFound) {
return ModelRules{}, err
}
}
return ModelRules{}, fmt.Errorf("model %s not found in any registry", name)
return ModelRules{}, fmt.Errorf("%w: %s (in any registry)", ErrModelNotFound, name)
}
// RegisterModelWithRules registers a model with specific rules in the default registry
func RegisterModelWithRules(model interface{}, name string, rules ModelRules) error {
return defaultRegistry.RegisterModelWithRules(name, model, rules)
return GetDefaultRegistry().RegisterModelWithRules(name, model, rules)
}
+161
View File
@@ -0,0 +1,161 @@
package modelregistry
import (
"errors"
"sync"
"testing"
)
type testModel struct{ ID int }
type recursivePtr *recursivePtr
func TestSentinelErrors(t *testing.T) {
r := NewModelRegistry()
if _, err := r.GetModel("nope"); !errors.Is(err, ErrModelNotFound) {
t.Fatalf("GetModel: want ErrModelNotFound, got %v", err)
}
if _, err := r.GetModelRules("nope"); !errors.Is(err, ErrModelNotFound) {
t.Fatalf("GetModelRules: want ErrModelNotFound, got %v", err)
}
if err := r.SetModelRules("nope", ModelRules{}); !errors.Is(err, ErrModelNotFound) {
t.Fatalf("SetModelRules: want ErrModelNotFound, got %v", err)
}
if err := r.RegisterModel("a", testModel{}); err != nil {
t.Fatal(err)
}
if err := r.RegisterModel("a", testModel{}); !errors.Is(err, ErrModelExists) {
t.Fatalf("want ErrModelExists, got %v", err)
}
if err := r.RegisterModel("b", nil); !errors.Is(err, ErrInvalidModel) {
t.Fatalf("want ErrInvalidModel, got %v", err)
}
if err := r.RegisterModel("c", 42); !errors.Is(err, ErrInvalidModel) {
t.Fatalf("want ErrInvalidModel, got %v", err)
}
}
func TestRecursivePointerTypeRejected(t *testing.T) {
var x recursivePtr
if err := NewModelRegistry().RegisterModel("r", x); !errors.Is(err, ErrInvalidModel) {
t.Fatalf("want ErrInvalidModel, got %v", err)
}
}
func TestPointerNormalised(t *testing.T) {
r := NewModelRegistry()
if err := r.RegisterModel("p", &testModel{}); err != nil {
t.Fatal(err)
}
m, _ := r.GetModel("p")
if _, ok := m.(testModel); !ok {
t.Fatalf("want testModel value, got %T", m)
}
}
// Rules must never be observable as permissive for a restrictively registered model.
func TestRegisterModelWithRulesAtomic(t *testing.T) {
for i := 0; i < 200; i++ {
r := NewModelRegistry()
var wg sync.WaitGroup
stop := make(chan struct{})
wg.Add(1)
go func() {
defer wg.Done()
for {
select {
case <-stop:
return
default:
}
if rules, err := r.GetModelRules("m"); err == nil && rules.CanDelete {
t.Error("observed permissive rules for restrictive model")
return
}
}
}()
if err := r.RegisterModelWithRules("m", testModel{}, ModelRules{CanRead: true}); err != nil {
t.Fatal(err)
}
close(stop)
wg.Wait()
}
}
func TestIterateModelsCallbackMayReenter(t *testing.T) {
prev := GetDefaultRegistry()
reg := NewModelRegistry()
SetDefaultRegistry(reg)
defer SetDefaultRegistry(prev)
if err := RegisterModel(testModel{}, "iter.a"); err != nil {
t.Fatal(err)
}
done := make(chan struct{})
go func() {
defer close(done)
IterateModels(func(name string, _ interface{}) {
_ = RegisterModel(testModel{}, name+".copy") // would deadlock if lock held
})
}()
<-done
if _, err := reg.GetModel("iter.a.copy"); err != nil {
t.Fatal(err)
}
}
func TestGetModelRulesByNameAcrossRegistries(t *testing.T) {
extra := NewModelRegistry()
if err := extra.RegisterModelWithRules("x.only", testModel{}, ModelRules{CanRead: true}); err != nil {
t.Fatal(err)
}
AddRegistry(extra)
rules, err := GetModelRulesByName("x.only")
if err != nil || rules.CanDelete {
t.Fatalf("rules=%+v err=%v", rules, err)
}
if _, err := GetModelRulesByName("x.missing"); !errors.Is(err, ErrModelNotFound) {
t.Fatalf("want ErrModelNotFound, got %v", err)
}
}
func TestConcurrentAccessRace(t *testing.T) {
r := NewModelRegistry()
_ = r.RegisterModel("seed", testModel{})
var wg sync.WaitGroup
for i := 0; i < 16; i++ {
wg.Add(1)
go func(i int) {
defer wg.Done()
for j := 0; j < 100; j++ {
_ = r.SetModelRules("seed", ModelRules{CanRead: j%2 == 0})
_, _ = r.GetModelRules("seed")
_ = r.GetAllModels()
_ = GetDefaultRegistry()
_ = GetModels()
}
}(i)
}
wg.Wait()
}
func TestIterateModelsRecoversCallbackPanic(t *testing.T) {
prev := GetDefaultRegistry()
SetDefaultRegistry(NewModelRegistry())
defer SetDefaultRegistry(prev)
_ = RegisterModel(testModel{}, "p.a")
_ = RegisterModel(testModel{}, "p.b")
calls := 0
IterateModels(func(string, interface{}) {
calls++
panic("boom")
})
if calls != 2 {
t.Fatalf("want both models visited, got %d", calls)
}
// Registry must still be usable (no lock left held).
if err := RegisterModel(testModel{}, "p.c"); err != nil {
t.Fatal(err)
}
}
+1 -1
View File
@@ -324,7 +324,7 @@ func (ebc *ExternalBrokerClient) Stop(ctx context.Context) error {
}
if ebc.client != nil && ebc.client.IsConnected() {
ebc.client.Disconnect(uint(ebc.config.ConnectTimeout.Milliseconds()))
ebc.client.Disconnect(uint(ebc.config.ConnectTimeout.Milliseconds())) //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
}
ebc.connected = false
+94 -64
View File
@@ -41,6 +41,12 @@ func setupTestHandler(t *testing.T) (*Handler, *gorm.DB) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
// Each connection to ":memory:" gets its own database; pin to one so
// concurrent requests all see the migrated schema.
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)
// Auto-migrate test model
err = db.AutoMigrate(&TestUser{})
require.NoError(t, err)
@@ -93,9 +99,9 @@ func TestHandler_HandleRead_Single(t *testing.T) {
// Insert test data
user := &TestUser{
ID: 1,
Name: "John Doe",
Email: "john@example.com",
ID: 1,
Name: "John Doe",
Email: "john@example.com",
Status: "active",
}
db.Create(user)
@@ -115,13 +121,16 @@ func TestHandler_HandleRead_Single(t *testing.T) {
// Create hook context
hookCtx := &HookContext{
Context: context.Background(),
Handler: nil,
Schema: "public",
Entity: "users",
ID: "1",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
Context: context.Background(),
TableName: "users",
Model: &TestUser{},
ModelPtr: &TestUser{},
Handler: nil,
Schema: "public",
Entity: "users",
ID: "1",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
}
// Handle read
@@ -164,12 +173,15 @@ func TestHandler_HandleRead_Multiple(t *testing.T) {
// Create hook context
hookCtx := &HookContext{
Context: context.Background(),
Handler: nil,
Schema: "public",
Entity: "users",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
Context: context.Background(),
TableName: "users",
Model: &TestUser{},
ModelPtr: &TestUser{},
Handler: nil,
Schema: "public",
Entity: "users",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
}
// Handle read
@@ -208,13 +220,16 @@ func TestHandler_HandleCreate(t *testing.T) {
// Create hook context
hookCtx := &HookContext{
Context: context.Background(),
Handler: nil,
Schema: "public",
Entity: "users",
Data: newUser,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
Context: context.Background(),
TableName: "users",
Model: &TestUser{},
ModelPtr: &TestUser{},
Handler: nil,
Schema: "public",
Entity: "users",
Data: newUser,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
}
// Handle create
@@ -266,14 +281,17 @@ func TestHandler_HandleUpdate(t *testing.T) {
// Create hook context
hookCtx := &HookContext{
Context: context.Background(),
Handler: nil,
Schema: "public",
Entity: "users",
ID: "1",
Data: updateData,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
Context: context.Background(),
TableName: "users",
Model: &TestUser{},
ModelPtr: &TestUser{},
Handler: nil,
Schema: "public",
Entity: "users",
ID: "1",
Data: updateData,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
}
// Handle update
@@ -318,13 +336,16 @@ func TestHandler_HandleDelete(t *testing.T) {
// Create hook context
hookCtx := &HookContext{
Context: context.Background(),
Handler: nil,
Schema: "public",
Entity: "users",
ID: "1",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
Context: context.Background(),
TableName: "users",
Model: &TestUser{},
ModelPtr: &TestUser{},
Handler: nil,
Schema: "public",
Entity: "users",
ID: "1",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
}
// Handle delete
@@ -388,13 +409,13 @@ func TestHandler_HandleUnsubscribe(t *testing.T) {
sub := handler.subscriptionManager.Subscribe("sub-1", client.ID, "public", "users", &common.RequestOptions{})
client.AddSubscription(sub)
// Create unsubscribe message with subscription ID in Data
// Create unsubscribe message with the subscription ID
msg := &Message{
ID: "msg-7",
Type: MessageTypeSubscription,
Operation: OperationUnsubscribe,
Data: map[string]interface{}{"subscription_id": "sub-1"},
Options: &common.RequestOptions{},
ID: "msg-7",
Type: MessageTypeSubscription,
Operation: OperationUnsubscribe,
SubscriptionID: "sub-1",
Options: &common.RequestOptions{},
}
// Handle unsubscribe
@@ -490,12 +511,15 @@ func TestHandler_Hooks_BeforeRead(t *testing.T) {
// Create hook context
hookCtx := &HookContext{
Context: context.Background(),
Handler: nil,
Schema: "public",
Entity: "users",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
Context: context.Background(),
TableName: "users",
Model: &TestUser{},
ModelPtr: &TestUser{},
Handler: nil,
Schema: "public",
Entity: "users",
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
}
// Handle read
@@ -544,13 +568,16 @@ func TestHandler_Hooks_BeforeCreate(t *testing.T) {
}
hookCtx := &HookContext{
Context: context.Background(),
Handler: nil,
Schema: "public",
Entity: "users",
Data: newUser,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
Context: context.Background(),
TableName: "users",
Model: &TestUser{},
ModelPtr: &TestUser{},
Handler: nil,
Schema: "public",
Entity: "users",
Data: newUser,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
}
// Handle create
@@ -600,13 +627,16 @@ func TestHandler_ConcurrentRequests(t *testing.T) {
}
hookCtx := &HookContext{
Context: context.Background(),
Handler: nil,
Schema: "public",
Entity: "users",
Data: newUser,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
Context: context.Background(),
TableName: "users",
Model: &TestUser{},
ModelPtr: &TestUser{},
Handler: nil,
Schema: "public",
Entity: "users",
Data: newUser,
Options: msg.Options,
Metadata: map[string]interface{}{"mqtt_client": client},
}
handler.handleCreate(client, msg, hookCtx)
+4 -4
View File
@@ -149,7 +149,7 @@ func generateSwaggerUI(config UIConfig) (string, error) {
data := templateData{
UIConfig: config,
SafeCustomCSS: template.CSS(config.CustomCSS),
SafeCustomCSS: template.CSS(config.CustomCSS), //nolint:gosec // G203: CSS from trusted server config
}
var buf strings.Builder
@@ -200,7 +200,7 @@ func generateRapiDoc(config UIConfig) (string, error) {
data := templateData{
UIConfig: config,
SafeCustomCSS: template.CSS(config.CustomCSS),
SafeCustomCSS: template.CSS(config.CustomCSS), //nolint:gosec // G203: CSS from trusted server config
}
var buf strings.Builder
@@ -238,7 +238,7 @@ func generateRedoc(config UIConfig) (string, error) {
data := templateData{
UIConfig: config,
SafeCustomCSS: template.CSS(config.CustomCSS),
SafeCustomCSS: template.CSS(config.CustomCSS), //nolint:gosec // G203: CSS from trusted server config
}
var buf strings.Builder
@@ -276,7 +276,7 @@ func generateScalar(config UIConfig) (string, error) {
data := templateData{
UIConfig: config,
SafeCustomCSS: template.CSS(config.CustomCSS),
SafeCustomCSS: template.CSS(config.CustomCSS), //nolint:gosec // G203: CSS from trusted server config
}
var buf strings.Builder
+70 -13
View File
@@ -439,6 +439,63 @@ func GetSQLModelColumns(model any) []string {
return columns
}
// HasColumn reports whether the model has a struct field that bun/gorm would
// scan a column named columnName into. Unlike GetSQLModelColumns, this
// includes scanonly fields (e.g. a `bun:"jsonvalue_product_cost,scanonly"`
// field added specifically to receive a computed/JSON-path SELECT expression)
// since those are legitimate scan targets even though they are not writable.
// Matching is case-insensitive against the resolved bun/gorm/json column name
// and against the bare Go field name.
func HasColumn(model any, columnName string) bool {
if columnName == "" {
return false
}
modelType := reflect.TypeOf(model)
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
modelType = modelType.Elem()
}
if modelType == nil || modelType.Kind() != reflect.Struct {
return false
}
return hasColumnInType(modelType, columnName)
}
func hasColumnInType(typ reflect.Type, columnName string) bool {
for i := 0; i < typ.NumField(); i++ {
field := typ.Field(i)
if !field.IsExported() {
continue
}
bunTag := field.Tag.Get("bun")
gormTag := field.Tag.Get("gorm")
if field.Anonymous {
fieldType := field.Type
if fieldType.Kind() == reflect.Pointer {
fieldType = fieldType.Elem()
}
if fieldType.Kind() == reflect.Struct {
if hasColumnInType(fieldType, columnName) {
return true
}
continue
}
}
if bunTag == "-" || gormTag == "-" {
continue
}
if strings.EqualFold(getColumnNameFromField(field), columnName) || strings.EqualFold(field.Name, columnName) {
return true
}
}
return false
}
// collectSQLColumnsFromType recursively collects SQL column names from a struct type
// scanOnlyEmbedded indicates if we're inside a scan-only embedded struct
func collectSQLColumnsFromType(typ reflect.Type, columns *[]string, scanOnlyEmbedded bool) {
@@ -795,11 +852,11 @@ func ConvertToNumericType(value string, kind reflect.Kind) (interface{}, error)
case reflect.Int:
return int(intVal), nil
case reflect.Int8:
return int8(intVal), nil
return int8(intVal), nil //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case reflect.Int16:
return int16(intVal), nil
return int16(intVal), nil //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case reflect.Int32:
return int32(intVal), nil
return int32(intVal), nil //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case reflect.Int64:
return intVal, nil
}
@@ -826,11 +883,11 @@ func ConvertToNumericType(value string, kind reflect.Kind) (interface{}, error)
case reflect.Uint:
return uint(uintVal), nil
case reflect.Uint8:
return uint8(uintVal), nil
return uint8(uintVal), nil //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case reflect.Uint16:
return uint16(uintVal), nil
return uint16(uintVal), nil //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case reflect.Uint32:
return uint32(uintVal), nil
return uint32(uintVal), nil //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case reflect.Uint64:
return uintVal, nil
}
@@ -1489,7 +1546,7 @@ func convertToInt64(value interface{}) (int64, bool) {
case int64:
return v, true
case uint:
return int64(v), true
return int64(v), true //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case uint8:
return int64(v), true
case uint16:
@@ -1497,7 +1554,7 @@ func convertToInt64(value interface{}) (int64, bool) {
case uint32:
return int64(v), true
case uint64:
return int64(v), true
return int64(v), true //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case float32:
return int64(v), true
case float64:
@@ -1514,15 +1571,15 @@ func convertToInt64(value interface{}) (int64, bool) {
func convertToUint64(value interface{}) (uint64, bool) {
switch v := value.(type) {
case int:
return uint64(v), true
return uint64(v), true //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case int8:
return uint64(v), true
return uint64(v), true //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case int16:
return uint64(v), true
return uint64(v), true //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case int32:
return uint64(v), true
return uint64(v), true //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case int64:
return uint64(v), true
return uint64(v), true //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
case uint:
return uint64(v), true
case uint8:
+17
View File
@@ -411,6 +411,23 @@ func TestGetModelColumnsWithEmbedded(t *testing.T) {
}
}
func TestHasColumn(t *testing.T) {
m := ModelWithEmbedded{}
for _, col := range []string{"name", "description", "rid_base", "created_at", "cql1", "cql2"} {
if !HasColumn(m, col) {
t.Errorf("HasColumn(%q) = false, want true", col)
}
}
if HasColumn(m, "nonexistent_column") {
t.Error("HasColumn(nonexistent_column) = true, want false")
}
if HasColumn(m, "") {
t.Error("HasColumn(\"\") = true, want false")
}
}
func TestIsColumnWritableWithEmbedded(t *testing.T) {
tests := []struct {
name string
+1 -1
View File
@@ -213,7 +213,7 @@ func OAuth2CallbackHandler(auth *security.DatabaseAuthenticator, providerName, a
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(loginResp) //nolint:errcheck
json.NewEncoder(w).Encode(loginResp) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
}
+4
View File
@@ -349,6 +349,10 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
logger.Debug("Selecting columns: %v", options.Columns)
for _, col := range options.Columns {
if expr, jargs, alias, ok := common.ResolveJSONColumnExpr(model, "", col); ok {
if !reflection.HasColumn(model, alias) {
logger.Warn("Skipping JSON select column %q: model has no scan target for alias %q", col, alias)
continue
}
query = query.ColumnExpr(expr+" AS "+common.QuoteIdent(alias), jargs...)
continue
}
+4
View File
@@ -530,6 +530,10 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
// JSON sub-field selection (data->>'x', data.x, data#>>'{a,b}'):
// emit a parameterised expression aliased to a stable name.
if expr, jargs, alias, ok := common.ResolveJSONColumnExpr(model, selectAlias, col); ok {
if !reflection.HasColumn(model, alias) {
logger.Warn("Skipping JSON select column %q: model has no scan target for alias %q", col, alias)
continue
}
query = query.ColumnExpr(expr+" AS "+common.QuoteIdent(alias), jargs...)
continue
}
+2 -2
View File
@@ -529,7 +529,7 @@ func ExampleBunRouterWithBunDB(bunDB *bun.DB) {
SetupBunRouterRoutes(bunRouter, handler, nil)
// Start server
if err := http.ListenAndServe(":8080", bunRouter); err != nil {
if err := http.ListenAndServe(":8080", bunRouter); err != nil { //nolint:gosec // G114: example code only
logger.Error("Server failed to start: %v", err)
}
}
@@ -549,7 +549,7 @@ func ExampleBunRouterWithGroup(bunDB *bun.DB) {
SetupBunRouterRoutes(apiGroup, handler, nil)
// Start server
if err := http.ListenAndServe(":8080", bunRouter); err != nil {
if err := http.ListenAndServe(":8080", bunRouter); err != nil { //nolint:gosec // G114: example code only
logger.Error("Server failed to start: %v", err)
}
}
+196
View File
@@ -0,0 +1,196 @@
package security
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"reflect"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
)
// slowProvider embeds a nil SecurityProvider; only the two load methods are used.
type slowProvider struct {
SecurityProvider
calls atomic.Int32
active atomic.Int32
maxSeen atomic.Int32
delay time.Duration
}
func (p *slowProvider) enter() {
p.calls.Add(1)
n := p.active.Add(1)
for {
m := p.maxSeen.Load()
if n <= m || p.maxSeen.CompareAndSwap(m, n) {
break
}
}
time.Sleep(p.delay)
p.active.Add(-1)
}
func (p *slowProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error) {
p.enter()
return []ColumnSecurity{{Schema: schema, Tablename: table}}, nil
}
func (p *slowProvider) GetRowSecurity(ctx context.Context, ref any, schema, table string) (RowSecurity, error) {
p.enter()
return RowSecurity{Schema: schema, Tablename: table}, nil
}
func TestLoadDoesNotHoldLockAcrossProvider(t *testing.T) {
p := &slowProvider{delay: 100 * time.Millisecond}
sl, _ := NewSecurityList(p)
var wg sync.WaitGroup
start := time.Now()
for i := 0; i < 8; i++ {
wg.Add(2)
go func(i int) { defer wg.Done(); _ = sl.LoadColumnSecurity(context.Background(), i, "s", "t", false) }(i)
go func(i int) { defer wg.Done(); _, _ = sl.LoadRowSecurity(context.Background(), i, "s", "t", false) }(i)
}
wg.Wait()
if el := time.Since(start); el > 500*time.Millisecond {
t.Fatalf("loads serialised: %v", el)
}
if p.maxSeen.Load() < 2 {
t.Fatal("provider calls never overlapped")
}
}
func TestLoadCachesAndHonoursOverwrite(t *testing.T) {
p := &slowProvider{}
sl, _ := NewSecurityList(p)
ctx := context.Background()
for i := 0; i < 3; i++ {
_ = sl.LoadColumnSecurity(ctx, 1, "s", "t", false)
_, _ = sl.LoadRowSecurity(ctx, 1, "s", "t", false)
}
if got := p.calls.Load(); got != 2 {
t.Fatalf("expected 2 provider calls (cached), got %d", got)
}
_ = sl.LoadColumnSecurity(ctx, 1, "s", "t", true)
_, _ = sl.LoadRowSecurity(ctx, 1, "s", "t", true)
if got := p.calls.Load(); got != 4 {
t.Fatalf("overwrite should reload: got %d calls", got)
}
}
func TestLoadExpiryAndPrune(t *testing.T) {
p := &slowProvider{}
sl, _ := NewSecurityList(p)
ctx := context.Background()
_ = sl.LoadColumnSecurity(ctx, 1, "s", "t", false)
sl.ColumnSecurityMutex.Lock()
sl.colSecExpiry["s.t@1"] = time.Now().Add(-time.Second)
sl.ColumnSecurityMutex.Unlock()
_ = sl.LoadColumnSecurity(ctx, 1, "s", "t", false)
if p.calls.Load() != 2 {
t.Fatal("expired entry should reload")
}
sl.ColumnSecurityMutex.Lock()
sl.colSecExpiry["s.old@9"] = time.Now().Add(-time.Hour)
sl.ColumnSecurity["s.old@9"] = nil
sl.lastColPrune = time.Time{}
sl.ColumnSecurityMutex.Unlock()
_ = sl.LoadColumnSecurity(ctx, 2, "s", "t", false)
sl.ColumnSecurityMutex.RLock()
_, ok := sl.ColumnSecurity["s.old@9"]
sl.ColumnSecurityMutex.RUnlock()
if ok {
t.Fatal("stale entry not pruned")
}
}
func TestAuthenticateRejectsTooManyTokens(t *testing.T) {
db, _, err := sqlmock.New()
if err != nil {
t.Fatal(err)
}
defer db.Close()
a := NewDatabaseAuthenticator(db)
r := httptest.NewRequest(http.MethodGet, "/", nil)
r.Header.Set("Authorization", "a,b,c,d,e,f,g,h")
if _, err := a.Authenticate(r); err == nil || err.Error() != "too many authorization tokens" {
t.Fatalf("got %v", err)
}
}
func TestOAuth2CleanupStopsOnClose(t *testing.T) {
db, _, _ := sqlmock.New()
defer db.Close()
a := NewDatabaseAuthenticator(db).WithOAuth2(OAuth2Config{ClientID: "x", ProviderName: "p"})
p := a.oauth2Providers["p"]
if err := a.Close(); err != nil {
t.Fatal(err)
}
_ = a.Close() // idempotent
select {
case <-p.stopCh:
default:
t.Fatal("stop channel not closed")
}
}
func TestSplitTagDropsEmpty(t *testing.T) {
got := splitTag("a,,b,c,", ',')
if len(got) != 3 || got[0] != "a" || got[2] != "c" {
t.Fatalf("got %v", got)
}
}
func TestColumnSecurityPanicFailsClosed(t *testing.T) {
type Rec struct {
JSONCol string `json:"json_col" bun:"json_col"`
}
sl, _ := NewSecurityList(&slowProvider{})
sl.ColumnSecurity["public.t@1"] = []ColumnSecurity{{
Schema: "public", Tablename: "t", Path: []string{"JSONCol"}, Accesstype: "mask", UserID: 1,
}}
// A struct boxed in an interface is not addressable, so SetString panics.
recs := []any{Rec{JSONCol: "secret"}}
out, err := sl.ApplyColumnSecurity(reflect.ValueOf(recs), reflect.TypeOf(Rec{}), 1, "public", "t")
if err == nil {
t.Fatalf("panic must be returned as an error, got out=%v", out)
}
if errors.Is(err, ErrNoColumnSecurity) {
t.Fatal("a panic must not look like 'no rules'")
}
}
func TestNoRulesIsNotAnError(t *testing.T) {
sl, _ := NewSecurityList(&slowProvider{})
if _, err := sl.GetRowSecurityTemplate(1, "s", "t"); !errors.Is(err, ErrNoRowSecurity) {
t.Fatalf("got %v", err)
}
if _, err := sl.ApplyColumnSecurity(reflect.ValueOf([]int{}), reflect.TypeOf(0), 1, "s", "t"); !errors.Is(err, ErrNoColumnSecurity) {
t.Fatalf("got %v", err)
}
}
func TestApplyColumnSecurityHookFailsClosedOnPanic(t *testing.T) {
type Rec struct {
JSONCol string `bun:"json_col"`
}
sl, _ := NewSecurityList(&slowProvider{})
sl.ColumnSecurity["public.t@1"] = []ColumnSecurity{{
Schema: "public", Tablename: "t", Path: []string{"JSONCol"}, Accesstype: "mask", UserID: 1,
}}
secCtx := &mockSecurityContext{
ctx: context.Background(), userID: 1, hasUser: true, schema: "public", entity: "t",
model: &Rec{}, result: []any{Rec{JSONCol: "secret"}},
}
if err := ApplyColumnSecurity(secCtx, sl); err == nil {
t.Fatal("a panic during masking must fail the request, not return unmasked data")
}
}
+73 -29
View File
@@ -1,12 +1,16 @@
-- Database Schema for DatabaseAuthenticator
-- ============================================
-- pgcrypto provides gen_random_bytes(), crypt() and gen_salt(); it is required
-- for session token generation and for password hashing/verification below.
CREATE EXTENSION IF NOT EXISTS pgcrypto;
-- Users table
CREATE TABLE IF NOT EXISTS users (
id SERIAL PRIMARY KEY,
username VARCHAR(255) NOT NULL UNIQUE,
email VARCHAR(255) NOT NULL UNIQUE,
password VARCHAR(255), -- bcrypt hashed password (nullable for OAuth2 users)
password VARCHAR(255), -- bcrypt hash (nullable for OAuth2 users); legacy cleartext is accepted at login (upgrade to bcrypt is opt-in)
user_level INTEGER DEFAULT 0,
roles VARCHAR(500), -- Comma-separated roles: "admin,manager,user"
is_active BOOLEAN DEFAULT true,
@@ -98,6 +102,8 @@ DECLARE
v_user_level INTEGER;
v_roles TEXT;
v_password_hash TEXT;
v_supplied_password TEXT;
v_password_ok BOOLEAN := false;
v_session_token TEXT;
v_expires_at TIMESTAMP;
v_ip_address TEXT;
@@ -107,6 +113,7 @@ DECLARE
BEGIN
-- Extract login request fields
v_username := p_request->>'username';
v_supplied_password := p_request->>'password';
v_ip_address := p_request->'claims'->>'ip_address';
v_user_agent := p_request->'claims'->>'user_agent';
@@ -121,12 +128,30 @@ BEGIN
RETURN;
END IF;
-- TODO: Verify password hash using pgcrypto extension
-- Enable pgcrypto: CREATE EXTENSION IF NOT EXISTS pgcrypto;
-- IF NOT (crypt(p_request->>'password', v_password_hash) = v_password_hash) THEN
-- RETURN QUERY SELECT false, 'Invalid credentials'::text, NULL::jsonb;
-- RETURN;
-- END IF;
-- Verify the password. bcrypt hashes are checked with crypt(); a legacy
-- cleartext value is still accepted (and only rewritten as bcrypt if the
-- upgrade is explicitly enabled).
-- bcrypt only uses the first 72 bytes, so longer input is rejected.
IF v_password_hash IS NOT NULL AND v_password_hash <> ''
AND v_supplied_password IS NOT NULL AND v_supplied_password <> ''
AND octet_length(v_supplied_password) <= 72 THEN
IF v_password_hash ~ '^\$2[aby]\$' THEN
v_password_ok := (crypt(v_supplied_password, v_password_hash) = v_password_hash);
ELSE
v_password_ok := (v_password_hash = v_supplied_password);
-- Upgrading the stored value is opt-in:
-- ALTER DATABASE <db> SET resolvespec.upgrade_password_hash = 'on';
IF v_password_ok AND COALESCE(current_setting('resolvespec.upgrade_password_hash', true), 'off') = 'on' THEN
UPDATE users SET password = crypt(v_supplied_password, gen_salt('bf')), updated_at = now()
WHERE id = v_user_id;
END IF;
END IF;
END IF;
IF NOT v_password_ok THEN
RETURN QUERY SELECT false, 'Invalid credentials'::text, NULL::jsonb;
RETURN;
END IF;
-- Generate session token
v_session_token := 'sess_' || encode(gen_random_bytes(32), 'hex') || '_' || extract(epoch from now())::bigint::text;
@@ -336,6 +361,7 @@ DECLARE
v_username TEXT;
v_email TEXT;
v_password TEXT;
v_password_ok BOOLEAN := false;
v_user_level INTEGER;
v_roles TEXT;
BEGIN
@@ -350,11 +376,26 @@ BEGIN
RETURN;
END IF;
-- TODO: Verify password hash
-- IF NOT (crypt(p_password, v_password) = v_password) THEN
-- RETURN QUERY SELECT false, 'Invalid credentials'::text, NULL::jsonb;
-- RETURN;
-- END IF;
-- Verify the password (bcrypt, or legacy cleartext).
IF v_password IS NOT NULL AND v_password <> ''
AND p_password IS NOT NULL AND p_password <> ''
AND octet_length(p_password) <= 72 THEN
IF v_password ~ '^\$2[aby]\$' THEN
v_password_ok := (crypt(p_password, v_password) = v_password);
ELSE
v_password_ok := (v_password = p_password);
-- Upgrading the stored value is opt-in (see resolvespec_login).
IF v_password_ok AND COALESCE(current_setting('resolvespec.upgrade_password_hash', true), 'off') = 'on' THEN
UPDATE users SET password = crypt(p_password, gen_salt('bf')), updated_at = now()
WHERE id = v_user_id;
END IF;
END IF;
END IF;
IF NOT v_password_ok THEN
RETURN QUERY SELECT false, 'Invalid credentials'::text, NULL::jsonb;
RETURN;
END IF;
-- Return user data for JWT token generation
RETURN QUERY SELECT
@@ -364,7 +405,6 @@ BEGIN
'id', v_user_id,
'username', v_username,
'email', v_email,
'password', v_password,
'user_level', v_user_level,
'roles', v_roles
);
@@ -442,7 +482,8 @@ END;
$$ LANGUAGE plpgsql;
-- 10. resolvespec_register - Registers a new user and creates session
-- Input: RegisterRequest as jsonb {username: string, password: string, email: string, user_level: int, roles: array, claims: object, meta: object}
-- Input: RegisterRequest as jsonb {username: string, password: string, email: string, claims: object, meta: object}
-- (user_level / roles in the request are ignored; new users are unprivileged)
-- Output: p_success (bool), p_error (text), p_data (LoginResponse as jsonb)
CREATE OR REPLACE FUNCTION resolvespec_register(p_request jsonb)
RETURNS TABLE(p_success boolean, p_error text, p_data jsonb) AS $$
@@ -465,15 +506,14 @@ BEGIN
v_username := p_request->>'username';
v_email := p_request->>'email';
v_password := p_request->>'password';
v_user_level := COALESCE((p_request->>'user_level')::integer, 0);
-- Privileges are never taken from the request: self-registration always
-- creates an unprivileged user (level 0, no roles, no program user link).
v_user_level := 0;
v_roles := '';
v_ip_address := p_request->'claims'->>'ip_address';
v_user_agent := p_request->'claims'->>'user_agent';
v_program_user_id := COALESCE((p_request->>'program_user_id')::integer, 0);
v_program_user_table := COALESCE(p_request->>'program_user_table', '');
-- Convert roles array from JSON to comma-separated string
SELECT array_to_string(ARRAY(SELECT jsonb_array_elements_text(p_request->'roles')), ',')
INTO v_roles;
v_program_user_id := 0;
v_program_user_table := '';
-- Validate required fields
IF v_username IS NULL OR v_username = '' THEN
@@ -491,6 +531,11 @@ BEGIN
RETURN;
END IF;
IF octet_length(v_password) > 72 THEN
RETURN QUERY SELECT false, 'Password must be at most 72 bytes'::text, NULL::jsonb;
RETURN;
END IF;
-- Check if username already exists
IF EXISTS (SELECT 1 FROM users WHERE username = v_username) THEN
RETURN QUERY SELECT false, 'Username already exists'::text, NULL::jsonb;
@@ -503,9 +548,7 @@ BEGIN
RETURN;
END IF;
-- TODO: Hash password using pgcrypto extension
-- Enable pgcrypto: CREATE EXTENSION IF NOT EXISTS pgcrypto;
-- v_password := crypt(v_password, gen_salt('bf'));
v_password := crypt(v_password, gen_salt('bf'));
-- Create new user
INSERT INTO users (username, email, password, user_level, roles, is_active, created_at, updated_at, program_user_id, program_user_table)
@@ -1520,8 +1563,7 @@ $$ LANGUAGE plpgsql;
-- 2. resolvespec_password_reset - Validates the token and updates the user's password
-- Input: p_request jsonb {token: string, new_password: string}
-- Output: p_success (bool), p_error (text)
-- NOTE: Hash the new_password with bcrypt before storing (pgcrypto crypt/gen_salt).
-- The TODO below mirrors the convention used in resolvespec_register.
-- NOTE: The new password is hashed with bcrypt (pgcrypto crypt/gen_salt) before storing.
CREATE OR REPLACE FUNCTION resolvespec_password_reset(p_request jsonb)
RETURNS TABLE(p_success boolean, p_error text) AS $$
DECLARE
@@ -1563,9 +1605,11 @@ BEGIN
RETURN;
END IF;
-- TODO: Hash new password with pgcrypto before storing
-- Enable pgcrypto: CREATE EXTENSION IF NOT EXISTS pgcrypto;
-- v_new_pw := crypt(v_new_pw, gen_salt('bf'));
IF octet_length(v_new_pw) > 72 THEN
RETURN QUERY SELECT false, 'new_password must be at most 72 bytes'::text;
RETURN;
END IF;
v_new_pw := crypt(v_new_pw, gen_salt('bf'));
-- Update password and invalidate all sessions
UPDATE users SET password = v_new_pw, updated_at = now() WHERE id = v_user_id;
+1 -1
View File
@@ -7,7 +7,7 @@ CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username VARCHAR(255) NOT NULL UNIQUE,
email VARCHAR(255) NOT NULL UNIQUE,
password VARCHAR(255),
password VARCHAR(255), -- bcrypt hash (nullable for OAuth2 users); legacy cleartext is accepted at login (upgrade to bcrypt is opt-in)
user_level INTEGER DEFAULT 0,
roles VARCHAR(500),
is_active BOOLEAN DEFAULT 1,
+75 -7
View File
@@ -53,11 +53,12 @@ func TestDirectMode_RegisterThenLogin(t *testing.T) {
ctx := context.Background()
regResp, err := auth.Register(ctx, RegisterRequest{
Username: "alice",
Password: "hunter2",
Email: "alice@example.com",
Roles: []string{"user", "admin"},
Claims: map[string]any{"ip_address": "127.0.0.1", "user_agent": "test-agent"},
Username: "alice",
Password: "hunter2",
Email: "alice@example.com",
Roles: []string{"user", "admin"},
UserLevel: 99,
Claims: map[string]any{"ip_address": "127.0.0.1", "user_agent": "test-agent"},
})
if err != nil {
t.Fatalf("Register() error = %v", err)
@@ -65,8 +66,27 @@ func TestDirectMode_RegisterThenLogin(t *testing.T) {
if regResp.Token == "" || regResp.User == nil {
t.Fatalf("Register() returned incomplete response: %+v", regResp)
}
if len(regResp.User.Roles) != 2 {
t.Errorf("expected 2 roles, got %v", regResp.User.Roles)
if len(regResp.User.Roles) != 0 || regResp.User.UserLevel != 0 {
t.Errorf("client-supplied privileges must be ignored, got level=%d roles=%v", regResp.User.UserLevel, regResp.User.Roles)
}
// Password must be stored as a bcrypt hash, not cleartext.
var stored string
if err := db.QueryRow(`SELECT password FROM users WHERE username = 'alice'`).Scan(&stored); err != nil {
t.Fatal(err)
}
if !isBcryptHash(stored) || stored == "hunter2" {
t.Errorf("password not hashed: %q", stored)
}
if _, err := auth.Login(ctx, LoginRequest{Username: "alice", Password: "wrong"}); err == nil {
t.Error("login with wrong password must fail")
}
if _, err := auth.Login(ctx, LoginRequest{Username: "alice"}); err == nil {
t.Error("login with empty password must fail")
}
if _, err := auth.Login(ctx, LoginRequest{Username: "nobody", Password: "hunter2"}); err == nil {
t.Error("login for unknown user must fail")
}
loginResp, err := auth.Login(ctx, LoginRequest{Username: "alice", Password: "hunter2"})
@@ -465,3 +485,51 @@ func TestDirectMode_OAuthServerClientAndCode(t *testing.T) {
t.Error("expected token to be inactive after revoke")
}
}
func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) {
for _, enabled := range []bool{false, true} {
db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect, UpgradePasswordHash: enabled})
ctx := context.Background()
if _, err := db.Exec(`DELETE FROM users`); err != nil {
t.Fatal(err)
}
if _, err := db.Exec(`INSERT INTO users (username, email, password, user_level, roles, is_active, created_at, updated_at) VALUES ('legacy', 'l@example.com', 'oldpass', 0, '', 1, datetime('now'), datetime('now'))`); err != nil {
t.Fatal(err)
}
if _, err := auth.Login(ctx, LoginRequest{Username: "legacy", Password: "nope"}); err == nil {
t.Fatal("wrong password must fail for legacy row")
}
if _, err := auth.Login(ctx, LoginRequest{Username: "legacy", Password: "oldpass"}); err != nil {
t.Fatalf("legacy login failed (upgrade=%v): %v", enabled, err)
}
var stored string
_ = db.QueryRow(`SELECT password FROM users WHERE username = 'legacy'`).Scan(&stored)
if enabled && !isBcryptHash(stored) {
t.Fatalf("upgrade enabled but password not upgraded: %q", stored)
}
if !enabled && stored != "oldpass" {
t.Fatalf("upgrade must not happen unless enabled, stored=%q", stored)
}
if _, err := auth.Login(ctx, LoginRequest{Username: "legacy", Password: "oldpass"}); err != nil {
t.Fatalf("second login failed (upgrade=%v): %v", enabled, err)
}
}
}
func TestVerifyPasswordEdgeCases(t *testing.T) {
h, _ := hashPassword("pw")
if ok, _ := verifyPassword(h, "pw"); !ok {
t.Error("bcrypt match failed")
}
if ok, _ := verifyPassword("", "pw"); ok {
t.Error("empty stored must not match")
}
if ok, _ := verifyPassword("pw", ""); ok {
t.Error("empty supplied must not match")
}
if _, err := hashPassword(string(make([]byte, 73))); err == nil {
t.Error("73-byte password must be rejected")
}
}
+41 -38
View File
@@ -2,9 +2,12 @@ package security
import (
"context"
"errors"
"fmt"
"reflect"
"strings"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
)
@@ -83,9 +86,13 @@ func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error
// Get row security template
rowSec, err := securityList.GetRowSecurityTemplate(userRef, schema, tablename)
if err != nil {
// No row security defined, allow query to proceed
logger.Debug("No row security for %s.%s@%v: %v", schema, tablename, userRef, err)
return nil
if errors.Is(err, ErrNoRowSecurity) {
// No row security defined for this user/table: nothing to apply.
logger.Debug("No row security for %s.%s", schema, tablename)
return nil
}
// Anything else (including a recovered panic) fails closed.
return fmt.Errorf("row security failed for %s.%s: %w", schema, tablename, err)
}
// Check if user has a blocking rule
@@ -123,21 +130,21 @@ func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error
}
}
// Generate the WHERE clause from template
whereClause := rowSec.GetTemplate(pkName, modelType)
logger.Info("Applying row security filter for user %v on %s.%s: %s",
userRef, schema, tablename, whereClause)
// Apply the WHERE clause to the query
query := secCtx.GetQuery()
if selectQuery, ok := query.(interface {
Where(string, ...interface{}) interface{}
}); ok {
secCtx.SetQuery(selectQuery.Where(whereClause))
} else {
logger.Debug("Query doesn't support Where method, skipping row security")
// Generate the WHERE clause and bind arguments from the template
whereClause, whereArgs, err := rowSec.GetTemplate(pkName, modelType)
if err != nil {
return fmt.Errorf("row security failed for %s.%s: %w", schema, tablename, err)
}
logger.Debug("Applying row security filter on %s.%s: %s", schema, tablename, whereClause)
// A filter that cannot be attached must fail the request; silently
// skipping it would expose every row.
selectQuery, ok := secCtx.GetQuery().(common.SelectQuery)
if !ok {
return fmt.Errorf("row security: query type %T on %s.%s does not support Where", secCtx.GetQuery(), schema, tablename)
}
secCtx.SetQuery(selectQuery.Where(whereClause, whereArgs...))
}
return nil
@@ -181,9 +188,14 @@ func applyColumnSecurity(secCtx SecurityContext, securityList *SecurityList) err
maskedResult, err := securityList.ApplyColumnSecurity(resultValue, modelType, userID, schema, tablename)
if err != nil {
logger.Warn("Column security error: %v", err)
// Don't fail the request, just log the issue
return nil
if errors.Is(err, ErrNoColumnSecurity) {
// No rules for this user/table: nothing to mask.
logger.Debug("No column security for %s.%s", schema, tablename)
return nil
}
// Anything else (including a recovered panic) fails closed rather
// than returning unmasked data.
return fmt.Errorf("column security failed for %s.%s: %w", schema, tablename, err)
}
// Update the result with masked data
@@ -284,7 +296,10 @@ func checkModelUpdateAllowed(secCtx SecurityContext) error {
rules, err = modelregistry.GetModelRulesByName(entity)
}
if err != nil {
return nil // model not registered, allow by default
if errors.Is(err, modelregistry.ErrModelNotFound) {
return nil // model not registered, allow by default
}
return err
}
}
if !rules.CanUpdate {
@@ -308,7 +323,10 @@ func checkModelDeleteAllowed(secCtx SecurityContext) error {
rules, err = modelregistry.GetModelRulesByName(entity)
}
if err != nil {
return nil // model not registered, allow by default
if errors.Is(err, modelregistry.ErrModelNotFound) {
return nil // model not registered, allow by default
}
return err
}
}
if !rules.CanDelete {
@@ -422,20 +440,5 @@ func extractSQLName(tag string) string {
}
func splitTag(tag string, sep rune) []string {
var parts []string
var current string
for _, ch := range tag {
if ch == sep {
if current != "" {
parts = append(parts, current)
current = ""
}
} else {
current += string(ch)
}
}
if current != "" {
parts = append(parts, current)
}
return parts
return strings.FieldsFunc(tag, func(r rune) bool { return r == sep })
}
+99 -18
View File
@@ -3,19 +3,23 @@ package security
import (
"context"
"reflect"
"strings"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/common"
)
// Mock SecurityContext for testing hooks
type mockSecurityContext struct {
ctx context.Context
userID int
hasUser bool
schema string
entity string
model interface{}
query interface{}
result interface{}
ctx context.Context
userID int
hasUser bool
schema string
entity string
model interface{}
query interface{}
result interface{}
userRef any
}
func (m *mockSecurityContext) GetContext() context.Context {
@@ -27,6 +31,9 @@ func (m *mockSecurityContext) GetUserID() (int, bool) {
}
func (m *mockSecurityContext) GetUserRef() (any, bool) {
if m.userRef != nil {
return m.userRef, m.hasUser
}
return m.userID, m.hasUser
}
@@ -194,6 +201,19 @@ func TestLoadSecurityRules(t *testing.T) {
})
}
// recordingQuery is a common.SelectQuery that records Where calls.
type recordingQuery struct {
common.SelectQuery
clauses []string
args [][]any
}
func (q *recordingQuery) Where(query string, args ...interface{}) common.SelectQuery {
q.clauses = append(q.clauses, query)
q.args = append(q.args, args)
return q
}
// Test applyRowSecurity
func TestApplyRowSecurity(t *testing.T) {
type TestModel struct {
@@ -207,6 +227,7 @@ func TestApplyRowSecurity(t *testing.T) {
Tablename: "orders",
Template: "user_id = {UserID}",
HasBlock: false,
UserID: 1,
},
}
secList, _ := NewSecurityList(provider)
@@ -215,11 +236,7 @@ func TestApplyRowSecurity(t *testing.T) {
// Load row security
_, _ = secList.LoadRowSecurity(ctx, 1, "public", "orders", false)
// Mock query that supports Where
type MockQuery struct {
whereClause string
}
mockQuery := &MockQuery{}
mockQuery := &recordingQuery{}
secCtx := &mockSecurityContext{
ctx: ctx,
@@ -236,8 +253,61 @@ func TestApplyRowSecurity(t *testing.T) {
t.Fatalf("expected no error, got %v", err)
}
// Note: The actual WHERE clause application requires a query type that implements Where()
// In a real scenario, this would be a bun.SelectQuery or similar
if len(mockQuery.clauses) != 1 || mockQuery.clauses[0] != "user_id = ?" {
t.Fatalf("expected filter to be attached as %q, got %v", "user_id = ?", mockQuery.clauses)
}
if len(mockQuery.args[0]) != 1 || mockQuery.args[0][0] != 1 {
t.Fatalf("expected bound arg [1], got %v", mockQuery.args[0])
}
})
t.Run("fails closed when filter cannot be attached", func(t *testing.T) {
provider := &mockSecurityProvider{rowSecurity: RowSecurity{
Schema: "public", Tablename: "orders", Template: "user_id = {UserID}", UserID: 1,
}}
secList, _ := NewSecurityList(provider)
ctx := context.Background()
_, _ = secList.LoadRowSecurity(ctx, 1, "public", "orders", false)
secCtx := &mockSecurityContext{
ctx: ctx, userID: 1, hasUser: true, schema: "public", entity: "orders",
model: &TestModel{}, query: struct{}{},
}
if err := ApplyRowSecurity(secCtx, secList); err == nil {
t.Fatal("expected an error when the query does not support Where")
}
})
t.Run("user context is bound as its id, never rendered into SQL", func(t *testing.T) {
uc := &UserContext{UserID: 7, SessionID: "sess_secret", UserName: "x' OR '1'='1"}
provider := &mockSecurityProvider{rowSecurity: RowSecurity{
Schema: "public", Tablename: "orders", Template: "user_id = {UserID}", UserID: uc,
}}
secList, _ := NewSecurityList(provider)
ctx := context.Background()
_, _ = secList.LoadRowSecurity(ctx, uc, "public", "orders", false)
q := &recordingQuery{}
secCtx := &mockSecurityContext{
ctx: ctx, userID: 7, hasUser: true, schema: "public", entity: "orders",
model: &TestModel{}, query: q, userRef: uc,
}
if err := ApplyRowSecurity(secCtx, secList); err != nil {
t.Fatal(err)
}
if len(q.clauses) != 1 || strings.Contains(q.clauses[0], "sess_secret") || strings.Contains(q.clauses[0], "OR") {
t.Fatalf("user data leaked into SQL: %v", q.clauses)
}
if q.args[0][0] != 7 {
t.Fatalf("expected bound user id 7, got %v", q.args[0])
}
})
t.Run("invalid identifier is rejected", func(t *testing.T) {
rs := RowSecurity{Schema: "public", Tablename: "orders; DROP TABLE x", Template: "{TableName}.uid = 1"}
if _, _, err := rs.GetTemplate("id", nil); err == nil {
t.Fatal("expected invalid identifier error")
}
})
t.Run("block access", func(t *testing.T) {
@@ -472,6 +542,7 @@ func TestSecurityIntegration(t *testing.T) {
Tablename: "orders",
Template: "user_id = {UserID}",
HasBlock: false,
UserID: 1,
},
}
@@ -486,6 +557,7 @@ func TestSecurityIntegration(t *testing.T) {
schema: "public",
entity: "orders",
model: &Order{},
query: &recordingQuery{},
}
// Step 1: Load security rules
@@ -549,6 +621,7 @@ func TestRowSecurityGetTemplateIntegration(t *testing.T) {
rowSec RowSecurity
pkName string
expectedPart string // Part of the expected output
expectedArgs []any
}{
{
name: "with all placeholders",
@@ -559,7 +632,8 @@ func TestRowSecurityGetTemplateIntegration(t *testing.T) {
Template: "{PrimaryKeyName} IN (SELECT {PrimaryKeyName} FROM {SchemaName}.{TableName}_access WHERE user_id = {UserID})",
},
pkName: "order_id",
expectedPart: "order_id IN (SELECT order_id FROM sales.orders_access WHERE user_id = 42)",
expectedPart: "order_id IN (SELECT order_id FROM sales.orders_access WHERE user_id = ?)",
expectedArgs: []any{42},
},
{
name: "simple user filter",
@@ -570,18 +644,25 @@ func TestRowSecurityGetTemplateIntegration(t *testing.T) {
Template: "user_id = {UserID}",
},
pkName: "id",
expectedPart: "user_id = 1",
expectedPart: "user_id = ?",
expectedArgs: []any{1},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
modelType := reflect.TypeOf(Model{})
result := tt.rowSec.GetTemplate(tt.pkName, modelType)
result, args, err := tt.rowSec.GetTemplate(tt.pkName, modelType)
if err != nil {
t.Fatalf("GetTemplate() error = %v", err)
}
if result != tt.expectedPart {
t.Errorf("GetTemplate() = %q, want %q", result, tt.expectedPart)
}
if !reflect.DeepEqual(args, tt.expectedArgs) {
t.Errorf("GetTemplate() args = %v, want %v", args, tt.expectedArgs)
}
})
}
}
+2 -2
View File
@@ -520,7 +520,7 @@ func SetSessionCookie(w http.ResponseWriter, loginResp *LoginResponse, opts ...S
maxAge = int(loginResp.ExpiresIn)
}
http.SetCookie(w, &http.Cookie{
http.SetCookie(w, &http.Cookie{ //nolint:gosec // G124: Secure/HttpOnly/SameSite set from options with secure defaults
Name: o.name(),
Value: loginResp.Token,
Path: o.path(),
@@ -563,7 +563,7 @@ func ClearSessionCookie(w http.ResponseWriter, opts ...SessionCookieOptions) {
o = opts[0]
}
http.SetCookie(w, &http.Cookie{
http.SetCookie(w, &http.Cookie{ //nolint:gosec // G124: Secure/HttpOnly/SameSite set from options with secure defaults
Name: o.name(),
Value: "",
Path: o.path(),
+26 -26
View File
@@ -54,10 +54,10 @@ func ExampleOAuth2Google() {
})
// Return user info as JSON
_ = json.NewEncoder(w).Encode(loginResp)
_ = json.NewEncoder(w).Encode(loginResp) //nolint:gosec // G117: intentional: field must be serialized
})
_ = http.ListenAndServe(":8080", router)
_ = http.ListenAndServe(":8080", router) //nolint:gosec // G114: example code only
}
// Example: OAuth2 Authentication with GitHub
@@ -89,10 +89,10 @@ func ExampleOAuth2GitHub() {
return
}
_ = json.NewEncoder(w).Encode(loginResp)
_ = json.NewEncoder(w).Encode(loginResp) //nolint:gosec // G117: intentional: field must be serialized
})
_ = http.ListenAndServe(":8080", router)
_ = http.ListenAndServe(":8080", router) //nolint:gosec // G114: example code only
}
// Example: Custom OAuth2 Provider
@@ -100,7 +100,7 @@ func ExampleOAuth2Custom() {
db, _ := sql.Open("postgres", "connection-string")
// Custom OAuth2 provider configuration
oauth2Auth := NewDatabaseAuthenticator(db).WithOAuth2(OAuth2Config{
oauth2Auth := NewDatabaseAuthenticator(db).WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
ClientID: "your-client-id",
ClientSecret: "your-client-secret",
RedirectURL: "http://localhost:8080/auth/callback",
@@ -142,10 +142,10 @@ func ExampleOAuth2Custom() {
return
}
_ = json.NewEncoder(w).Encode(loginResp)
_ = json.NewEncoder(w).Encode(loginResp) //nolint:gosec // G117: intentional: field must be serialized
})
_ = http.ListenAndServe(":8080", router)
_ = http.ListenAndServe(":8080", router) //nolint:gosec // G114: example code only
}
// Example: Multi-Provider OAuth2 with Security Integration
@@ -190,7 +190,7 @@ func ExampleOAuth2MultiProvider() {
return
}
http.SetCookie(w, &http.Cookie{
http.SetCookie(w, &http.Cookie{ //nolint:gosec // G124: Secure/HttpOnly/SameSite set from options with secure defaults
Name: "session_token",
Value: loginResp.Token,
Path: "/",
@@ -218,7 +218,7 @@ func ExampleOAuth2MultiProvider() {
return
}
http.SetCookie(w, &http.Cookie{
http.SetCookie(w, &http.Cookie{ //nolint:gosec // G124: Secure/HttpOnly/SameSite set from options with secure defaults
Name: "session_token",
Value: loginResp.Token,
Path: "/",
@@ -243,7 +243,7 @@ func ExampleOAuth2MultiProvider() {
_ = json.NewEncoder(w).Encode(userCtx)
})
_ = http.ListenAndServe(":8080", router)
_ = http.ListenAndServe(":8080", router) //nolint:gosec // G114: example code only
}
// Example: OAuth2 with Token Refresh
@@ -294,10 +294,10 @@ func ExampleOAuth2TokenRefresh() {
SameSite: http.SameSiteLaxMode,
})
_ = json.NewEncoder(w).Encode(loginResp)
_ = json.NewEncoder(w).Encode(loginResp) //nolint:gosec // G117: intentional: field must be serialized
})
_ = http.ListenAndServe(":8080", router)
_ = http.ListenAndServe(":8080", router) //nolint:gosec // G114: example code only
}
// Example: OAuth2 Logout
@@ -334,7 +334,7 @@ func ExampleOAuth2Logout() {
}
// Clear cookie
http.SetCookie(w, &http.Cookie{
http.SetCookie(w, &http.Cookie{ //nolint:gosec // G124: Secure/HttpOnly/SameSite set from options with secure defaults
Name: "session_token",
Value: "",
Path: "/",
@@ -346,7 +346,7 @@ func ExampleOAuth2Logout() {
_, _ = w.Write([]byte("Logged out successfully"))
})
_ = http.ListenAndServe(":8080", router)
_ = http.ListenAndServe(":8080", router) //nolint:gosec // G114: example code only
}
// Example: Complete OAuth2 Integration with Database Setup
@@ -393,7 +393,7 @@ func ExampleOAuth2Complete() {
return
}
http.SetCookie(w, &http.Cookie{
http.SetCookie(w, &http.Cookie{ //nolint:gosec // G124: Secure/HttpOnly/SameSite set from options with secure defaults
Name: "session_token",
Value: loginResp.Token,
Path: "/",
@@ -426,7 +426,7 @@ func ExampleOAuth2Complete() {
UserID: userCtx.UserID,
})
http.SetCookie(w, &http.Cookie{
http.SetCookie(w, &http.Cookie{ //nolint:gosec // G124: Secure/HttpOnly/SameSite set from options with secure defaults
Name: "session_token",
Value: "",
Path: "/",
@@ -437,7 +437,7 @@ func ExampleOAuth2Complete() {
http.Redirect(w, r, "/", http.StatusTemporaryRedirect)
})
_ = http.ListenAndServe(":8080", router)
_ = http.ListenAndServe(":8080", router) //nolint:gosec // G114: example code only
}
func setupOAuth2Tables(db *sql.DB) {
@@ -488,7 +488,7 @@ func ExampleOAuth2AllProviders() {
// Create authenticator with ALL OAuth2 providers
auth := NewDatabaseAuthenticator(db).
WithOAuth2(OAuth2Config{
WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
ClientID: "google-client-id",
ClientSecret: "google-client-secret",
RedirectURL: "http://localhost:8080/auth/google/callback",
@@ -498,7 +498,7 @@ func ExampleOAuth2AllProviders() {
UserInfoURL: "https://www.googleapis.com/oauth2/v2/userinfo",
ProviderName: "google",
}).
WithOAuth2(OAuth2Config{
WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
ClientID: "github-client-id",
ClientSecret: "github-client-secret",
RedirectURL: "http://localhost:8080/auth/github/callback",
@@ -508,7 +508,7 @@ func ExampleOAuth2AllProviders() {
UserInfoURL: "https://api.github.com/user",
ProviderName: "github",
}).
WithOAuth2(OAuth2Config{
WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
ClientID: "microsoft-client-id",
ClientSecret: "microsoft-client-secret",
RedirectURL: "http://localhost:8080/auth/microsoft/callback",
@@ -518,7 +518,7 @@ func ExampleOAuth2AllProviders() {
UserInfoURL: "https://graph.microsoft.com/v1.0/me",
ProviderName: "microsoft",
}).
WithOAuth2(OAuth2Config{
WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
ClientID: "facebook-client-id",
ClientSecret: "facebook-client-secret",
RedirectURL: "http://localhost:8080/auth/facebook/callback",
@@ -547,7 +547,7 @@ func ExampleOAuth2AllProviders() {
http.Error(w, err.Error(), http.StatusUnauthorized)
return
}
_ = json.NewEncoder(w).Encode(loginResp)
_ = json.NewEncoder(w).Encode(loginResp) //nolint:gosec // G117: intentional: field must be serialized
})
// GitHub routes
@@ -562,7 +562,7 @@ func ExampleOAuth2AllProviders() {
http.Error(w, err.Error(), http.StatusUnauthorized)
return
}
_ = json.NewEncoder(w).Encode(loginResp)
_ = json.NewEncoder(w).Encode(loginResp) //nolint:gosec // G117: intentional: field must be serialized
})
// Microsoft routes
@@ -577,7 +577,7 @@ func ExampleOAuth2AllProviders() {
http.Error(w, err.Error(), http.StatusUnauthorized)
return
}
_ = json.NewEncoder(w).Encode(loginResp)
_ = json.NewEncoder(w).Encode(loginResp) //nolint:gosec // G117: intentional: field must be serialized
})
// Facebook routes
@@ -592,7 +592,7 @@ func ExampleOAuth2AllProviders() {
http.Error(w, err.Error(), http.StatusUnauthorized)
return
}
_ = json.NewEncoder(w).Encode(loginResp)
_ = json.NewEncoder(w).Encode(loginResp) //nolint:gosec // G117: intentional: field must be serialized
})
// Create security list for protected routes
@@ -611,5 +611,5 @@ func ExampleOAuth2AllProviders() {
_ = json.NewEncoder(w).Encode(userCtx)
})
_ = http.ListenAndServe(":8080", router)
_ = http.ListenAndServe(":8080", router) //nolint:gosec // G114: example code only
}
+36 -5
View File
@@ -12,6 +12,8 @@ import (
"sync"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"golang.org/x/oauth2"
)
@@ -39,6 +41,8 @@ type OAuth2Provider struct {
providerName string
states map[string]time.Time // state -> expiry time
statesMutex sync.RWMutex
stopCh chan struct{} // closed to stop cleanupStates
stopOnce sync.Once
}
// WithOAuth2 configures OAuth2 support for the DatabaseAuthenticator
@@ -68,6 +72,7 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
userInfoParser: cfg.UserInfoParser,
providerName: cfg.ProviderName,
states: make(map[string]time.Time),
stopCh: make(chan struct{}),
}
// Initialize providers map if needed
@@ -77,6 +82,9 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
}
// Register provider
if old := a.oauth2Providers[cfg.ProviderName]; old != nil {
old.stop() // replaced provider: stop its cleanup goroutine
}
a.oauth2Providers[cfg.ProviderName] = provider
a.oauth2ProvidersMutex.Unlock()
@@ -335,10 +343,16 @@ func (p *OAuth2Provider) validateState(state string) bool {
// cleanupStates removes expired states periodically
func (p *OAuth2Provider) cleanupStates() {
defer logger.CatchPanic("OAuth2Provider.cleanupStates")()
ticker := time.NewTicker(5 * time.Minute)
defer ticker.Stop()
for range ticker.C {
for {
select {
case <-p.stopCh:
return
case <-ticker.C:
}
p.statesMutex.Lock()
now := time.Now()
for state, expiry := range p.states {
@@ -350,6 +364,23 @@ func (p *OAuth2Provider) cleanupStates() {
}
}
// stop terminates the cleanup goroutine; safe to call more than once.
func (p *OAuth2Provider) stop() {
p.stopOnce.Do(func() { close(p.stopCh) })
}
// Close stops the background OAuth2 state cleanup goroutines and waits for
// in-flight session activity updates. It is safe to call more than once.
func (a *DatabaseAuthenticator) Close() error {
a.oauth2ProvidersMutex.RLock()
for _, p := range a.oauth2Providers {
p.stop()
}
a.oauth2ProvidersMutex.RUnlock()
a.activityWG.Wait()
return nil
}
// defaultOAuth2UserInfoParser parses standard OAuth2 user info claims
func defaultOAuth2UserInfoParser(userInfo map[string]any) (*UserContext, error) {
ctx := &UserContext{
@@ -441,7 +472,7 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
// NewGoogleAuthenticator creates a DatabaseAuthenticator configured for Google OAuth2
func NewGoogleAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
auth := NewDatabaseAuthenticator(db)
return auth.WithOAuth2(OAuth2Config{
return auth.WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
ClientID: clientID,
ClientSecret: clientSecret,
RedirectURL: redirectURL,
@@ -456,7 +487,7 @@ func NewGoogleAuthenticator(clientID, clientSecret, redirectURL string, db *sql.
// NewGitHubAuthenticator creates a DatabaseAuthenticator configured for GitHub OAuth2
func NewGitHubAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
auth := NewDatabaseAuthenticator(db)
return auth.WithOAuth2(OAuth2Config{
return auth.WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
ClientID: clientID,
ClientSecret: clientSecret,
RedirectURL: redirectURL,
@@ -471,7 +502,7 @@ func NewGitHubAuthenticator(clientID, clientSecret, redirectURL string, db *sql.
// NewMicrosoftAuthenticator creates a DatabaseAuthenticator configured for Microsoft OAuth2
func NewMicrosoftAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
auth := NewDatabaseAuthenticator(db)
return auth.WithOAuth2(OAuth2Config{
return auth.WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
ClientID: clientID,
ClientSecret: clientSecret,
RedirectURL: redirectURL,
@@ -486,7 +517,7 @@ func NewMicrosoftAuthenticator(clientID, clientSecret, redirectURL string, db *s
// NewFacebookAuthenticator creates a DatabaseAuthenticator configured for Facebook OAuth2
func NewFacebookAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
auth := NewDatabaseAuthenticator(db)
return auth.WithOAuth2(OAuth2Config{
return auth.WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
ClientID: clientID,
ClientSecret: clientSecret,
RedirectURL: redirectURL,
+21 -21
View File
@@ -305,7 +305,7 @@ func (s *OAuthServer) serverMetadata() map[string]interface{} {
func (s *OAuthServer) metadataHandler(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(s.serverMetadata()) //nolint:errcheck
json.NewEncoder(w).Encode(s.serverMetadata()) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
// --------------------------------------------------------------------------
@@ -318,7 +318,7 @@ func (s *OAuthServer) openIDConfigurationHandler(w http.ResponseWriter, r *http.
meta["id_token_signing_alg_values_supported"] = []string{"RS256"}
meta["claims_supported"] = []string{"sub", "iss", "aud", "exp", "iat", "email", "preferred_username"}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(meta) //nolint:errcheck
json.NewEncoder(w).Encode(meta) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
// --------------------------------------------------------------------------
@@ -333,7 +333,7 @@ func (s *OAuthServer) protectedResourceHandler(w http.ResponseWriter, r *http.Re
"bearer_methods_supported": []string{"header"},
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(meta) //nolint:errcheck
json.NewEncoder(w).Encode(meta) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
// --------------------------------------------------------------------------
@@ -343,7 +343,7 @@ func (s *OAuthServer) protectedResourceHandler(w http.ResponseWriter, r *http.Re
func (s *OAuthServer) jwksHandler(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
if s.signingKey == nil {
json.NewEncoder(w).Encode(map[string]interface{}{"keys": []interface{}{}}) //nolint:errcheck
json.NewEncoder(w).Encode(map[string]interface{}{"keys": []interface{}{}}) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
return
}
pub := s.signingKey.PublicKey
@@ -355,7 +355,7 @@ func (s *OAuthServer) jwksHandler(w http.ResponseWriter, r *http.Request) {
"n": base64.RawURLEncoding.EncodeToString(pub.N.Bytes()),
"e": base64.RawURLEncoding.EncodeToString(bigEndianBytes(pub.E)),
}
json.NewEncoder(w).Encode(map[string]interface{}{"keys": []interface{}{jwk}}) //nolint:errcheck
json.NewEncoder(w).Encode(map[string]interface{}{"keys": []interface{}{jwk}}) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
// --------------------------------------------------------------------------
@@ -392,7 +392,7 @@ func (s *OAuthServer) userinfoHandler(w http.ResponseWriter, r *http.Request) {
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]interface{}{ //nolint:errcheck
json.NewEncoder(w).Encode(map[string]interface{}{ //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
"sub": info.Sub,
"preferred_username": info.Username,
"email": info.Email,
@@ -507,7 +507,7 @@ func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusCreated)
json.NewEncoder(w).Encode(resp) //nolint:errcheck
json.NewEncoder(w).Encode(resp) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
// --------------------------------------------------------------------------
@@ -995,7 +995,7 @@ func (s *OAuthServer) revokeHandler(w http.ResponseWriter, r *http.Request) {
}
if s.auth != nil {
s.auth.OAuthRevokeToken(r.Context(), token) //nolint:errcheck
s.auth.OAuthRevokeToken(r.Context(), token) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
} else {
// In external-provider-only mode, attempt revocation via the first provider's auth.
s.mu.RLock()
@@ -1005,7 +1005,7 @@ func (s *OAuthServer) revokeHandler(w http.ResponseWriter, r *http.Request) {
}
s.mu.RUnlock()
if providerAuth != nil {
providerAuth.OAuthRevokeToken(r.Context(), token) //nolint:errcheck
providerAuth.OAuthRevokeToken(r.Context(), token) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
}
w.WriteHeader(http.StatusOK)
@@ -1022,14 +1022,14 @@ func (s *OAuthServer) introspectHandler(w http.ResponseWriter, r *http.Request)
}
if err := r.ParseForm(); err != nil {
w.Header().Set("Content-Type", "application/json")
w.Write([]byte(`{"active":false}`)) //nolint:errcheck
w.Write([]byte(`{"active":false}`)) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
return
}
token := r.FormValue("token")
w.Header().Set("Content-Type", "application/json")
if token == "" {
w.Write([]byte(`{"active":false}`)) //nolint:errcheck
w.Write([]byte(`{"active":false}`)) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
return
}
@@ -1043,16 +1043,16 @@ func (s *OAuthServer) introspectHandler(w http.ResponseWriter, r *http.Request)
s.mu.RUnlock()
}
if authToUse == nil {
w.Write([]byte(`{"active":false}`)) //nolint:errcheck
w.Write([]byte(`{"active":false}`)) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
return
}
info, err := authToUse.OAuthIntrospectToken(r.Context(), token)
if err != nil {
w.Write([]byte(`{"active":false}`)) //nolint:errcheck
w.Write([]byte(`{"active":false}`)) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
return
}
json.NewEncoder(w).Encode(info) //nolint:errcheck
json.NewEncoder(w).Encode(info) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
// --------------------------------------------------------------------------
@@ -1063,13 +1063,13 @@ func (s *OAuthServer) renderLoginForm(w http.ResponseWriter, r *http.Request, cl
w.Header().Set("Content-Type", "text/html; charset=utf-8")
errHTML := ""
if errMsg != "" {
errHTML = `<p style="color:red">` + errMsg + `</p>`
errHTML = `<p style="color:red">` + htmlEscape(errMsg) + `</p>`
}
fmt.Fprintf(w, loginFormHTML,
s.cfg.LoginTitle,
s.cfg.LoginTitle,
fmt.Fprintf(w, loginFormHTML, //nolint:gosec // G705: output is HTML-escaped
htmlEscape(s.cfg.LoginTitle),
htmlEscape(s.cfg.LoginTitle),
errHTML,
clientID,
htmlEscape(clientID),
htmlEscape(redirectURI),
htmlEscape(clientState),
htmlEscape(codeChallenge),
@@ -1195,7 +1195,7 @@ func (s *OAuthServer) writeOAuthToken(w http.ResponseWriter, r *http.Request, ac
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Pragma", "no-cache")
json.NewEncoder(w).Encode(resp) //nolint:errcheck
json.NewEncoder(w).Encode(resp) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
// buildIDToken issues an OIDC id_token for the just-issued access token by reusing the
@@ -1294,7 +1294,7 @@ func writeOAuthError(w http.ResponseWriter, errCode, description string, status
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
json.NewEncoder(w).Encode(resp) //nolint:errcheck
json.NewEncoder(w).Encode(resp) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
func htmlEscape(s string) string {
+1 -1
View File
@@ -117,7 +117,7 @@ func (a *DatabaseAuthenticator) OAuthSaveCode(ctx context.Context, code *OAuthCo
return a.oauthSaveCodeDirect(ctx, code)
}
input, err := json.Marshal(code)
input, err := json.Marshal(code) //nolint:gosec // G117: intentional: field must be serialized
if err != nil {
return fmt.Errorf("failed to marshal code: %w", err)
}
+1 -1
View File
@@ -239,7 +239,7 @@ func PasskeyHTTPHandlersExample(auth *DatabaseAuthenticator) {
})
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(loginResponse)
_ = json.NewEncoder(w).Encode(loginResponse) //nolint:gosec // G117: intentional: field must be serialized
})
// List credentials endpoint
+67
View File
@@ -0,0 +1,67 @@
package security
import (
"crypto/subtle"
"errors"
"strings"
"sync"
"golang.org/x/crypto/bcrypt"
)
// bcrypt only considers the first 72 bytes of input; longer passwords are
// rejected rather than silently truncated.
const maxPasswordBytes = 72
var errPasswordTooLong = errors.New("password must be at most 72 bytes")
func hashPassword(password string) (string, error) {
if len(password) > maxPasswordBytes {
return "", errPasswordTooLong
}
h, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
if err != nil {
return "", err
}
return string(h), nil
}
func isBcryptHash(s string) bool {
return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$")
}
// verifyPassword checks supplied against the stored value. A stored bcrypt hash
// is compared with bcrypt. A legacy cleartext value (written before hashing was
// implemented) is compared in constant time and, on a match, needsRehash is true
// so the caller can upgrade the row to a bcrypt hash. An empty stored value
// (e.g. an OAuth2-only user) never matches.
func verifyPassword(stored, supplied string) (ok, needsRehash bool) {
if stored == "" || supplied == "" || len(supplied) > maxPasswordBytes {
return false, false
}
if isBcryptHash(stored) {
return bcrypt.CompareHashAndPassword([]byte(stored), []byte(supplied)) == nil, false
}
if subtle.ConstantTimeCompare([]byte(stored), []byte(supplied)) == 1 {
return true, true
}
return false, false
}
var (
dummyHashOnce sync.Once
dummyHash string
)
// burnPasswordCheck spends roughly one bcrypt comparison so an unknown username
// costs about the same as a wrong password.
func burnPasswordCheck(supplied string) {
dummyHashOnce.Do(func() {
h, _ := bcrypt.GenerateFromPassword([]byte("resolvespec-dummy"), bcrypt.DefaultCost)
dummyHash = string(h)
})
if len(supplied) > maxPasswordBytes {
supplied = supplied[:maxPasswordBytes]
}
_ = bcrypt.CompareHashAndPassword([]byte(dummyHash), []byte(supplied))
}
+195 -47
View File
@@ -2,10 +2,13 @@ package security
import (
"context"
"errors"
"fmt"
"reflect"
"regexp"
"strings"
"sync"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
@@ -40,13 +43,76 @@ type RowSecurity struct {
UserID any `json:"user_id"`
}
func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Type) string {
// safeIdentRe matches an unquoted SQL identifier. Identifiers substituted into a
// row-security template must match it; anything else is rejected.
var safeIdentRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
// ErrNoRowSecurity is returned by GetRowSecurityTemplate when no row security
// entry is loaded for the user and table. It means "no rules", as opposed to a
// failure, which callers must treat as fatal.
var ErrNoRowSecurity = errors.New("no row security data")
// ErrNoColumnSecurity is the column-security equivalent of ErrNoRowSecurity.
var ErrNoColumnSecurity = errors.New("no column security data")
// userIDScalar reduces the opaque user reference to a scalar that is safe to
// bind as a query argument. A *UserContext is reduced to its UserID; other
// structured values are rejected rather than stringified into SQL.
func userIDScalar(ref any) (any, error) {
switch v := ref.(type) {
case nil:
return nil, fmt.Errorf("row security: no user reference")
case *UserContext:
if v == nil {
return nil, fmt.Errorf("row security: nil user context")
}
return v.UserID, nil
case UserContext:
return v.UserID, nil
case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64:
return v, nil
case string:
return v, nil
default:
return nil, fmt.Errorf("row security: unsupported user reference type %T", ref)
}
}
// GetTemplate expands the row-security template into a WHERE clause and its
// bind arguments. {PrimaryKeyName}, {TableName} and {SchemaName} are validated
// identifiers substituted in place; every {UserID} becomes a `?` placeholder
// with the user reference bound as an argument, so user data never reaches the
// SQL text.
func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Type) (clause string, args []any, err error) {
str := m.Template
str = strings.ReplaceAll(str, "{PrimaryKeyName}", pPrimaryKeyName)
str = strings.ReplaceAll(str, "{TableName}", m.Tablename)
str = strings.ReplaceAll(str, "{SchemaName}", m.Schema)
str = strings.ReplaceAll(str, "{UserID}", fmt.Sprintf("%v", m.UserID))
return str
for placeholder, ident := range map[string]string{
"{PrimaryKeyName}": pPrimaryKeyName,
"{TableName}": m.Tablename,
"{SchemaName}": m.Schema,
} {
if !strings.Contains(str, placeholder) {
continue
}
if !safeIdentRe.MatchString(ident) {
return "", nil, fmt.Errorf("row security: invalid identifier %q for %s", ident, placeholder)
}
str = strings.ReplaceAll(str, placeholder, ident)
}
n := strings.Count(str, "{UserID}")
if n == 0 {
return str, nil, nil
}
uid, err := userIDScalar(m.UserID)
if err != nil {
return "", nil, err
}
args = make([]any, n)
for i := range args {
args[i] = uid
}
return strings.ReplaceAll(str, "{UserID}", "?"), args, nil
}
// SecurityList manages security state and caching
@@ -58,8 +124,25 @@ type SecurityList struct {
ColumnSecurity map[string][]ColumnSecurity
RowSecurityMutex sync.RWMutex
RowSecurity map[string]RowSecurity
// Expiry bookkeeping for the two caches above; guarded by the same mutexes.
colSecExpiry map[string]time.Time
rowSecExpiry map[string]time.Time
lastColPrune time.Time
lastRowPrune time.Time
}
const (
// securityCacheTTL is how long loaded rules are served without re-querying
// the provider. Revoked rules take up to this long to take effect.
securityCacheTTL = 30 * time.Second
// securityLoadTimeout bounds a single provider call.
securityLoadTimeout = 10 * time.Second
// securityPruneGrace keeps expired entries around long enough that a
// request that loaded them can still read them.
securityPruneGrace = securityCacheTTL
)
// NewSecurityList creates a new security list with the given provider
func NewSecurityList(provider SecurityProvider) (*SecurityList, error) {
if provider == nil {
@@ -85,7 +168,8 @@ const SECURITY_CONTEXT_KEY CONTEXT_KEY = "SecurityList"
func maskString(pString string, maskStart, maskEnd int, maskChar string, invert bool) string {
strLen := len(pString)
middleIndex := (strLen / 2)
newStr := ""
var newStr strings.Builder
newStr.Grow(strLen)
if maskStart == 0 && maskEnd == 0 {
maskStart = strLen
maskEnd = strLen
@@ -101,32 +185,29 @@ func maskString(pString string, maskStart, maskEnd int, maskChar string, invert
}
for index, char := range pString {
if invert && index >= middleIndex-maskStart && index <= middleIndex {
newStr += maskChar
newStr.WriteString(maskChar)
continue
}
if invert && index <= middleIndex+maskEnd && index >= middleIndex {
newStr += maskChar
newStr.WriteString(maskChar)
continue
}
if !invert && index <= maskStart {
newStr += maskChar
newStr.WriteString(maskChar)
continue
}
if !invert && index >= strLen-1-maskEnd {
newStr += maskChar
newStr.WriteString(maskChar)
continue
}
newStr += string(char)
newStr.WriteRune(char)
}
return newStr
return newStr.String()
}
func (m *SecurityList) ColumSecurityApplyOnRecord(prevRecord reflect.Value, newRecord reflect.Value, modelType reflect.Type, pUserID int, pSchema, pTablename string) ([]string, error) {
cols := make([]string, 0)
if m.ColumnSecurity == nil {
return cols, fmt.Errorf("security not initialized")
}
if prevRecord.Type() != newRecord.Type() {
logger.Error("prev:%s and new:%s record type mismatch", prevRecord.Type(), newRecord.Type())
@@ -136,9 +217,13 @@ func (m *SecurityList) ColumSecurityApplyOnRecord(prevRecord reflect.Value, newR
m.ColumnSecurityMutex.RLock()
defer m.ColumnSecurityMutex.RUnlock()
if m.ColumnSecurity == nil {
return cols, fmt.Errorf("security not initialized")
}
colsecList, ok := m.ColumnSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
if !ok || colsecList == nil {
return cols, fmt.Errorf("no column security data")
return cols, ErrNoColumnSecurity
}
for i := range colsecList {
@@ -298,19 +383,26 @@ func setColSecValue(fieldsrc reflect.Value, colsec ColumnSecurity, fieldTypeName
return 0, fieldsrc
}
func (m *SecurityList) ApplyColumnSecurity(records reflect.Value, modelType reflect.Type, pUserID int, pSchema, pTablename string) (reflect.Value, error) {
defer logger.CatchPanic("ApplyColumnSecurity")()
func (m *SecurityList) ApplyColumnSecurity(records reflect.Value, modelType reflect.Type, pUserID int, pSchema, pTablename string) (out reflect.Value, err error) {
// A panic must surface as an error: recovering into zero results would
// read as "success, nothing to mask" and let the response go out unmasked.
defer func() {
if r := recover(); r != nil {
out = reflect.Value{}
err = logger.HandlePanic("ApplyColumnSecurity", r)
}
}()
m.ColumnSecurityMutex.RLock()
defer m.ColumnSecurityMutex.RUnlock()
if m.ColumnSecurity == nil {
return records, fmt.Errorf("security not initialized")
}
m.ColumnSecurityMutex.RLock()
defer m.ColumnSecurityMutex.RUnlock()
colsecList, ok := m.ColumnSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
if !ok || colsecList == nil {
return records, fmt.Errorf("nocolumn security data")
return records, ErrNoColumnSecurity
}
for i := range colsecList {
@@ -372,25 +464,49 @@ func (m *SecurityList) LoadColumnSecurity(ctx context.Context, pUserID int, pSch
return fmt.Errorf("security provider not set")
}
m.ColumnSecurityMutex.Lock()
defer m.ColumnSecurityMutex.Unlock()
if m.ColumnSecurity == nil {
m.ColumnSecurity = make(map[string][]ColumnSecurity, 0)
}
secKey := fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)
if pOverwrite || m.ColumnSecurity[secKey] == nil {
m.ColumnSecurity[secKey] = make([]ColumnSecurity, 0)
if !pOverwrite {
m.ColumnSecurityMutex.RLock()
exp, ok := m.colSecExpiry[secKey]
fresh := ok && m.ColumnSecurity[secKey] != nil && time.Now().Before(exp)
m.ColumnSecurityMutex.RUnlock()
if fresh {
return nil
}
}
// Call the provider to load security rules
colSecList, err := m.provider.GetColumnSecurity(ctx, pUserID, pSchema, pTablename)
// Query the provider without holding any lock.
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
defer cancel()
colSecList, err := m.provider.GetColumnSecurity(loadCtx, pUserID, pSchema, pTablename)
if err != nil {
return fmt.Errorf("GetColumnSecurity failed: %v", err)
}
if colSecList == nil {
colSecList = make([]ColumnSecurity, 0)
}
now := time.Now()
m.ColumnSecurityMutex.Lock()
defer m.ColumnSecurityMutex.Unlock()
if m.ColumnSecurity == nil {
m.ColumnSecurity = make(map[string][]ColumnSecurity)
}
if m.colSecExpiry == nil {
m.colSecExpiry = make(map[string]time.Time)
}
m.ColumnSecurity[secKey] = colSecList
m.colSecExpiry[secKey] = now.Add(securityCacheTTL)
if now.Sub(m.lastColPrune) > securityCacheTTL {
m.lastColPrune = now
for k, exp := range m.colSecExpiry {
if now.Sub(exp) > securityPruneGrace {
delete(m.colSecExpiry, k)
delete(m.ColumnSecurity, k)
}
}
}
return nil
}
@@ -421,37 +537,69 @@ func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserRef any, pSchem
return RowSecurity{}, fmt.Errorf("security provider not set")
}
m.RowSecurityMutex.Lock()
defer m.RowSecurityMutex.Unlock()
if m.RowSecurity == nil {
m.RowSecurity = make(map[string]RowSecurity, 0)
}
secKey := fmt.Sprintf("%s.%s@%v", pSchema, pTablename, pUserRef)
// Call the provider to load security rules
record, err := m.provider.GetRowSecurity(ctx, pUserRef, pSchema, pTablename)
if !pOverwrite {
m.RowSecurityMutex.RLock()
exp, ok := m.rowSecExpiry[secKey]
cached, present := m.RowSecurity[secKey]
m.RowSecurityMutex.RUnlock()
if ok && present && time.Now().Before(exp) {
return cached, nil
}
}
// Query the provider without holding any lock.
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
defer cancel()
record, err := m.provider.GetRowSecurity(loadCtx, pUserRef, pSchema, pTablename)
if err != nil {
return RowSecurity{}, fmt.Errorf("GetRowSecurity failed: %v", err)
}
now := time.Now()
m.RowSecurityMutex.Lock()
defer m.RowSecurityMutex.Unlock()
if m.RowSecurity == nil {
m.RowSecurity = make(map[string]RowSecurity)
}
if m.rowSecExpiry == nil {
m.rowSecExpiry = make(map[string]time.Time)
}
m.RowSecurity[secKey] = record
m.rowSecExpiry[secKey] = now.Add(securityCacheTTL)
if now.Sub(m.lastRowPrune) > securityCacheTTL {
m.lastRowPrune = now
for k, exp := range m.rowSecExpiry {
if now.Sub(exp) > securityPruneGrace {
delete(m.rowSecExpiry, k)
delete(m.RowSecurity, k)
}
}
}
return record, nil
}
func (m *SecurityList) GetRowSecurityTemplate(pUserRef any, pSchema, pTablename string) (RowSecurity, error) {
defer logger.CatchPanic("GetRowSecurityTemplate")()
func (m *SecurityList) GetRowSecurityTemplate(pUserRef any, pSchema, pTablename string) (out RowSecurity, err error) {
// A panic must surface as an error: recovering into zero results would
// read as "no row security" and unblock the user.
defer func() {
if r := recover(); r != nil {
out = RowSecurity{}
err = logger.HandlePanic("GetRowSecurityTemplate", r)
}
}()
m.RowSecurityMutex.RLock()
defer m.RowSecurityMutex.RUnlock()
if m.RowSecurity == nil {
return RowSecurity{}, fmt.Errorf("security not initialized")
}
m.RowSecurityMutex.RLock()
defer m.RowSecurityMutex.RUnlock()
rowSec, ok := m.RowSecurity[fmt.Sprintf("%s.%s@%v", pSchema, pTablename, pUserRef)]
if !ok {
return RowSecurity{}, fmt.Errorf("no row security data")
return RowSecurity{}, ErrNoRowSecurity
}
return rowSec, nil
+17 -10
View File
@@ -38,7 +38,8 @@ func (m *mockSecurityProvider) Authenticate(r *http.Request) (*UserContext, erro
return m.authUser, m.authError
}
func (m *mockSecurityProvider) SetAuthenticateCallback(_ func(r *http.Request) (*UserContext, error)) {}
func (m *mockSecurityProvider) SetAuthenticateCallback(_ func(r *http.Request) (*UserContext, error)) {
}
func (m *mockSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error) {
return m.columnSecurity, nil
@@ -78,13 +79,13 @@ func TestNewSecurityList(t *testing.T) {
// Test maskString function
func TestMaskString(t *testing.T) {
tests := []struct {
name string
input string
maskStart int
maskEnd int
maskChar string
invert bool
expected string
name string
input string
maskStart int
maskEnd int
maskChar string
invert bool
expected string
}{
{
name: "mask first 3 characters",
@@ -299,12 +300,18 @@ func TestRowSecurityGetTemplate(t *testing.T) {
UserID: 42,
}
result := rowSec.GetTemplate("order_id", nil)
result, args, err := rowSec.GetTemplate("order_id", nil)
if err != nil {
t.Fatalf("GetTemplate() error = %v", err)
}
expected := "order_id IN (SELECT order_id FROM public.orders_access WHERE user_id = 42)"
expected := "order_id IN (SELECT order_id FROM public.orders_access WHERE user_id = ?)"
if result != expected {
t.Errorf("GetTemplate() = %q, want %q", result, expected)
}
if len(args) != 1 || args[0] != 42 {
t.Errorf("GetTemplate() args = %v, want [42]", args)
}
}
// Test ClearSecurity
+52 -11
View File
@@ -64,6 +64,13 @@ func (a *HeaderAuthenticator) Authenticate(r *http.Request) (*UserContext, error
}, nil
}
// maxAuthTokens caps the comma-separated credentials tried per request so one
// request cannot drive unbounded session lookups.
const maxAuthTokens = 4
// sessionActivityTimeout bounds the detached last-activity update.
const sessionActivityTimeout = 5 * time.Second
// DatabaseAuthenticator provides session-based authentication with database storage
// All database operations go through stored procedures for security and consistency
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
@@ -81,6 +88,13 @@ type DatabaseAuthenticator struct {
queryMode QueryMode
capability *dbCapability
// upgradePasswordHash enables rewriting legacy cleartext passwords as bcrypt
// on successful login (opt-in, see DatabaseAuthenticatorOptions).
upgradePasswordHash bool
// activityWG tracks in-flight asynchronous session activity updates
activityWG sync.WaitGroup
// Cookie session support (optional, gated by enableCookieSession)
enableCookieSession bool
cookieOptions SessionCookieOptions
@@ -120,6 +134,11 @@ type DatabaseAuthenticatorOptions struct {
// CookieOptions.Name (default "session_token") in addition to the Authorization header,
// and LoginWithCookie / LogoutWithCookie automatically set / clear the cookie.
EnableCookieSession bool
// UpgradePasswordHash, when true, rewrites a legacy cleartext password as a
// bcrypt hash after a successful login. It is off by default and is never
// enabled automatically: legacy cleartext values are still accepted at login,
// but stored rows are left untouched unless this is set.
UpgradePasswordHash bool
// CookieOptions configures the session cookie written by LoginWithCookie.
// Only used when EnableCookieSession is true.
CookieOptions SessionCookieOptions
@@ -159,6 +178,7 @@ func NewDatabaseAuthenticatorWithOptions(db *sql.DB, opts DatabaseAuthenticatorO
capability: newDBCapability(),
passkeyProvider: opts.PasskeyProvider,
enableCookieSession: opts.EnableCookieSession,
upgradePasswordHash: opts.UpgradePasswordHash,
cookieOptions: opts.CookieOptions,
authenticateCallback: opts.AuthenticateCallback,
}
@@ -212,7 +232,7 @@ func (a *DatabaseAuthenticator) Login(ctx context.Context, req LoginRequest) (*L
return a.loginDirect(ctx, req)
}
// Convert LoginRequest to JSON
reqJSON, err := json.Marshal(req)
reqJSON, err := json.Marshal(req) //nolint:gosec // G117: intentional: field must be serialized
if err != nil {
return nil, fmt.Errorf("failed to marshal login request: %w", err)
}
@@ -251,7 +271,7 @@ func (a *DatabaseAuthenticator) Register(ctx context.Context, req RegisterReques
return a.registerDirect(ctx, req)
}
// Convert RegisterRequest to JSON
reqJSON, err := json.Marshal(req)
reqJSON, err := json.Marshal(req) //nolint:gosec // G117: intentional: field must be serialized
if err != nil {
return nil, fmt.Errorf("failed to marshal register request: %w", err)
}
@@ -365,7 +385,10 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err
} else {
// Parse Authorization header which may contain multiple comma-separated tokens
// Format: "Token abc, Token def" or "Bearer abc" or just "abc"
rawTokens := strings.Split(sessionToken, ",")
rawTokens := strings.SplitN(sessionToken, ",", maxAuthTokens+2)
if len(rawTokens) > maxAuthTokens {
return nil, fmt.Errorf("too many authorization tokens")
}
for _, token := range rawTokens {
token = strings.TrimSpace(token)
// Remove "Bearer " prefix if present
@@ -411,7 +434,7 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.sqlNames.Session)
return db.QueryRowContext(r.Context(), query, token, reference).Scan(&success, &errorMsg, &userJSON)
return db.QueryRowContext(r.Context(), query, token, reference).Scan(&success, &errorMsg, &userJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
})
if err != nil {
return nil, fmt.Errorf("session query failed: %w", err)
@@ -444,7 +467,17 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err
// Authentication succeeded with this token
// Update last activity timestamp asynchronously
go a.updateSessionActivity(r.Context(), token, &userCtx)
activityCtx := userCtx
// Detach from the request (it is cancelled when the handler returns) but
// keep a deadline, and never let a panic here take the process down.
detached, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), sessionActivityTimeout)
a.activityWG.Add(1)
go func(ctx context.Context, token string) {
defer a.activityWG.Done()
defer cancel()
defer logger.CatchPanic("updateSessionActivity")()
a.updateSessionActivity(ctx, token, &activityCtx)
}(detached, token)
return &userCtx, nil
}
@@ -497,7 +530,7 @@ func (a *DatabaseAuthenticator) updateSessionActivity(ctx context.Context, sessi
_ = a.runDBOpWithReconnect(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2::jsonb)`, a.sqlNames.SessionUpdate)
return db.QueryRowContext(ctx, query, sessionToken, string(userJSON)).Scan(&success, &errorMsg, &updatedUserJSON)
return db.QueryRowContext(ctx, query, sessionToken, string(userJSON)).Scan(&success, &errorMsg, &updatedUserJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
})
}
@@ -587,6 +620,17 @@ type JWTAuthenticator struct {
tableNames *TableNames
queryMode QueryMode
capability *dbCapability
// upgradePasswordHash enables rewriting legacy cleartext passwords as bcrypt
// on successful login. Off by default; enable with WithPasswordHashUpgrade.
upgradePasswordHash bool
}
// WithPasswordHashUpgrade explicitly enables (or disables) upgrading legacy
// cleartext passwords to bcrypt after a successful login. Off by default.
func (a *JWTAuthenticator) WithPasswordHashUpgrade(enabled bool) *JWTAuthenticator {
a.upgradePasswordHash = enabled
return a
}
func NewJWTAuthenticator(secretKey string, db *sql.DB, names ...*SQLNames) *JWTAuthenticator {
@@ -675,7 +719,6 @@ func (a *JWTAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginR
ID int `json:"id"`
Username string `json:"username"`
Email string `json:"email"`
Password string `json:"password"`
UserLevel int `json:"user_level"`
Roles string `json:"roles"`
}
@@ -684,10 +727,8 @@ func (a *JWTAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginR
return nil, fmt.Errorf("failed to parse user data: %w", err)
}
// TODO: Verify password
// if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(req.Password)); err != nil {
// return nil, fmt.Errorf("invalid credentials")
// }
// The password is verified inside resolvespec_jwt_login; the hash is never
// returned to Go.
// Generate token (placeholder - implement JWT signing when library is available)
expiresAt := time.Now().Add(24 * time.Hour)
+80 -18
View File
@@ -10,6 +10,8 @@ import (
"fmt"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// Direct-mode implementations for DatabaseAuthenticator and JWTAuthenticator.
@@ -17,10 +19,10 @@ import (
// parameterized SQL against the configured TableNames, so they work on
// SQLite, MySQL, or Postgres without the resolvespec_* functions installed.
//
// Password verification is intentionally not implemented here: the stored
// procedures never verify the password hash either (see the TODOs in
// database_schema.sql), so Direct mode matches that behavior exactly rather
// than introducing a mismatch between modes.
// Passwords are verified with bcrypt (see password.go). Legacy cleartext rows
// are still accepted at login; they are only rewritten as bcrypt when the
// upgrade is explicitly enabled (UpgradePasswordHash). Registration
// never honours client-supplied user_level/roles.
var (
errUsernameExists = errors.New("username already exists")
@@ -29,22 +31,34 @@ var (
func (a *DatabaseAuthenticator) loginDirect(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
var userID int
var email, roles, programUserTable sql.NullString
var email, roles, programUserTable, storedPassword sql.NullString
var userLevel, programUserID sql.NullInt64
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`SELECT id, email, user_level, roles, program_user_id, program_user_table FROM %s WHERE username = ? AND is_active = ?`,
`SELECT id, email, user_level, roles, program_user_id, program_user_table, password FROM %s WHERE username = ? AND is_active = ?`,
a.tableNames.Users))
return db.QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles, &programUserID, &programUserTable)
return db.QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles, &programUserID, &programUserTable, &storedPassword)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
burnPasswordCheck(req.Password)
return nil, fmt.Errorf("invalid credentials")
}
return nil, fmt.Errorf("login query failed: %w", err)
}
ok, needsRehash := verifyPassword(storedPassword.String, req.Password)
if !ok {
if storedPassword.String == "" {
burnPasswordCheck(req.Password)
}
return nil, fmt.Errorf("invalid credentials")
}
if needsRehash && a.upgradePasswordHash {
a.upgradePasswordHashFor(ctx, userID, req.Password)
}
sessionToken, err := generateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
@@ -86,6 +100,24 @@ func (a *DatabaseAuthenticator) loginDirect(ctx context.Context, req LoginReques
}, nil
}
// upgradePasswordHashFor replaces a legacy cleartext password with a bcrypt hash.
// Only called when the upgrade has been explicitly enabled. Failure is logged
// and ignored: the login itself already succeeded.
func (a *DatabaseAuthenticator) upgradePasswordHashFor(ctx context.Context, userID int, password string) {
h, err := hashPassword(password)
if err != nil {
return
}
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
q := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET password = ?, updated_at = ? WHERE id = ?`, a.tableNames.Users))
_, err := db.ExecContext(ctx, q, h, time.Now(), userID)
return err
})
if err != nil {
logger.Warn("failed to upgrade legacy password hash for user %d: %v", userID, err)
}
}
func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req RegisterRequest) (*LoginResponse, error) {
if req.Username == "" {
return nil, fmt.Errorf("username is required")
@@ -97,12 +129,20 @@ func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req Register
return nil, fmt.Errorf("password is required")
}
rolesStr := strings.Join(req.Roles, ",")
passwordHash, err := hashPassword(req.Password)
if err != nil {
return nil, err
}
// Privileges are never taken from the request: self-registration always
// creates an unprivileged user.
const userLevel = 0
const rolesStr = ""
now := time.Now()
ipAddress, userAgent := claimStrings(req.Claims)
var userID int64
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
var count int
checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT COUNT(*) FROM %s WHERE username = ?`, a.tableNames.Users))
if err := db.QueryRowContext(ctx, checkQuery, req.Username).Scan(&count); err != nil {
@@ -122,7 +162,7 @@ func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req Register
insertQuery := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (username, email, password, user_level, roles, is_active, created_at, updated_at, program_user_id, program_user_table) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
a.tableNames.Users))
res, err := db.ExecContext(ctx, insertQuery, req.Username, req.Email, req.Password, req.UserLevel, rolesStr, true, now, now, 0, "")
res, err := db.ExecContext(ctx, insertQuery, req.Username, req.Email, passwordHash, userLevel, rolesStr, true, now, now, 0, "")
if err != nil {
return err
}
@@ -164,7 +204,7 @@ func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req Register
UserID: int(userID),
UserName: req.Username,
Email: req.Email,
UserLevel: req.UserLevel,
UserLevel: userLevel,
Roles: parseRoles(rolesStr),
SessionID: sessionToken,
ProgramUserID: 0,
@@ -218,7 +258,7 @@ func (a *DatabaseAuthenticator) sessionDirect(ctx context.Context, token string)
FROM %s s JOIN %s u ON s.user_id = u.id
WHERE s.session_token = ? AND s.expires_at > ? AND u.is_active = ?`,
a.tableNames.UserSessions, a.tableNames.Users))
return db.QueryRowContext(ctx, query, token, time.Now(), true).Scan(&userID, &username, &email, &userLevel, &roles, &programUserID, &programUserTable)
return db.QueryRowContext(ctx, query, token, time.Now(), true).Scan(&userID, &username, &email, &userLevel, &roles, &programUserID, &programUserTable) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
@@ -242,7 +282,7 @@ func (a *DatabaseAuthenticator) sessionDirect(ctx context.Context, token string)
func (a *DatabaseAuthenticator) updateSessionActivityDirect(ctx context.Context, sessionToken string) error {
return a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET last_activity_at = ? WHERE session_token = ? AND expires_at > ?`, a.tableNames.UserSessions))
_, err := db.ExecContext(ctx, query, time.Now(), sessionToken, time.Now())
_, err := db.ExecContext(ctx, query, time.Now(), sessionToken, time.Now()) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
return err
})
}
@@ -367,12 +407,17 @@ func (a *DatabaseAuthenticator) completePasswordResetDirect(ctx context.Context,
return fmt.Errorf("new_password is required")
}
newHash, err := hashPassword(req.NewPassword)
if err != nil {
return err
}
hash := sha256.Sum256([]byte(req.Token))
tokenHash := hex.EncodeToString(hash[:])
var resetID, userID int
var expiresAt time.Time
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT id, user_id, expires_at FROM %s WHERE token_hash = ? AND used = ?`, a.tableNames.UserPasswordResets))
return db.QueryRowContext(ctx, query, tokenHash, false).Scan(&resetID, &userID, &expiresAt)
})
@@ -389,7 +434,7 @@ func (a *DatabaseAuthenticator) completePasswordResetDirect(ctx context.Context,
now := time.Now()
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
updUser := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET password = ?, updated_at = ? WHERE id = ?`, a.tableNames.Users))
if _, err := db.ExecContext(ctx, updUser, req.NewPassword, now, userID); err != nil {
if _, err := db.ExecContext(ctx, updUser, newHash, now, userID); err != nil {
return err
}
delSessions := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, a.tableNames.UserSessions))
@@ -424,12 +469,12 @@ func claimStrings(claims map[string]any) (ipAddress, userAgent string) {
// jwtLoginDirect mirrors resolvespec_jwt_login.
func (a *JWTAuthenticator) jwtLoginDirect(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
var userID int
var email, roles sql.NullString
var email, roles, storedPassword sql.NullString
var userLevel sql.NullInt64
runQuery := func() error {
query := rewritePlaceholders(a.getDB(), fmt.Sprintf(`SELECT id, email, user_level, roles FROM %s WHERE username = ? AND is_active = ?`, a.tableNames.Users))
return a.getDB().QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles)
query := rewritePlaceholders(a.getDB(), fmt.Sprintf(`SELECT id, email, user_level, roles, password FROM %s WHERE username = ? AND is_active = ?`, a.tableNames.Users))
return a.getDB().QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles, &storedPassword)
}
err := runQuery()
if isDBClosed(err) {
@@ -439,11 +484,28 @@ func (a *JWTAuthenticator) jwtLoginDirect(ctx context.Context, req LoginRequest)
}
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
burnPasswordCheck(req.Password)
return nil, fmt.Errorf("invalid credentials")
}
return nil, fmt.Errorf("login query failed: %w", err)
}
ok, needsRehash := verifyPassword(storedPassword.String, req.Password)
if !ok {
if storedPassword.String == "" {
burnPasswordCheck(req.Password)
}
return nil, fmt.Errorf("invalid credentials")
}
if needsRehash && a.upgradePasswordHash {
if h, herr := hashPassword(req.Password); herr == nil {
q := rewritePlaceholders(a.getDB(), fmt.Sprintf(`UPDATE %s SET password = ?, updated_at = ? WHERE id = ?`, a.tableNames.Users))
if _, uerr := a.getDB().ExecContext(ctx, q, h, time.Now(), userID); uerr != nil {
logger.Warn("failed to upgrade legacy password hash for user %d: %v", userID, uerr)
}
}
}
expiresAt := time.Now().Add(24 * time.Hour)
tokenString := fmt.Sprintf("token_%d_%d", userID, expiresAt.Unix())
+26 -18
View File
@@ -194,7 +194,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("cached-token-123", "authenticate").
WillReturnRows(rows)
userCtx1, err := auth.Authenticate(req)
userCtx1, err := authenticateSync(auth, req)
if err != nil {
t.Fatalf("first authenticate failed: %v", err)
}
@@ -203,7 +203,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
}
// Second call - should use cache, no database call expected
userCtx2, err := auth.Authenticate(req)
userCtx2, err := authenticateSync(auth, req)
if err != nil {
t.Fatalf("second authenticate failed: %v", err)
}
@@ -229,7 +229,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("expire-token-456", "authenticate").
WillReturnRows(rows1)
_, err := auth.Authenticate(req)
_, err := authenticateSync(auth, req)
if err != nil {
t.Fatalf("first authenticate failed: %v", err)
}
@@ -245,7 +245,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("expire-token-456", "authenticate").
WillReturnRows(rows2)
_, err = auth.Authenticate(req)
_, err = authenticateSync(auth, req)
if err != nil {
t.Fatalf("second authenticate after expiration failed: %v", err)
}
@@ -267,7 +267,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("logout-token-789", "authenticate").
WillReturnRows(rows1)
_, err := auth.Authenticate(req)
_, err := authenticateSync(auth, req)
if err != nil {
t.Fatalf("authenticate failed: %v", err)
}
@@ -296,7 +296,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("logout-token-789", "authenticate").
WillReturnRows(rows2)
_, err = auth.Authenticate(req)
_, err = authenticateSync(auth, req)
if err != nil {
t.Fatalf("authenticate after logout failed: %v", err)
}
@@ -318,7 +318,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("manual-clear-token", "authenticate").
WillReturnRows(rows)
_, err := auth.Authenticate(req)
_, err := authenticateSync(auth, req)
if err != nil {
t.Fatalf("authenticate failed: %v", err)
}
@@ -334,7 +334,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("manual-clear-token", "authenticate").
WillReturnRows(rows2)
_, err = auth.Authenticate(req)
_, err = authenticateSync(auth, req)
if err != nil {
t.Fatalf("authenticate after cache clear failed: %v", err)
}
@@ -356,7 +356,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("user-token-1", "authenticate").
WillReturnRows(rows1)
_, err := auth.Authenticate(req1)
_, err := authenticateSync(auth, req1)
if err != nil {
t.Fatalf("first authenticate failed: %v", err)
}
@@ -371,7 +371,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("user-token-2", "authenticate").
WillReturnRows(rows2)
_, err = auth.Authenticate(req2)
_, err = authenticateSync(auth, req2)
if err != nil {
t.Fatalf("second authenticate failed: %v", err)
}
@@ -387,7 +387,7 @@ func TestDatabaseAuthenticatorCaching(t *testing.T) {
WithArgs("user-token-1", "authenticate").
WillReturnRows(rows3)
_, err = auth.Authenticate(req1)
_, err = authenticateSync(auth, req1)
if err != nil {
t.Fatalf("authenticate after user cache clear failed: %v", err)
}
@@ -496,7 +496,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
WithArgs("test-token-123", "authenticate").
WillReturnRows(rows)
userCtx, err := auth.Authenticate(req)
userCtx, err := authenticateSync(auth, req)
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
@@ -528,7 +528,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
WithArgs("cookie-token-456", "cookie").
WillReturnRows(rows)
userCtx, err := cookieAuth.Authenticate(req)
userCtx, err := authenticateSync(cookieAuth, req)
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
@@ -545,7 +545,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
t.Run("authenticate missing token", func(t *testing.T) {
req := httptest.NewRequest("GET", "/test", nil)
_, err := auth.Authenticate(req)
_, err := authenticateSync(auth, req)
if err == nil {
t.Fatal("expected error when token is missing")
}
@@ -571,7 +571,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
WithArgs("valid-token-123", "authenticate").
WillReturnRows(rows2)
userCtx, err := auth.Authenticate(req)
userCtx, err := authenticateSync(auth, req)
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
@@ -597,7 +597,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
WithArgs("968CA5AE-4F83-4D55-A3C6-51AE4410E03A", "authenticate").
WillReturnRows(rows)
userCtx, err := auth.Authenticate(req)
userCtx, err := authenticateSync(auth, req)
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
@@ -631,7 +631,7 @@ func TestDatabaseAuthenticator(t *testing.T) {
WithArgs("bad-token-2", "authenticate").
WillReturnRows(rows2)
_, err := auth.Authenticate(req)
_, err := authenticateSync(auth, req)
if err == nil {
t.Fatal("expected error when all tokens fail")
}
@@ -891,7 +891,7 @@ func TestDatabaseAuthenticatorReconnectsClosedDBPaths(t *testing.T) {
WithArgs("reconnect-auth-token", "authenticate").
WillReturnRows(reconnectRows)
userCtx, err := auth.Authenticate(req)
userCtx, err := authenticateSync(auth, req)
if err != nil {
t.Fatalf("expected authenticate to reconnect, got %v", err)
}
@@ -1328,3 +1328,11 @@ func TestConfigRowSecurityProvider(t *testing.T) {
}
})
}
// authenticateSync authenticates and waits for the asynchronous session
// activity update so sqlmock expectations are never touched concurrently.
func authenticateSync(auth *DatabaseAuthenticator, req *http.Request) (*UserContext, error) {
userCtx, err := auth.Authenticate(req)
auth.activityWG.Wait()
return userCtx, err
}
+1 -1
View File
@@ -69,7 +69,7 @@ type SQLNames struct {
// DefaultSQLNames returns an SQLNames with all default resolvespec_* values.
func DefaultSQLNames() *SQLNames {
return &SQLNames{
return &SQLNames{ //nolint:gosec // G101: false positive: identifier/example, not a credential
Login: "resolvespec_login",
Register: "resolvespec_register",
Logout: "resolvespec_logout",
+1 -1
View File
@@ -31,7 +31,7 @@ type TableNames struct {
// DefaultTableNames returns a TableNames with all default table names.
func DefaultTableNames() *TableNames {
return &TableNames{
return &TableNames{ //nolint:gosec // G101: false positive: identifier/example, not a credential
Users: "users",
UserSessions: "user_sessions",
TokenBlacklist: "token_blacklist",
+2 -2
View File
@@ -3,7 +3,7 @@ package security
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha1"
"crypto/sha1" //nolint:gosec // G505: SHA-1 is required by RFC 6238/4226 HMAC-TOTP
"crypto/sha256"
"crypto/sha512"
"encoding/base32"
@@ -117,7 +117,7 @@ func (t *TOTPGenerator) GenerateCode(secret string, timestamp time.Time) (string
}
// Calculate counter (time steps since Unix epoch)
counter := uint64(timestamp.Unix()) / uint64(t.config.Period)
counter := uint64(timestamp.Unix()) / uint64(t.config.Period) //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
// Generate HMAC
h := t.getHashFunc()
+2 -2
View File
@@ -493,9 +493,9 @@ func newInstance(cfg Config) (*serverInstance, error) {
if cfg.HTTP2 {
if existing := os.Getenv("GODEBUG"); !strings.Contains(existing, "http2xconnect=1") {
if existing == "" {
os.Setenv("GODEBUG", "http2xconnect=1")
os.Setenv("GODEBUG", "http2xconnect=1") //nolint:gosec // G104: best-effort call, error intentionally ignored
} else {
os.Setenv("GODEBUG", existing+",http2xconnect=1")
os.Setenv("GODEBUG", existing+",http2xconnect=1") //nolint:gosec // G104: best-effort call, error intentionally ignored
}
}
if httpServer.HTTP2 == nil {
+1 -1
View File
@@ -217,7 +217,7 @@ func (s *Service) Handler(fallback http.Handler) http.Handler {
// attempt fails; see ErrorHandler above.
if r.Body != nil && r.Body != http.NoBody {
bodyBytes, err := io.ReadAll(r.Body)
r.Body.Close()
r.Body.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
if err != nil {
http.Error(w, "failed to read request body", http.StatusInternalServerError)
return
+1 -1
View File
@@ -116,7 +116,7 @@ func getCertDirectory() (string, error) {
// isCertificateValid checks if a certificate file exists and is not expired.
func isCertificateValid(certFile string) bool {
// Check if file exists
certData, err := os.ReadFile(certFile)
certData, err := os.ReadFile(certFile) //nolint:gosec // G304: path from trusted server config
if err != nil {
return false
}
+3 -3
View File
@@ -60,7 +60,7 @@ func (f *ZipFile) Read(b []byte) (int, error) {
n, err := f.rc.Read(b)
f.offset += int64(n)
if err == io.EOF {
f.rc.Close()
f.rc.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
f.rc = nil
}
return n, err
@@ -68,7 +68,7 @@ func (f *ZipFile) Read(b []byte) (int, error) {
}
func (f *ZipFile) Seek(offset int64, whence int) (int64, error) {
if f.rc != nil {
f.rc.Close()
f.rc.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
f.rc = nil
}
switch whence {
@@ -83,7 +83,7 @@ func (f *ZipFile) Seek(offset int64, whence int) (int64, error) {
}
f.offset += offset
case io.SeekEnd:
size := int64(f.UncompressedSize64)
size := int64(f.UncompressedSize64) //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
if size+offset < 0 {
return 0, &fs.PathError{Op: "seek", Path: f.Name, Err: fmt.Errorf("negative position")}
}

Some files were not shown because too many files have changed in this diff Show More