mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-09-30 12:01:59 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d3a99550d9 | ||
|
|
8a94d884e7 | ||
|
|
f9c948ca4e | ||
|
|
164ba2b240 | ||
|
|
97fe88b3a6 | ||
|
|
a4e1abc1df | ||
|
|
d7cb111496 | ||
|
|
f66930c3c9 | ||
|
|
9533c3a0ed | ||
|
|
652621a70e | ||
|
|
c7b4530689 | ||
|
|
e1cf72834e | ||
|
|
3657aa94cc | ||
|
|
da1af1487e | ||
|
|
bc8bff7955 | ||
|
|
a74eebc7f3 | ||
|
|
6687a7a5cd | ||
|
|
e8fbbede7e | ||
|
|
7f8982fa35 | ||
|
|
b587cbd3c4 | ||
|
|
a220338eea | ||
|
|
20c67166d0 | ||
|
|
749dad4ed1 | ||
|
|
d6c5740f9c | ||
|
|
817b781c88 | ||
|
|
87eaa9e18c | ||
|
|
4f6878099b | ||
|
|
0d8b136b91 | ||
|
|
6de9be0ae7 | ||
|
|
82e923b16e | ||
|
|
9a664593f0 | ||
|
|
6e3124e4e0 | ||
|
|
d5de48011b | ||
|
|
6bd6a6f164 | ||
|
|
1885ce016b | ||
|
|
eeb7ba04d8 | ||
|
|
e957753ce4 | ||
|
|
206edd4bfd | ||
|
|
f841d58c59 | ||
|
|
cbac47e052 | ||
|
|
3d4f6faa8e | ||
|
|
f259df1258 | ||
|
|
798bb47e71 | ||
|
|
105a5e1b87 | ||
|
|
a68cf83be6 | ||
|
|
dab4940ace | ||
|
|
c7178e0a2b | ||
|
|
0261f121e8 | ||
|
|
93dc1008ee | ||
|
|
c60565e4e0 | ||
|
|
16cc7d350e | ||
|
|
a172c73ab0 | ||
|
|
ef28959c4d | ||
|
|
7c737afc5a | ||
|
|
a70e3e02d0 | ||
|
|
cec8eb5c0f | ||
|
|
06fa3198f2 | ||
|
|
52d3dca1fa | ||
|
|
873e8925d4 | ||
|
|
b23916048a | ||
|
|
47708fc87a | ||
|
|
a85e572732 | ||
|
|
598fd687f6 | ||
|
|
eee83f9dc6 | ||
|
|
8a06aacfb2 | ||
|
|
705c4f8001 | ||
|
|
d648614611 | ||
|
|
3f86eb0f06 | ||
|
|
3dac55cb19 | ||
|
|
bbb2c6d127 | ||
|
|
3fec7b1a90 | ||
|
|
910390f62d | ||
|
|
b9bed67bd7 | ||
|
|
11ef16f75a | ||
|
|
48b72a7631 | ||
|
|
4c512acf25 | ||
|
|
07a402634e | ||
|
|
0e8f8925c6 | ||
|
|
5a359a160b | ||
|
|
a2799fa224 | ||
|
|
1419542650 | ||
|
|
c120b49529 | ||
|
|
66348dac97 | ||
|
|
a87cd18b1b | ||
|
|
29449c93d5 | ||
|
|
3b6e5c75be | ||
|
|
549ccb8468 | ||
|
|
1af9c76337 | ||
|
|
938a2ef3d9 | ||
|
|
69cc3e2839 | ||
|
|
4018af0636 | ||
|
|
c4e79d6950 | ||
|
|
982a0e62ac | ||
|
|
5d459c95a7 | ||
|
|
e9f7726e43 | ||
|
|
3d2251317a | ||
|
|
1ce0ab1ab4 | ||
|
|
1f9b230f7f | ||
|
|
c42c6b28e3 | ||
|
|
57e7503389 | ||
|
|
0308644075 | ||
|
|
e5984f5205 | ||
|
|
76909ae869 | ||
|
|
c90c2984ac | ||
|
|
1ab4ae33e7 | ||
|
|
905457964c | ||
|
|
c42d09238f | ||
|
|
0647a88aba | ||
|
|
3d2e11eeed | ||
|
|
4493bfa40f | ||
|
|
b157379ff8 | ||
|
|
52752d9c8b | ||
|
|
baca5ad29e | ||
|
|
53ab22ce02 | ||
|
|
09a3dc92b9 | ||
|
|
6590cd789a | ||
|
|
4244e838b1 | ||
|
|
c42fa11c1a | ||
|
|
85bb0f7874 | ||
|
|
cd65946191 | ||
|
|
cb416d49c4 | ||
|
|
cb921f2c5e | ||
|
|
1ebe0d7ac3 | ||
|
|
ae9e06c98b | ||
|
|
2ae4d07544 | ||
|
|
49639b6c19 | ||
|
|
8733176cba | ||
|
|
bce27f7ed2 | ||
|
|
987a2a7faf | ||
|
|
157788b73b | ||
|
|
fb051b5577 | ||
|
|
cc9c4337fd | ||
|
|
0aaeff63a2 | ||
|
|
325769be4e | ||
|
|
f79a400772 | ||
|
|
aef1f96c10 | ||
|
|
354ed2a8dc | ||
|
|
dfb63c3328 | ||
|
|
e8d0ab28c3 |
@@ -27,6 +27,17 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
name: coverage-report
|
name: coverage-report
|
||||||
path: coverage.html
|
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:
|
integration-tests:
|
||||||
name: Integration Tests
|
name: Integration Tests
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
|
|||||||
@@ -30,6 +30,7 @@
|
|||||||
"linters": {
|
"linters": {
|
||||||
"enable": [
|
"enable": [
|
||||||
"gocritic",
|
"gocritic",
|
||||||
|
"gosec",
|
||||||
"misspell",
|
"misspell",
|
||||||
"revive"
|
"revive"
|
||||||
],
|
],
|
||||||
|
|||||||
@@ -1,9 +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
|
# Run all unit tests
|
||||||
test-unit:
|
test-unit:
|
||||||
@echo "Running unit tests..."
|
@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)
|
# Run all integration tests (requires PostgreSQL)
|
||||||
test-integration:
|
test-integration:
|
||||||
@@ -11,7 +20,7 @@ test-integration:
|
|||||||
@go test -tags=integration ./pkg/resolvespec ./pkg/restheadspec -v
|
@go test -tags=integration ./pkg/resolvespec ./pkg/restheadspec -v
|
||||||
|
|
||||||
# Run all tests (unit + integration)
|
# 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)
|
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 \
|
@if [ -z "$(VERSION)" ]; then \
|
||||||
@@ -49,7 +58,9 @@ release-version: ## Create and push a release with specific version (use: make r
|
|||||||
|
|
||||||
lint: ## Run linter
|
lint: ## Run linter
|
||||||
@echo "Running linter..."
|
@echo "Running linter..."
|
||||||
@if command -v golangci-lint > /dev/null; then \
|
@if [ -x "$(GOLANGCI_LINT)" ]; then \
|
||||||
|
"$(GOLANGCI_LINT)" run --config=.golangci.json; \
|
||||||
|
elif command -v golangci-lint > /dev/null; then \
|
||||||
golangci-lint run --config=.golangci.json; \
|
golangci-lint run --config=.golangci.json; \
|
||||||
else \
|
else \
|
||||||
echo "golangci-lint not installed. Install with: go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest"; \
|
echo "golangci-lint not installed. Install with: go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest"; \
|
||||||
@@ -58,7 +69,9 @@ lint: ## Run linter
|
|||||||
|
|
||||||
lintfix: ## Run linter
|
lintfix: ## Run linter
|
||||||
@echo "Running linter..."
|
@echo "Running linter..."
|
||||||
@if command -v golangci-lint > /dev/null; then \
|
@if [ -x "$(GOLANGCI_LINT)" ]; then \
|
||||||
|
"$(GOLANGCI_LINT)" run --config=.golangci.json --fix; \
|
||||||
|
elif command -v golangci-lint > /dev/null; then \
|
||||||
golangci-lint run --config=.golangci.json --fix; \
|
golangci-lint run --config=.golangci.json --fix; \
|
||||||
else \
|
else \
|
||||||
echo "golangci-lint not installed. Install with: go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest"; \
|
echo "golangci-lint not installed. Install with: go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest"; \
|
||||||
@@ -107,7 +120,8 @@ coverage-integration:
|
|||||||
|
|
||||||
help:
|
help:
|
||||||
@echo "Available targets:"
|
@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-integration - Run integration tests (requires PostgreSQL)"
|
||||||
@echo " test - Run all tests"
|
@echo " test - Run all tests"
|
||||||
@echo " docker-up - Start PostgreSQL container"
|
@echo " docker-up - Start PostgreSQL container"
|
||||||
|
|||||||
@@ -135,6 +135,73 @@ For complete documentation including setup, headers, lifecycle hooks, cursor pag
|
|||||||
|
|
||||||
For detailed examples of reading data, cursor pagination, recursive CRUD operations, filtering, sorting, and more, see [pkg/resolvespec/README.md](pkg/resolvespec/README.md).
|
For detailed examples of reading data, cursor pagination, recursive CRUD operations, filtering, sorting, and more, see [pkg/resolvespec/README.md](pkg/resolvespec/README.md).
|
||||||
|
|
||||||
|
## PostGIS & Vector (PostgreSQL only)
|
||||||
|
|
||||||
|
First-class support for PostGIS geometry/geography and pgvector columns in `resolvespec` + `restheadspec`. No extra dependencies. On non-Postgres databases the spatial/vector operators simply don't match.
|
||||||
|
|
||||||
|
### Column types (`pkg/spectypes`)
|
||||||
|
|
||||||
|
| Go type | SQL type | Wire / JSON |
|
||||||
|
|--------------------|--------------|--------------------------------------------------------|
|
||||||
|
| `SqlGeometry` | `geometry` | JSON in/out = **GeoJSON**; also accepts EWKT / hex-EWKB |
|
||||||
|
| `SqlGeography` | `geography` | same as `SqlGeometry` |
|
||||||
|
| `SqlVector` | `vector` | `[]float32` ⇄ `[1,2,3]` |
|
||||||
|
| `SqlHalfVector` | `halfvec` | `[]float32` ⇄ `[1,2,3]` |
|
||||||
|
| `SqlSparseVector` | `sparsevec` | `{"dim":8,"indices":[1,4],"values":[0.5,0.2]}` |
|
||||||
|
| `SqlBitVector` | `bit`/`varbit` | bool array or `"1011"` string |
|
||||||
|
|
||||||
|
- Geometry `Value()` emits `SRID=<n>;<WKT>` (PostGIS implicit text→geometry cast; no wrapper function needed).
|
||||||
|
- Declare dimensioned types with a tag: `gorm:"type:vector(1536)"` — the tag wins over the canonical name in metadata/OpenAPI.
|
||||||
|
- Metadata endpoint and OpenAPI schema report `geometry`/`vector`/`halfvec`/`sparsevec`/`bit`.
|
||||||
|
|
||||||
|
### Spatial filter operators
|
||||||
|
|
||||||
|
`value` is a geometry (GeoJSON object, EWKT string, or hex-EWKB) unless noted.
|
||||||
|
|
||||||
|
| Operator | Value shape |
|
||||||
|
|----------|-------------|
|
||||||
|
| `st_intersects`, `st_contains`, `st_within`, `st_covers`, `st_coveredby`, `st_overlaps`, `st_touches`, `st_crosses`, `st_equals`, `st_disjoint` | geometry |
|
||||||
|
| `st_dwithin` | `{"geom": <geometry>, "distance": <meters>}` |
|
||||||
|
| `bbox` (alias `&&`) | geometry, or `{"bbox":[minx,miny,maxx,maxy],"srid":4326}` |
|
||||||
|
|
||||||
|
### Vector similarity filter operators
|
||||||
|
|
||||||
|
| Operator | pgvector op | Value shape |
|
||||||
|
|----------|-------------|-------------|
|
||||||
|
| `l2_within` / `euclidean_within` | `<->` | `{"vector":[...], "distance": <n>}` |
|
||||||
|
| `cosine_within` | `<=>` | same (also `"lt"`/`"lte"`/`"gt"`/`"gte"` instead of `"distance"`) |
|
||||||
|
| `ip_within` / `inner_within` | `<#>` | same |
|
||||||
|
|
||||||
|
### KNN search (ordering + distance column)
|
||||||
|
|
||||||
|
**resolvespec** — `options.vector_search`:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{ "options": { "vector_search": {
|
||||||
|
"column": "embedding",
|
||||||
|
"vector": [0.1, 0.2, 0.3],
|
||||||
|
"metric": "cosine",
|
||||||
|
"as": "_distance",
|
||||||
|
"direction": "asc"
|
||||||
|
}}}
|
||||||
|
```
|
||||||
|
|
||||||
|
Orders rows by distance; when `as` is set, returns the distance as an extra column (all model columns are auto-selected).
|
||||||
|
`metric`: `l2` (default) | `cosine` | `ip`.
|
||||||
|
|
||||||
|
**restheadspec** — headers:
|
||||||
|
|
||||||
|
```HTTP
|
||||||
|
X-Vector-Search-embedding: cosine
|
||||||
|
X-Vector-Search-Vector: [0.1,0.2,0.3]
|
||||||
|
X-Vector-Search-As: _distance
|
||||||
|
X-Vector-Search-Dir: asc
|
||||||
|
```
|
||||||
|
|
||||||
|
Spatial/vector filters via headers: `X-SpatialFilter-<col>` / `X-VectorFilter-<col>` with a JSON operator object, e.g.
|
||||||
|
`X-SpatialFilter-geom: {"op":"st_dwithin","geom":"SRID=4326;POINT(0 0)","distance":1000}`
|
||||||
|
(optional `"logic":"or"`).
|
||||||
|
|
||||||
## Installation
|
## Installation
|
||||||
|
|
||||||
```Shell
|
```Shell
|
||||||
@@ -509,11 +576,32 @@ Centralized management of multiple database connections with support for Postgre
|
|||||||
- Multiple named database connections
|
- Multiple named database connections
|
||||||
- Multi-ORM access (Bun, GORM, Native SQL) sharing the same connection pool
|
- Multi-ORM access (Bun, GORM, Native SQL) sharing the same connection pool
|
||||||
- Automatic SQLite schema translation (`schema.table` → `schema_table`)
|
- Automatic SQLite schema translation (`schema.table` → `schema_table`)
|
||||||
- Health checks with auto-reconnect
|
- Background health checks (report status; they never close the pool)
|
||||||
- Prometheus metrics for monitoring
|
- Prometheus metrics for monitoring
|
||||||
- Configuration-driven via YAML
|
- Configuration-driven via YAML
|
||||||
- Per-connection statistics and management
|
- Per-connection statistics and management
|
||||||
|
|
||||||
|
**How to use it correctly**:
|
||||||
|
|
||||||
|
```go
|
||||||
|
mgr, err := dbmanager.NewManager(cfg) // or dbmanager.SetupManager(cfg) + GetInstance()
|
||||||
|
if err != nil { /* handle */ }
|
||||||
|
if err := mgr.Connect(ctx); err != nil { /* handle */ } // SetupManager does NOT connect
|
||||||
|
defer mgr.Close() // once, at shutdown
|
||||||
|
|
||||||
|
conn, _ := mgr.GetDefault()
|
||||||
|
db, _ := conn.Bun() // or conn.GORM() / conn.Native() / conn.Database()
|
||||||
|
handler := restheadspec.NewHandlerWithBun(db)
|
||||||
|
```
|
||||||
|
|
||||||
|
- **Fetch a handle once and keep it.** `Bun()`, `GORM()`, `Native()` and `Database()` return handles over one long-lived `*sql.DB`. You do not need to re-fetch them per request, and they stay valid for the life of the connection.
|
||||||
|
- **Never close a handle yourself.** Closing a `*bun.DB`, `*gorm.DB` or the `*sql.DB` closes the shared pool for everyone. Only `mgr.Close()` (at shutdown) should close it. After `Close`, the handles are dead.
|
||||||
|
- **Don't reconnect to recover from errors.** `database/sql` already discards bad connections and dials new ones. The manager does not close the pool on errors or failed health checks. `conn.Reconnect(ctx)` is for explicit operator use only (for example after rotating credentials): on PostgreSQL it retires pooled connections without closing the pool, so held handles keep working. Other databases close and reopen the pool, which invalidates handles you already hold.
|
||||||
|
- **Bring your own `*sql.DB`.** `dbmanager.NewConnectionFromDB(name, type, db)` wraps a pool you opened. The manager never closes it (`Close` only logs a warning); you own it and must close it.
|
||||||
|
- **Set deadlines on request contexts.** `query_timeout` is applied to PostgreSQL as `statement_timeout` (server side) and TCP timeouts detect dead sockets, but pass a context with a deadline to your queries so callers fail fast.
|
||||||
|
- **Pool tuning.** Keep `conn_max_idle_time` below the shortest idle timeout of any NAT, load balancer or pgbouncer between you and the database (typically 60-240s). SQLite `:memory:` is pinned to a single connection.
|
||||||
|
- **Health checks** run every `health_check_interval` (default 15s; a negative value disables them) and publish Prometheus metrics. `enable_auto_reconnect` is deprecated and ignored.
|
||||||
|
|
||||||
For documentation, see [pkg/dbmanager/README.md](pkg/dbmanager/README.md).
|
For documentation, see [pkg/dbmanager/README.md](pkg/dbmanager/README.md).
|
||||||
|
|
||||||
#### Cache
|
#### Cache
|
||||||
@@ -524,9 +612,9 @@ For documentation, see [pkg/cache/README.md](pkg/cache/README.md).
|
|||||||
|
|
||||||
#### Security
|
#### Security
|
||||||
|
|
||||||
Authentication and authorization framework with hooks integration.
|
Authentication and authorization framework with hooks integration. Database-backed providers use PostgreSQL stored procedures by default, with a portable Direct mode (plain Go/SQL) for SQLite, MySQL, or Postgres without the procedures installed.
|
||||||
|
|
||||||
For documentation, see [pkg/security/README.md](pkg/security/README.md).
|
For documentation, see [pkg/security/README.md](pkg/security/README.md) (see "Direct Mode" for the SQLite/portable-SQL path).
|
||||||
|
|
||||||
#### Middleware
|
#### Middleware
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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.
|
||||||
@@ -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).
|
||||||
@@ -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).
|
||||||
@@ -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.
|
||||||
@@ -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
@@ -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
@@ -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.
|
||||||
@@ -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.
|
||||||
@@ -24,6 +24,7 @@ import (
|
|||||||
func main() {
|
func main() {
|
||||||
// Load configuration
|
// Load configuration
|
||||||
cfgMgr := config.NewManager()
|
cfgMgr := config.NewManager()
|
||||||
|
config.SetConfigManager(cfgMgr)
|
||||||
if err := cfgMgr.Load(); err != nil {
|
if err := cfgMgr.Load(); err != nil {
|
||||||
log.Fatalf("Failed to load configuration: %v", err)
|
log.Fatalf("Failed to load configuration: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,48 +1,49 @@
|
|||||||
module github.com/bitechdev/ResolveSpec
|
module github.com/bitechdev/ResolveSpec
|
||||||
|
|
||||||
go 1.24.0
|
go 1.25.7
|
||||||
|
|
||||||
toolchain go1.24.6
|
|
||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/DATA-DOG/go-sqlmock v1.5.2
|
github.com/DATA-DOG/go-sqlmock v1.5.2
|
||||||
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf
|
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c
|
||||||
github.com/eclipse/paho.mqtt.golang v1.5.1
|
github.com/eclipse/paho.mqtt.golang v1.5.1
|
||||||
github.com/getsentry/sentry-go v0.40.0
|
github.com/getsentry/sentry-go v0.46.2
|
||||||
github.com/glebarez/sqlite v1.11.0
|
github.com/glebarez/sqlite v1.11.0
|
||||||
|
github.com/golang-jwt/jwt/v5 v5.3.1
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/gorilla/mux v1.8.1
|
github.com/gorilla/mux v1.8.1
|
||||||
github.com/gorilla/websocket v1.5.3
|
github.com/gorilla/websocket v1.5.3
|
||||||
github.com/jackc/pgx/v5 v5.8.0
|
github.com/jackc/pgx/v5 v5.9.2
|
||||||
github.com/klauspost/compress v1.18.2
|
github.com/klauspost/compress v1.18.6
|
||||||
github.com/mark3labs/mcp-go v0.46.0
|
github.com/mark3labs/mcp-go v0.54.0
|
||||||
github.com/mattn/go-sqlite3 v1.14.33
|
github.com/mattn/go-sqlite3 v1.14.44
|
||||||
github.com/microsoft/go-mssqldb v1.9.5
|
github.com/microsoft/go-mssqldb v1.10.0
|
||||||
github.com/mochi-mqtt/server/v2 v2.7.9
|
github.com/mochi-mqtt/server/v2 v2.7.9
|
||||||
github.com/nats-io/nats.go v1.48.0
|
github.com/nats-io/nats.go v1.52.0
|
||||||
github.com/prometheus/client_golang v1.23.2
|
github.com/prometheus/client_golang v1.23.2
|
||||||
github.com/redis/go-redis/v9 v9.17.2
|
github.com/redis/go-redis/v9 v9.19.0
|
||||||
github.com/spf13/viper v1.21.0
|
github.com/spf13/viper v1.21.0
|
||||||
github.com/stretchr/testify v1.11.1
|
github.com/stretchr/testify v1.11.1
|
||||||
github.com/testcontainers/testcontainers-go v0.40.0
|
github.com/testcontainers/testcontainers-go v0.40.0
|
||||||
github.com/tidwall/gjson v1.18.0
|
github.com/tidwall/gjson v1.19.0
|
||||||
github.com/tidwall/sjson v1.2.5
|
github.com/tidwall/sjson v1.2.5
|
||||||
github.com/uptrace/bun v1.2.16
|
github.com/uptrace/bun v1.2.18
|
||||||
github.com/uptrace/bun/dialect/mssqldialect v1.2.16
|
github.com/uptrace/bun/dialect/mssqldialect v1.2.16
|
||||||
github.com/uptrace/bun/dialect/pgdialect v1.2.16
|
github.com/uptrace/bun/dialect/pgdialect v1.2.16
|
||||||
github.com/uptrace/bun/dialect/sqlitedialect v1.2.16
|
github.com/uptrace/bun/dialect/sqlitedialect v1.2.16
|
||||||
github.com/uptrace/bun/driver/sqliteshim v1.2.16
|
github.com/uptrace/bun/driver/sqliteshim v1.2.16
|
||||||
github.com/uptrace/bunrouter v1.0.23
|
github.com/uptrace/bunrouter v1.0.23
|
||||||
go.mongodb.org/mongo-driver v1.17.6
|
go.mongodb.org/mongo-driver v1.17.9
|
||||||
go.opentelemetry.io/otel v1.38.0
|
go.opentelemetry.io/otel v1.44.0
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.38.0
|
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.43.0
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.38.0
|
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.43.0
|
||||||
go.opentelemetry.io/otel/sdk v1.38.0
|
go.opentelemetry.io/otel/sdk v1.44.0
|
||||||
go.opentelemetry.io/otel/trace v1.38.0
|
go.opentelemetry.io/otel/trace v1.44.0
|
||||||
go.uber.org/zap v1.27.1
|
go.uber.org/zap v1.28.0
|
||||||
golang.org/x/crypto v0.46.0
|
golang.org/x/crypto v0.55.0
|
||||||
golang.org/x/oauth2 v0.34.0
|
golang.org/x/oauth2 v0.36.0
|
||||||
golang.org/x/time v0.14.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/postgres v1.6.0
|
||||||
gorm.io/driver/sqlite v1.6.0
|
gorm.io/driver/sqlite v1.6.0
|
||||||
gorm.io/driver/sqlserver v1.6.3
|
gorm.io/driver/sqlserver v1.6.3
|
||||||
@@ -62,8 +63,7 @@ require (
|
|||||||
github.com/containerd/log v0.1.0 // indirect
|
github.com/containerd/log v0.1.0 // indirect
|
||||||
github.com/containerd/platforms v0.2.1 // indirect
|
github.com/containerd/platforms v0.2.1 // indirect
|
||||||
github.com/cpuguy83/dockercfg v0.3.2 // indirect
|
github.com/cpuguy83/dockercfg v0.3.2 // indirect
|
||||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
|
||||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
|
|
||||||
github.com/distribution/reference v0.6.0 // indirect
|
github.com/distribution/reference v0.6.0 // indirect
|
||||||
github.com/docker/docker v28.5.1+incompatible // indirect
|
github.com/docker/docker v28.5.1+incompatible // indirect
|
||||||
github.com/docker/go-connections v0.6.0 // indirect
|
github.com/docker/go-connections v0.6.0 // indirect
|
||||||
@@ -71,17 +71,17 @@ require (
|
|||||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||||
github.com/ebitengine/purego v0.8.4 // indirect
|
github.com/ebitengine/purego v0.8.4 // indirect
|
||||||
github.com/felixge/httpsnoop v1.0.4 // indirect
|
github.com/felixge/httpsnoop v1.0.4 // indirect
|
||||||
github.com/fsnotify/fsnotify v1.9.0 // indirect
|
github.com/fsnotify/fsnotify v1.10.1 // indirect
|
||||||
github.com/glebarez/go-sqlite v1.22.0 // indirect
|
github.com/glebarez/go-sqlite v1.22.0 // indirect
|
||||||
github.com/go-logr/logr v1.4.3 // indirect
|
github.com/go-logr/logr v1.4.3 // indirect
|
||||||
github.com/go-logr/stdr v1.2.2 // indirect
|
github.com/go-logr/stdr v1.2.2 // indirect
|
||||||
github.com/go-ole/go-ole v1.2.6 // indirect
|
github.com/go-ole/go-ole v1.2.6 // indirect
|
||||||
github.com/go-viper/mapstructure/v2 v2.4.0 // indirect
|
github.com/go-viper/mapstructure/v2 v2.5.0 // indirect
|
||||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 // indirect
|
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 // indirect
|
||||||
github.com/golang-sql/sqlexp v0.1.0 // indirect
|
github.com/golang-sql/sqlexp v0.1.0 // indirect
|
||||||
github.com/golang/snappy v1.0.0 // indirect
|
github.com/golang/snappy v1.0.0 // indirect
|
||||||
github.com/google/jsonschema-go v0.4.2 // indirect
|
github.com/google/jsonschema-go v0.4.3 // indirect
|
||||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.2 // indirect
|
github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 // indirect
|
||||||
github.com/jackc/pgpassfile v1.0.0 // indirect
|
github.com/jackc/pgpassfile v1.0.0 // indirect
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
||||||
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
||||||
@@ -89,7 +89,7 @@ require (
|
|||||||
github.com/jinzhu/now v1.1.5 // indirect
|
github.com/jinzhu/now v1.1.5 // indirect
|
||||||
github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0 // indirect
|
github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0 // indirect
|
||||||
github.com/magiconair/properties v1.8.10 // indirect
|
github.com/magiconair/properties v1.8.10 // indirect
|
||||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
github.com/mattn/go-isatty v0.0.22 // indirect
|
||||||
github.com/moby/docker-image-spec v1.3.1 // indirect
|
github.com/moby/docker-image-spec v1.3.1 // indirect
|
||||||
github.com/moby/go-archive v0.1.0 // indirect
|
github.com/moby/go-archive v0.1.0 // indirect
|
||||||
github.com/moby/patternmatcher v0.6.0 // indirect
|
github.com/moby/patternmatcher v0.6.0 // indirect
|
||||||
@@ -97,25 +97,26 @@ require (
|
|||||||
github.com/moby/sys/user v0.4.0 // indirect
|
github.com/moby/sys/user v0.4.0 // indirect
|
||||||
github.com/moby/sys/userns v0.1.0 // indirect
|
github.com/moby/sys/userns v0.1.0 // indirect
|
||||||
github.com/moby/term v0.5.0 // indirect
|
github.com/moby/term v0.5.0 // indirect
|
||||||
github.com/montanaflynn/stats v0.7.1 // indirect
|
github.com/montanaflynn/stats v0.9.0 // indirect
|
||||||
github.com/morikuni/aec v1.0.0 // indirect
|
github.com/morikuni/aec v1.0.0 // indirect
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
||||||
github.com/nats-io/nkeys v0.4.11 // indirect
|
github.com/nats-io/nkeys v0.4.15 // indirect
|
||||||
github.com/nats-io/nuid v1.0.1 // indirect
|
github.com/nats-io/nuid v1.0.1 // indirect
|
||||||
github.com/ncruces/go-strftime v1.0.0 // indirect
|
github.com/ncruces/go-strftime v1.0.0 // indirect
|
||||||
github.com/opencontainers/go-digest v1.0.0 // indirect
|
github.com/opencontainers/go-digest v1.0.0 // indirect
|
||||||
github.com/opencontainers/image-spec v1.1.1 // indirect
|
github.com/opencontainers/image-spec v1.1.1 // indirect
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
|
github.com/pelletier/go-toml/v2 v2.3.1 // indirect
|
||||||
github.com/pkg/errors v0.9.1 // indirect
|
github.com/pkg/errors v0.9.1 // indirect
|
||||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||||
github.com/power-devops/perfstat v0.0.0-20210106213030-5aafc221ea8c // indirect
|
github.com/power-devops/perfstat v0.0.0-20210106213030-5aafc221ea8c // indirect
|
||||||
github.com/prometheus/client_model v0.6.2 // indirect
|
github.com/prometheus/client_model v0.6.2 // indirect
|
||||||
github.com/prometheus/common v0.67.4 // indirect
|
github.com/prometheus/common v0.67.5 // indirect
|
||||||
github.com/prometheus/procfs v0.19.2 // indirect
|
github.com/prometheus/procfs v0.20.1 // indirect
|
||||||
github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect
|
github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||||
github.com/rs/xid v1.4.0 // indirect
|
github.com/rs/xid v1.6.0 // indirect
|
||||||
github.com/sagikazarmark/locafero v0.12.0 // indirect
|
github.com/sagikazarmark/locafero v0.12.0 // indirect
|
||||||
|
github.com/santhosh-tekuri/jsonschema/v6 v6.0.2 // indirect
|
||||||
github.com/shirou/gopsutil/v4 v4.25.6 // indirect
|
github.com/shirou/gopsutil/v4 v4.25.6 // indirect
|
||||||
github.com/shopspring/decimal v1.4.0 // indirect
|
github.com/shopspring/decimal v1.4.0 // indirect
|
||||||
github.com/sirupsen/logrus v1.9.3 // indirect
|
github.com/sirupsen/logrus v1.9.3 // indirect
|
||||||
@@ -137,28 +138,26 @@ require (
|
|||||||
github.com/yosida95/uritemplate/v3 v3.0.2 // indirect
|
github.com/yosida95/uritemplate/v3 v3.0.2 // indirect
|
||||||
github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 // indirect
|
github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 // indirect
|
||||||
github.com/yusufpapurcu/wmi v1.2.4 // indirect
|
github.com/yusufpapurcu/wmi v1.2.4 // indirect
|
||||||
go.opentelemetry.io/auto/sdk v1.1.0 // indirect
|
go.opentelemetry.io/auto/sdk v1.2.1 // indirect
|
||||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 // indirect
|
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0 // indirect
|
||||||
go.opentelemetry.io/otel/metric v1.38.0 // indirect
|
go.opentelemetry.io/otel/metric v1.44.0 // indirect
|
||||||
go.opentelemetry.io/proto/otlp v1.7.1 // indirect
|
go.opentelemetry.io/proto/otlp v1.10.0 // indirect
|
||||||
|
go.uber.org/atomic v1.11.0 // indirect
|
||||||
go.uber.org/multierr v1.11.0 // indirect
|
go.uber.org/multierr v1.11.0 // indirect
|
||||||
go.yaml.in/yaml/v2 v2.4.3 // indirect
|
go.yaml.in/yaml/v2 v2.4.4 // indirect
|
||||||
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
||||||
golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93 // indirect
|
golang.org/x/mod v0.38.0 // indirect
|
||||||
golang.org/x/mod v0.31.0 // indirect
|
golang.org/x/net v0.58.0 // indirect
|
||||||
golang.org/x/net v0.48.0 // indirect
|
golang.org/x/sync v0.22.0 // indirect
|
||||||
golang.org/x/sync v0.19.0 // indirect
|
golang.org/x/text v0.41.0 // indirect
|
||||||
golang.org/x/sys v0.39.0 // indirect
|
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||||
golang.org/x/text v0.32.0 // indirect
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||||
google.golang.org/genproto/googleapis/api v0.0.0-20250825161204-c5933d9347a5 // indirect
|
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20250825161204-c5933d9347a5 // indirect
|
|
||||||
google.golang.org/grpc v1.75.0 // indirect
|
|
||||||
google.golang.org/protobuf v1.36.11 // indirect
|
google.golang.org/protobuf v1.36.11 // indirect
|
||||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||||
modernc.org/libc v1.67.4 // indirect
|
modernc.org/libc v1.72.3 // indirect
|
||||||
modernc.org/mathutil v1.7.1 // indirect
|
modernc.org/mathutil v1.7.1 // indirect
|
||||||
modernc.org/memory v1.11.0 // indirect
|
modernc.org/memory v1.11.0 // indirect
|
||||||
modernc.org/sqlite v1.42.2 // indirect
|
modernc.org/sqlite v1.50.1 // indirect
|
||||||
)
|
)
|
||||||
|
|
||||||
replace github.com/uptrace/bun => github.com/warkanum/bun v1.2.17
|
replace github.com/uptrace/bun => github.com/warkanum/bun v1.2.17
|
||||||
|
|||||||
@@ -5,37 +5,37 @@ github.com/AdaLogics/go-fuzz-headers v0.0.0-20240806141605-e8a1dd7889d6/go.mod h
|
|||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.0/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.0/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.1/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.1/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.11.1/go.mod h1:a6xsAQUZg+VsS3TJ05SRp524Hs4pZ/AeFSr5ENf0Yjo=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.11.1/go.mod h1:a6xsAQUZg+VsS3TJ05SRp524Hs4pZ/AeFSr5ENf0Yjo=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.18.0 h1:Gt0j3wceWMwPmiazCa8MzMA0MfhmPIz0Qp0FJ6qcM0U=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1 h1:jHb/wfvRikGdxMXYV3QG/SzUOPYN9KEUUuC0Yd0/vC0=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.18.0/go.mod h1:Ot/6aikWnKWi4l9QB7qVSwa8iMphQNqkWALMoNT3rzM=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1/go.mod h1:pzBXCYn05zvYIrwLgtK8Ap8QcjRg+0i76tMQdWN6wOk=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.3.1/go.mod h1:uE9zaUfEQT/nbQjVi2IblCG9iaLtZsuYZ8ne+PuQ02M=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.3.1/go.mod h1:uE9zaUfEQT/nbQjVi2IblCG9iaLtZsuYZ8ne+PuQ02M=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.6.0/go.mod h1:9kIvujWAA58nmPmWB1m23fyWic1kYZMxD9CxaWn4Qpg=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.6.0/go.mod h1:9kIvujWAA58nmPmWB1m23fyWic1kYZMxD9CxaWn4Qpg=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.10.1 h1:B+blDbyVIG3WaikNxPnhPiJ1MThR03b3vKGtER95TP4=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1 h1:Hk5QBxZQC1jb2Fwj6mpzme37xbCDdNTxU7O9eb5+LB4=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.10.1/go.mod h1:JdM5psgjfBf5fo2uWOZhflPWyDBZ/O/CNAH9CtsuZE4=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1/go.mod h1:IYus9qsFobWIc2YVwe/WPjcnyCkPKtnHAqUYeebc8z0=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.3.0/go.mod h1:okt5dMMTOFjX/aovMlrjvvXoPMBVSPzk9185BT0+eZM=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.3.0/go.mod h1:okt5dMMTOFjX/aovMlrjvvXoPMBVSPzk9185BT0+eZM=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.5.2/go.mod h1:yInRyqWXAuaPrgI7p70+lDDgh3mlBohis29jGMISnmc=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.5.2/go.mod h1:yInRyqWXAuaPrgI7p70+lDDgh3mlBohis29jGMISnmc=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.8.0/go.mod h1:4OG6tQ9EOP/MT0NMjDlRzWoVFxfu9rN9B2X+tlSVktg=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.8.0/go.mod h1:4OG6tQ9EOP/MT0NMjDlRzWoVFxfu9rN9B2X+tlSVktg=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.11.1 h1:FPKJS1T+clwv+OLGt13a8UjqeRuh0O4SJ3lUriThc+4=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.12.0 h1:fhqpLE3UEXi9lPaBRpQ6XuRW0nU7hgg4zlmZZa+a9q4=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.11.1/go.mod h1:j2chePtV91HrC22tGoRX3sGY42uF13WzmmV80/OdVAA=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.12.0/go.mod h1:7dCRMLwisfRH3dBupKeNCioWYUZ4SS09Z14H+7i8ZoY=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.0.1/go.mod h1:GpPjLhVR9dnUoJMyHWSPy71xY9/lcmpzIPZXmF0FCVY=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.0.1/go.mod h1:GpPjLhVR9dnUoJMyHWSPy71xY9/lcmpzIPZXmF0FCVY=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.3.1 h1:Wgf5rZba3YZqeTNJPtvqZoBu1sBN/L4sry+u2U3Y75w=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.4.0 h1:E4MgwLBGeVB5f2MdcIVD3ELVAWpr+WD6MUe1i+tM/PA=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.3.1/go.mod h1:xxCBG/f/4Vbmh2XQJBsOmNdxWUY5j/s27jujKPbQf14=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.4.0/go.mod h1:Y2b/1clN4zsAoUd/pgNAQHjLDnTis/6ROkUfyob6psM=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.0.0/go.mod h1:bTSOgj05NGRuHHhQwAdPnYr9TOdNmKlZTgGLL6nyAdI=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.0.0/go.mod h1:bTSOgj05NGRuHHhQwAdPnYr9TOdNmKlZTgGLL6nyAdI=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.1.1 h1:bFWuoEKg+gImo7pvkiQEFAc8ocibADgXeiLAxWhWmkI=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.2.0 h1:nCYfgcSyHZXJI8J0IWE5MsCGlb2xp9fJiXyxWgmOFg4=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.1.1/go.mod h1:Vih/3yc6yac2JzU4hzpaDupBJP0Flaia9rXXrU8xyww=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.2.0/go.mod h1:ucUjca2JtSZboY8IoUqyQyuuXvwbMBVwFOm0vdQPNhA=
|
||||||
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1 h1:UQHMgLO+TxOElx5B5HZ4hJQsoJ/PvUvKRhJHDQXO8P8=
|
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1 h1:UQHMgLO+TxOElx5B5HZ4hJQsoJ/PvUvKRhJHDQXO8P8=
|
||||||
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E=
|
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.1.1/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
github.com/AzureAD/microsoft-authentication-library-for-go v1.1.1/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.2.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
github.com/AzureAD/microsoft-authentication-library-for-go v1.2.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.4.2 h1:oygO0locgZJe7PpYPXT5A29ZkwJaPqcva7BVeemZOZs=
|
github.com/AzureAD/microsoft-authentication-library-for-go v1.6.0 h1:XRzhVemXdgvJqCH0sFfrBUTnUJSBrBf7++ypk+twtRs=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.4.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
github.com/AzureAD/microsoft-authentication-library-for-go v1.6.0/go.mod h1:HKpQxkWaGLJ+D/5H8QRpyQXA1eKjxkFlOMwck5+33Jk=
|
||||||
github.com/DATA-DOG/go-sqlmock v1.5.2 h1:OcvFkGmslmlZibjAjaHm3L//6LiuBgolP7OputlJIzU=
|
github.com/DATA-DOG/go-sqlmock v1.5.2 h1:OcvFkGmslmlZibjAjaHm3L//6LiuBgolP7OputlJIzU=
|
||||||
github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU=
|
github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU=
|
||||||
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
|
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
|
||||||
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
|
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
|
||||||
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
||||||
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
||||||
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf h1:TqhNAT4zKbTdLa62d2HDBFdvgSbIGB3eJE8HqhgiL9I=
|
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c h1:6Gpm9YYUEQx2T9zMsYolQhr6sjwwGtFitSA0pQsa7a8=
|
||||||
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf/go.mod h1:r5xuitiExdLAJ09PR7vBVENGvp4ZuTBeWTGtxuX3K+c=
|
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c/go.mod h1:r5xuitiExdLAJ09PR7vBVENGvp4ZuTBeWTGtxuX3K+c=
|
||||||
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
|
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
|
||||||
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
|
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
|
||||||
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
|
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
|
||||||
@@ -60,12 +60,13 @@ github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ3
|
|||||||
github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY=
|
github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY=
|
||||||
github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4=
|
github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4=
|
||||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
|
||||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78=
|
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM=
|
||||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc=
|
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk=
|
github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk=
|
||||||
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
|
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
|
||||||
|
github.com/dlclark/regexp2 v1.11.0 h1:G/nrcoOa7ZXlpoa/91N3X7mM3r8eIlMBBJZvsz/mxKI=
|
||||||
|
github.com/dlclark/regexp2 v1.11.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8=
|
||||||
github.com/dnaeon/go-vcr v1.1.0/go.mod h1:M7tiix8f0r6mKKJ3Yq/kqU1OYf3MnfmBWVbPx/yU9ko=
|
github.com/dnaeon/go-vcr v1.1.0/go.mod h1:M7tiix8f0r6mKKJ3Yq/kqU1OYf3MnfmBWVbPx/yU9ko=
|
||||||
github.com/dnaeon/go-vcr v1.2.0/go.mod h1:R4UdLID7HZT3taECzJs4YgbbH6PIGXB6W/sc5OLb6RQ=
|
github.com/dnaeon/go-vcr v1.2.0/go.mod h1:R4UdLID7HZT3taECzJs4YgbbH6PIGXB6W/sc5OLb6RQ=
|
||||||
github.com/docker/docker v28.5.1+incompatible h1:Bm8DchhSD2J6PsFzxC35TZo4TLGR2PdW/E69rU45NhM=
|
github.com/docker/docker v28.5.1+incompatible h1:Bm8DchhSD2J6PsFzxC35TZo4TLGR2PdW/E69rU45NhM=
|
||||||
@@ -84,10 +85,10 @@ github.com/felixge/httpsnoop v1.0.4 h1:NFTV2Zj1bL4mc9sqWACXbQFVBBg2W3GPvqp8/ESS2
|
|||||||
github.com/felixge/httpsnoop v1.0.4/go.mod h1:m8KPJKqk1gH5J9DgRY2ASl2lWCfGKXixSwevea8zH2U=
|
github.com/felixge/httpsnoop v1.0.4/go.mod h1:m8KPJKqk1gH5J9DgRY2ASl2lWCfGKXixSwevea8zH2U=
|
||||||
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
|
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
|
||||||
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
||||||
github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k=
|
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
|
||||||
github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0=
|
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
|
||||||
github.com/getsentry/sentry-go v0.40.0 h1:VTJMN9zbTvqDqPwheRVLcp0qcUcM+8eFivvGocAaSbo=
|
github.com/getsentry/sentry-go v0.46.2 h1:1jhYwrKGa3sIpo/y5iDNXS5wDoT7I1KNzMHrnK6ojns=
|
||||||
github.com/getsentry/sentry-go v0.40.0/go.mod h1:eRXCoh3uvmjQLY6qu63BjUZnaBu5L5WhMV1RwYO8W5s=
|
github.com/getsentry/sentry-go v0.46.2/go.mod h1:evVbw2qotNUdYG8KxXbAdjOQWWvWIwKxpjdZZIvcIPw=
|
||||||
github.com/glebarez/go-sqlite v1.22.0 h1:uAcMJhaA6r3LHMTFgP0SifzgXg46yJkgxqyuyec+ruQ=
|
github.com/glebarez/go-sqlite v1.22.0 h1:uAcMJhaA6r3LHMTFgP0SifzgXg46yJkgxqyuyec+ruQ=
|
||||||
github.com/glebarez/go-sqlite v1.22.0/go.mod h1:PlBIdHe0+aUEFn+r2/uthrWq4FxbzugL0L8Li6yQJbc=
|
github.com/glebarez/go-sqlite v1.22.0/go.mod h1:PlBIdHe0+aUEFn+r2/uthrWq4FxbzugL0L8Li6yQJbc=
|
||||||
github.com/glebarez/sqlite v1.11.0 h1:wSG0irqzP6VurnMEpFGer5Li19RpIRi2qvQz++w0GMw=
|
github.com/glebarez/sqlite v1.11.0 h1:wSG0irqzP6VurnMEpFGer5Li19RpIRi2qvQz++w0GMw=
|
||||||
@@ -101,13 +102,13 @@ github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
|
|||||||
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
|
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
|
||||||
github.com/go-ole/go-ole v1.2.6 h1:/Fpf6oFPoeFik9ty7siob0G6Ke8QvQEuVcuChpwXzpY=
|
github.com/go-ole/go-ole v1.2.6 h1:/Fpf6oFPoeFik9ty7siob0G6Ke8QvQEuVcuChpwXzpY=
|
||||||
github.com/go-ole/go-ole v1.2.6/go.mod h1:pprOEPIfldk/42T2oK7lQ4v4JSDwmV0As9GaiUsvbm0=
|
github.com/go-ole/go-ole v1.2.6/go.mod h1:pprOEPIfldk/42T2oK7lQ4v4JSDwmV0As9GaiUsvbm0=
|
||||||
github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs=
|
github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro=
|
||||||
github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
|
github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
|
||||||
github.com/golang-jwt/jwt/v5 v5.0.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
github.com/golang-jwt/jwt/v5 v5.0.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||||
github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||||
github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo=
|
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
||||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
|
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
|
||||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0=
|
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0=
|
||||||
github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A=
|
github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A=
|
||||||
@@ -120,8 +121,8 @@ github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/
|
|||||||
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||||
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||||
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
||||||
github.com/google/jsonschema-go v0.4.2 h1:tmrUohrwoLZZS/P3x7ex0WAVknEkBZM46iALbcqoRA8=
|
github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0=
|
||||||
github.com/google/jsonschema-go v0.4.2/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE=
|
github.com/google/jsonschema-go v0.4.3/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE=
|
||||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
|
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
|
||||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
|
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
|
||||||
github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||||
@@ -133,8 +134,8 @@ github.com/gorilla/securecookie v1.1.1/go.mod h1:ra0sb63/xPlUeL+yeDciTfxMRAA+MP+
|
|||||||
github.com/gorilla/sessions v1.2.1/go.mod h1:dk2InVEVJ0sfLlnXv9EAgkf6ecYs/i80K/zI+bUmuGM=
|
github.com/gorilla/sessions v1.2.1/go.mod h1:dk2InVEVJ0sfLlnXv9EAgkf6ecYs/i80K/zI+bUmuGM=
|
||||||
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
||||||
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
||||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.2 h1:8Tjv8EJ+pM1xP8mK6egEbD1OgnVTyacbefKhmbLhIhU=
|
github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 h1:5VipnvEpbqr2gA2VbM+nYVbkIF28c5ZQfqCBQ5g2xfk=
|
||||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.2/go.mod h1:pkJQ2tZHJ0aFOVEEot6oZmaVEZcRme73eIFmhiVuRWs=
|
github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0/go.mod h1:Hyl3n6Twe1hvtd9XUXDec4pTvgMSEixRuQKPTMH2bNs=
|
||||||
github.com/hashicorp/go-uuid v1.0.2/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro=
|
github.com/hashicorp/go-uuid v1.0.2/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro=
|
||||||
github.com/hashicorp/go-uuid v1.0.3/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro=
|
github.com/hashicorp/go-uuid v1.0.3/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro=
|
||||||
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
|
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
|
||||||
@@ -143,8 +144,8 @@ github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsI
|
|||||||
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
||||||
github.com/jackc/pgx/v5 v5.8.0 h1:TYPDoleBBme0xGSAX3/+NujXXtpZn9HBONkQC7IEZSo=
|
github.com/jackc/pgx/v5 v5.9.2 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw=
|
||||||
github.com/jackc/pgx/v5 v5.8.0/go.mod h1:QVeDInX2m9VyzvNeiCJVjCkNFqzsNb43204HshNSZKw=
|
github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
|
||||||
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
||||||
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
||||||
github.com/jcmturner/aescts/v2 v2.0.0/go.mod h1:AiaICIRyfYg35RUkr8yESTqvSy7csK90qZ5xfvvsoNs=
|
github.com/jcmturner/aescts/v2 v2.0.0/go.mod h1:AiaICIRyfYg35RUkr8yESTqvSy7csK90qZ5xfvvsoNs=
|
||||||
@@ -160,8 +161,10 @@ github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkr
|
|||||||
github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
|
github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
|
||||||
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
|
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
|
||||||
github.com/kisielk/sqlstruct v0.0.0-20201105191214-5f3e10d3ab46/go.mod h1:yyMNCyc/Ib3bDTKd379tNMpB/7/H5TjM2Y9QJ5THLbE=
|
github.com/kisielk/sqlstruct v0.0.0-20201105191214-5f3e10d3ab46/go.mod h1:yyMNCyc/Ib3bDTKd379tNMpB/7/H5TjM2Y9QJ5THLbE=
|
||||||
github.com/klauspost/compress v1.18.2 h1:iiPHWW0YrcFgpBYhsA6D1+fqHssJscY/Tm/y2Uqnapk=
|
github.com/klauspost/compress v1.18.6 h1:2jupLlAwFm95+YDR+NwD2MEfFO9d4z4Prjl1XXDjuao=
|
||||||
github.com/klauspost/compress v1.18.2/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4=
|
github.com/klauspost/compress v1.18.6/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||||
|
github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE=
|
||||||
|
github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
|
||||||
github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
|
github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
|
||||||
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
||||||
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
||||||
@@ -175,15 +178,15 @@ github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0 h1:6E+4a0GO5zZEnZ
|
|||||||
github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0/go.mod h1:zJYVVT2jmtg6P3p1VtQj7WsuWi/y4VnjVBn7F8KPB3I=
|
github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0/go.mod h1:zJYVVT2jmtg6P3p1VtQj7WsuWi/y4VnjVBn7F8KPB3I=
|
||||||
github.com/magiconair/properties v1.8.10 h1:s31yESBquKXCV9a/ScB3ESkOjUYYv+X0rg8SYxI99mE=
|
github.com/magiconair/properties v1.8.10 h1:s31yESBquKXCV9a/ScB3ESkOjUYYv+X0rg8SYxI99mE=
|
||||||
github.com/magiconair/properties v1.8.10/go.mod h1:Dhd985XPs7jluiymwWYZ0G4Z61jb3vdS329zhj2hYo0=
|
github.com/magiconair/properties v1.8.10/go.mod h1:Dhd985XPs7jluiymwWYZ0G4Z61jb3vdS329zhj2hYo0=
|
||||||
github.com/mark3labs/mcp-go v0.46.0 h1:8KRibF4wcKejbLsHxCA/QBVUr5fQ9nwz/n8lGqmaALo=
|
github.com/mark3labs/mcp-go v0.54.0 h1:PZhQvd+5xrT43cUoiaKn/hDcvLUhcLc1twSEKYPTcTA=
|
||||||
github.com/mark3labs/mcp-go v0.46.0/go.mod h1:JKTC7R2LLVagkEWK7Kwu7DbmA6iIvnNAod6yrHiQMag=
|
github.com/mark3labs/mcp-go v0.54.0/go.mod h1:+8WclSK1ZUweCP3hvktSji8n8ABG/95QaEkeVE/Uwas=
|
||||||
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
github.com/mattn/go-isatty v0.0.22 h1:j8l17JJ9i6VGPUFUYoTUKPSgKe/83EYU2zBC7YNKMw4=
|
||||||
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
github.com/mattn/go-isatty v0.0.22/go.mod h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4=
|
||||||
github.com/mattn/go-sqlite3 v1.14.33 h1:A5blZ5ulQo2AtayQ9/limgHEkFreKj1Dv226a1K73s0=
|
github.com/mattn/go-sqlite3 v1.14.44 h1:3VSe+xafpbzsLbdr2AWlAZk9yRHiBhTBakioXaCKTF8=
|
||||||
github.com/mattn/go-sqlite3 v1.14.33/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y=
|
github.com/mattn/go-sqlite3 v1.14.44/go.mod h1:pjEuOr8IwzLJP2MfGeTb0A35jauH+C2kbHKBr7yXKVQ=
|
||||||
github.com/microsoft/go-mssqldb v1.8.2/go.mod h1:vp38dT33FGfVotRiTmDo3bFyaHq+p3LektQrjTULowo=
|
github.com/microsoft/go-mssqldb v1.8.2/go.mod h1:vp38dT33FGfVotRiTmDo3bFyaHq+p3LektQrjTULowo=
|
||||||
github.com/microsoft/go-mssqldb v1.9.5 h1:orwya0X/5bsL1o+KasupTkk2eNTNFkTQG0BEe/HxCn0=
|
github.com/microsoft/go-mssqldb v1.10.0 h1:pHEt+Qz6YFPWqREq10mqSE524QQo+/QremwTCQht7TY=
|
||||||
github.com/microsoft/go-mssqldb v1.9.5/go.mod h1:VCP2a0KEZZtGLRHd1PsLavLFYy/3xX2yJUPycv3Sr2Q=
|
github.com/microsoft/go-mssqldb v1.10.0/go.mod h1:mnG7lGa9iYJbzJqGCXyuQCegStKMr3kogDLD6+bmggg=
|
||||||
github.com/moby/docker-image-spec v1.3.1 h1:jMKff3w6PgbfSa69GfNg+zN/XLhfXJGnEx3Nl2EsFP0=
|
github.com/moby/docker-image-spec v1.3.1 h1:jMKff3w6PgbfSa69GfNg+zN/XLhfXJGnEx3Nl2EsFP0=
|
||||||
github.com/moby/docker-image-spec v1.3.1/go.mod h1:eKmb5VW8vQEh/BAr2yvVNvuiJuY6UIocYsFu/DxxRpo=
|
github.com/moby/docker-image-spec v1.3.1/go.mod h1:eKmb5VW8vQEh/BAr2yvVNvuiJuY6UIocYsFu/DxxRpo=
|
||||||
github.com/moby/go-archive v0.1.0 h1:Kk/5rdW/g+H8NHdJW2gsXyZ7UnzvJNOy6VKJqueWdcQ=
|
github.com/moby/go-archive v0.1.0 h1:Kk/5rdW/g+H8NHdJW2gsXyZ7UnzvJNOy6VKJqueWdcQ=
|
||||||
@@ -204,16 +207,16 @@ github.com/mochi-mqtt/server/v2 v2.7.9 h1:y0g4vrSLAag7T07l2oCzOa/+nKVLoazKEWAArw
|
|||||||
github.com/mochi-mqtt/server/v2 v2.7.9/go.mod h1:lZD3j35AVNqJL5cezlnSkuG05c0FCHSsfAKSPBOSbqc=
|
github.com/mochi-mqtt/server/v2 v2.7.9/go.mod h1:lZD3j35AVNqJL5cezlnSkuG05c0FCHSsfAKSPBOSbqc=
|
||||||
github.com/modocache/gover v0.0.0-20171022184752-b58185e213c5/go.mod h1:caMODM3PzxT8aQXRPkAt8xlV/e7d7w8GM5g0fa5F0D8=
|
github.com/modocache/gover v0.0.0-20171022184752-b58185e213c5/go.mod h1:caMODM3PzxT8aQXRPkAt8xlV/e7d7w8GM5g0fa5F0D8=
|
||||||
github.com/montanaflynn/stats v0.7.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
github.com/montanaflynn/stats v0.7.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
||||||
github.com/montanaflynn/stats v0.7.1 h1:etflOAAHORrCC44V+aR6Ftzort912ZU+YLiSTuV8eaE=
|
github.com/montanaflynn/stats v0.9.0 h1:tsBJ0RXwph9BmAuFoCmqGv6e8xa0MENQ8m0ptKq29mQ=
|
||||||
github.com/montanaflynn/stats v0.7.1/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
github.com/montanaflynn/stats v0.9.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
||||||
github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A=
|
github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A=
|
||||||
github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc=
|
github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc=
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
||||||
github.com/nats-io/nats.go v1.48.0 h1:pSFyXApG+yWU/TgbKCjmm5K4wrHu86231/w84qRVR+U=
|
github.com/nats-io/nats.go v1.52.0 h1:n3avV4VBsCgsdwh71TppsTwtv+QdPs7ntSKM8qJLGsc=
|
||||||
github.com/nats-io/nats.go v1.48.0/go.mod h1:iRWIPokVIFbVijxuMQq4y9ttaBTMe0SFdlZfMDd+33g=
|
github.com/nats-io/nats.go v1.52.0/go.mod h1:26HypzazeOkyO3/mqd1zZd53STJN0EjCYF9Uy2ZOBno=
|
||||||
github.com/nats-io/nkeys v0.4.11 h1:q44qGV008kYd9W1b1nEBkNzvnWxtRSQ7A8BoqRrcfa0=
|
github.com/nats-io/nkeys v0.4.15 h1:JACV5jRVO9V856KOapQ7x+EY8Jo3qw1vJt/9Jpwzkk4=
|
||||||
github.com/nats-io/nkeys v0.4.11/go.mod h1:szDimtgmfOi9n25JpfIdGw12tZFYXqhGxjhVxsatHVE=
|
github.com/nats-io/nkeys v0.4.15/go.mod h1:CpMchTXC9fxA5zrMo4KpySxNjiDVvr8ANOSZdiNfUrs=
|
||||||
github.com/nats-io/nuid v1.0.1 h1:5iA8DT8V7q8WK2EScv2padNa/rTESc1KdnPw4TC2paw=
|
github.com/nats-io/nuid v1.0.1 h1:5iA8DT8V7q8WK2EScv2padNa/rTESc1KdnPw4TC2paw=
|
||||||
github.com/nats-io/nuid v1.0.1/go.mod h1:19wcPz3Ph3q0Jbyiqsd0kePYG7A95tJPxeL+1OSON2c=
|
github.com/nats-io/nuid v1.0.1/go.mod h1:19wcPz3Ph3q0Jbyiqsd0kePYG7A95tJPxeL+1OSON2c=
|
||||||
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
|
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
|
||||||
@@ -222,8 +225,8 @@ github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8
|
|||||||
github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
|
github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
|
||||||
github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040=
|
github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040=
|
||||||
github.com/opencontainers/image-spec v1.1.1/go.mod h1:qpqAh3Dmcf36wStyyWU+kCeDgrGnAve2nCC8+7h8Q0M=
|
github.com/opencontainers/image-spec v1.1.1/go.mod h1:qpqAh3Dmcf36wStyyWU+kCeDgrGnAve2nCC8+7h8Q0M=
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4=
|
github.com/pelletier/go-toml/v2 v2.3.1 h1:MYEvvGnQjeNkRF1qUuGolNtNExTDwct51yp7olPtrEc=
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
|
github.com/pelletier/go-toml/v2 v2.3.1/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
|
||||||
github.com/pingcap/errors v0.11.4 h1:lFuQV/oaUMGcD2tqt+01ROSmJs75VG1ToEOkZIZ4nE4=
|
github.com/pingcap/errors v0.11.4 h1:lFuQV/oaUMGcD2tqt+01ROSmJs75VG1ToEOkZIZ4nE4=
|
||||||
github.com/pingcap/errors v0.11.4/go.mod h1:Oi8TUi2kEtXXLMJk9l1cGmz20kV3TaQ0usTwv5KuLY8=
|
github.com/pingcap/errors v0.11.4/go.mod h1:Oi8TUi2kEtXXLMJk9l1cGmz20kV3TaQ0usTwv5KuLY8=
|
||||||
github.com/pkg/browser v0.0.0-20210911075715-681adbf594b8/go.mod h1:HKlIX3XHQyzLZPlr7++PzdhaXEj94dEiJgZDTsxEqUI=
|
github.com/pkg/browser v0.0.0-20210911075715-681adbf594b8/go.mod h1:HKlIX3XHQyzLZPlr7++PzdhaXEj94dEiJgZDTsxEqUI=
|
||||||
@@ -240,24 +243,26 @@ github.com/prometheus/client_golang v1.23.2 h1:Je96obch5RDVy3FDMndoUsjAhG5Edi49h
|
|||||||
github.com/prometheus/client_golang v1.23.2/go.mod h1:Tb1a6LWHB3/SPIzCoaDXI4I8UHKeFTEQ1YCr+0Gyqmg=
|
github.com/prometheus/client_golang v1.23.2/go.mod h1:Tb1a6LWHB3/SPIzCoaDXI4I8UHKeFTEQ1YCr+0Gyqmg=
|
||||||
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
||||||
github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
|
github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
|
||||||
github.com/prometheus/common v0.67.4 h1:yR3NqWO1/UyO1w2PhUvXlGQs/PtFmoveVO0KZ4+Lvsc=
|
github.com/prometheus/common v0.67.5 h1:pIgK94WWlQt1WLwAC5j2ynLaBRDiinoAb86HZHTUGI4=
|
||||||
github.com/prometheus/common v0.67.4/go.mod h1:gP0fq6YjjNCLssJCQp0yk4M8W6ikLURwkdd/YKtTbyI=
|
github.com/prometheus/common v0.67.5/go.mod h1:SjE/0MzDEEAyrdr5Gqc6G+sXI67maCxzaT3A2+HqjUw=
|
||||||
github.com/prometheus/procfs v0.19.2 h1:zUMhqEW66Ex7OXIiDkll3tl9a1ZdilUOd/F6ZXw4Vws=
|
github.com/prometheus/procfs v0.20.1 h1:XwbrGOIplXW/AU3YhIhLODXMJYyC1isLFfYCsTEycfc=
|
||||||
github.com/prometheus/procfs v0.19.2/go.mod h1:M0aotyiemPhBCM0z5w87kL22CxfcH05ZpYlu+b4J7mw=
|
github.com/prometheus/procfs v0.20.1/go.mod h1:o9EMBZGRyvDrSPH1RqdxhojkuXstoe4UlK79eF5TGGo=
|
||||||
github.com/puzpuzpuz/xsync/v3 v3.5.1 h1:GJYJZwO6IdxN/IKbneznS6yPkVC+c3zyY/j19c++5Fg=
|
github.com/puzpuzpuz/xsync/v3 v3.5.1 h1:GJYJZwO6IdxN/IKbneznS6yPkVC+c3zyY/j19c++5Fg=
|
||||||
github.com/puzpuzpuz/xsync/v3 v3.5.1/go.mod h1:VjzYrABPabuM4KyBh1Ftq6u8nhwY5tBPKP9jpmh0nnA=
|
github.com/puzpuzpuz/xsync/v3 v3.5.1/go.mod h1:VjzYrABPabuM4KyBh1Ftq6u8nhwY5tBPKP9jpmh0nnA=
|
||||||
github.com/redis/go-redis/v9 v9.17.2 h1:P2EGsA4qVIM3Pp+aPocCJ7DguDHhqrXNhVcEp4ViluI=
|
github.com/redis/go-redis/v9 v9.19.0 h1:XPVaaPSnG6RhYf7p+rmSa9zZfeVAnWsH5h3lxthOm/k=
|
||||||
github.com/redis/go-redis/v9 v9.17.2/go.mod h1:u410H11HMLoB+TP67dz8rL9s6QW2j76l0//kSOd3370=
|
github.com/redis/go-redis/v9 v9.19.0/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA=
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
||||||
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
||||||
github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4=
|
github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4=
|
||||||
github.com/rogpeppe/go-internal v1.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII=
|
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
|
||||||
github.com/rogpeppe/go-internal v1.13.1/go.mod h1:uMEvuHeurkdAXX61udpOXGD/AzZDWNMNyH2VO9fmH0o=
|
github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc=
|
||||||
github.com/rs/xid v1.4.0 h1:qd7wPTDkN6KQx2VmMBLrpHkiyQwgFXRnkOLacUiaSNY=
|
github.com/rs/xid v1.6.0 h1:fV591PaemRlL6JfRxGDEPl69wICngIQ3shQtzfy2gxU=
|
||||||
github.com/rs/xid v1.4.0/go.mod h1:trrq9SKmegXys3aeAKXMUTdJsYXVwGY3RLcfgqegfbg=
|
github.com/rs/xid v1.6.0/go.mod h1:7XoLgs4eV+QndskICGsho+ADou8ySMSjJKDIan90Nz0=
|
||||||
github.com/sagikazarmark/locafero v0.12.0 h1:/NQhBAkUb4+fH1jivKHWusDYFjMOOKU88eegjfxfHb4=
|
github.com/sagikazarmark/locafero v0.12.0 h1:/NQhBAkUb4+fH1jivKHWusDYFjMOOKU88eegjfxfHb4=
|
||||||
github.com/sagikazarmark/locafero v0.12.0/go.mod h1:sZh36u/YSZ918v0Io+U9ogLYQJ9tLLBmM4eneO6WwsI=
|
github.com/sagikazarmark/locafero v0.12.0/go.mod h1:sZh36u/YSZ918v0Io+U9ogLYQJ9tLLBmM4eneO6WwsI=
|
||||||
|
github.com/santhosh-tekuri/jsonschema/v6 v6.0.2 h1:KRzFb2m7YtdldCEkzs6KqmJw4nqEVZGK7IN2kJkjTuQ=
|
||||||
|
github.com/santhosh-tekuri/jsonschema/v6 v6.0.2/go.mod h1:JXeL+ps8p7/KNMjDQk3TCwPpBy0wYklyWTfbkIzdIFU=
|
||||||
github.com/shirou/gopsutil/v4 v4.25.6 h1:kLysI2JsKorfaFPcYmcJqbzROzsBWEOAtw6A7dIfqXs=
|
github.com/shirou/gopsutil/v4 v4.25.6 h1:kLysI2JsKorfaFPcYmcJqbzROzsBWEOAtw6A7dIfqXs=
|
||||||
github.com/shirou/gopsutil/v4 v4.25.6/go.mod h1:PfybzyydfZcN+JMMjkF6Zb8Mq1A/VcogFFg7hj50W9c=
|
github.com/shirou/gopsutil/v4 v4.25.6/go.mod h1:PfybzyydfZcN+JMMjkF6Zb8Mq1A/VcogFFg7hj50W9c=
|
||||||
github.com/shopspring/decimal v1.4.0 h1:bxl37RwXBklmTi0C79JfXCEBD1cqqHt0bbgBAGFp81k=
|
github.com/shopspring/decimal v1.4.0 h1:bxl37RwXBklmTi0C79JfXCEBD1cqqHt0bbgBAGFp81k=
|
||||||
@@ -292,8 +297,8 @@ github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSW
|
|||||||
github.com/testcontainers/testcontainers-go v0.40.0 h1:pSdJYLOVgLE8YdUY2FHQ1Fxu+aMnb6JfVz1mxk7OeMU=
|
github.com/testcontainers/testcontainers-go v0.40.0 h1:pSdJYLOVgLE8YdUY2FHQ1Fxu+aMnb6JfVz1mxk7OeMU=
|
||||||
github.com/testcontainers/testcontainers-go v0.40.0/go.mod h1:FSXV5KQtX2HAMlm7U3APNyLkkap35zNLxukw9oBi/MY=
|
github.com/testcontainers/testcontainers-go v0.40.0/go.mod h1:FSXV5KQtX2HAMlm7U3APNyLkkap35zNLxukw9oBi/MY=
|
||||||
github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
||||||
github.com/tidwall/gjson v1.18.0 h1:FIDeeyB800efLX89e5a8Y0BNH+LOngJyGrIWxG2FKQY=
|
github.com/tidwall/gjson v1.19.0 h1:xwxm7n691Uf3u5OFjzngavjGTh55KX5q/9w9xHW88JU=
|
||||||
github.com/tidwall/gjson v1.18.0/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
github.com/tidwall/gjson v1.19.0/go.mod h1:V37/opeE/JbLUOfH0QTXiNez2l0RUjYUhpT4szFQAfc=
|
||||||
github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
||||||
github.com/tidwall/match v1.2.0 h1:0pt8FlkOwjN2fPt4bIl4BoNxb98gGHN2ObFEDkrfZnM=
|
github.com/tidwall/match v1.2.0 h1:0pt8FlkOwjN2fPt4bIl4BoNxb98gGHN2ObFEDkrfZnM=
|
||||||
github.com/tidwall/match v1.2.0/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
github.com/tidwall/match v1.2.0/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
||||||
@@ -337,38 +342,42 @@ github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78/go.mod h1:aL8wCCfTfS
|
|||||||
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
|
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
|
||||||
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
|
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
|
||||||
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
||||||
go.mongodb.org/mongo-driver v1.17.6 h1:87JUG1wZfWsr6rIz3ZmpH90rL5tea7O3IHuSwHUpsss=
|
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
|
||||||
go.mongodb.org/mongo-driver v1.17.6/go.mod h1:Hy04i7O2kC4RS06ZrhPRqj/u4DTYkFDAAccj+rVKqgQ=
|
github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s=
|
||||||
go.opentelemetry.io/auto/sdk v1.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA=
|
go.mongodb.org/mongo-driver v1.17.9 h1:IexDdCuuNJ3BHrELgBlyaH9p60JXAvdzWR128q+U5tU=
|
||||||
go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A=
|
go.mongodb.org/mongo-driver v1.17.9/go.mod h1:LlOhpH5NUEfhxcAwG0UEkMqwYcc4JU18gtCdGudk/tQ=
|
||||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 h1:jq9TW8u3so/bN+JPT166wjOI6/vQPF6Xe7nMNIltagk=
|
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
|
||||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0/go.mod h1:p8pYQP+m5XfbZm9fxtSKAbM6oIllS7s2AfxrChvc7iw=
|
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
|
||||||
go.opentelemetry.io/otel v1.38.0 h1:RkfdswUDRimDg0m2Az18RKOsnI8UDzppJAtj01/Ymk8=
|
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0 h1:F7Jx+6hwnZ41NSFTO5q4LYDtJRXBf2PD0rNBkeB/lus=
|
||||||
go.opentelemetry.io/otel v1.38.0/go.mod h1:zcmtmQ1+YmQM9wrNsTGV/q/uyusom3P8RxwExxkZhjM=
|
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0/go.mod h1:UHB22Z8QsdRDrnAtX4PntOl36ajSxcdUMt1sF7Y6E7Q=
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.38.0 h1:GqRJVj7UmLjCVyVJ3ZFLdPRmhDUp2zFmQe3RHIOsw24=
|
go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU=
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.38.0/go.mod h1:ri3aaHSmCTVYu2AWv44YMauwAQc0aqI9gHKIcSbI1pU=
|
go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc=
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.38.0 h1:lwI4Dc5leUqENgGuQImwLo4WnuXFPetmPpkLi2IrX54=
|
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.43.0 h1:88Y4s2C8oTui1LGM6bTWkw0ICGcOLCAI5l6zsD1j20k=
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.38.0/go.mod h1:Kz/oCE7z5wuyhPxsXDuaPteSWqjSBD5YaSdbxZYGbGk=
|
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.43.0/go.mod h1:Vl1/iaggsuRlrHf/hfPJPvVag77kKyvrLeD10kpMl+A=
|
||||||
|
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.43.0 h1:RAE+JPfvEmvy+0LzyUA25/SGawPwIUbZ6u0Wug54sLc=
|
||||||
|
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.43.0/go.mod h1:AGmbycVGEsRx9mXMZ75CsOyhSP6MFIcj/6dnG+vhVjk=
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.19.0 h1:IeMeyr1aBvBiPVYihXIaeIZba6b8E1bYp7lbdxK8CQg=
|
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.19.0 h1:IeMeyr1aBvBiPVYihXIaeIZba6b8E1bYp7lbdxK8CQg=
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.19.0/go.mod h1:oVdCUtjq9MK9BlS7TtucsQwUcXcymNiEDjgDD2jMtZU=
|
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.19.0/go.mod h1:oVdCUtjq9MK9BlS7TtucsQwUcXcymNiEDjgDD2jMtZU=
|
||||||
go.opentelemetry.io/otel/metric v1.38.0 h1:Kl6lzIYGAh5M159u9NgiRkmoMKjvbsKtYRwgfrA6WpA=
|
go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc=
|
||||||
go.opentelemetry.io/otel/metric v1.38.0/go.mod h1:kB5n/QoRM8YwmUahxvI3bO34eVtQf2i4utNVLr9gEmI=
|
go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo=
|
||||||
go.opentelemetry.io/otel/sdk v1.38.0 h1:l48sr5YbNf2hpCUj/FoGhW9yDkl+Ma+LrVl8qaM5b+E=
|
go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58=
|
||||||
go.opentelemetry.io/otel/sdk v1.38.0/go.mod h1:ghmNdGlVemJI3+ZB5iDEuk4bWA3GkTpW+DOoZMYBVVg=
|
go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0=
|
||||||
go.opentelemetry.io/otel/sdk/metric v1.38.0 h1:aSH66iL0aZqo//xXzQLYozmWrXxyFkBJ6qT5wthqPoM=
|
go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI=
|
||||||
go.opentelemetry.io/otel/sdk/metric v1.38.0/go.mod h1:dg9PBnW9XdQ1Hd6ZnRz689CbtrUp0wMMs9iPcgT9EZA=
|
go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA=
|
||||||
go.opentelemetry.io/otel/trace v1.38.0 h1:Fxk5bKrDZJUH+AMyyIXGcFAPah0oRcT+LuNtJrmcNLE=
|
go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk=
|
||||||
go.opentelemetry.io/otel/trace v1.38.0/go.mod h1:j1P9ivuFsTceSWe1oY+EeW3sc+Pp42sO++GHkg4wwhs=
|
go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE=
|
||||||
go.opentelemetry.io/proto/otlp v1.7.1 h1:gTOMpGDb0WTBOP8JaO72iL3auEZhVmAQg4ipjOVAtj4=
|
go.opentelemetry.io/proto/otlp v1.10.0 h1:IQRWgT5srOCYfiWnpqUYz9CVmbO8bFmKcwYxpuCSL2g=
|
||||||
go.opentelemetry.io/proto/otlp v1.7.1/go.mod h1:b2rVh6rfI/s2pHWNlB7ILJcRALpcNDzKhACevjI+ZnE=
|
go.opentelemetry.io/proto/otlp v1.10.0/go.mod h1:/CV4QoCR/S9yaPj8utp3lvQPoqMtxXdzn7ozvvozVqk=
|
||||||
|
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
|
||||||
|
go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
|
||||||
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||||
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||||
go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0=
|
go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0=
|
||||||
go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
||||||
go.uber.org/zap v1.27.1 h1:08RqriUEv8+ArZRYSTXy1LeBScaMpVSTBhCeaZYfMYc=
|
go.uber.org/zap v1.28.0 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo=
|
||||||
go.uber.org/zap v1.27.1/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E=
|
go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q=
|
||||||
go.yaml.in/yaml/v2 v2.4.3 h1:6gvOSjQoTB3vt1l+CU+tSyi/HOjfOjRLJ4YwYZGwRO0=
|
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
||||||
go.yaml.in/yaml/v2 v2.4.3/go.mod h1:zSxWcmIDjOzPXpjlTTbAsKokqkDNAVtZO0WOMiT90s8=
|
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
||||||
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
|
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
|
||||||
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
||||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||||
@@ -383,18 +392,16 @@ golang.org/x/crypto v0.21.0/go.mod h1:0BP7YvVV9gBbVKyeTG0Gyn+gZm94bibOW5BjDEYAOM
|
|||||||
golang.org/x/crypto v0.22.0/go.mod h1:vr6Su+7cTlO45qkww3VDJlzDn0ctJvRgYbC2NvXHt+M=
|
golang.org/x/crypto v0.22.0/go.mod h1:vr6Su+7cTlO45qkww3VDJlzDn0ctJvRgYbC2NvXHt+M=
|
||||||
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
||||||
golang.org/x/crypto v0.24.0/go.mod h1:Z1PMYSOR5nyMcyAVAIQSKCDwalqy85Aqn1x3Ws4L5DM=
|
golang.org/x/crypto v0.24.0/go.mod h1:Z1PMYSOR5nyMcyAVAIQSKCDwalqy85Aqn1x3Ws4L5DM=
|
||||||
golang.org/x/crypto v0.46.0 h1:cKRW/pmt1pKAfetfu+RCEvjvZkA9RimPbh7bhFjGVBU=
|
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
||||||
golang.org/x/crypto v0.46.0/go.mod h1:Evb/oLKmMraqjZ2iQTwDwvCtJkczlDuTmdJXoZVzqU0=
|
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
||||||
golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93 h1:fQsdNF2N+/YewlRZiricy4P1iimyPKZ/xwniHj8Q2a0=
|
|
||||||
golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93/go.mod h1:EPRbTFwzwjXj9NpYyyrvenVh9Y+GFeEvMNh7Xuz7xgU=
|
|
||||||
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
|
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
|
||||||
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||||
golang.org/x/mod v0.9.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
golang.org/x/mod v0.9.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||||
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||||
golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||||
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||||
golang.org/x/mod v0.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
|
golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk=
|
||||||
golang.org/x/mod v0.31.0/go.mod h1:43JraMp9cGx1Rx3AqioxrbrhNsLl2l/iNAvuBkrezpg=
|
golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40=
|
||||||
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||||
golang.org/x/net v0.0.0-20200114155413-6afb5195e5aa/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
golang.org/x/net v0.0.0-20200114155413-6afb5195e5aa/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||||
@@ -412,10 +419,10 @@ golang.org/x/net v0.22.0/go.mod h1:JKghWKKOSdJwpW2GEx0Ja7fmaKnMsbu+MWVZTokSYmg=
|
|||||||
golang.org/x/net v0.24.0/go.mod h1:2Q7sJY5mzlzWjKtYUEXSlBWCdyaioyXzRB2RtU8KVE8=
|
golang.org/x/net v0.24.0/go.mod h1:2Q7sJY5mzlzWjKtYUEXSlBWCdyaioyXzRB2RtU8KVE8=
|
||||||
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
||||||
golang.org/x/net v0.26.0/go.mod h1:5YKkiSynbBIh3p6iOc/vibscux0x38BZDkn8sCUPxHE=
|
golang.org/x/net v0.26.0/go.mod h1:5YKkiSynbBIh3p6iOc/vibscux0x38BZDkn8sCUPxHE=
|
||||||
golang.org/x/net v0.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU=
|
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
|
||||||
golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY=
|
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
|
||||||
golang.org/x/oauth2 v0.34.0 h1:hqK/t4AKgbqWkdkcAeI8XLmbK+4m4G5YeQRrmiotGlw=
|
golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
|
||||||
golang.org/x/oauth2 v0.34.0/go.mod h1:lzm5WQJQwKZ3nwavOZ3IS5Aulzxi68dUSgRHujetwEA=
|
golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q=
|
||||||
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
@@ -423,8 +430,8 @@ golang.org/x/sync v0.3.0/go.mod h1:FU7BRWz2tNW+3quACPkgCx/L+uEAv1htQ0V83Z9Rj+Y=
|
|||||||
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||||
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||||
golang.org/x/sync v0.9.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
golang.org/x/sync v0.9.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||||
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
|
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||||
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
@@ -448,8 +455,8 @@ golang.org/x/sys v0.18.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
|||||||
golang.org/x/sys v0.19.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
golang.org/x/sys v0.19.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||||
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||||
golang.org/x/sys v0.21.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
golang.org/x/sys v0.21.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||||
golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk=
|
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||||
golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||||
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
|
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
|
||||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||||
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
|
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
|
||||||
@@ -465,8 +472,8 @@ golang.org/x/term v0.18.0/go.mod h1:ILwASektA3OnRv7amZ1xhE/KTR+u50pbXfZ03+6Nx58=
|
|||||||
golang.org/x/term v0.19.0/go.mod h1:2CuTdWZ7KHSQwUzKva0cbMg6q2DMI3Mmxp+gKJbskEk=
|
golang.org/x/term v0.19.0/go.mod h1:2CuTdWZ7KHSQwUzKva0cbMg6q2DMI3Mmxp+gKJbskEk=
|
||||||
golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
|
golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
|
||||||
golang.org/x/term v0.21.0/go.mod h1:ooXLefLobQVslOqselCNF4SxFAaoS6KujMbsGzSDmX0=
|
golang.org/x/term v0.21.0/go.mod h1:ooXLefLobQVslOqselCNF4SxFAaoS6KujMbsGzSDmX0=
|
||||||
golang.org/x/term v0.38.0 h1:PQ5pkm/rLO6HnxFR7N2lJHOZX6Kez5Y1gDSJla6jo7Q=
|
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
||||||
golang.org/x/term v0.38.0/go.mod h1:bSEAKrOT1W+VSu9TSCMtoGEOUcKxOKgl3LE5QEF/xVg=
|
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
|
||||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||||
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
||||||
@@ -481,28 +488,28 @@ golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
|||||||
golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||||
golang.org/x/text v0.16.0/go.mod h1:GhwF1Be+LQoKShO3cGOHzqOgRrGaYc9AvblQOmPVHnI=
|
golang.org/x/text v0.16.0/go.mod h1:GhwF1Be+LQoKShO3cGOHzqOgRrGaYc9AvblQOmPVHnI=
|
||||||
golang.org/x/text v0.20.0/go.mod h1:D4IsuqiFMhST5bX19pQ9ikHC2GsaKyk/oF+pn3ducp4=
|
golang.org/x/text v0.20.0/go.mod h1:D4IsuqiFMhST5bX19pQ9ikHC2GsaKyk/oF+pn3ducp4=
|
||||||
golang.org/x/text v0.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU=
|
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
|
||||||
golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY=
|
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
|
||||||
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
|
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
|
||||||
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
|
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
|
||||||
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
||||||
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
|
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
|
||||||
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
|
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
|
||||||
golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
|
golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
|
||||||
golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58=
|
golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58=
|
||||||
golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
|
golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
|
||||||
golang.org/x/tools v0.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA=
|
golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE=
|
||||||
golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc=
|
golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk=
|
||||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk=
|
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
|
||||||
gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E=
|
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
|
||||||
google.golang.org/genproto/googleapis/api v0.0.0-20250825161204-c5933d9347a5 h1:BIRfGDEjiHRrk0QKZe3Xv2ieMhtgRGeLcZQ0mIVn4EY=
|
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa h1:Kjn0N0tCrDgiAFW+lGO4JZ3ck44CehvJQMAwj9QF0G8=
|
||||||
google.golang.org/genproto/googleapis/api v0.0.0-20250825161204-c5933d9347a5/go.mod h1:j3QtIyytwqGr1JUDtYXwtMXWPKsEa5LtzIFN1Wn5WvE=
|
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:q4lMZS6kskjT5HvCPrnnypcDPVJqT/f4nfxmkE7gryY=
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20250825161204-c5933d9347a5 h1:eaY8u2EuxbRv7c3NiGK0/NedzVsCcV6hDuU5qPX5EGE=
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk=
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20250825161204-c5933d9347a5/go.mod h1:M4/wBTSeyLxupu3W3tJtOgB14jILAS/XWPSSa3TAlJc=
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
|
||||||
google.golang.org/grpc v1.75.0 h1:+TW+dqTd2Biwe6KKfhE5JpiYIBWq865PhKGSXiivqt4=
|
google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU=
|
||||||
google.golang.org/grpc v1.75.0/go.mod h1:JtPAzKiq4v1xcAB2hydNlWI2RnF85XXcV0mhKXr2ecQ=
|
google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8=
|
||||||
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
||||||
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||||
@@ -526,30 +533,30 @@ gorm.io/gorm v1.31.1 h1:7CA8FTFz/gRfgqgpeKIBcervUn3xSyPUmr6B2WXJ7kg=
|
|||||||
gorm.io/gorm v1.31.1/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs=
|
gorm.io/gorm v1.31.1/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs=
|
||||||
gotest.tools/v3 v3.5.2 h1:7koQfIKdy+I8UTetycgUqXWSDwpgv193Ka+qRsmBY8Q=
|
gotest.tools/v3 v3.5.2 h1:7koQfIKdy+I8UTetycgUqXWSDwpgv193Ka+qRsmBY8Q=
|
||||||
gotest.tools/v3 v3.5.2/go.mod h1:LtdLGcnqToBH83WByAAi/wiwSFCArdFIUV/xxN4pcjA=
|
gotest.tools/v3 v3.5.2/go.mod h1:LtdLGcnqToBH83WByAAi/wiwSFCArdFIUV/xxN4pcjA=
|
||||||
modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis=
|
modernc.org/cc/v4 v4.28.2 h1:3tQ0lf2ADtoby2EtSP+J7IE2SHwEJdP8ioR59wx7XpY=
|
||||||
modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
|
modernc.org/cc/v4 v4.28.2/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI=
|
||||||
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
|
modernc.org/ccgo/v4 v4.34.0 h1:yRLPFZieg532OT4rp4JFNIVcquwalMX26G95WQDqwCQ=
|
||||||
modernc.org/ccgo/v4 v4.30.1/go.mod h1:bIOeI1JL54Utlxn+LwrFyjCx2n2RDiYEaJVSrgdrRfM=
|
modernc.org/ccgo/v4 v4.34.0/go.mod h1:AS5WYMyBakQ+fhsHhtP8mWB82KTGPkNNJDGfGQCe0/A=
|
||||||
modernc.org/fileutil v1.3.40 h1:ZGMswMNc9JOCrcrakF1HrvmergNLAmxOPjizirpfqBA=
|
modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM=
|
||||||
modernc.org/fileutil v1.3.40/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc=
|
modernc.org/fileutil v1.4.0/go.mod h1:EqdKFDxiByqxLk8ozOxObDSfcVOv/54xDs/DUHdvCUU=
|
||||||
modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI=
|
modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI=
|
||||||
modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
|
modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
|
||||||
modernc.org/gc/v3 v3.1.1 h1:k8T3gkXWY9sEiytKhcgyiZ2L0DTyCQ/nvX+LoCljoRE=
|
modernc.org/gc/v3 v3.1.2 h1:ZtDCnhonXSZexk/AYsegNRV1lJGgaNZJuKjJSWKyEqo=
|
||||||
modernc.org/gc/v3 v3.1.1/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY=
|
modernc.org/gc/v3 v3.1.2/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY=
|
||||||
modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks=
|
modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks=
|
||||||
modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI=
|
modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI=
|
||||||
modernc.org/libc v1.67.4 h1:zZGmCMUVPORtKv95c2ReQN5VDjvkoRm9GWPTEPuvlWg=
|
modernc.org/libc v1.72.3 h1:ZnDF4tXn4NBXFutMMQC4vtbTFSXhhKzR73fv0beZEAU=
|
||||||
modernc.org/libc v1.67.4/go.mod h1:QvvnnJ5P7aitu0ReNpVIEyesuhmDLQ8kaEoyMjIFZJA=
|
modernc.org/libc v1.72.3/go.mod h1:dn0dZNnnn1clLyvRxLxYExxiKRZIRENOfqQ8XEeg4Qs=
|
||||||
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
|
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
|
||||||
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
|
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
|
||||||
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
|
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
|
||||||
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
|
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
|
||||||
modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8=
|
modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg=
|
||||||
modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
|
modernc.org/opt v0.2.0/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
|
||||||
modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=
|
modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=
|
||||||
modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
|
modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
|
||||||
modernc.org/sqlite v1.42.2 h1:7hkZUNJvJFN2PgfUdjni9Kbvd4ef4mNLOu0B9FGxM74=
|
modernc.org/sqlite v1.50.1 h1:l+cQvn0sd0zJJtfygGHuQJ5AjlrwXmWPw4KP3ZMwr9w=
|
||||||
modernc.org/sqlite v1.42.2/go.mod h1:+VkC6v3pLOAE0A0uVucQEcbVW0I5nHCeDaBf+DpsQT8=
|
modernc.org/sqlite v1.50.1/go.mod h1:tcNzv5p84E0skkmJn038y+hWJbLQXQqEnQfeh5r2JLM=
|
||||||
modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
|
modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
|
||||||
modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A=
|
modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A=
|
||||||
modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=
|
modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=
|
||||||
|
|||||||
Vendored
+28
-15
@@ -3,23 +3,29 @@ package cache
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var defaultCache atomic.Pointer[Cache]
|
||||||
defaultCache *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.
|
// Initialize initializes the cache with a provider.
|
||||||
// If not called, the package will use an in-memory provider by default.
|
// If not called, the package will use an in-memory provider by default.
|
||||||
func Initialize(provider Provider) {
|
func Initialize(provider Provider) {
|
||||||
defaultCache = NewCache(provider)
|
swapOwned(NewCache(provider))
|
||||||
}
|
}
|
||||||
|
|
||||||
// UseMemory configures the cache to use in-memory storage.
|
// UseMemory configures the cache to use in-memory storage.
|
||||||
func UseMemory(opts *Options) error {
|
func UseMemory(opts *Options) error {
|
||||||
provider := NewMemoryProvider(opts)
|
swapOwned(NewCache(NewMemoryProvider(opts)))
|
||||||
defaultCache = NewCache(provider)
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -29,7 +35,7 @@ func UseRedis(config *RedisConfig) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to initialize Redis provider: %w", err)
|
return fmt.Errorf("failed to initialize Redis provider: %w", err)
|
||||||
}
|
}
|
||||||
defaultCache = NewCache(provider)
|
swapOwned(NewCache(provider))
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -39,26 +45,33 @@ func UseMemcache(config *MemcacheConfig) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to initialize Memcache provider: %w", err)
|
return fmt.Errorf("failed to initialize Memcache provider: %w", err)
|
||||||
}
|
}
|
||||||
defaultCache = NewCache(provider)
|
swapOwned(NewCache(provider))
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetDefaultCache returns the default cache instance.
|
// GetDefaultCache returns the default cache instance.
|
||||||
// Initializes with in-memory provider if not already initialized.
|
// Initializes with in-memory provider if not already initialized.
|
||||||
|
// Safe for concurrent use.
|
||||||
func GetDefaultCache() *Cache {
|
func GetDefaultCache() *Cache {
|
||||||
if defaultCache == nil {
|
if c := defaultCache.Load(); c != nil {
|
||||||
_ = UseMemory(&Options{
|
return c
|
||||||
|
}
|
||||||
|
fresh := NewCache(NewMemoryProvider(&Options{
|
||||||
DefaultTTL: 5 * time.Minute,
|
DefaultTTL: 5 * time.Minute,
|
||||||
MaxSize: 10000,
|
MaxSize: 10000,
|
||||||
})
|
}))
|
||||||
|
if defaultCache.CompareAndSwap(nil, fresh) {
|
||||||
|
return fresh
|
||||||
}
|
}
|
||||||
return defaultCache
|
_ = fresh.Close() // lost the race; discard our provider
|
||||||
|
return defaultCache.Load()
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetDefaultCache sets a custom cache instance as the default cache.
|
// 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.
|
// 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) {
|
func SetDefaultCache(cache *Cache) {
|
||||||
defaultCache = cache
|
defaultCache.Store(cache)
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetStats returns cache statistics.
|
// GetStats returns cache statistics.
|
||||||
@@ -69,8 +82,8 @@ func GetStats(ctx context.Context) (*CacheStats, error) {
|
|||||||
|
|
||||||
// Close closes the cache and releases resources.
|
// Close closes the cache and releases resources.
|
||||||
func Close() error {
|
func Close() error {
|
||||||
if defaultCache != nil {
|
if c := defaultCache.Load(); c != nil {
|
||||||
return defaultCache.Close()
|
return c.Close()
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
Vendored
+19
-6
@@ -3,10 +3,17 @@ package cache
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"time"
|
"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.
|
// Cache is the main cache manager that wraps a Provider.
|
||||||
type Cache struct {
|
type Cache struct {
|
||||||
provider Provider
|
provider Provider
|
||||||
@@ -23,7 +30,7 @@ func NewCache(provider Provider) *Cache {
|
|||||||
func (c *Cache) Get(ctx context.Context, key string, dest interface{}) error {
|
func (c *Cache) Get(ctx context.Context, key string, dest interface{}) error {
|
||||||
data, exists := c.provider.Get(ctx, key)
|
data, exists := c.provider.Get(ctx, key)
|
||||||
if !exists {
|
if !exists {
|
||||||
return fmt.Errorf("key not found: %s", key)
|
return ErrNotFound
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := json.Unmarshal(data, dest); err != nil {
|
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) {
|
func (c *Cache) GetBytes(ctx context.Context, key string) ([]byte, error) {
|
||||||
data, exists := c.provider.Get(ctx, key)
|
data, exists := c.provider.Get(ctx, key)
|
||||||
if !exists {
|
if !exists {
|
||||||
return nil, fmt.Errorf("key not found: %s", key)
|
return nil, ErrNotFound
|
||||||
}
|
}
|
||||||
return data, nil
|
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)
|
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 {
|
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
|
// 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.
|
// Remember is a convenience function that caches the result of a function call.
|
||||||
// It's similar to GetOrSet but returns the value directly.
|
// 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) {
|
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
|
// Try to get from cache first as bytes
|
||||||
data, err := c.GetBytes(ctx, key)
|
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)
|
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 {
|
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
|
return value, nil
|
||||||
|
|||||||
Vendored
+4
@@ -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
|
package cache
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
|||||||
Vendored
+154
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
Vendored
+10
@@ -2,9 +2,14 @@ package cache
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"time"
|
"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.
|
// Provider defines the interface that all cache providers must implement.
|
||||||
type Provider interface {
|
type Provider interface {
|
||||||
// Get retrieves a value from the cache by key.
|
// Get retrieves a value from the cache by key.
|
||||||
@@ -58,8 +63,13 @@ type Options struct {
|
|||||||
DefaultTTL time.Duration
|
DefaultTTL time.Duration
|
||||||
|
|
||||||
// MaxSize is the maximum number of items (for in-memory provider).
|
// MaxSize is the maximum number of items (for in-memory provider).
|
||||||
|
// 0 selects the default (10000); a negative value means unbounded.
|
||||||
MaxSize int
|
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 determines how items are evicted (LRU, LFU, etc).
|
||||||
EvictionPolicy string
|
EvictionPolicy string
|
||||||
}
|
}
|
||||||
|
|||||||
Vendored
+224
-105
@@ -2,17 +2,34 @@ package cache
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bradfitz/gomemcache/memcache"
|
"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.
|
// MemcacheProvider is a Memcache implementation of the Provider interface.
|
||||||
type MemcacheProvider struct {
|
type MemcacheProvider struct {
|
||||||
client *memcache.Client
|
client *memcache.Client
|
||||||
options *Options
|
options *Options
|
||||||
|
allowFlush bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// MemcacheConfig contains Memcache-specific configuration.
|
// MemcacheConfig contains Memcache-specific configuration.
|
||||||
@@ -28,37 +45,46 @@ type MemcacheConfig struct {
|
|||||||
|
|
||||||
// Options contains general cache options
|
// Options contains general cache options
|
||||||
Options *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.
|
// NewMemcacheProvider creates a new Memcache cache provider.
|
||||||
func NewMemcacheProvider(config *MemcacheConfig) (*MemcacheProvider, error) {
|
func NewMemcacheProvider(config *MemcacheConfig) (*MemcacheProvider, error) {
|
||||||
if config == nil {
|
// Work on a copy so the caller's struct is not mutated
|
||||||
config = &MemcacheConfig{
|
var cfg MemcacheConfig
|
||||||
Servers: []string{"localhost:11211"},
|
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 {
|
if len(cfg.Servers) == 0 {
|
||||||
config.Servers = []string{"localhost:11211"}
|
cfg.Servers = []string{"localhost:11211"}
|
||||||
}
|
}
|
||||||
|
|
||||||
if config.MaxIdleConns == 0 {
|
if cfg.MaxIdleConns == 0 {
|
||||||
config.MaxIdleConns = 2
|
cfg.MaxIdleConns = 2
|
||||||
}
|
}
|
||||||
|
|
||||||
if config.Timeout == 0 {
|
if cfg.Timeout == 0 {
|
||||||
config.Timeout = 1 * time.Second
|
cfg.Timeout = 1 * time.Second
|
||||||
}
|
}
|
||||||
|
|
||||||
if config.Options == nil {
|
if cfg.Options == nil {
|
||||||
config.Options = &Options{
|
cfg.Options = &Options{
|
||||||
DefaultTTL: 5 * time.Minute,
|
DefaultTTL: 5 * time.Minute,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
client := memcache.New(config.Servers...)
|
client := memcache.New(cfg.Servers...)
|
||||||
client.MaxIdleConns = config.MaxIdleConns
|
client.MaxIdleConns = cfg.MaxIdleConns
|
||||||
client.Timeout = config.Timeout
|
client.Timeout = cfg.Timeout
|
||||||
|
|
||||||
// Test connection
|
// Test connection
|
||||||
if err := client.Ping(); err != nil {
|
if err := client.Ping(); err != nil {
|
||||||
@@ -67,17 +93,71 @@ func NewMemcacheProvider(config *MemcacheConfig) (*MemcacheProvider, error) {
|
|||||||
|
|
||||||
return &MemcacheProvider{
|
return &MemcacheProvider{
|
||||||
client: client,
|
client: client,
|
||||||
options: config.Options,
|
options: cfg.Options,
|
||||||
|
allowFlush: cfg.AllowFlush,
|
||||||
}, nil
|
}, 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.
|
// Get retrieves a value from the cache by key.
|
||||||
func (m *MemcacheProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
func (m *MemcacheProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
||||||
item, err := m.client.Get(key)
|
if ctx.Err() != nil {
|
||||||
if err == memcache.ErrCacheMiss {
|
return nil, false
|
||||||
|
}
|
||||||
|
item, err := m.client.Get(memcacheKey(key))
|
||||||
|
if errors.Is(err, memcache.ErrCacheMiss) {
|
||||||
return nil, false
|
return nil, false
|
||||||
}
|
}
|
||||||
if err != nil {
|
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 nil, false
|
||||||
}
|
}
|
||||||
return item.Value, true
|
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.
|
// 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 {
|
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 {
|
if ttl == 0 {
|
||||||
ttl = m.options.DefaultTTL
|
ttl = m.options.DefaultTTL
|
||||||
}
|
}
|
||||||
|
|
||||||
item := &memcache.Item{
|
return m.client.Set(&memcache.Item{
|
||||||
Key: key,
|
Key: memcacheKey(key),
|
||||||
Value: value,
|
Value: value,
|
||||||
Expiration: int32(ttl.Seconds()),
|
Expiration: memcacheExpiry(ttl),
|
||||||
}
|
})
|
||||||
|
|
||||||
return m.client.Set(item)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetWithTags stores a value in the cache with the specified TTL and tags.
|
// 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 {
|
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 {
|
if ttl == 0 {
|
||||||
ttl = m.options.DefaultTTL
|
ttl = m.options.DefaultTTL
|
||||||
}
|
}
|
||||||
|
|
||||||
expiration := int32(ttl.Seconds())
|
expiration := memcacheExpiry(ttl)
|
||||||
|
mkey := memcacheKey(key)
|
||||||
|
|
||||||
// Set the main value
|
if err := m.client.Set(&memcache.Item{Key: mkey, Value: value, Expiration: expiration}); err != nil {
|
||||||
item := &memcache.Item{
|
return err
|
||||||
Key: key,
|
|
||||||
Value: value,
|
|
||||||
Expiration: expiration,
|
|
||||||
}
|
}
|
||||||
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
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Store tags for this key
|
|
||||||
if len(tags) > 0 {
|
|
||||||
tagsData, err := json.Marshal(tags)
|
tagsData, err := json.Marshal(tags)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to marshal tags: %w", err)
|
return fail(fmt.Errorf("failed to marshal tags: %w", err))
|
||||||
}
|
}
|
||||||
|
if err := m.client.Set(&memcache.Item{
|
||||||
tagsItem := &memcache.Item{
|
Key: memcacheTagKey("cache:tags:", key),
|
||||||
Key: fmt.Sprintf("cache:tags:%s", key),
|
|
||||||
Value: tagsData,
|
Value: tagsData,
|
||||||
Expiration: expiration,
|
Expiration: expiration,
|
||||||
}
|
}); err != nil {
|
||||||
if err := m.client.Set(tagsItem); err != nil {
|
return fail(err)
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add key to each tag's key list
|
// Tag lists live longer than the entries they index
|
||||||
|
tagExpiry := memcacheExpiry(ttl + time.Hour)
|
||||||
for _, tag := range tags {
|
for _, tag := range tags {
|
||||||
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
if err := m.updateTagKeys(memcacheTagKey("cache:tag:", tag), tagExpiry, func(keys []string) ([]string, error) {
|
||||||
|
|
||||||
// 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
|
|
||||||
for _, k := range keys {
|
for _, k := range keys {
|
||||||
if k == key {
|
if k == key {
|
||||||
found = true
|
return keys, nil
|
||||||
break
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if !found {
|
if len(keys) >= memcacheMaxTagKeys {
|
||||||
keys = append(keys, key)
|
return nil, fmt.Errorf("tag index for %q is full (%d keys)", tag, memcacheMaxTagKeys)
|
||||||
}
|
}
|
||||||
|
return append(keys, key), nil
|
||||||
// Store updated key list
|
}); err != nil {
|
||||||
keysData, err := json.Marshal(keys)
|
return fail(err)
|
||||||
if err != nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
tagItem := &memcache.Item{
|
|
||||||
Key: tagKey,
|
|
||||||
Value: keysData,
|
|
||||||
Expiration: expiration + 3600, // Give tag lists longer TTL
|
|
||||||
}
|
|
||||||
_ = m.client.Set(tagItem)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
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.
|
// Delete removes a key from the cache.
|
||||||
func (m *MemcacheProvider) Delete(ctx context.Context, key string) error {
|
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
|
// 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 {
|
if item, err := m.client.Get(tagsKey); err == nil {
|
||||||
var tags []string
|
var tags []string
|
||||||
if err := json.Unmarshal(item.Value, &tags); err == nil {
|
if err := json.Unmarshal(item.Value, &tags); err == nil {
|
||||||
// Remove key from each tag's key list
|
|
||||||
for _, tag := range tags {
|
for _, tag := range tags {
|
||||||
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
err := m.updateTagKeys(memcacheTagKey("cache:tag:", tag), memcacheExpiry(m.options.DefaultTTL+time.Hour), func(keys []string) ([]string, error) {
|
||||||
if tagItem, err := m.client.Get(tagKey); err == nil {
|
out := make([]string, 0, len(keys))
|
||||||
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 {
|
for _, k := range keys {
|
||||||
if k != key {
|
if k != key {
|
||||||
newKeys = append(newKeys, k)
|
out = append(out, k)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// Update the tag's key list
|
return out, nil
|
||||||
if keysData, err := json.Marshal(newKeys); err == nil {
|
})
|
||||||
tagItem.Value = keysData
|
if err != nil {
|
||||||
_ = m.client.Set(tagItem)
|
logger.Warn("cache: failed to update memcache tag index on delete: %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
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 tags key
|
|
||||||
_ = m.client.Delete(tagsKey)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Delete the actual key
|
// Delete the actual key
|
||||||
err := m.client.Delete(key)
|
err := m.client.Delete(mkey)
|
||||||
if err == memcache.ErrCacheMiss {
|
if errors.Is(err, memcache.ErrCacheMiss) {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return err
|
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.
|
// DeleteByTag removes all keys associated with the given tag.
|
||||||
func (m *MemcacheProvider) DeleteByTag(ctx context.Context, tag string) error {
|
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)
|
item, err := m.client.Get(tagKey)
|
||||||
if err == memcache.ErrCacheMiss {
|
if errors.Is(err, memcache.ErrCacheMiss) {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if err != 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)
|
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 {
|
for _, key := range keys {
|
||||||
_ = m.client.Delete(key)
|
note(m.client.Delete(memcacheKey(key)))
|
||||||
// Also delete the tags key for this cache key
|
note(m.client.Delete(memcacheTagKey("cache:tags:", key)))
|
||||||
tagsKey := fmt.Sprintf("cache:tags:%s", key)
|
|
||||||
_ = m.client.Delete(tagsKey)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Delete the tag key itself
|
if firstErr != nil {
|
||||||
_ = m.client.Delete(tagKey)
|
return firstErr // keep the tag index so the invalidation can be retried
|
||||||
|
}
|
||||||
return nil
|
note(m.client.Delete(tagKey))
|
||||||
|
return firstErr
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeleteByPattern removes all keys matching the pattern.
|
// DeleteByPattern is not supported by Memcache; it always returns an error.
|
||||||
// Note: Memcache does not support pattern-based deletion natively.
|
// Use tags (SetWithTags / DeleteByTag) for group invalidation instead.
|
||||||
// This is a no-op for memcache and returns an error.
|
|
||||||
func (m *MemcacheProvider) DeleteByPattern(ctx context.Context, pattern string) error {
|
func (m *MemcacheProvider) DeleteByPattern(ctx context.Context, pattern string) error {
|
||||||
return fmt.Errorf("pattern-based deletion is not supported by Memcache")
|
return fmt.Errorf("pattern-based deletion is not supported by Memcache")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Clear removes all items from the cache.
|
// 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 {
|
func (m *MemcacheProvider) Clear(ctx context.Context) error {
|
||||||
|
if !m.allowFlush {
|
||||||
|
return ErrFlushNotAllowed
|
||||||
|
}
|
||||||
return m.client.FlushAll()
|
return m.client.FlushAll()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Exists checks if a key exists in the cache.
|
// Exists checks if a key exists in the cache.
|
||||||
func (m *MemcacheProvider) Exists(ctx context.Context, key string) bool {
|
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
|
return err == nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close closes the provider and releases any resources.
|
// Close closes the provider and releases idle connections.
|
||||||
func (m *MemcacheProvider) Close() error {
|
func (m *MemcacheProvider) Close() error {
|
||||||
// Memcache client doesn't have a close method
|
return m.client.Close()
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Stats returns statistics about the cache provider.
|
// Stats returns statistics about the cache provider.
|
||||||
|
|||||||
Vendored
+112
-103
@@ -2,20 +2,39 @@ package cache
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"regexp"
|
"regexp"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"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.
|
// memoryItem represents a cached item in memory.
|
||||||
type memoryItem struct {
|
type memoryItem struct {
|
||||||
Value []byte
|
Value []byte
|
||||||
Expiration time.Time
|
Expiration time.Time
|
||||||
LastAccess time.Time
|
|
||||||
HitCount int64
|
|
||||||
Tags []string
|
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.
|
// isExpired checks if the item has expired.
|
||||||
@@ -34,30 +53,72 @@ type MemoryProvider struct {
|
|||||||
options *Options
|
options *Options
|
||||||
hits atomic.Int64
|
hits atomic.Int64
|
||||||
misses atomic.Int64
|
misses atomic.Int64
|
||||||
|
closed bool
|
||||||
|
done chan struct{}
|
||||||
|
closeOnce sync.Once
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewMemoryProvider creates a new in-memory cache provider.
|
// 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 {
|
func NewMemoryProvider(opts *Options) *MemoryProvider {
|
||||||
if opts == nil {
|
var o Options
|
||||||
opts = &Options{
|
if opts != nil {
|
||||||
DefaultTTL: 5 * time.Minute,
|
o = *opts // do not mutate the caller's struct
|
||||||
MaxSize: 10000,
|
} 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),
|
items: make(map[string]*memoryItem),
|
||||||
tagToKeys: make(map[string]map[string]struct{}),
|
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.
|
// Get retrieves a value from the cache by key.
|
||||||
func (m *MemoryProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
func (m *MemoryProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
||||||
// First try with read lock for fast path
|
|
||||||
m.mu.RLock()
|
m.mu.RLock()
|
||||||
item, exists := m.items[key]
|
item, exists := m.items[key]
|
||||||
if !exists {
|
if !exists || m.closed {
|
||||||
m.mu.RUnlock()
|
m.mu.RUnlock()
|
||||||
m.misses.Add(1)
|
m.misses.Add(1)
|
||||||
return nil, false
|
return nil, false
|
||||||
@@ -65,56 +126,29 @@ func (m *MemoryProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
|||||||
|
|
||||||
if item.isExpired() {
|
if item.isExpired() {
|
||||||
m.mu.RUnlock()
|
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()
|
m.mu.Lock()
|
||||||
delete(m.items, key)
|
if cur, ok := m.items[key]; ok && cur == item {
|
||||||
|
m.removeLocked(key)
|
||||||
|
}
|
||||||
m.mu.Unlock()
|
m.mu.Unlock()
|
||||||
m.misses.Add(1)
|
m.misses.Add(1)
|
||||||
return nil, false
|
return nil, false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Update stats and access time with write lock
|
item.lastAccess.Store(time.Now().UnixNano())
|
||||||
value := item.Value
|
item.hitCount.Add(1)
|
||||||
|
out := make([]byte, len(item.Value))
|
||||||
|
copy(out, item.Value)
|
||||||
m.mu.RUnlock()
|
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)
|
m.hits.Add(1)
|
||||||
return value, true
|
return out, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// Set stores a value in the cache with the specified TTL.
|
// 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 {
|
func (m *MemoryProvider) Set(ctx context.Context, key string, value []byte, ttl time.Duration) error {
|
||||||
m.mu.Lock()
|
return m.SetWithTags(ctx, key, value, ttl, nil)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetWithTags stores a value in the cache with the specified TTL and tags.
|
// 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()
|
m.mu.Lock()
|
||||||
defer m.mu.Unlock()
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
if m.closed {
|
||||||
|
return ErrClosed
|
||||||
|
}
|
||||||
|
|
||||||
if ttl == 0 {
|
if ttl == 0 {
|
||||||
ttl = m.options.DefaultTTL
|
ttl = m.options.DefaultTTL
|
||||||
}
|
}
|
||||||
@@ -131,34 +169,14 @@ func (m *MemoryProvider) SetWithTags(ctx context.Context, key string, value []by
|
|||||||
expiration = time.Now().Add(ttl)
|
expiration = time.Now().Add(ttl)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check max size and evict if necessary
|
if _, exists := m.items[key]; exists {
|
||||||
if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
|
m.removeLocked(key) // drops old tag associations
|
||||||
if _, exists := m.items[key]; !exists {
|
} else if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
|
||||||
m.evictOne()
|
m.evictOne()
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
// Remove old tag associations if key exists
|
m.items[key] = newMemoryItem(value, expiration, tags)
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Store the item
|
|
||||||
m.items[key] = &memoryItem{
|
|
||||||
Value: value,
|
|
||||||
Expiration: expiration,
|
|
||||||
LastAccess: time.Now(),
|
|
||||||
Tags: tags,
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add new tag associations
|
|
||||||
for _, tag := range tags {
|
for _, tag := range tags {
|
||||||
if m.tagToKeys[tag] == nil {
|
if m.tagToKeys[tag] == nil {
|
||||||
m.tagToKeys[tag] = make(map[string]struct{})
|
m.tagToKeys[tag] = make(map[string]struct{})
|
||||||
@@ -174,19 +192,7 @@ func (m *MemoryProvider) Delete(ctx context.Context, key string) error {
|
|||||||
m.mu.Lock()
|
m.mu.Lock()
|
||||||
defer m.mu.Unlock()
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
// Remove tag associations
|
m.removeLocked(key)
|
||||||
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)
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -195,16 +201,13 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
|
|||||||
m.mu.Lock()
|
m.mu.Lock()
|
||||||
defer m.mu.Unlock()
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
// Get all keys associated with this tag
|
|
||||||
keySet, exists := m.tagToKeys[tag]
|
keySet, exists := m.tagToKeys[tag]
|
||||||
if !exists {
|
if !exists {
|
||||||
return nil // No keys with this tag
|
return nil // No keys with this tag
|
||||||
}
|
}
|
||||||
|
|
||||||
// Delete all items with this tag
|
|
||||||
for key := range keySet {
|
for key := range keySet {
|
||||||
if item, ok := m.items[key]; ok {
|
if item, ok := m.items[key]; ok {
|
||||||
// Remove this tag from the item's tag list
|
|
||||||
newTags := make([]string, 0, len(item.Tags))
|
newTags := make([]string, 0, len(item.Tags))
|
||||||
for _, t := range item.Tags {
|
for _, t := range item.Tags {
|
||||||
if t != tag {
|
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
|
// If item has no more tags, delete it; otherwise update its tags
|
||||||
// Otherwise update its tags
|
|
||||||
if len(newTags) == 0 {
|
if len(newTags) == 0 {
|
||||||
delete(m.items, key)
|
delete(m.items, key)
|
||||||
} else {
|
} else {
|
||||||
@@ -222,24 +224,24 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Remove the tag mapping
|
|
||||||
delete(m.tagToKeys, tag)
|
delete(m.tagToKeys, tag)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeleteByPattern removes all keys matching the pattern.
|
// 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 {
|
func (m *MemoryProvider) DeleteByPattern(ctx context.Context, pattern string) error {
|
||||||
m.mu.Lock()
|
|
||||||
defer m.mu.Unlock()
|
|
||||||
|
|
||||||
re, err := regexp.Compile(pattern)
|
re, err := regexp.Compile(pattern)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid pattern: %w", err)
|
return fmt.Errorf("invalid pattern: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
for key := range m.items {
|
for key := range m.items {
|
||||||
if re.MatchString(key) {
|
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()
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
m.items = make(map[string]*memoryItem)
|
m.items = make(map[string]*memoryItem)
|
||||||
|
m.tagToKeys = make(map[string]map[string]struct{})
|
||||||
m.hits.Store(0)
|
m.hits.Store(0)
|
||||||
m.misses.Store(0)
|
m.misses.Store(0)
|
||||||
return nil
|
return nil
|
||||||
@@ -270,12 +273,17 @@ func (m *MemoryProvider) Exists(ctx context.Context, key string) bool {
|
|||||||
return !item.isExpired()
|
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 {
|
func (m *MemoryProvider) Close() error {
|
||||||
|
m.closeOnce.Do(func() { close(m.done) })
|
||||||
|
|
||||||
m.mu.Lock()
|
m.mu.Lock()
|
||||||
defer m.mu.Unlock()
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -284,7 +292,7 @@ func (m *MemoryProvider) Stats(ctx context.Context) (*CacheStats, error) {
|
|||||||
m.mu.RLock()
|
m.mu.RLock()
|
||||||
defer m.mu.RUnlock()
|
defer m.mu.RUnlock()
|
||||||
|
|
||||||
// Clean expired items first
|
// Count non-expired items (read-only)
|
||||||
validKeys := 0
|
validKeys := 0
|
||||||
for _, item := range m.items {
|
for _, item := range m.items {
|
||||||
if !item.isExpired() {
|
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.
|
// 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() {
|
func (m *MemoryProvider) evictOne() {
|
||||||
var oldestKey string
|
var oldestKey string
|
||||||
var oldestTime time.Time
|
var oldest int64
|
||||||
|
|
||||||
for key, item := range m.items {
|
for key, item := range m.items {
|
||||||
if item.isExpired() {
|
if item.isExpired() {
|
||||||
delete(m.items, key)
|
m.removeLocked(key)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if oldestKey == "" || item.LastAccess.Before(oldestTime) {
|
if la := item.lastAccess.Load(); oldestKey == "" || la < oldest {
|
||||||
oldestKey = key
|
oldestKey = key
|
||||||
oldestTime = item.LastAccess
|
oldest = la
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if oldestKey != "" {
|
if oldestKey != "" {
|
||||||
delete(m.items, oldestKey)
|
m.removeLocked(oldestKey)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -333,7 +342,7 @@ func (m *MemoryProvider) CleanExpired(ctx context.Context) int {
|
|||||||
count := 0
|
count := 0
|
||||||
for key, item := range m.items {
|
for key, item := range m.items {
|
||||||
if item.isExpired() {
|
if item.isExpired() {
|
||||||
delete(m.items, key)
|
m.removeLocked(key)
|
||||||
count++
|
count++
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Vendored
+47
-10
@@ -3,15 +3,20 @@ package cache
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/redis/go-redis/v9"
|
"github.com/redis/go-redis/v9"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
// RedisProvider is a Redis implementation of the Provider interface.
|
// RedisProvider is a Redis implementation of the Provider interface.
|
||||||
type RedisProvider struct {
|
type RedisProvider struct {
|
||||||
client *redis.Client
|
client *redis.Client
|
||||||
options *Options
|
options *Options
|
||||||
|
allowFlush bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// RedisConfig contains Redis-specific configuration.
|
// RedisConfig contains Redis-specific configuration.
|
||||||
@@ -33,16 +38,25 @@ type RedisConfig struct {
|
|||||||
|
|
||||||
// Options contains general cache options
|
// Options contains general cache options
|
||||||
Options *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.
|
// NewRedisProvider creates a new Redis cache provider.
|
||||||
func NewRedisProvider(config *RedisConfig) (*RedisProvider, error) {
|
func NewRedisProvider(config *RedisConfig) (*RedisProvider, error) {
|
||||||
if config == nil {
|
// Work on a copy so the caller's struct is not mutated
|
||||||
config = &RedisConfig{
|
var cfg RedisConfig
|
||||||
Host: "localhost",
|
if config != nil {
|
||||||
Port: 6379,
|
cfg = *config
|
||||||
DB: 0,
|
} else {
|
||||||
|
cfg = RedisConfig{Host: "localhost", Port: 6379, DB: 0}
|
||||||
}
|
}
|
||||||
|
config = &cfg
|
||||||
|
if config.Options != nil {
|
||||||
|
o := *config.Options
|
||||||
|
config.Options = &o
|
||||||
}
|
}
|
||||||
|
|
||||||
if config.Host == "" {
|
if config.Host == "" {
|
||||||
@@ -79,6 +93,7 @@ func NewRedisProvider(config *RedisConfig) (*RedisProvider, error) {
|
|||||||
return &RedisProvider{
|
return &RedisProvider{
|
||||||
client: client,
|
client: client,
|
||||||
options: config.Options,
|
options: config.Options,
|
||||||
|
allowFlush: config.AllowFlush,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -89,6 +104,8 @@ func (r *RedisProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
|||||||
return nil, false
|
return nil, false
|
||||||
}
|
}
|
||||||
if err != nil {
|
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 nil, false
|
||||||
}
|
}
|
||||||
return val, true
|
return val, true
|
||||||
@@ -194,7 +211,7 @@ func (r *RedisProvider) DeleteByTag(ctx context.Context, tag string) error {
|
|||||||
|
|
||||||
// DeleteByPattern removes all keys matching the pattern.
|
// DeleteByPattern removes all keys matching the pattern.
|
||||||
func (r *RedisProvider) DeleteByPattern(ctx context.Context, pattern string) error {
|
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()
|
pipe := r.client.Pipeline()
|
||||||
|
|
||||||
count := 0
|
count := 0
|
||||||
@@ -225,7 +242,11 @@ func (r *RedisProvider) DeleteByPattern(ctx context.Context, pattern string) err
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Clear removes all items from the cache.
|
// Clear removes all items from the cache.
|
||||||
|
// It runs FLUSHDB and therefore requires RedisConfig.AllowFlush.
|
||||||
func (r *RedisProvider) Clear(ctx context.Context) error {
|
func (r *RedisProvider) Clear(ctx context.Context) error {
|
||||||
|
if !r.allowFlush {
|
||||||
|
return ErrFlushNotAllowed
|
||||||
|
}
|
||||||
return r.client.FlushDB(ctx).Err()
|
return r.client.FlushDB(ctx).Err()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -244,8 +265,9 @@ func (r *RedisProvider) Close() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Stats returns statistics about the cache provider.
|
// 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) {
|
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 {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to get Redis stats: %w", err)
|
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)
|
return nil, fmt.Errorf("failed to get DB size: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse stats from INFO command
|
counters := map[string]int64{}
|
||||||
// This is a simplified version - you may want to parse more detailed stats
|
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{
|
stats := &CacheStats{
|
||||||
|
Hits: counters["keyspace_hits"],
|
||||||
|
Misses: counters["keyspace_misses"],
|
||||||
Keys: dbSize,
|
Keys: dbSize,
|
||||||
ProviderType: "redis",
|
ProviderType: "redis",
|
||||||
ProviderStats: map[string]any{
|
ProviderStats: map[string]any{
|
||||||
"info": info,
|
"evicted_keys": counters["evicted_keys"],
|
||||||
|
"expired_keys": counters["expired_keys"],
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -39,7 +39,7 @@ func (h *QueryDebugHook) AfterQuery(ctx context.Context, event *bun.QueryEvent)
|
|||||||
// This helps identify which specific field is causing scanning issues
|
// This helps identify which specific field is causing scanning issues
|
||||||
func debugScanIntoStruct(rows interface{}, dest interface{}) error {
|
func debugScanIntoStruct(rows interface{}, dest interface{}) error {
|
||||||
v := reflect.ValueOf(dest)
|
v := reflect.ValueOf(dest)
|
||||||
if v.Kind() != reflect.Ptr {
|
if v.Kind() != reflect.Pointer {
|
||||||
return fmt.Errorf("dest must be a pointer")
|
return fmt.Errorf("dest must be a pointer")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -59,7 +59,7 @@ func debugScanIntoStruct(rows interface{}, dest interface{}) error {
|
|||||||
logger.Debug(" Slice element type: %s", elemType)
|
logger.Debug(" Slice element type: %s", elemType)
|
||||||
|
|
||||||
// If slice of pointers, get the underlying type
|
// If slice of pointers, get the underlying type
|
||||||
if elemType.Kind() == reflect.Ptr {
|
if elemType.Kind() == reflect.Pointer {
|
||||||
structType = elemType.Elem()
|
structType = elemType.Elem()
|
||||||
} else {
|
} else {
|
||||||
structType = elemType
|
structType = elemType
|
||||||
@@ -99,11 +99,12 @@ type BunAdapter struct {
|
|||||||
dbMu sync.RWMutex
|
dbMu sync.RWMutex
|
||||||
dbFactory func() (*bun.DB, error)
|
dbFactory func() (*bun.DB, error)
|
||||||
driverName string
|
driverName string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewBunAdapter creates a new Bun adapter
|
// NewBunAdapter creates a new Bun adapter
|
||||||
func NewBunAdapter(db *bun.DB) *BunAdapter {
|
func NewBunAdapter(db *bun.DB) *BunAdapter {
|
||||||
adapter := &BunAdapter{db: db}
|
adapter := &BunAdapter{db: db, metricsEnabled: true}
|
||||||
// Initialize driver name
|
// Initialize driver name
|
||||||
adapter.driverName = adapter.DriverName()
|
adapter.driverName = adapter.DriverName()
|
||||||
return adapter
|
return adapter
|
||||||
@@ -115,6 +116,12 @@ func (b *BunAdapter) WithDBFactory(factory func() (*bun.DB, error)) *BunAdapter
|
|||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetMetricsEnabled enables or disables query metrics for this adapter.
|
||||||
|
func (b *BunAdapter) SetMetricsEnabled(enabled bool) *BunAdapter {
|
||||||
|
b.metricsEnabled = enabled
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
func (b *BunAdapter) getDB() *bun.DB {
|
func (b *BunAdapter) getDB() *bun.DB {
|
||||||
b.dbMu.RLock()
|
b.dbMu.RLock()
|
||||||
defer b.dbMu.RUnlock()
|
defer b.dbMu.RUnlock()
|
||||||
@@ -162,19 +169,20 @@ func (b *BunAdapter) NewSelect() common.SelectQuery {
|
|||||||
query: b.getDB().NewSelect(),
|
query: b.getDB().NewSelect(),
|
||||||
db: b.db,
|
db: b.db,
|
||||||
driverName: b.driverName,
|
driverName: b.driverName,
|
||||||
|
metricsEnabled: b.metricsEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunAdapter) NewInsert() common.InsertQuery {
|
func (b *BunAdapter) NewInsert() common.InsertQuery {
|
||||||
return &BunInsertQuery{query: b.getDB().NewInsert()}
|
return &BunInsertQuery{query: b.getDB().NewInsert(), driverName: b.driverName, metricsEnabled: b.metricsEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunAdapter) NewUpdate() common.UpdateQuery {
|
func (b *BunAdapter) NewUpdate() common.UpdateQuery {
|
||||||
return &BunUpdateQuery{query: b.getDB().NewUpdate()}
|
return &BunUpdateQuery{query: b.getDB().NewUpdate(), driverName: b.driverName, metricsEnabled: b.metricsEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunAdapter) NewDelete() common.DeleteQuery {
|
func (b *BunAdapter) NewDelete() common.DeleteQuery {
|
||||||
return &BunDeleteQuery{query: b.getDB().NewDelete()}
|
return &BunDeleteQuery{query: b.getDB().NewDelete(), driverName: b.driverName, metricsEnabled: b.metricsEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunAdapter) Exec(ctx context.Context, query string, args ...interface{}) (res common.Result, err error) {
|
func (b *BunAdapter) Exec(ctx context.Context, query string, args ...interface{}) (res common.Result, err error) {
|
||||||
@@ -183,6 +191,8 @@ func (b *BunAdapter) Exec(ctx context.Context, query string, args ...interface{}
|
|||||||
err = logger.HandlePanic("BunAdapter.Exec", r)
|
err = logger.HandlePanic("BunAdapter.Exec", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
|
operation, schema, entity, table := metricTargetFromRawQuery(query, b.driverName)
|
||||||
var result sql.Result
|
var result sql.Result
|
||||||
run := func() error { var e error; result, e = b.getDB().ExecContext(ctx, query, args...); return e }
|
run := func() error { var e error; result, e = b.getDB().ExecContext(ctx, query, args...); return e }
|
||||||
err = run()
|
err = run()
|
||||||
@@ -191,6 +201,7 @@ func (b *BunAdapter) Exec(ctx context.Context, query string, args ...interface{}
|
|||||||
err = run()
|
err = run()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(b.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
return &BunResult{result: result}, err
|
return &BunResult{result: result}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -200,12 +211,15 @@ func (b *BunAdapter) Query(ctx context.Context, dest interface{}, query string,
|
|||||||
err = logger.HandlePanic("BunAdapter.Query", r)
|
err = logger.HandlePanic("BunAdapter.Query", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
|
operation, schema, entity, table := metricTargetFromRawQuery(query, b.driverName)
|
||||||
err = b.getDB().NewRaw(query, args...).Scan(ctx, dest)
|
err = b.getDB().NewRaw(query, args...).Scan(ctx, dest)
|
||||||
if isDBClosed(err) {
|
if isDBClosed(err) {
|
||||||
if reconnErr := b.reconnectDB(); reconnErr == nil {
|
if reconnErr := b.reconnectDB(); reconnErr == nil {
|
||||||
err = b.getDB().NewRaw(query, args...).Scan(ctx, dest)
|
err = b.getDB().NewRaw(query, args...).Scan(ctx, dest)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(b.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -219,7 +233,7 @@ func (b *BunAdapter) BeginTx(ctx context.Context) (common.Database, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return &BunTxAdapter{tx: tx, driverName: b.driverName}, nil
|
return &BunTxAdapter{tx: tx, driverName: b.driverName, metricsEnabled: b.metricsEnabled}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunAdapter) CommitTx(ctx context.Context) error {
|
func (b *BunAdapter) CommitTx(ctx context.Context) error {
|
||||||
@@ -242,7 +256,7 @@ func (b *BunAdapter) RunInTransaction(ctx context.Context, fn func(common.Databa
|
|||||||
}()
|
}()
|
||||||
run := func() error {
|
run := func() error {
|
||||||
return b.getDB().RunInTx(ctx, &sql.TxOptions{}, func(ctx context.Context, tx bun.Tx) error {
|
return b.getDB().RunInTx(ctx, &sql.TxOptions{}, func(ctx context.Context, tx bun.Tx) error {
|
||||||
adapter := &BunTxAdapter{tx: tx, driverName: b.driverName}
|
adapter := &BunTxAdapter{tx: tx, driverName: b.driverName, metricsEnabled: b.metricsEnabled}
|
||||||
return fn(adapter)
|
return fn(adapter)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -280,25 +294,25 @@ type BunSelectQuery struct {
|
|||||||
hasModel bool // Track if Model() was called
|
hasModel bool // Track if Model() was called
|
||||||
schema string // Separated schema name
|
schema string // Separated schema name
|
||||||
tableName string // Just the table name, without schema
|
tableName string // Just the table name, without schema
|
||||||
|
entity string
|
||||||
tableAlias string
|
tableAlias string
|
||||||
driverName string // Database driver name (postgres, sqlite, mssql)
|
driverName string // Database driver name (postgres, sqlite, mssql)
|
||||||
inJoinContext bool // Track if we're in a JOIN relation context
|
inJoinContext bool // Track if we're in a JOIN relation context
|
||||||
joinTableAlias string // Alias to use for JOIN conditions
|
joinTableAlias string // Alias to use for JOIN conditions
|
||||||
skipAutoDetect bool // Skip auto-detection to prevent circular calls
|
skipAutoDetect bool // Skip auto-detection to prevent circular calls
|
||||||
|
preloadRelationAlias string // Relation alias used in separate-query preloads (e.g. "tprp" for relation "TPRP")
|
||||||
customPreloads map[string][]func(common.SelectQuery) common.SelectQuery // Relations to load with custom implementation
|
customPreloads map[string][]func(common.SelectQuery) common.SelectQuery // Relations to load with custom implementation
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunSelectQuery) Model(model interface{}) common.SelectQuery {
|
func (b *BunSelectQuery) Model(model interface{}) common.SelectQuery {
|
||||||
b.query = b.query.Model(model)
|
b.query = b.query.Model(model)
|
||||||
b.hasModel = true // Mark that we have a model
|
b.hasModel = true // Mark that we have a model
|
||||||
|
b.schema, b.tableName = schemaAndTableFromModel(model, b.driverName)
|
||||||
// Try to get table name from model if it implements TableNameProvider
|
if b.tableName == "" {
|
||||||
if provider, ok := model.(common.TableNameProvider); ok {
|
b.schema, b.tableName = parseTableName(b.query.GetTableName(), b.driverName)
|
||||||
fullTableName := provider.TableName()
|
|
||||||
// Check if the table name contains schema (e.g., "schema.table")
|
|
||||||
// For SQLite, this will convert "schema.table" to "schema_table"
|
|
||||||
b.schema, b.tableName = parseTableName(fullTableName, b.driverName)
|
|
||||||
}
|
}
|
||||||
|
b.entity = entityNameFromModel(model, b.tableName)
|
||||||
|
|
||||||
if provider, ok := model.(common.TableAliasProvider); ok {
|
if provider, ok := model.(common.TableAliasProvider); ok {
|
||||||
b.tableAlias = provider.TableAlias()
|
b.tableAlias = provider.TableAlias()
|
||||||
@@ -312,6 +326,9 @@ func (b *BunSelectQuery) Table(table string) common.SelectQuery {
|
|||||||
// Check if the table name contains schema (e.g., "schema.table")
|
// Check if the table name contains schema (e.g., "schema.table")
|
||||||
// For SQLite, this will convert "schema.table" to "schema_table"
|
// For SQLite, this will convert "schema.table" to "schema_table"
|
||||||
b.schema, b.tableName = parseTableName(table, b.driverName)
|
b.schema, b.tableName = parseTableName(table, b.driverName)
|
||||||
|
if b.entity == "" {
|
||||||
|
b.entity = cleanMetricIdentifier(b.tableName)
|
||||||
|
}
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -322,7 +339,7 @@ func (b *BunSelectQuery) Column(columns ...string) common.SelectQuery {
|
|||||||
|
|
||||||
func (b *BunSelectQuery) ColumnExpr(query string, args ...interface{}) common.SelectQuery {
|
func (b *BunSelectQuery) ColumnExpr(query string, args ...interface{}) common.SelectQuery {
|
||||||
if len(args) > 0 {
|
if len(args) > 0 {
|
||||||
b.query = b.query.ColumnExpr(query, args)
|
b.query = b.query.ColumnExpr(query, args...)
|
||||||
} else {
|
} else {
|
||||||
b.query = b.query.ColumnExpr(query)
|
b.query = b.query.ColumnExpr(query)
|
||||||
}
|
}
|
||||||
@@ -330,12 +347,14 @@ func (b *BunSelectQuery) ColumnExpr(query string, args ...interface{}) common.Se
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunSelectQuery) Where(query string, args ...interface{}) common.SelectQuery {
|
func (b *BunSelectQuery) Where(query string, args ...interface{}) common.SelectQuery {
|
||||||
// If we're in a JOIN context, add table prefix to unqualified columns
|
|
||||||
if b.inJoinContext && b.joinTableAlias != "" {
|
if b.inJoinContext && b.joinTableAlias != "" {
|
||||||
query = addTablePrefix(query, b.joinTableAlias)
|
query = addTablePrefix(query, b.joinTableAlias)
|
||||||
|
} else if b.preloadRelationAlias != "" && b.tableName != "" {
|
||||||
|
// Separate-query preload: the caller may have written conditions using the
|
||||||
|
// relation name as a prefix (e.g. "TPRP.col"). Bun uses the real table name
|
||||||
|
// as the alias, so rewrite any such references to use tableName instead.
|
||||||
|
query = replaceRelationAlias(query, b.preloadRelationAlias, b.tableName)
|
||||||
} else if b.tableAlias != "" && b.tableName != "" {
|
} else if b.tableAlias != "" && b.tableName != "" {
|
||||||
// If we have a table alias defined, check if the query references a different alias
|
|
||||||
// This can happen in preloads where the user expects a certain alias but Bun generates another
|
|
||||||
query = normalizeTableAlias(query, b.tableAlias, b.tableName)
|
query = normalizeTableAlias(query, b.tableAlias, b.tableName)
|
||||||
}
|
}
|
||||||
b.query = b.query.Where(query, args...)
|
b.query = b.query.Where(query, args...)
|
||||||
@@ -471,6 +490,38 @@ func normalizeTableAlias(query, expectedAlias, tableName string) string {
|
|||||||
return modified
|
return modified
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// replaceRelationAlias rewrites WHERE conditions written with a relation alias prefix
|
||||||
|
// (e.g. "TPRP.col") to use the real table name that bun uses in separate queries
|
||||||
|
// (e.g. "t_proposalinstance.col"). Only called for separate-query preload wrappers.
|
||||||
|
func replaceRelationAlias(query, relationAlias, tableName string) string {
|
||||||
|
if relationAlias == "" || tableName == "" || query == "" {
|
||||||
|
return query
|
||||||
|
}
|
||||||
|
parts := strings.FieldsFunc(query, func(r rune) bool {
|
||||||
|
return r == ' ' || r == '(' || r == ')' || r == ','
|
||||||
|
})
|
||||||
|
modified := query
|
||||||
|
for _, part := range parts {
|
||||||
|
if dotIndex := strings.Index(part, "."); dotIndex > 0 {
|
||||||
|
prefix := part[:dotIndex]
|
||||||
|
column := part[dotIndex+1:]
|
||||||
|
if strings.EqualFold(prefix, relationAlias) {
|
||||||
|
logger.Debug("Replacing relation alias '%s' with table name '%s' in preload WHERE condition", prefix, tableName)
|
||||||
|
modified = strings.ReplaceAll(modified, part, tableName+"."+column)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return modified
|
||||||
|
}
|
||||||
|
|
||||||
|
func isJoinKeyword(word string) bool {
|
||||||
|
switch strings.ToUpper(word) {
|
||||||
|
case "JOIN", "INNER", "LEFT", "RIGHT", "FULL", "OUTER", "CROSS":
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
func (b *BunSelectQuery) WhereOr(query string, args ...interface{}) common.SelectQuery {
|
func (b *BunSelectQuery) WhereOr(query string, args ...interface{}) common.SelectQuery {
|
||||||
b.query = b.query.WhereOr(query, args...)
|
b.query = b.query.WhereOr(query, args...)
|
||||||
return b
|
return b
|
||||||
@@ -501,7 +552,7 @@ func (b *BunSelectQuery) Join(query string, args ...interface{}) common.SelectQu
|
|||||||
if prefix != "" && !strings.Contains(strings.ToUpper(query), " AS ") {
|
if prefix != "" && !strings.Contains(strings.ToUpper(query), " AS ") {
|
||||||
// If query doesn't already have AS, check if it's a simple table name
|
// If query doesn't already have AS, check if it's a simple table name
|
||||||
parts := strings.Fields(query)
|
parts := strings.Fields(query)
|
||||||
if len(parts) > 0 && !strings.HasPrefix(strings.ToUpper(parts[0]), "JOIN") {
|
if len(parts) > 0 && !isJoinKeyword(parts[0]) {
|
||||||
// Simple table name, add prefix: "table AS prefix"
|
// Simple table name, add prefix: "table AS prefix"
|
||||||
joinClause = fmt.Sprintf("%s AS %s", parts[0], prefix)
|
joinClause = fmt.Sprintf("%s AS %s", parts[0], prefix)
|
||||||
if len(parts) > 1 {
|
if len(parts) > 1 {
|
||||||
@@ -536,7 +587,7 @@ func (b *BunSelectQuery) LeftJoin(query string, args ...interface{}) common.Sele
|
|||||||
joinClause := query
|
joinClause := query
|
||||||
if prefix != "" && !strings.Contains(strings.ToUpper(query), " AS ") {
|
if prefix != "" && !strings.Contains(strings.ToUpper(query), " AS ") {
|
||||||
parts := strings.Fields(query)
|
parts := strings.Fields(query)
|
||||||
if len(parts) > 0 && !strings.HasPrefix(strings.ToUpper(parts[0]), "LEFT") && !strings.HasPrefix(strings.ToUpper(parts[0]), "JOIN") {
|
if len(parts) > 0 && !isJoinKeyword(parts[0]) {
|
||||||
joinClause = fmt.Sprintf("%s AS %s", parts[0], prefix)
|
joinClause = fmt.Sprintf("%s AS %s", parts[0], prefix)
|
||||||
if len(parts) > 1 {
|
if len(parts) > 1 {
|
||||||
joinClause += " " + strings.Join(parts[1:], " ")
|
joinClause += " " + strings.Join(parts[1:], " ")
|
||||||
@@ -581,6 +632,19 @@ func (b *BunSelectQuery) PreloadRelation(relation string, apply ...func(common.S
|
|||||||
if !b.skipAutoDetect {
|
if !b.skipAutoDetect {
|
||||||
model := b.query.GetModel()
|
model := b.query.GetModel()
|
||||||
if model != nil && model.Value() != nil {
|
if model != nil && model.Value() != nil {
|
||||||
|
// Guard against relations that don't exist on the model. Without this,
|
||||||
|
// bun panics inside Count/Scan with `model=X does not have relation="Y"`.
|
||||||
|
// Only validate the root segment so nested paths (e.g. "PRM.CHILD") still
|
||||||
|
// fall through to bun's native resolution.
|
||||||
|
rootRelation := relation
|
||||||
|
if idx := strings.Index(rootRelation, "."); idx >= 0 {
|
||||||
|
rootRelation = rootRelation[:idx]
|
||||||
|
}
|
||||||
|
if reflection.GetRelationType(model.Value(), rootRelation) == reflection.RelationUnknown {
|
||||||
|
logger.Warn("Skipping preload '%s': relation '%s' is not declared on model %T", relation, rootRelation, model.Value())
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
relType := reflection.GetRelationType(model.Value(), relation)
|
relType := reflection.GetRelationType(model.Value(), relation)
|
||||||
|
|
||||||
// Log the detected relationship type
|
// Log the detected relationship type
|
||||||
@@ -605,10 +669,7 @@ func (b *BunSelectQuery) PreloadRelation(relation string, apply ...func(common.S
|
|||||||
b.query = b.query.Relation(relation, func(sq *bun.SelectQuery) *bun.SelectQuery {
|
b.query = b.query.Relation(relation, func(sq *bun.SelectQuery) *bun.SelectQuery {
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
err := logger.HandlePanic("BunSelectQuery.PreloadRelation", r)
|
_ = logger.HandlePanic("BunSelectQuery.PreloadRelation", r)
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
if len(apply) == 0 {
|
if len(apply) == 0 {
|
||||||
@@ -620,6 +681,7 @@ func (b *BunSelectQuery) PreloadRelation(relation string, apply ...func(common.S
|
|||||||
query: sq,
|
query: sq,
|
||||||
db: b.db,
|
db: b.db,
|
||||||
driverName: b.driverName,
|
driverName: b.driverName,
|
||||||
|
metricsEnabled: b.metricsEnabled,
|
||||||
}
|
}
|
||||||
|
|
||||||
// Try to extract table name and alias from the preload model
|
// Try to extract table name and alias from the preload model
|
||||||
@@ -638,8 +700,20 @@ func (b *BunSelectQuery) PreloadRelation(relation string, apply ...func(common.S
|
|||||||
wrapper.tableAlias = provider.TableAlias()
|
wrapper.tableAlias = provider.TableAlias()
|
||||||
logger.Debug("Preload relation '%s' using table alias: %s", relation, wrapper.tableAlias)
|
logger.Debug("Preload relation '%s' using table alias: %s", relation, wrapper.tableAlias)
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Fallback: if the model didn't provide a table name, ask bun directly.
|
||||||
|
if wrapper.tableName == "" {
|
||||||
|
wrapper.schema, wrapper.tableName = parseTableName(sq.GetTableName(), b.driverName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// For separate-query preloads (has-many), bun aliases the related table using
|
||||||
|
// the actual table name, not the relation name. Record the relation alias so
|
||||||
|
// Where() can rewrite conditions like "TPRP.col" to "t_proposalinstance.col".
|
||||||
|
wrapper.preloadRelationAlias = strings.ToLower(relation)
|
||||||
|
logger.Debug("Preload relation '%s' registered alias '%s' for separate-query WHERE rewriting", relation, wrapper.preloadRelationAlias)
|
||||||
|
|
||||||
// Start with the interface value (not pointer)
|
// Start with the interface value (not pointer)
|
||||||
current := common.SelectQuery(wrapper)
|
current := common.SelectQuery(wrapper)
|
||||||
|
|
||||||
@@ -670,7 +744,7 @@ func checkIfRelationAlreadyLoaded(parents reflect.Value, relationName string) (r
|
|||||||
|
|
||||||
// Get the first parent to check the relation field
|
// Get the first parent to check the relation field
|
||||||
firstParent := parents.Index(0)
|
firstParent := parents.Index(0)
|
||||||
if firstParent.Kind() == reflect.Ptr {
|
if firstParent.Kind() == reflect.Pointer {
|
||||||
firstParent = firstParent.Elem()
|
firstParent = firstParent.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -685,7 +759,7 @@ func checkIfRelationAlreadyLoaded(parents reflect.Value, relationName string) (r
|
|||||||
// Check if any parent has a non-empty slice
|
// Check if any parent has a non-empty slice
|
||||||
for i := 0; i < parents.Len(); i++ {
|
for i := 0; i < parents.Len(); i++ {
|
||||||
parent := parents.Index(i)
|
parent := parents.Index(i)
|
||||||
if parent.Kind() == reflect.Ptr {
|
if parent.Kind() == reflect.Pointer {
|
||||||
parent = parent.Elem()
|
parent = parent.Elem()
|
||||||
}
|
}
|
||||||
field := parent.FieldByName(relationName)
|
field := parent.FieldByName(relationName)
|
||||||
@@ -694,7 +768,7 @@ func checkIfRelationAlreadyLoaded(parents reflect.Value, relationName string) (r
|
|||||||
allRelated := reflect.MakeSlice(field.Type(), 0, field.Len()*parents.Len())
|
allRelated := reflect.MakeSlice(field.Type(), 0, field.Len()*parents.Len())
|
||||||
for j := 0; j < parents.Len(); j++ {
|
for j := 0; j < parents.Len(); j++ {
|
||||||
p := parents.Index(j)
|
p := parents.Index(j)
|
||||||
if p.Kind() == reflect.Ptr {
|
if p.Kind() == reflect.Pointer {
|
||||||
p = p.Elem()
|
p = p.Elem()
|
||||||
}
|
}
|
||||||
f := p.FieldByName(relationName)
|
f := p.FieldByName(relationName)
|
||||||
@@ -707,7 +781,7 @@ func checkIfRelationAlreadyLoaded(parents reflect.Value, relationName string) (r
|
|||||||
return allRelated, true
|
return allRelated, true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
} else if relationField.Kind() == reflect.Ptr {
|
} else if relationField.Kind() == reflect.Pointer {
|
||||||
// Check if it's a pointer (has-one/belongs-to)
|
// Check if it's a pointer (has-one/belongs-to)
|
||||||
if !relationField.IsNil() {
|
if !relationField.IsNil() {
|
||||||
// Already loaded! Collect all related records from all parents
|
// Already loaded! Collect all related records from all parents
|
||||||
@@ -715,7 +789,7 @@ func checkIfRelationAlreadyLoaded(parents reflect.Value, relationName string) (r
|
|||||||
allRelated := reflect.MakeSlice(reflect.SliceOf(relatedType), 0, parents.Len())
|
allRelated := reflect.MakeSlice(reflect.SliceOf(relatedType), 0, parents.Len())
|
||||||
for j := 0; j < parents.Len(); j++ {
|
for j := 0; j < parents.Len(); j++ {
|
||||||
p := parents.Index(j)
|
p := parents.Index(j)
|
||||||
if p.Kind() == reflect.Ptr {
|
if p.Kind() == reflect.Pointer {
|
||||||
p = p.Elem()
|
p = p.Elem()
|
||||||
}
|
}
|
||||||
f := p.FieldByName(relationName)
|
f := p.FieldByName(relationName)
|
||||||
@@ -739,7 +813,7 @@ func (b *BunSelectQuery) loadCustomPreloads(ctx context.Context) error {
|
|||||||
|
|
||||||
// Get the actual data from the model
|
// Get the actual data from the model
|
||||||
modelValue := reflect.ValueOf(model.Value())
|
modelValue := reflect.ValueOf(model.Value())
|
||||||
if modelValue.Kind() == reflect.Ptr {
|
if modelValue.Kind() == reflect.Pointer {
|
||||||
modelValue = modelValue.Elem()
|
modelValue = modelValue.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -807,7 +881,7 @@ func (b *BunSelectQuery) loadRelationLevel(ctx context.Context, parentRecords re
|
|||||||
|
|
||||||
// Get the first record to inspect the struct type
|
// Get the first record to inspect the struct type
|
||||||
firstRecord := parentRecords.Index(0)
|
firstRecord := parentRecords.Index(0)
|
||||||
if firstRecord.Kind() == reflect.Ptr {
|
if firstRecord.Kind() == reflect.Pointer {
|
||||||
firstRecord = firstRecord.Elem()
|
firstRecord = firstRecord.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -853,7 +927,7 @@ func (b *BunSelectQuery) loadRelationLevel(ctx context.Context, parentRecords re
|
|||||||
if isSlice {
|
if isSlice {
|
||||||
relatedType = relatedType.Elem()
|
relatedType = relatedType.Elem()
|
||||||
}
|
}
|
||||||
if relatedType.Kind() == reflect.Ptr {
|
if relatedType.Kind() == reflect.Pointer {
|
||||||
relatedType = relatedType.Elem()
|
relatedType = relatedType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -870,7 +944,7 @@ func (b *BunSelectQuery) loadRelationLevel(ctx context.Context, parentRecords re
|
|||||||
|
|
||||||
// Apply user's functions (if any)
|
// Apply user's functions (if any)
|
||||||
if isLast && len(applyFuncs) > 0 {
|
if isLast && len(applyFuncs) > 0 {
|
||||||
wrapper := &BunSelectQuery{query: query, db: b.db, driverName: b.driverName}
|
wrapper := &BunSelectQuery{query: query, db: b.db, driverName: b.driverName, metricsEnabled: b.metricsEnabled}
|
||||||
for _, fn := range applyFuncs {
|
for _, fn := range applyFuncs {
|
||||||
if fn != nil {
|
if fn != nil {
|
||||||
wrapper = fn(wrapper).(*BunSelectQuery)
|
wrapper = fn(wrapper).(*BunSelectQuery)
|
||||||
@@ -941,7 +1015,7 @@ func extractForeignKeyValues(records reflect.Value, fkFieldName string) ([]inter
|
|||||||
|
|
||||||
for i := 0; i < records.Len(); i++ {
|
for i := 0; i < records.Len(); i++ {
|
||||||
record := records.Index(i)
|
record := records.Index(i)
|
||||||
if record.Kind() == reflect.Ptr {
|
if record.Kind() == reflect.Pointer {
|
||||||
record = record.Elem()
|
record = record.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1006,7 +1080,7 @@ func associateRelatedRecords(parents, related reflect.Value, fieldName string, r
|
|||||||
for i := 0; i < related.Len(); i++ {
|
for i := 0; i < related.Len(); i++ {
|
||||||
relRecord := related.Index(i)
|
relRecord := related.Index(i)
|
||||||
relRecordElem := relRecord
|
relRecordElem := relRecord
|
||||||
if relRecordElem.Kind() == reflect.Ptr {
|
if relRecordElem.Kind() == reflect.Pointer {
|
||||||
relRecordElem = relRecordElem.Elem()
|
relRecordElem = relRecordElem.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1032,7 +1106,7 @@ func associateRelatedRecords(parents, related reflect.Value, fieldName string, r
|
|||||||
for i := 0; i < parents.Len(); i++ {
|
for i := 0; i < parents.Len(); i++ {
|
||||||
parentPtr := parents.Index(i)
|
parentPtr := parents.Index(i)
|
||||||
parent := parentPtr
|
parent := parentPtr
|
||||||
if parent.Kind() == reflect.Ptr {
|
if parent.Kind() == reflect.Pointer {
|
||||||
parent = parent.Elem()
|
parent = parent.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1222,27 +1296,29 @@ func (b *BunSelectQuery) Having(having string, args ...interface{}) common.Selec
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunSelectQuery) Scan(ctx context.Context, dest interface{}) (err error) {
|
func (b *BunSelectQuery) Scan(ctx context.Context, dest interface{}) (err error) {
|
||||||
|
startedAt := time.Now()
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("BunSelectQuery.Scan", r)
|
err = logger.HandlePanic("BunSelectQuery.Scan", r)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(b.metricsEnabled, "SELECT", b.schema, b.entity, b.tableName, startedAt, err)
|
||||||
}()
|
}()
|
||||||
if dest == nil {
|
if dest == nil {
|
||||||
return fmt.Errorf("destination cannot be nil")
|
err = fmt.Errorf("destination cannot be nil")
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
err = b.query.Scan(ctx, dest)
|
err = b.query.Scan(ctx, dest)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Log SQL string for debugging
|
|
||||||
sqlStr := b.query.String()
|
sqlStr := b.query.String()
|
||||||
logger.Error("BunSelectQuery.Scan failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("BunSelectQuery.Scan failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
return err
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
|
return err
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunSelectQuery) ScanModel(ctx context.Context) (err error) {
|
func (b *BunSelectQuery) ScanModel(ctx context.Context) (err error) {
|
||||||
|
startedAt := time.Now()
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
// Enhanced panic recovery with model information
|
// Enhanced panic recovery with model information
|
||||||
@@ -1252,13 +1328,12 @@ func (b *BunSelectQuery) ScanModel(ctx context.Context) (err error) {
|
|||||||
modelValue := model.Value()
|
modelValue := model.Value()
|
||||||
modelInfo = fmt.Sprintf("Model type: %T", modelValue)
|
modelInfo = fmt.Sprintf("Model type: %T", modelValue)
|
||||||
|
|
||||||
// Try to get the model's underlying struct type
|
|
||||||
v := reflect.ValueOf(modelValue)
|
v := reflect.ValueOf(modelValue)
|
||||||
if v.Kind() == reflect.Ptr {
|
if v.Kind() == reflect.Pointer {
|
||||||
v = v.Elem()
|
v = v.Elem()
|
||||||
}
|
}
|
||||||
if v.Kind() == reflect.Slice {
|
if v.Kind() == reflect.Slice {
|
||||||
if v.Type().Elem().Kind() == reflect.Ptr {
|
if v.Type().Elem().Kind() == reflect.Pointer {
|
||||||
modelInfo += fmt.Sprintf(", Slice of: %s", v.Type().Elem().Elem().Name())
|
modelInfo += fmt.Sprintf(", Slice of: %s", v.Type().Elem().Elem().Name())
|
||||||
} else {
|
} else {
|
||||||
modelInfo += fmt.Sprintf(", Slice of: %s", v.Type().Elem().Name())
|
modelInfo += fmt.Sprintf(", Slice of: %s", v.Type().Elem().Name())
|
||||||
@@ -1272,9 +1347,11 @@ func (b *BunSelectQuery) ScanModel(ctx context.Context) (err error) {
|
|||||||
logger.Error("Panic in BunSelectQuery.ScanModel: %v. %s. SQL: %s", r, modelInfo, sqlStr)
|
logger.Error("Panic in BunSelectQuery.ScanModel: %v. %s. SQL: %s", r, modelInfo, sqlStr)
|
||||||
err = logger.HandlePanic("BunSelectQuery.ScanModel", r)
|
err = logger.HandlePanic("BunSelectQuery.ScanModel", r)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(b.metricsEnabled, "SELECT", b.schema, b.entity, b.tableName, startedAt, err)
|
||||||
}()
|
}()
|
||||||
if b.query.GetModel() == nil {
|
if b.query.GetModel() == nil {
|
||||||
return fmt.Errorf("model is nil")
|
err = fmt.Errorf("model is nil")
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Optional: Enable detailed field-level debugging (set to true to debug)
|
// Optional: Enable detailed field-level debugging (set to true to debug)
|
||||||
@@ -1290,16 +1367,15 @@ func (b *BunSelectQuery) ScanModel(ctx context.Context) (err error) {
|
|||||||
|
|
||||||
err = b.query.Scan(ctx)
|
err = b.query.Scan(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Log SQL string for debugging
|
|
||||||
sqlStr := b.query.String()
|
sqlStr := b.query.String()
|
||||||
logger.Error("BunSelectQuery.ScanModel failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("BunSelectQuery.ScanModel failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
return err
|
return common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
|
|
||||||
// After main query, load custom preloads using separate queries
|
// After main query, load custom preloads using separate queries
|
||||||
if len(b.customPreloads) > 0 {
|
if len(b.customPreloads) > 0 {
|
||||||
logger.Info("Loading %d custom preload(s) with separate queries", len(b.customPreloads))
|
logger.Info("Loading %d custom preload(s) with separate queries", len(b.customPreloads))
|
||||||
if err := b.loadCustomPreloads(ctx); err != nil {
|
if err = b.loadCustomPreloads(ctx); err != nil {
|
||||||
logger.Error("Failed to load custom preloads: %v", err)
|
logger.Error("Failed to load custom preloads: %v", err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -1309,21 +1385,23 @@ func (b *BunSelectQuery) ScanModel(ctx context.Context) (err error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunSelectQuery) Count(ctx context.Context) (count int, err error) {
|
func (b *BunSelectQuery) Count(ctx context.Context) (count int, err error) {
|
||||||
|
startedAt := time.Now()
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("BunSelectQuery.Count", r)
|
err = logger.HandlePanic("BunSelectQuery.Count", r)
|
||||||
count = 0
|
count = 0
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(b.metricsEnabled, "COUNT", b.schema, b.entity, b.tableName, startedAt, err)
|
||||||
}()
|
}()
|
||||||
// If Model() was set, use bun's native Count() which works properly
|
// If Model() was set, use bun's native Count() which works properly
|
||||||
if b.hasModel {
|
if b.hasModel {
|
||||||
count, err := b.query.Count(ctx)
|
count, err = b.query.Count(ctx) // assign to named returns, not shadow vars
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Log SQL string for debugging
|
|
||||||
sqlStr := b.query.String()
|
sqlStr := b.query.String()
|
||||||
logger.Error("BunSelectQuery.Count failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("BunSelectQuery.Count failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
return count, err
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Otherwise, wrap as subquery to avoid "Model(nil)" error
|
// Otherwise, wrap as subquery to avoid "Model(nil)" error
|
||||||
@@ -1333,27 +1411,29 @@ func (b *BunSelectQuery) Count(ctx context.Context) (count int, err error) {
|
|||||||
ColumnExpr("COUNT(*)")
|
ColumnExpr("COUNT(*)")
|
||||||
err = countQuery.Scan(ctx, &count)
|
err = countQuery.Scan(ctx, &count)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Log SQL string for debugging
|
|
||||||
sqlStr := countQuery.String()
|
sqlStr := countQuery.String()
|
||||||
logger.Error("BunSelectQuery.Count (subquery) failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("BunSelectQuery.Count (subquery) failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
return count, err
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunSelectQuery) Exists(ctx context.Context) (exists bool, err error) {
|
func (b *BunSelectQuery) Exists(ctx context.Context) (exists bool, err error) {
|
||||||
|
startedAt := time.Now()
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("BunSelectQuery.Exists", r)
|
err = logger.HandlePanic("BunSelectQuery.Exists", r)
|
||||||
exists = false
|
exists = false
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(b.metricsEnabled, "EXISTS", b.schema, b.entity, b.tableName, startedAt, err)
|
||||||
}()
|
}()
|
||||||
exists, err = b.query.Exists(ctx)
|
exists, err = b.query.Exists(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Log SQL string for debugging
|
|
||||||
sqlStr := b.query.String()
|
sqlStr := b.query.String()
|
||||||
logger.Error("BunSelectQuery.Exists failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("BunSelectQuery.Exists failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
return exists, err
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// BunInsertQuery implements InsertQuery for Bun
|
// BunInsertQuery implements InsertQuery for Bun
|
||||||
@@ -1361,11 +1441,21 @@ type BunInsertQuery struct {
|
|||||||
query *bun.InsertQuery
|
query *bun.InsertQuery
|
||||||
values map[string]interface{}
|
values map[string]interface{}
|
||||||
hasModel bool
|
hasModel bool
|
||||||
|
driverName string
|
||||||
|
schema string
|
||||||
|
tableName string
|
||||||
|
entity string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunInsertQuery) Model(model interface{}) common.InsertQuery {
|
func (b *BunInsertQuery) Model(model interface{}) common.InsertQuery {
|
||||||
b.query = b.query.Model(model)
|
b.query = b.query.Model(model)
|
||||||
b.hasModel = true
|
b.hasModel = true
|
||||||
|
b.schema, b.tableName = schemaAndTableFromModel(model, b.driverName)
|
||||||
|
if b.tableName == "" {
|
||||||
|
b.schema, b.tableName = parseTableName(b.query.GetTableName(), b.driverName)
|
||||||
|
}
|
||||||
|
b.entity = entityNameFromModel(model, b.tableName)
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1374,6 +1464,10 @@ func (b *BunInsertQuery) Table(table string) common.InsertQuery {
|
|||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
b.query = b.query.Table(table)
|
b.query = b.query.Table(table)
|
||||||
|
b.schema, b.tableName = parseTableName(table, b.driverName)
|
||||||
|
if b.entity == "" {
|
||||||
|
b.entity = cleanMetricIdentifier(b.tableName)
|
||||||
|
}
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1392,53 +1486,84 @@ func (b *BunInsertQuery) OnConflict(action string) common.InsertQuery {
|
|||||||
|
|
||||||
func (b *BunInsertQuery) Returning(columns ...string) common.InsertQuery {
|
func (b *BunInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||||
if len(columns) > 0 {
|
if len(columns) > 0 {
|
||||||
b.query = b.query.Returning(columns[0])
|
b.query = b.query.Returning(strings.Join(columns, ", "))
|
||||||
}
|
}
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (b *BunInsertQuery) prepareValues() {
|
||||||
|
if len(b.values) > 0 {
|
||||||
|
if !b.hasModel {
|
||||||
|
b.query = b.query.Model(&b.values)
|
||||||
|
} else {
|
||||||
|
for k, v := range b.values {
|
||||||
|
b.query = b.query.Value(k, "?", v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (b *BunInsertQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
func (b *BunInsertQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("BunInsertQuery.Exec", r)
|
err = logger.HandlePanic("BunInsertQuery.Exec", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
if len(b.values) > 0 {
|
startedAt := time.Now()
|
||||||
if !b.hasModel {
|
b.prepareValues()
|
||||||
// If no model was set, use the values map as the model
|
|
||||||
// Bun can insert map[string]interface{} directly
|
|
||||||
b.query = b.query.Model(&b.values)
|
|
||||||
} else {
|
|
||||||
// If model was set, use Value() to add individual values
|
|
||||||
for k, v := range b.values {
|
|
||||||
b.query = b.query.Value(k, "?", v)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
result, err := b.query.Exec(ctx)
|
result, err := b.query.Exec(ctx)
|
||||||
|
recordQueryMetrics(b.metricsEnabled, "INSERT", b.schema, b.entity, b.tableName, startedAt, err)
|
||||||
return &BunResult{result: result}, err
|
return &BunResult{result: result}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (b *BunInsertQuery) Scan(ctx context.Context, dest interface{}) (err error) {
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
err = logger.HandlePanic("BunInsertQuery.Scan", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
|
b.prepareValues()
|
||||||
|
err = b.query.Scan(ctx, dest)
|
||||||
|
recordQueryMetrics(b.metricsEnabled, "INSERT", b.schema, b.entity, b.tableName, startedAt, err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
// BunUpdateQuery implements UpdateQuery for Bun
|
// BunUpdateQuery implements UpdateQuery for Bun
|
||||||
type BunUpdateQuery struct {
|
type BunUpdateQuery struct {
|
||||||
query *bun.UpdateQuery
|
query *bun.UpdateQuery
|
||||||
model interface{}
|
model interface{}
|
||||||
|
driverName string
|
||||||
|
schema string
|
||||||
|
tableName string
|
||||||
|
entity string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunUpdateQuery) Model(model interface{}) common.UpdateQuery {
|
func (b *BunUpdateQuery) Model(model interface{}) common.UpdateQuery {
|
||||||
b.query = b.query.Model(model)
|
b.query = b.query.Model(model)
|
||||||
b.model = model
|
b.model = model
|
||||||
|
b.schema, b.tableName = schemaAndTableFromModel(model, b.driverName)
|
||||||
|
if b.tableName == "" {
|
||||||
|
b.schema, b.tableName = parseTableName(b.query.GetTableName(), b.driverName)
|
||||||
|
}
|
||||||
|
b.entity = entityNameFromModel(model, b.tableName)
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunUpdateQuery) Table(table string) common.UpdateQuery {
|
func (b *BunUpdateQuery) Table(table string) common.UpdateQuery {
|
||||||
b.query = b.query.Table(table)
|
b.query = b.query.Table(table)
|
||||||
|
b.schema, b.tableName = parseTableName(table, b.driverName)
|
||||||
|
if b.entity == "" {
|
||||||
|
b.entity = cleanMetricIdentifier(b.tableName)
|
||||||
|
}
|
||||||
if b.model == nil {
|
if b.model == nil {
|
||||||
// Try to get table name from table string if model is not set
|
// Try to get table name from table string if model is not set
|
||||||
|
|
||||||
model, err := modelregistry.GetModelByName(table)
|
model, err := modelregistry.GetModelByName(table)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
b.model = model
|
b.model = model
|
||||||
|
b.entity = entityNameFromModel(model, b.tableName)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return b
|
return b
|
||||||
@@ -1466,7 +1591,7 @@ func (b *BunUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuer
|
|||||||
// Skip primary key updates
|
// Skip primary key updates
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
b.query = b.query.Set(column+" = ?", value)
|
b.query = b.query.Set(column+" = ?", common.ConvertSliceForBun(value))
|
||||||
}
|
}
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
@@ -1478,7 +1603,7 @@ func (b *BunUpdateQuery) Where(query string, args ...interface{}) common.UpdateQ
|
|||||||
|
|
||||||
func (b *BunUpdateQuery) Returning(columns ...string) common.UpdateQuery {
|
func (b *BunUpdateQuery) Returning(columns ...string) common.UpdateQuery {
|
||||||
if len(columns) > 0 {
|
if len(columns) > 0 {
|
||||||
b.query = b.query.Returning(columns[0])
|
b.query = b.query.Returning(strings.Join(columns, ", "))
|
||||||
}
|
}
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
@@ -1489,27 +1614,44 @@ func (b *BunUpdateQuery) Exec(ctx context.Context) (res common.Result, err error
|
|||||||
err = logger.HandlePanic("BunUpdateQuery.Exec", r)
|
err = logger.HandlePanic("BunUpdateQuery.Exec", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
result, err := b.query.Exec(ctx)
|
result, err := b.query.Exec(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Log SQL string for debugging
|
// Log SQL string for debugging
|
||||||
sqlStr := b.query.String()
|
sqlStr := b.query.String()
|
||||||
logger.Error("BunUpdateQuery.Exec failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("BunUpdateQuery.Exec failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(b.metricsEnabled, "UPDATE", b.schema, b.entity, b.tableName, startedAt, err)
|
||||||
return &BunResult{result: result}, err
|
return &BunResult{result: result}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// BunDeleteQuery implements DeleteQuery for Bun
|
// BunDeleteQuery implements DeleteQuery for Bun
|
||||||
type BunDeleteQuery struct {
|
type BunDeleteQuery struct {
|
||||||
query *bun.DeleteQuery
|
query *bun.DeleteQuery
|
||||||
|
driverName string
|
||||||
|
schema string
|
||||||
|
tableName string
|
||||||
|
entity string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunDeleteQuery) Model(model interface{}) common.DeleteQuery {
|
func (b *BunDeleteQuery) Model(model interface{}) common.DeleteQuery {
|
||||||
b.query = b.query.Model(model)
|
b.query = b.query.Model(model)
|
||||||
|
b.schema, b.tableName = schemaAndTableFromModel(model, b.driverName)
|
||||||
|
if b.tableName == "" {
|
||||||
|
b.schema, b.tableName = parseTableName(b.query.GetTableName(), b.driverName)
|
||||||
|
}
|
||||||
|
b.entity = entityNameFromModel(model, b.tableName)
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunDeleteQuery) Table(table string) common.DeleteQuery {
|
func (b *BunDeleteQuery) Table(table string) common.DeleteQuery {
|
||||||
b.query = b.query.Table(table)
|
b.query = b.query.Table(table)
|
||||||
|
b.schema, b.tableName = parseTableName(table, b.driverName)
|
||||||
|
if b.entity == "" {
|
||||||
|
b.entity = cleanMetricIdentifier(b.tableName)
|
||||||
|
}
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1524,12 +1666,15 @@ func (b *BunDeleteQuery) Exec(ctx context.Context) (res common.Result, err error
|
|||||||
err = logger.HandlePanic("BunDeleteQuery.Exec", r)
|
err = logger.HandlePanic("BunDeleteQuery.Exec", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
result, err := b.query.Exec(ctx)
|
result, err := b.query.Exec(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Log SQL string for debugging
|
// Log SQL string for debugging
|
||||||
sqlStr := b.query.String()
|
sqlStr := b.query.String()
|
||||||
logger.Error("BunDeleteQuery.Exec failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("BunDeleteQuery.Exec failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(b.metricsEnabled, "DELETE", b.schema, b.entity, b.tableName, startedAt, err)
|
||||||
return &BunResult{result: result}, err
|
return &BunResult{result: result}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1557,6 +1702,7 @@ func (b *BunResult) LastInsertId() (int64, error) {
|
|||||||
type BunTxAdapter struct {
|
type BunTxAdapter struct {
|
||||||
tx bun.Tx
|
tx bun.Tx
|
||||||
driverName string
|
driverName string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunTxAdapter) NewSelect() common.SelectQuery {
|
func (b *BunTxAdapter) NewSelect() common.SelectQuery {
|
||||||
@@ -1564,28 +1710,36 @@ func (b *BunTxAdapter) NewSelect() common.SelectQuery {
|
|||||||
query: b.tx.NewSelect(),
|
query: b.tx.NewSelect(),
|
||||||
db: b.tx,
|
db: b.tx,
|
||||||
driverName: b.driverName,
|
driverName: b.driverName,
|
||||||
|
metricsEnabled: b.metricsEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunTxAdapter) NewInsert() common.InsertQuery {
|
func (b *BunTxAdapter) NewInsert() common.InsertQuery {
|
||||||
return &BunInsertQuery{query: b.tx.NewInsert()}
|
return &BunInsertQuery{query: b.tx.NewInsert(), driverName: b.driverName, metricsEnabled: b.metricsEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunTxAdapter) NewUpdate() common.UpdateQuery {
|
func (b *BunTxAdapter) NewUpdate() common.UpdateQuery {
|
||||||
return &BunUpdateQuery{query: b.tx.NewUpdate()}
|
return &BunUpdateQuery{query: b.tx.NewUpdate(), driverName: b.driverName, metricsEnabled: b.metricsEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunTxAdapter) NewDelete() common.DeleteQuery {
|
func (b *BunTxAdapter) NewDelete() common.DeleteQuery {
|
||||||
return &BunDeleteQuery{query: b.tx.NewDelete()}
|
return &BunDeleteQuery{query: b.tx.NewDelete(), driverName: b.driverName, metricsEnabled: b.metricsEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunTxAdapter) Exec(ctx context.Context, query string, args ...interface{}) (common.Result, error) {
|
func (b *BunTxAdapter) Exec(ctx context.Context, query string, args ...interface{}) (common.Result, error) {
|
||||||
|
startedAt := time.Now()
|
||||||
|
operation, schema, entity, table := metricTargetFromRawQuery(query, b.driverName)
|
||||||
result, err := b.tx.ExecContext(ctx, query, args...)
|
result, err := b.tx.ExecContext(ctx, query, args...)
|
||||||
|
recordQueryMetrics(b.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
return &BunResult{result: result}, err
|
return &BunResult{result: result}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunTxAdapter) Query(ctx context.Context, dest interface{}, query string, args ...interface{}) error {
|
func (b *BunTxAdapter) Query(ctx context.Context, dest interface{}, query string, args ...interface{}) error {
|
||||||
return b.tx.NewRaw(query, args...).Scan(ctx, dest)
|
startedAt := time.Now()
|
||||||
|
operation, schema, entity, table := metricTargetFromRawQuery(query, b.driverName)
|
||||||
|
err := b.tx.NewRaw(query, args...).Scan(ctx, dest)
|
||||||
|
recordQueryMetrics(b.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunTxAdapter) BeginTx(ctx context.Context) (common.Database, error) {
|
func (b *BunTxAdapter) BeginTx(ctx context.Context) (common.Database, error) {
|
||||||
|
|||||||
@@ -0,0 +1,31 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestBunSelectQuery_ColumnExpr_SpreadsArgs is a regression test for a bug
|
||||||
|
// where ColumnExpr passed its variadic args slice as a single argument
|
||||||
|
// (b.query.ColumnExpr(query, args) instead of args...), causing bun to
|
||||||
|
// serialize the arg slice itself (e.g. producing `'["{product,cost}"]'`
|
||||||
|
// instead of `'{product,cost}'` for a JSON path parameter).
|
||||||
|
func TestBunSelectQuery_ColumnExpr_SpreadsArgs(t *testing.T) {
|
||||||
|
db := setupBunTestDB(t)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
adapter := NewBunAdapter(db)
|
||||||
|
|
||||||
|
sq := adapter.NewSelect().
|
||||||
|
Table("test_inserts").
|
||||||
|
ColumnExpr("(jsonvalue #>> ?::text[]) AS jsonvalue_product_cost", "{product,cost}")
|
||||||
|
|
||||||
|
bsq, ok := sq.(*BunSelectQuery)
|
||||||
|
require.True(t, ok, "expected *BunSelectQuery")
|
||||||
|
|
||||||
|
sqlStr := bsq.query.String()
|
||||||
|
require.NotContains(t, sqlStr, `["{product,cost}"]`, "arg slice must not be serialized as a JSON array: %s", sqlStr)
|
||||||
|
require.True(t, strings.Contains(sqlStr, `'{product,cost}'`), "expected the bound text[] literal in SQL: %s", sqlStr)
|
||||||
|
}
|
||||||
@@ -3,10 +3,13 @@ package database
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"reflect"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
"gorm.io/gorm/clause"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
@@ -20,11 +23,12 @@ type GormAdapter struct {
|
|||||||
db *gorm.DB
|
db *gorm.DB
|
||||||
dbFactory func() (*gorm.DB, error)
|
dbFactory func() (*gorm.DB, error)
|
||||||
driverName string
|
driverName string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewGormAdapter creates a new GORM adapter
|
// NewGormAdapter creates a new GORM adapter
|
||||||
func NewGormAdapter(db *gorm.DB) *GormAdapter {
|
func NewGormAdapter(db *gorm.DB) *GormAdapter {
|
||||||
adapter := &GormAdapter{db: db}
|
adapter := &GormAdapter{db: db, metricsEnabled: true}
|
||||||
// Initialize driver name
|
// Initialize driver name
|
||||||
adapter.driverName = adapter.DriverName()
|
adapter.driverName = adapter.DriverName()
|
||||||
return adapter
|
return adapter
|
||||||
@@ -36,6 +40,12 @@ func (g *GormAdapter) WithDBFactory(factory func() (*gorm.DB, error)) *GormAdapt
|
|||||||
return g
|
return g
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetMetricsEnabled enables or disables query metrics for this adapter.
|
||||||
|
func (g *GormAdapter) SetMetricsEnabled(enabled bool) *GormAdapter {
|
||||||
|
g.metricsEnabled = enabled
|
||||||
|
return g
|
||||||
|
}
|
||||||
|
|
||||||
func (g *GormAdapter) getDB() *gorm.DB {
|
func (g *GormAdapter) getDB() *gorm.DB {
|
||||||
g.dbMu.RLock()
|
g.dbMu.RLock()
|
||||||
defer g.dbMu.RUnlock()
|
defer g.dbMu.RUnlock()
|
||||||
@@ -109,19 +119,19 @@ func (g *GormAdapter) DisableQueryDebug() *GormAdapter {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormAdapter) NewSelect() common.SelectQuery {
|
func (g *GormAdapter) NewSelect() common.SelectQuery {
|
||||||
return &GormSelectQuery{db: g.getDB(), driverName: g.driverName, reconnect: g.reconnectDB}
|
return &GormSelectQuery{db: g.getDB(), driverName: g.driverName, reconnect: g.reconnectDB, metricsEnabled: g.metricsEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormAdapter) NewInsert() common.InsertQuery {
|
func (g *GormAdapter) NewInsert() common.InsertQuery {
|
||||||
return &GormInsertQuery{db: g.getDB(), reconnect: g.reconnectDB}
|
return &GormInsertQuery{db: g.getDB(), reconnect: g.reconnectDB, driverName: g.driverName, metricsEnabled: g.metricsEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormAdapter) NewUpdate() common.UpdateQuery {
|
func (g *GormAdapter) NewUpdate() common.UpdateQuery {
|
||||||
return &GormUpdateQuery{db: g.getDB(), reconnect: g.reconnectDB}
|
return &GormUpdateQuery{db: g.getDB(), reconnect: g.reconnectDB, driverName: g.driverName, metricsEnabled: g.metricsEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormAdapter) NewDelete() common.DeleteQuery {
|
func (g *GormAdapter) NewDelete() common.DeleteQuery {
|
||||||
return &GormDeleteQuery{db: g.getDB(), reconnect: g.reconnectDB}
|
return &GormDeleteQuery{db: g.getDB(), reconnect: g.reconnectDB, driverName: g.driverName, metricsEnabled: g.metricsEnabled}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormAdapter) Exec(ctx context.Context, query string, args ...interface{}) (res common.Result, err error) {
|
func (g *GormAdapter) Exec(ctx context.Context, query string, args ...interface{}) (res common.Result, err error) {
|
||||||
@@ -130,6 +140,8 @@ func (g *GormAdapter) Exec(ctx context.Context, query string, args ...interface{
|
|||||||
err = logger.HandlePanic("GormAdapter.Exec", r)
|
err = logger.HandlePanic("GormAdapter.Exec", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
|
operation, schema, entity, table := metricTargetFromRawQuery(query, g.driverName)
|
||||||
run := func() *gorm.DB {
|
run := func() *gorm.DB {
|
||||||
return g.getDB().WithContext(ctx).Exec(query, args...)
|
return g.getDB().WithContext(ctx).Exec(query, args...)
|
||||||
}
|
}
|
||||||
@@ -139,6 +151,7 @@ func (g *GormAdapter) Exec(ctx context.Context, query string, args ...interface{
|
|||||||
result = run()
|
result = run()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(g.metricsEnabled, operation, schema, entity, table, startedAt, result.Error)
|
||||||
return &GormResult{result: result}, result.Error
|
return &GormResult{result: result}, result.Error
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -148,6 +161,8 @@ func (g *GormAdapter) Query(ctx context.Context, dest interface{}, query string,
|
|||||||
err = logger.HandlePanic("GormAdapter.Query", r)
|
err = logger.HandlePanic("GormAdapter.Query", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
|
operation, schema, entity, table := metricTargetFromRawQuery(query, g.driverName)
|
||||||
run := func() error {
|
run := func() error {
|
||||||
return g.getDB().WithContext(ctx).Raw(query, args...).Find(dest).Error
|
return g.getDB().WithContext(ctx).Raw(query, args...).Find(dest).Error
|
||||||
}
|
}
|
||||||
@@ -157,6 +172,7 @@ func (g *GormAdapter) Query(ctx context.Context, dest interface{}, query string,
|
|||||||
err = run()
|
err = run()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(g.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -173,7 +189,7 @@ func (g *GormAdapter) BeginTx(ctx context.Context) (common.Database, error) {
|
|||||||
if tx.Error != nil {
|
if tx.Error != nil {
|
||||||
return nil, tx.Error
|
return nil, tx.Error
|
||||||
}
|
}
|
||||||
return &GormAdapter{db: tx, dbFactory: g.dbFactory, driverName: g.driverName}, nil
|
return &GormAdapter{db: tx, dbFactory: g.dbFactory, driverName: g.driverName, metricsEnabled: g.metricsEnabled}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormAdapter) CommitTx(ctx context.Context) error {
|
func (g *GormAdapter) CommitTx(ctx context.Context) error {
|
||||||
@@ -192,7 +208,7 @@ func (g *GormAdapter) RunInTransaction(ctx context.Context, fn func(common.Datab
|
|||||||
}()
|
}()
|
||||||
run := func() error {
|
run := func() error {
|
||||||
return g.getDB().WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
return g.getDB().WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||||
adapter := &GormAdapter{db: tx, dbFactory: g.dbFactory, driverName: g.driverName}
|
adapter := &GormAdapter{db: tx, dbFactory: g.dbFactory, driverName: g.driverName, metricsEnabled: g.metricsEnabled}
|
||||||
return fn(adapter)
|
return fn(adapter)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -236,22 +252,18 @@ type GormSelectQuery struct {
|
|||||||
reconnect func(...*gorm.DB) error
|
reconnect func(...*gorm.DB) error
|
||||||
schema string // Separated schema name
|
schema string // Separated schema name
|
||||||
tableName string // Just the table name, without schema
|
tableName string // Just the table name, without schema
|
||||||
|
entity string
|
||||||
tableAlias string
|
tableAlias string
|
||||||
driverName string // Database driver name (postgres, sqlite, mssql)
|
driverName string // Database driver name (postgres, sqlite, mssql)
|
||||||
inJoinContext bool // Track if we're in a JOIN relation context
|
inJoinContext bool // Track if we're in a JOIN relation context
|
||||||
joinTableAlias string // Alias to use for JOIN conditions
|
joinTableAlias string // Alias to use for JOIN conditions
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormSelectQuery) Model(model interface{}) common.SelectQuery {
|
func (g *GormSelectQuery) Model(model interface{}) common.SelectQuery {
|
||||||
g.db = g.db.Model(model)
|
g.db = g.db.Model(model)
|
||||||
|
g.schema, g.tableName = schemaAndTableFromModel(model, g.driverName)
|
||||||
// Try to get table name from model if it implements TableNameProvider
|
g.entity = entityNameFromModel(model, g.tableName)
|
||||||
if provider, ok := model.(common.TableNameProvider); ok {
|
|
||||||
fullTableName := provider.TableName()
|
|
||||||
// Check if the table name contains schema (e.g., "schema.table")
|
|
||||||
// For SQLite, this will convert "schema.table" to "schema_table"
|
|
||||||
g.schema, g.tableName = parseTableName(fullTableName, g.driverName)
|
|
||||||
}
|
|
||||||
|
|
||||||
if provider, ok := model.(common.TableAliasProvider); ok {
|
if provider, ok := model.(common.TableAliasProvider); ok {
|
||||||
g.tableAlias = provider.TableAlias()
|
g.tableAlias = provider.TableAlias()
|
||||||
@@ -265,6 +277,9 @@ func (g *GormSelectQuery) Table(table string) common.SelectQuery {
|
|||||||
// Check if the table name contains schema (e.g., "schema.table")
|
// Check if the table name contains schema (e.g., "schema.table")
|
||||||
// For SQLite, this will convert "schema.table" to "schema_table"
|
// For SQLite, this will convert "schema.table" to "schema_table"
|
||||||
g.schema, g.tableName = parseTableName(table, g.driverName)
|
g.schema, g.tableName = parseTableName(table, g.driverName)
|
||||||
|
if g.entity == "" {
|
||||||
|
g.entity = cleanMetricIdentifier(g.tableName)
|
||||||
|
}
|
||||||
|
|
||||||
return g
|
return g
|
||||||
}
|
}
|
||||||
@@ -453,6 +468,7 @@ func (g *GormSelectQuery) PreloadRelation(relation string, apply ...func(common.
|
|||||||
db: db,
|
db: db,
|
||||||
reconnect: g.reconnect,
|
reconnect: g.reconnect,
|
||||||
driverName: g.driverName,
|
driverName: g.driverName,
|
||||||
|
metricsEnabled: g.metricsEnabled,
|
||||||
}
|
}
|
||||||
|
|
||||||
current := common.SelectQuery(wrapper)
|
current := common.SelectQuery(wrapper)
|
||||||
@@ -494,6 +510,7 @@ func (g *GormSelectQuery) JoinRelation(relation string, apply ...func(common.Sel
|
|||||||
driverName: g.driverName,
|
driverName: g.driverName,
|
||||||
inJoinContext: true, // Mark as JOIN context
|
inJoinContext: true, // Mark as JOIN context
|
||||||
joinTableAlias: strings.ToLower(relation), // Use relation name as alias
|
joinTableAlias: strings.ToLower(relation), // Use relation name as alias
|
||||||
|
metricsEnabled: g.metricsEnabled,
|
||||||
}
|
}
|
||||||
current := common.SelectQuery(wrapper)
|
current := common.SelectQuery(wrapper)
|
||||||
|
|
||||||
@@ -550,6 +567,7 @@ func (g *GormSelectQuery) Scan(ctx context.Context, dest interface{}) (err error
|
|||||||
err = logger.HandlePanic("GormSelectQuery.Scan", r)
|
err = logger.HandlePanic("GormSelectQuery.Scan", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
run := func() error {
|
run := func() error {
|
||||||
return g.db.WithContext(ctx).Find(dest).Error
|
return g.db.WithContext(ctx).Find(dest).Error
|
||||||
}
|
}
|
||||||
@@ -565,7 +583,9 @@ func (g *GormSelectQuery) Scan(ctx context.Context, dest interface{}) (err error
|
|||||||
return tx.Find(dest)
|
return tx.Find(dest)
|
||||||
})
|
})
|
||||||
logger.Error("GormSelectQuery.Scan failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("GormSelectQuery.Scan failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(g.metricsEnabled, "SELECT", g.schema, g.entity, g.tableName, startedAt, err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -578,6 +598,7 @@ func (g *GormSelectQuery) ScanModel(ctx context.Context) (err error) {
|
|||||||
if g.db.Statement.Model == nil {
|
if g.db.Statement.Model == nil {
|
||||||
return fmt.Errorf("ScanModel requires Model() to be set before scanning")
|
return fmt.Errorf("ScanModel requires Model() to be set before scanning")
|
||||||
}
|
}
|
||||||
|
startedAt := time.Now()
|
||||||
run := func() error {
|
run := func() error {
|
||||||
return g.db.WithContext(ctx).Find(g.db.Statement.Model).Error
|
return g.db.WithContext(ctx).Find(g.db.Statement.Model).Error
|
||||||
}
|
}
|
||||||
@@ -593,7 +614,9 @@ func (g *GormSelectQuery) ScanModel(ctx context.Context) (err error) {
|
|||||||
return tx.Find(g.db.Statement.Model)
|
return tx.Find(g.db.Statement.Model)
|
||||||
})
|
})
|
||||||
logger.Error("GormSelectQuery.ScanModel failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("GormSelectQuery.ScanModel failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(g.metricsEnabled, "SELECT", g.schema, g.entity, g.tableName, startedAt, err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -604,6 +627,7 @@ func (g *GormSelectQuery) Count(ctx context.Context) (count int, err error) {
|
|||||||
count = 0
|
count = 0
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
var count64 int64
|
var count64 int64
|
||||||
run := func() error {
|
run := func() error {
|
||||||
return g.db.WithContext(ctx).Count(&count64).Error
|
return g.db.WithContext(ctx).Count(&count64).Error
|
||||||
@@ -620,7 +644,9 @@ func (g *GormSelectQuery) Count(ctx context.Context) (count int, err error) {
|
|||||||
return tx.Count(&count64)
|
return tx.Count(&count64)
|
||||||
})
|
})
|
||||||
logger.Error("GormSelectQuery.Count failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("GormSelectQuery.Count failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(g.metricsEnabled, "COUNT", g.schema, g.entity, g.tableName, startedAt, err)
|
||||||
return int(count64), err
|
return int(count64), err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -631,6 +657,7 @@ func (g *GormSelectQuery) Exists(ctx context.Context) (exists bool, err error) {
|
|||||||
exists = false
|
exists = false
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
var count int64
|
var count int64
|
||||||
run := func() error {
|
run := func() error {
|
||||||
return g.db.WithContext(ctx).Limit(1).Count(&count).Error
|
return g.db.WithContext(ctx).Limit(1).Count(&count).Error
|
||||||
@@ -647,7 +674,9 @@ func (g *GormSelectQuery) Exists(ctx context.Context) (exists bool, err error) {
|
|||||||
return tx.Limit(1).Count(&count)
|
return tx.Limit(1).Count(&count)
|
||||||
})
|
})
|
||||||
logger.Error("GormSelectQuery.Exists failed. SQL: %s. Error: %v", sqlStr, err)
|
logger.Error("GormSelectQuery.Exists failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(g.metricsEnabled, "EXISTS", g.schema, g.entity, g.tableName, startedAt, err)
|
||||||
return count > 0, err
|
return count > 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -657,16 +686,28 @@ type GormInsertQuery struct {
|
|||||||
reconnect func(...*gorm.DB) error
|
reconnect func(...*gorm.DB) error
|
||||||
model interface{}
|
model interface{}
|
||||||
values map[string]interface{}
|
values map[string]interface{}
|
||||||
|
schema string
|
||||||
|
tableName string
|
||||||
|
entity string
|
||||||
|
driverName string
|
||||||
|
metricsEnabled bool
|
||||||
|
returningColumns []string
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormInsertQuery) Model(model interface{}) common.InsertQuery {
|
func (g *GormInsertQuery) Model(model interface{}) common.InsertQuery {
|
||||||
g.model = model
|
g.model = model
|
||||||
g.db = g.db.Model(model)
|
g.db = g.db.Model(model)
|
||||||
|
g.schema, g.tableName = schemaAndTableFromModel(model, g.driverName)
|
||||||
|
g.entity = entityNameFromModel(model, g.tableName)
|
||||||
return g
|
return g
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormInsertQuery) Table(table string) common.InsertQuery {
|
func (g *GormInsertQuery) Table(table string) common.InsertQuery {
|
||||||
g.db = g.db.Table(table)
|
g.db = g.db.Table(table)
|
||||||
|
g.schema, g.tableName = parseTableName(table, g.driverName)
|
||||||
|
if g.entity == "" {
|
||||||
|
g.entity = cleanMetricIdentifier(g.tableName)
|
||||||
|
}
|
||||||
return g
|
return g
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -684,7 +725,7 @@ func (g *GormInsertQuery) OnConflict(action string) common.InsertQuery {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormInsertQuery) Returning(columns ...string) common.InsertQuery {
|
func (g *GormInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||||
// GORM doesn't have explicit RETURNING, but updates the model
|
g.returningColumns = columns
|
||||||
return g
|
return g
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -694,6 +735,7 @@ func (g *GormInsertQuery) Exec(ctx context.Context) (res common.Result, err erro
|
|||||||
err = logger.HandlePanic("GormInsertQuery.Exec", r)
|
err = logger.HandlePanic("GormInsertQuery.Exec", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
run := func() *gorm.DB {
|
run := func() *gorm.DB {
|
||||||
switch {
|
switch {
|
||||||
case g.model != nil:
|
case g.model != nil:
|
||||||
@@ -710,30 +752,113 @@ func (g *GormInsertQuery) Exec(ctx context.Context) (res common.Result, err erro
|
|||||||
result = run()
|
result = run()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(g.metricsEnabled, "INSERT", g.schema, g.entity, g.tableName, startedAt, result.Error)
|
||||||
return &GormResult{result: result}, result.Error
|
return &GormResult{result: result}, result.Error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (g *GormInsertQuery) Scan(ctx context.Context, dest interface{}) (err error) {
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
err = logger.HandlePanic("GormInsertQuery.Scan", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
|
|
||||||
|
var returningCols []clause.Column
|
||||||
|
for _, col := range g.returningColumns {
|
||||||
|
returningCols = append(returningCols, clause.Column{Name: col})
|
||||||
|
}
|
||||||
|
|
||||||
|
db := g.db.WithContext(ctx)
|
||||||
|
if len(returningCols) > 0 {
|
||||||
|
db = db.Clauses(clause.Returning{Columns: returningCols})
|
||||||
|
}
|
||||||
|
|
||||||
|
var result *gorm.DB
|
||||||
|
switch {
|
||||||
|
case g.model != nil:
|
||||||
|
result = db.Create(g.model)
|
||||||
|
case g.values != nil:
|
||||||
|
result = db.Create(g.values)
|
||||||
|
default:
|
||||||
|
result = db.Create(map[string]interface{}{})
|
||||||
|
}
|
||||||
|
|
||||||
|
if isDBClosed(result.Error) && g.reconnect != nil {
|
||||||
|
if reconnErr := g.reconnect(g.db); reconnErr == nil {
|
||||||
|
result = db.Create(g.model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
recordQueryMetrics(g.metricsEnabled, "INSERT", g.schema, g.entity, g.tableName, startedAt, result.Error)
|
||||||
|
if result.Error != nil {
|
||||||
|
return result.Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract the returning column value from the model or values map
|
||||||
|
if len(g.returningColumns) == 1 {
|
||||||
|
col := g.returningColumns[0]
|
||||||
|
if g.model != nil {
|
||||||
|
val := reflect.ValueOf(g.model)
|
||||||
|
if val.Kind() == reflect.Pointer {
|
||||||
|
val = val.Elem()
|
||||||
|
}
|
||||||
|
if val.Kind() == reflect.Struct {
|
||||||
|
for i := 0; i < val.NumField(); i++ {
|
||||||
|
f := val.Type().Field(i)
|
||||||
|
dbTag := strings.Split(f.Tag.Get("bun"), ",")[0]
|
||||||
|
jsonTag := strings.Split(f.Tag.Get("json"), ",")[0]
|
||||||
|
if strings.EqualFold(f.Name, col) || dbTag == col || jsonTag == col {
|
||||||
|
reflect.ValueOf(dest).Elem().Set(val.Field(i))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if g.values != nil {
|
||||||
|
if v, ok := g.values[col]; ok {
|
||||||
|
reflect.ValueOf(dest).Elem().Set(reflect.ValueOf(v))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// GormUpdateQuery implements UpdateQuery for GORM
|
// GormUpdateQuery implements UpdateQuery for GORM
|
||||||
type GormUpdateQuery struct {
|
type GormUpdateQuery struct {
|
||||||
db *gorm.DB
|
db *gorm.DB
|
||||||
reconnect func(...*gorm.DB) error
|
reconnect func(...*gorm.DB) error
|
||||||
model interface{}
|
model interface{}
|
||||||
updates interface{}
|
updates interface{}
|
||||||
|
schema string
|
||||||
|
tableName string
|
||||||
|
entity string
|
||||||
|
driverName string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormUpdateQuery) Model(model interface{}) common.UpdateQuery {
|
func (g *GormUpdateQuery) Model(model interface{}) common.UpdateQuery {
|
||||||
g.model = model
|
g.model = model
|
||||||
g.db = g.db.Model(model)
|
g.db = g.db.Model(model)
|
||||||
|
g.schema, g.tableName = schemaAndTableFromModel(model, g.driverName)
|
||||||
|
g.entity = entityNameFromModel(model, g.tableName)
|
||||||
return g
|
return g
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormUpdateQuery) Table(table string) common.UpdateQuery {
|
func (g *GormUpdateQuery) Table(table string) common.UpdateQuery {
|
||||||
g.db = g.db.Table(table)
|
g.db = g.db.Table(table)
|
||||||
|
g.schema, g.tableName = parseTableName(table, g.driverName)
|
||||||
|
if g.entity == "" {
|
||||||
|
g.entity = cleanMetricIdentifier(g.tableName)
|
||||||
|
}
|
||||||
if g.model == nil {
|
if g.model == nil {
|
||||||
// Try to get table name from table string if model is not set
|
// Try to get table name from table string if model is not set
|
||||||
model, err := modelregistry.GetModelByName(table)
|
model, err := modelregistry.GetModelByName(table)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
g.model = model
|
g.model = model
|
||||||
|
g.entity = entityNameFromModel(model, g.tableName)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return g
|
return g
|
||||||
@@ -794,6 +919,7 @@ func (g *GormUpdateQuery) Exec(ctx context.Context) (res common.Result, err erro
|
|||||||
err = logger.HandlePanic("GormUpdateQuery.Exec", r)
|
err = logger.HandlePanic("GormUpdateQuery.Exec", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
run := func() *gorm.DB {
|
run := func() *gorm.DB {
|
||||||
return g.db.WithContext(ctx).Updates(g.updates)
|
return g.db.WithContext(ctx).Updates(g.updates)
|
||||||
}
|
}
|
||||||
@@ -809,7 +935,9 @@ func (g *GormUpdateQuery) Exec(ctx context.Context) (res common.Result, err erro
|
|||||||
return tx.Updates(g.updates)
|
return tx.Updates(g.updates)
|
||||||
})
|
})
|
||||||
logger.Error("GormUpdateQuery.Exec failed. SQL: %s. Error: %v", sqlStr, result.Error)
|
logger.Error("GormUpdateQuery.Exec failed. SQL: %s. Error: %v", sqlStr, result.Error)
|
||||||
|
return &GormResult{result: result}, common.WrapSQLError(result.Error, sqlStr)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(g.metricsEnabled, "UPDATE", g.schema, g.entity, g.tableName, startedAt, result.Error)
|
||||||
return &GormResult{result: result}, result.Error
|
return &GormResult{result: result}, result.Error
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -818,16 +946,27 @@ type GormDeleteQuery struct {
|
|||||||
db *gorm.DB
|
db *gorm.DB
|
||||||
reconnect func(...*gorm.DB) error
|
reconnect func(...*gorm.DB) error
|
||||||
model interface{}
|
model interface{}
|
||||||
|
schema string
|
||||||
|
tableName string
|
||||||
|
entity string
|
||||||
|
driverName string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormDeleteQuery) Model(model interface{}) common.DeleteQuery {
|
func (g *GormDeleteQuery) Model(model interface{}) common.DeleteQuery {
|
||||||
g.model = model
|
g.model = model
|
||||||
g.db = g.db.Model(model)
|
g.db = g.db.Model(model)
|
||||||
|
g.schema, g.tableName = schemaAndTableFromModel(model, g.driverName)
|
||||||
|
g.entity = entityNameFromModel(model, g.tableName)
|
||||||
return g
|
return g
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GormDeleteQuery) Table(table string) common.DeleteQuery {
|
func (g *GormDeleteQuery) Table(table string) common.DeleteQuery {
|
||||||
g.db = g.db.Table(table)
|
g.db = g.db.Table(table)
|
||||||
|
g.schema, g.tableName = parseTableName(table, g.driverName)
|
||||||
|
if g.entity == "" {
|
||||||
|
g.entity = cleanMetricIdentifier(g.tableName)
|
||||||
|
}
|
||||||
return g
|
return g
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -842,6 +981,7 @@ func (g *GormDeleteQuery) Exec(ctx context.Context) (res common.Result, err erro
|
|||||||
err = logger.HandlePanic("GormDeleteQuery.Exec", r)
|
err = logger.HandlePanic("GormDeleteQuery.Exec", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
run := func() *gorm.DB {
|
run := func() *gorm.DB {
|
||||||
return g.db.WithContext(ctx).Delete(g.model)
|
return g.db.WithContext(ctx).Delete(g.model)
|
||||||
}
|
}
|
||||||
@@ -857,7 +997,9 @@ func (g *GormDeleteQuery) Exec(ctx context.Context) (res common.Result, err erro
|
|||||||
return tx.Delete(g.model)
|
return tx.Delete(g.model)
|
||||||
})
|
})
|
||||||
logger.Error("GormDeleteQuery.Exec failed. SQL: %s. Error: %v", sqlStr, result.Error)
|
logger.Error("GormDeleteQuery.Exec failed. SQL: %s. Error: %v", sqlStr, result.Error)
|
||||||
|
return &GormResult{result: result}, common.WrapSQLError(result.Error, sqlStr)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(g.metricsEnabled, "DELETE", g.schema, g.entity, g.tableName, startedAt, result.Error)
|
||||||
return &GormResult{result: result}, result.Error
|
return &GormResult{result: result}, result.Error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -5,8 +5,10 @@ import (
|
|||||||
"database/sql"
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
@@ -21,6 +23,7 @@ type PgSQLAdapter struct {
|
|||||||
dbMu sync.RWMutex
|
dbMu sync.RWMutex
|
||||||
dbFactory func() (*sql.DB, error)
|
dbFactory func() (*sql.DB, error)
|
||||||
driverName string
|
driverName string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewPgSQLAdapter creates a new adapter wrapping a standard sql.DB.
|
// NewPgSQLAdapter creates a new adapter wrapping a standard sql.DB.
|
||||||
@@ -31,7 +34,7 @@ func NewPgSQLAdapter(db *sql.DB, driverName ...string) *PgSQLAdapter {
|
|||||||
if len(driverName) > 0 && driverName[0] != "" {
|
if len(driverName) > 0 && driverName[0] != "" {
|
||||||
name = driverName[0]
|
name = driverName[0]
|
||||||
}
|
}
|
||||||
return &PgSQLAdapter{db: db, driverName: name}
|
return &PgSQLAdapter{db: db, driverName: name, metricsEnabled: true}
|
||||||
}
|
}
|
||||||
|
|
||||||
// WithDBFactory configures a factory used to reopen the database connection if it is closed.
|
// WithDBFactory configures a factory used to reopen the database connection if it is closed.
|
||||||
@@ -40,6 +43,12 @@ func (p *PgSQLAdapter) WithDBFactory(factory func() (*sql.DB, error)) *PgSQLAdap
|
|||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetMetricsEnabled enables or disables query metrics for this adapter.
|
||||||
|
func (p *PgSQLAdapter) SetMetricsEnabled(enabled bool) *PgSQLAdapter {
|
||||||
|
p.metricsEnabled = enabled
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
func (p *PgSQLAdapter) getDB() *sql.DB {
|
func (p *PgSQLAdapter) getDB() *sql.DB {
|
||||||
p.dbMu.RLock()
|
p.dbMu.RLock()
|
||||||
defer p.dbMu.RUnlock()
|
defer p.dbMu.RUnlock()
|
||||||
@@ -75,6 +84,7 @@ func (p *PgSQLAdapter) NewSelect() common.SelectQuery {
|
|||||||
driverName: p.driverName,
|
driverName: p.driverName,
|
||||||
columns: []string{"*"},
|
columns: []string{"*"},
|
||||||
args: make([]interface{}, 0),
|
args: make([]interface{}, 0),
|
||||||
|
metricsEnabled: p.metricsEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -83,6 +93,7 @@ func (p *PgSQLAdapter) NewInsert() common.InsertQuery {
|
|||||||
db: p.getDB(),
|
db: p.getDB(),
|
||||||
driverName: p.driverName,
|
driverName: p.driverName,
|
||||||
values: make(map[string]interface{}),
|
values: make(map[string]interface{}),
|
||||||
|
metricsEnabled: p.metricsEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -93,6 +104,7 @@ func (p *PgSQLAdapter) NewUpdate() common.UpdateQuery {
|
|||||||
sets: make(map[string]interface{}),
|
sets: make(map[string]interface{}),
|
||||||
args: make([]interface{}, 0),
|
args: make([]interface{}, 0),
|
||||||
whereClauses: make([]string, 0),
|
whereClauses: make([]string, 0),
|
||||||
|
metricsEnabled: p.metricsEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -102,6 +114,7 @@ func (p *PgSQLAdapter) NewDelete() common.DeleteQuery {
|
|||||||
driverName: p.driverName,
|
driverName: p.driverName,
|
||||||
args: make([]interface{}, 0),
|
args: make([]interface{}, 0),
|
||||||
whereClauses: make([]string, 0),
|
whereClauses: make([]string, 0),
|
||||||
|
metricsEnabled: p.metricsEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -111,6 +124,8 @@ func (p *PgSQLAdapter) Exec(ctx context.Context, query string, args ...interface
|
|||||||
err = logger.HandlePanic("PgSQLAdapter.Exec", r)
|
err = logger.HandlePanic("PgSQLAdapter.Exec", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
|
operation, schema, entity, table := metricTargetFromRawQuery(query, p.driverName)
|
||||||
logger.Debug("PgSQL Exec: %s [args: %v]", query, args)
|
logger.Debug("PgSQL Exec: %s [args: %v]", query, args)
|
||||||
var result sql.Result
|
var result sql.Result
|
||||||
run := func() error { var e error; result, e = p.getDB().ExecContext(ctx, query, args...); return e }
|
run := func() error { var e error; result, e = p.getDB().ExecContext(ctx, query, args...); return e }
|
||||||
@@ -122,8 +137,10 @@ func (p *PgSQLAdapter) Exec(ctx context.Context, query string, args ...interface
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL Exec failed: %v", err)
|
logger.Error("PgSQL Exec failed: %v", err)
|
||||||
return nil, err
|
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
|
return nil, common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, nil)
|
||||||
return &PgSQLResult{result: result}, nil
|
return &PgSQLResult{result: result}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -133,6 +150,8 @@ func (p *PgSQLAdapter) Query(ctx context.Context, dest interface{}, query string
|
|||||||
err = logger.HandlePanic("PgSQLAdapter.Query", r)
|
err = logger.HandlePanic("PgSQLAdapter.Query", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
|
operation, schema, entity, table := metricTargetFromRawQuery(query, p.driverName)
|
||||||
logger.Debug("PgSQL Query: %s [args: %v]", query, args)
|
logger.Debug("PgSQL Query: %s [args: %v]", query, args)
|
||||||
var rows *sql.Rows
|
var rows *sql.Rows
|
||||||
run := func() error { var e error; rows, e = p.getDB().QueryContext(ctx, query, args...); return e }
|
run := func() error { var e error; rows, e = p.getDB().QueryContext(ctx, query, args...); return e }
|
||||||
@@ -144,11 +163,14 @@ func (p *PgSQLAdapter) Query(ctx context.Context, dest interface{}, query string
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL Query failed: %v", err)
|
logger.Error("PgSQL Query failed: %v", err)
|
||||||
return err
|
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
|
return common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
defer rows.Close()
|
defer rows.Close()
|
||||||
|
|
||||||
return scanRows(rows, dest)
|
err = scanRows(rows, dest)
|
||||||
|
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLAdapter) BeginTx(ctx context.Context) (common.Database, error) {
|
func (p *PgSQLAdapter) BeginTx(ctx context.Context) (common.Database, error) {
|
||||||
@@ -156,7 +178,7 @@ func (p *PgSQLAdapter) BeginTx(ctx context.Context) (common.Database, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return &PgSQLTxAdapter{tx: tx, driverName: p.driverName}, nil
|
return &PgSQLTxAdapter{tx: tx, driverName: p.driverName, metricsEnabled: p.metricsEnabled}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLAdapter) CommitTx(ctx context.Context) error {
|
func (p *PgSQLAdapter) CommitTx(ctx context.Context) error {
|
||||||
@@ -179,7 +201,7 @@ func (p *PgSQLAdapter) RunInTransaction(ctx context.Context, fn func(common.Data
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
adapter := &PgSQLTxAdapter{tx: tx, driverName: p.driverName}
|
adapter := &PgSQLTxAdapter{tx: tx, driverName: p.driverName, metricsEnabled: p.metricsEnabled}
|
||||||
|
|
||||||
defer func() {
|
defer func() {
|
||||||
if p := recover(); p != nil {
|
if p := recover(); p != nil {
|
||||||
@@ -225,7 +247,9 @@ type PgSQLSelectQuery struct {
|
|||||||
db *sql.DB
|
db *sql.DB
|
||||||
tx *sql.Tx
|
tx *sql.Tx
|
||||||
model interface{}
|
model interface{}
|
||||||
|
entity string
|
||||||
tableName string
|
tableName string
|
||||||
|
schema string
|
||||||
tableAlias string
|
tableAlias string
|
||||||
driverName string // Database driver name (postgres, sqlite, mssql)
|
driverName string // Database driver name (postgres, sqlite, mssql)
|
||||||
columns []string
|
columns []string
|
||||||
@@ -241,15 +265,13 @@ type PgSQLSelectQuery struct {
|
|||||||
args []interface{}
|
args []interface{}
|
||||||
paramCounter int
|
paramCounter int
|
||||||
preloads []preloadConfig
|
preloads []preloadConfig
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLSelectQuery) Model(model interface{}) common.SelectQuery {
|
func (p *PgSQLSelectQuery) Model(model interface{}) common.SelectQuery {
|
||||||
p.model = model
|
p.model = model
|
||||||
if provider, ok := model.(common.TableNameProvider); ok {
|
p.schema, p.tableName = schemaAndTableFromModel(model, p.driverName)
|
||||||
fullTableName := provider.TableName()
|
p.entity = entityNameFromModel(model, p.tableName)
|
||||||
// For SQLite, convert "schema.table" to "schema_table"
|
|
||||||
_, p.tableName = parseTableName(fullTableName, p.driverName)
|
|
||||||
}
|
|
||||||
if provider, ok := model.(common.TableAliasProvider); ok {
|
if provider, ok := model.(common.TableAliasProvider); ok {
|
||||||
p.tableAlias = provider.TableAlias()
|
p.tableAlias = provider.TableAlias()
|
||||||
}
|
}
|
||||||
@@ -258,7 +280,10 @@ func (p *PgSQLSelectQuery) Model(model interface{}) common.SelectQuery {
|
|||||||
|
|
||||||
func (p *PgSQLSelectQuery) Table(table string) common.SelectQuery {
|
func (p *PgSQLSelectQuery) Table(table string) common.SelectQuery {
|
||||||
// For SQLite, convert "schema.table" to "schema_table"
|
// For SQLite, convert "schema.table" to "schema_table"
|
||||||
_, p.tableName = parseTableName(table, p.driverName)
|
p.schema, p.tableName = parseTableName(table, p.driverName)
|
||||||
|
if p.entity == "" {
|
||||||
|
p.entity = cleanMetricIdentifier(p.tableName)
|
||||||
|
}
|
||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -468,6 +493,7 @@ func (p *PgSQLSelectQuery) Scan(ctx context.Context, dest interface{}) (err erro
|
|||||||
err = logger.HandlePanic("PgSQLSelectQuery.Scan", r)
|
err = logger.HandlePanic("PgSQLSelectQuery.Scan", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
|
|
||||||
// Apply preloads that use JOINs
|
// Apply preloads that use JOINs
|
||||||
p.applyJoinPreloads()
|
p.applyJoinPreloads()
|
||||||
@@ -484,17 +510,21 @@ func (p *PgSQLSelectQuery) Scan(ctx context.Context, dest interface{}) (err erro
|
|||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL SELECT failed: %v", err)
|
logger.Error("PgSQL SELECT failed: %v", err)
|
||||||
return err
|
recordQueryMetrics(p.metricsEnabled, "SELECT", p.schema, p.entity, p.tableName, startedAt, err)
|
||||||
|
return common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
defer rows.Close()
|
defer rows.Close()
|
||||||
|
|
||||||
err = scanRows(rows, dest)
|
err = scanRows(rows, dest)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
recordQueryMetrics(p.metricsEnabled, "SELECT", p.schema, p.entity, p.tableName, startedAt, err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Apply preloads that use separate queries
|
// Apply preloads that use separate queries
|
||||||
return p.applySubqueryPreloads(ctx, dest)
|
err = p.applySubqueryPreloads(ctx, dest)
|
||||||
|
recordQueryMetrics(p.metricsEnabled, "SELECT", p.schema, p.entity, p.tableName, startedAt, err)
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLSelectQuery) ScanModel(ctx context.Context) error {
|
func (p *PgSQLSelectQuery) ScanModel(ctx context.Context) error {
|
||||||
@@ -504,15 +534,8 @@ func (p *PgSQLSelectQuery) ScanModel(ctx context.Context) error {
|
|||||||
return p.Scan(ctx, p.model)
|
return p.Scan(ctx, p.model)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLSelectQuery) Count(ctx context.Context) (count int, err error) {
|
// countInternal executes the COUNT query and returns the result and the SQL string without recording metrics.
|
||||||
defer func() {
|
func (p *PgSQLSelectQuery) countInternal(ctx context.Context) (rowCount int, querySQL string, retErr error) {
|
||||||
if r := recover(); r != nil {
|
|
||||||
err = logger.HandlePanic("PgSQLSelectQuery.Count", r)
|
|
||||||
count = 0
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
// Build a COUNT query
|
|
||||||
var sb strings.Builder
|
var sb strings.Builder
|
||||||
sb.WriteString("SELECT COUNT(*) FROM ")
|
sb.WriteString("SELECT COUNT(*) FROM ")
|
||||||
sb.WriteString(p.tableName)
|
sb.WriteString(p.tableName)
|
||||||
@@ -546,10 +569,28 @@ func (p *PgSQLSelectQuery) Count(ctx context.Context) (count int, err error) {
|
|||||||
row = p.db.QueryRowContext(ctx, query, p.args...)
|
row = p.db.QueryRowContext(ctx, query, p.args...)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = row.Scan(&count)
|
var count int
|
||||||
|
if err := row.Scan(&count); err != nil {
|
||||||
|
return 0, query, err
|
||||||
|
}
|
||||||
|
return count, query, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *PgSQLSelectQuery) Count(ctx context.Context) (count int, err error) {
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
err = logger.HandlePanic("PgSQLSelectQuery.Count", r)
|
||||||
|
count = 0
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
|
var sqlStr string
|
||||||
|
count, sqlStr, err = p.countInternal(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL COUNT failed: %v", err)
|
logger.Error("PgSQL COUNT failed: %v", err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(p.metricsEnabled, "COUNT", p.schema, p.entity, p.tableName, startedAt, err)
|
||||||
return count, err
|
return count, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -560,8 +601,14 @@ func (p *PgSQLSelectQuery) Exists(ctx context.Context) (exists bool, err error)
|
|||||||
exists = false
|
exists = false
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
startedAt := time.Now()
|
||||||
count, err := p.Count(ctx)
|
var sqlStr string
|
||||||
|
count, sqlStr, err := p.countInternal(ctx)
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("PgSQL EXISTS failed: %v", err)
|
||||||
|
err = common.WrapSQLError(err, sqlStr)
|
||||||
|
}
|
||||||
|
recordQueryMetrics(p.metricsEnabled, "EXISTS", p.schema, p.entity, p.tableName, startedAt, err)
|
||||||
return count > 0, err
|
return count > 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -569,18 +616,19 @@ func (p *PgSQLSelectQuery) Exists(ctx context.Context) (exists bool, err error)
|
|||||||
type PgSQLInsertQuery struct {
|
type PgSQLInsertQuery struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
tx *sql.Tx
|
tx *sql.Tx
|
||||||
|
schema string
|
||||||
tableName string
|
tableName string
|
||||||
|
entity string
|
||||||
driverName string
|
driverName string
|
||||||
values map[string]interface{}
|
values map[string]interface{}
|
||||||
|
valueOrder []string
|
||||||
returning []string
|
returning []string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLInsertQuery) Model(model interface{}) common.InsertQuery {
|
func (p *PgSQLInsertQuery) Model(model interface{}) common.InsertQuery {
|
||||||
if provider, ok := model.(common.TableNameProvider); ok {
|
p.schema, p.tableName = schemaAndTableFromModel(model, p.driverName)
|
||||||
fullTableName := provider.TableName()
|
p.entity = entityNameFromModel(model, p.tableName)
|
||||||
// For SQLite, convert "schema.table" to "schema_table"
|
|
||||||
_, p.tableName = parseTableName(fullTableName, p.driverName)
|
|
||||||
}
|
|
||||||
// Extract values from model using reflection
|
// Extract values from model using reflection
|
||||||
// This is a simplified implementation
|
// This is a simplified implementation
|
||||||
return p
|
return p
|
||||||
@@ -588,11 +636,17 @@ func (p *PgSQLInsertQuery) Model(model interface{}) common.InsertQuery {
|
|||||||
|
|
||||||
func (p *PgSQLInsertQuery) Table(table string) common.InsertQuery {
|
func (p *PgSQLInsertQuery) Table(table string) common.InsertQuery {
|
||||||
// For SQLite, convert "schema.table" to "schema_table"
|
// For SQLite, convert "schema.table" to "schema_table"
|
||||||
_, p.tableName = parseTableName(table, p.driverName)
|
p.schema, p.tableName = parseTableName(table, p.driverName)
|
||||||
|
if p.entity == "" {
|
||||||
|
p.entity = cleanMetricIdentifier(p.tableName)
|
||||||
|
}
|
||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLInsertQuery) Value(column string, value interface{}) common.InsertQuery {
|
func (p *PgSQLInsertQuery) Value(column string, value interface{}) common.InsertQuery {
|
||||||
|
if _, exists := p.values[column]; !exists {
|
||||||
|
p.valueOrder = append(p.valueOrder, column)
|
||||||
|
}
|
||||||
p.values[column] = value
|
p.values[column] = value
|
||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
@@ -608,29 +662,31 @@ func (p *PgSQLInsertQuery) Returning(columns ...string) common.InsertQuery {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLInsertQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
func (p *PgSQLInsertQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
||||||
|
startedAt := time.Now()
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("PgSQLInsertQuery.Exec", r)
|
err = logger.HandlePanic("PgSQLInsertQuery.Exec", r)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(p.metricsEnabled, "INSERT", p.schema, p.entity, p.tableName, startedAt, err)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
if len(p.values) == 0 {
|
if len(p.values) == 0 {
|
||||||
return nil, fmt.Errorf("no values to insert")
|
err = fmt.Errorf("no values to insert")
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
columns := make([]string, 0, len(p.values))
|
columns := make([]string, 0, len(p.values))
|
||||||
placeholders := make([]string, 0, len(p.values))
|
placeholders := make([]string, 0, len(p.values))
|
||||||
args := make([]interface{}, 0, len(p.values))
|
args := make([]interface{}, 0, len(p.values))
|
||||||
|
|
||||||
i := 1
|
i := 1
|
||||||
for col, val := range p.values {
|
for _, col := range p.valueOrder {
|
||||||
columns = append(columns, col)
|
columns = append(columns, col)
|
||||||
placeholders = append(placeholders, fmt.Sprintf("$%d", i))
|
placeholders = append(placeholders, fmt.Sprintf("$%d", i))
|
||||||
args = append(args, val)
|
args = append(args, p.values[col])
|
||||||
i++
|
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,
|
p.tableName,
|
||||||
strings.Join(columns, ", "),
|
strings.Join(columns, ", "),
|
||||||
strings.Join(placeholders, ", "))
|
strings.Join(placeholders, ", "))
|
||||||
@@ -650,43 +706,96 @@ func (p *PgSQLInsertQuery) Exec(ctx context.Context) (res common.Result, err err
|
|||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL INSERT failed: %v", err)
|
logger.Error("PgSQL INSERT failed: %v", err)
|
||||||
return nil, err
|
return nil, common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
|
|
||||||
return &PgSQLResult{result: result}, nil
|
return &PgSQLResult{result: result}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (p *PgSQLInsertQuery) Scan(ctx context.Context, dest interface{}) (err error) {
|
||||||
|
startedAt := time.Now()
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
err = logger.HandlePanic("PgSQLInsertQuery.Scan", r)
|
||||||
|
}
|
||||||
|
recordQueryMetrics(p.metricsEnabled, "INSERT", p.schema, p.entity, p.tableName, startedAt, err)
|
||||||
|
}()
|
||||||
|
|
||||||
|
if len(p.values) == 0 {
|
||||||
|
return fmt.Errorf("no values to insert")
|
||||||
|
}
|
||||||
|
|
||||||
|
columns := make([]string, 0, len(p.values))
|
||||||
|
placeholders := make([]string, 0, len(p.values))
|
||||||
|
args := make([]interface{}, 0, len(p.values))
|
||||||
|
i := 1
|
||||||
|
for _, col := range p.valueOrder {
|
||||||
|
columns = append(columns, col)
|
||||||
|
placeholders = append(placeholders, fmt.Sprintf("$%d", i))
|
||||||
|
args = append(args, p.values[col])
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
|
||||||
|
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, ", "))
|
||||||
|
|
||||||
|
if len(p.returning) > 0 {
|
||||||
|
query += " RETURNING " + strings.Join(p.returning, ", ")
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Debug("PgSQL INSERT (Scan): %s [args: %v]", query, args)
|
||||||
|
|
||||||
|
var row *sql.Row
|
||||||
|
if p.tx != nil {
|
||||||
|
row = p.tx.QueryRowContext(ctx, query, args...)
|
||||||
|
} else {
|
||||||
|
row = p.db.QueryRowContext(ctx, query, args...)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := row.Scan(dest); err != nil {
|
||||||
|
return common.WrapSQLError(err, query)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// PgSQLUpdateQuery implements UpdateQuery for PostgreSQL
|
// PgSQLUpdateQuery implements UpdateQuery for PostgreSQL
|
||||||
type PgSQLUpdateQuery struct {
|
type PgSQLUpdateQuery struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
tx *sql.Tx
|
tx *sql.Tx
|
||||||
|
schema string
|
||||||
tableName string
|
tableName string
|
||||||
|
entity string
|
||||||
driverName string
|
driverName string
|
||||||
model interface{}
|
model interface{}
|
||||||
sets map[string]interface{}
|
sets map[string]interface{}
|
||||||
|
setOrder []string
|
||||||
whereClauses []string
|
whereClauses []string
|
||||||
args []interface{}
|
args []interface{}
|
||||||
paramCounter int
|
paramCounter int
|
||||||
returning []string
|
returning []string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLUpdateQuery) Model(model interface{}) common.UpdateQuery {
|
func (p *PgSQLUpdateQuery) Model(model interface{}) common.UpdateQuery {
|
||||||
p.model = model
|
p.model = model
|
||||||
if provider, ok := model.(common.TableNameProvider); ok {
|
p.schema, p.tableName = schemaAndTableFromModel(model, p.driverName)
|
||||||
fullTableName := provider.TableName()
|
p.entity = entityNameFromModel(model, p.tableName)
|
||||||
// For SQLite, convert "schema.table" to "schema_table"
|
|
||||||
_, p.tableName = parseTableName(fullTableName, p.driverName)
|
|
||||||
}
|
|
||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLUpdateQuery) Table(table string) common.UpdateQuery {
|
func (p *PgSQLUpdateQuery) Table(table string) common.UpdateQuery {
|
||||||
// For SQLite, convert "schema.table" to "schema_table"
|
// For SQLite, convert "schema.table" to "schema_table"
|
||||||
_, p.tableName = parseTableName(table, p.driverName)
|
p.schema, p.tableName = parseTableName(table, p.driverName)
|
||||||
|
if p.entity == "" {
|
||||||
|
p.entity = cleanMetricIdentifier(p.tableName)
|
||||||
|
}
|
||||||
if p.model == nil {
|
if p.model == nil {
|
||||||
model, err := modelregistry.GetModelByName(table)
|
model, err := modelregistry.GetModelByName(table)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
p.model = model
|
p.model = model
|
||||||
|
p.entity = entityNameFromModel(model, p.tableName)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return p
|
return p
|
||||||
@@ -696,6 +805,9 @@ func (p *PgSQLUpdateQuery) Set(column string, value interface{}) common.UpdateQu
|
|||||||
if p.model != nil && !reflection.IsColumnWritable(p.model, column) {
|
if p.model != nil && !reflection.IsColumnWritable(p.model, column) {
|
||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
|
if _, exists := p.sets[column]; !exists {
|
||||||
|
p.setOrder = append(p.setOrder, column)
|
||||||
|
}
|
||||||
p.sets[column] = value
|
p.sets[column] = value
|
||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
@@ -706,13 +818,23 @@ func (p *PgSQLUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQu
|
|||||||
pkName = reflection.GetPrimaryKeyName(p.model)
|
pkName = reflection.GetPrimaryKeyName(p.model)
|
||||||
}
|
}
|
||||||
|
|
||||||
for column, value := range values {
|
orderedColumns := make([]string, 0, len(values))
|
||||||
|
for column := range values {
|
||||||
|
orderedColumns = append(orderedColumns, column)
|
||||||
|
}
|
||||||
|
sort.Strings(orderedColumns)
|
||||||
|
|
||||||
|
for _, column := range orderedColumns {
|
||||||
|
value := values[column]
|
||||||
if pkName != "" && column == pkName {
|
if pkName != "" && column == pkName {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
if p.model != nil && !reflection.IsColumnWritable(p.model, column) {
|
if p.model != nil && !reflection.IsColumnWritable(p.model, column) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
if _, exists := p.sets[column]; !exists {
|
||||||
|
p.setOrder = append(p.setOrder, column)
|
||||||
|
}
|
||||||
p.sets[column] = value
|
p.sets[column] = value
|
||||||
}
|
}
|
||||||
return p
|
return p
|
||||||
@@ -741,28 +863,30 @@ func (p *PgSQLUpdateQuery) replacePlaceholders(query string, argCount int) strin
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
||||||
|
startedAt := time.Now()
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("PgSQLUpdateQuery.Exec", r)
|
err = logger.HandlePanic("PgSQLUpdateQuery.Exec", r)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(p.metricsEnabled, "UPDATE", p.schema, p.entity, p.tableName, startedAt, err)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
if len(p.sets) == 0 {
|
if len(p.sets) == 0 {
|
||||||
return nil, fmt.Errorf("no values to update")
|
err = fmt.Errorf("no values to update")
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
setClauses := make([]string, 0, len(p.sets))
|
setClauses := make([]string, 0, len(p.sets))
|
||||||
setArgs := make([]interface{}, 0, len(p.sets))
|
setArgs := make([]interface{}, 0, len(p.sets))
|
||||||
|
|
||||||
// SET parameters start at $1
|
// SET parameters start at $1
|
||||||
i := 1
|
i := 1
|
||||||
for col, val := range p.sets {
|
for _, col := range p.setOrder {
|
||||||
setClauses = append(setClauses, fmt.Sprintf("%s = $%d", col, i))
|
setClauses = append(setClauses, fmt.Sprintf("%s = $%d", col, i))
|
||||||
setArgs = append(setArgs, val)
|
setArgs = append(setArgs, p.sets[col])
|
||||||
i++
|
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,
|
p.tableName,
|
||||||
strings.Join(setClauses, ", "))
|
strings.Join(setClauses, ", "))
|
||||||
|
|
||||||
@@ -812,7 +936,7 @@ func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err err
|
|||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL UPDATE failed: %v", err)
|
logger.Error("PgSQL UPDATE failed: %v", err)
|
||||||
return nil, err
|
return nil, common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
|
|
||||||
return &PgSQLResult{result: result}, nil
|
return &PgSQLResult{result: result}, nil
|
||||||
@@ -822,25 +946,28 @@ func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err err
|
|||||||
type PgSQLDeleteQuery struct {
|
type PgSQLDeleteQuery struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
tx *sql.Tx
|
tx *sql.Tx
|
||||||
|
schema string
|
||||||
tableName string
|
tableName string
|
||||||
|
entity string
|
||||||
driverName string
|
driverName string
|
||||||
whereClauses []string
|
whereClauses []string
|
||||||
args []interface{}
|
args []interface{}
|
||||||
paramCounter int
|
paramCounter int
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLDeleteQuery) Model(model interface{}) common.DeleteQuery {
|
func (p *PgSQLDeleteQuery) Model(model interface{}) common.DeleteQuery {
|
||||||
if provider, ok := model.(common.TableNameProvider); ok {
|
p.schema, p.tableName = schemaAndTableFromModel(model, p.driverName)
|
||||||
fullTableName := provider.TableName()
|
p.entity = entityNameFromModel(model, p.tableName)
|
||||||
// For SQLite, convert "schema.table" to "schema_table"
|
|
||||||
_, p.tableName = parseTableName(fullTableName, p.driverName)
|
|
||||||
}
|
|
||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLDeleteQuery) Table(table string) common.DeleteQuery {
|
func (p *PgSQLDeleteQuery) Table(table string) common.DeleteQuery {
|
||||||
// For SQLite, convert "schema.table" to "schema_table"
|
// For SQLite, convert "schema.table" to "schema_table"
|
||||||
_, p.tableName = parseTableName(table, p.driverName)
|
p.schema, p.tableName = parseTableName(table, p.driverName)
|
||||||
|
if p.entity == "" {
|
||||||
|
p.entity = cleanMetricIdentifier(p.tableName)
|
||||||
|
}
|
||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -862,13 +989,15 @@ func (p *PgSQLDeleteQuery) replacePlaceholders(query string, argCount int) strin
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLDeleteQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
func (p *PgSQLDeleteQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
||||||
|
startedAt := time.Now()
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("PgSQLDeleteQuery.Exec", r)
|
err = logger.HandlePanic("PgSQLDeleteQuery.Exec", r)
|
||||||
}
|
}
|
||||||
|
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 {
|
if len(p.whereClauses) > 0 {
|
||||||
query += " WHERE " + strings.Join(p.whereClauses, " AND ")
|
query += " WHERE " + strings.Join(p.whereClauses, " AND ")
|
||||||
@@ -885,7 +1014,7 @@ func (p *PgSQLDeleteQuery) Exec(ctx context.Context) (res common.Result, err err
|
|||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL DELETE failed: %v", err)
|
logger.Error("PgSQL DELETE failed: %v", err)
|
||||||
return nil, err
|
return nil, common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
|
|
||||||
return &PgSQLResult{result: result}, nil
|
return &PgSQLResult{result: result}, nil
|
||||||
@@ -915,6 +1044,7 @@ func (p *PgSQLResult) LastInsertId() (int64, error) {
|
|||||||
type PgSQLTxAdapter struct {
|
type PgSQLTxAdapter struct {
|
||||||
tx *sql.Tx
|
tx *sql.Tx
|
||||||
driverName string
|
driverName string
|
||||||
|
metricsEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLTxAdapter) NewSelect() common.SelectQuery {
|
func (p *PgSQLTxAdapter) NewSelect() common.SelectQuery {
|
||||||
@@ -923,6 +1053,7 @@ func (p *PgSQLTxAdapter) NewSelect() common.SelectQuery {
|
|||||||
driverName: p.driverName,
|
driverName: p.driverName,
|
||||||
columns: []string{"*"},
|
columns: []string{"*"},
|
||||||
args: make([]interface{}, 0),
|
args: make([]interface{}, 0),
|
||||||
|
metricsEnabled: p.metricsEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -931,6 +1062,7 @@ func (p *PgSQLTxAdapter) NewInsert() common.InsertQuery {
|
|||||||
tx: p.tx,
|
tx: p.tx,
|
||||||
driverName: p.driverName,
|
driverName: p.driverName,
|
||||||
values: make(map[string]interface{}),
|
values: make(map[string]interface{}),
|
||||||
|
metricsEnabled: p.metricsEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -941,6 +1073,7 @@ func (p *PgSQLTxAdapter) NewUpdate() common.UpdateQuery {
|
|||||||
sets: make(map[string]interface{}),
|
sets: make(map[string]interface{}),
|
||||||
args: make([]interface{}, 0),
|
args: make([]interface{}, 0),
|
||||||
whereClauses: make([]string, 0),
|
whereClauses: make([]string, 0),
|
||||||
|
metricsEnabled: p.metricsEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -950,29 +1083,39 @@ func (p *PgSQLTxAdapter) NewDelete() common.DeleteQuery {
|
|||||||
driverName: p.driverName,
|
driverName: p.driverName,
|
||||||
args: make([]interface{}, 0),
|
args: make([]interface{}, 0),
|
||||||
whereClauses: make([]string, 0),
|
whereClauses: make([]string, 0),
|
||||||
|
metricsEnabled: p.metricsEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLTxAdapter) Exec(ctx context.Context, query string, args ...interface{}) (common.Result, error) {
|
func (p *PgSQLTxAdapter) Exec(ctx context.Context, query string, args ...interface{}) (common.Result, error) {
|
||||||
|
startedAt := time.Now()
|
||||||
|
operation, schema, entity, table := metricTargetFromRawQuery(query, p.driverName)
|
||||||
logger.Debug("PgSQL Tx Exec: %s [args: %v]", query, args)
|
logger.Debug("PgSQL Tx Exec: %s [args: %v]", query, args)
|
||||||
result, err := p.tx.ExecContext(ctx, query, args...)
|
result, err := p.tx.ExecContext(ctx, query, args...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL Tx Exec failed: %v", err)
|
logger.Error("PgSQL Tx Exec failed: %v", err)
|
||||||
return nil, err
|
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
|
return nil, common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
|
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, nil)
|
||||||
return &PgSQLResult{result: result}, nil
|
return &PgSQLResult{result: result}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLTxAdapter) Query(ctx context.Context, dest interface{}, query string, args ...interface{}) error {
|
func (p *PgSQLTxAdapter) Query(ctx context.Context, dest interface{}, query string, args ...interface{}) error {
|
||||||
|
startedAt := time.Now()
|
||||||
|
operation, schema, entity, table := metricTargetFromRawQuery(query, p.driverName)
|
||||||
logger.Debug("PgSQL Tx Query: %s [args: %v]", query, args)
|
logger.Debug("PgSQL Tx Query: %s [args: %v]", query, args)
|
||||||
rows, err := p.tx.QueryContext(ctx, query, args...)
|
rows, err := p.tx.QueryContext(ctx, query, args...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL Tx Query failed: %v", err)
|
logger.Error("PgSQL Tx Query failed: %v", err)
|
||||||
return err
|
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
|
return common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
defer rows.Close()
|
defer rows.Close()
|
||||||
|
|
||||||
return scanRows(rows, dest)
|
err = scanRows(rows, dest)
|
||||||
|
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PgSQLTxAdapter) BeginTx(ctx context.Context) (common.Database, error) {
|
func (p *PgSQLTxAdapter) BeginTx(ctx context.Context) (common.Database, error) {
|
||||||
@@ -1052,7 +1195,7 @@ func (p *PgSQLSelectQuery) applySubqueryPreloads(ctx context.Context, dest inter
|
|||||||
|
|
||||||
// Use reflection to process the destination
|
// Use reflection to process the destination
|
||||||
destValue := reflect.ValueOf(dest)
|
destValue := reflect.ValueOf(dest)
|
||||||
if destValue.Kind() != reflect.Ptr {
|
if destValue.Kind() != reflect.Pointer {
|
||||||
return fmt.Errorf("dest must be a pointer")
|
return fmt.Errorf("dest must be a pointer")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1079,7 +1222,7 @@ func (p *PgSQLSelectQuery) applySubqueryPreloads(ctx context.Context, dest inter
|
|||||||
|
|
||||||
// loadPreloadsForRecord loads all preload relationships for a single record
|
// loadPreloadsForRecord loads all preload relationships for a single record
|
||||||
func (p *PgSQLSelectQuery) loadPreloadsForRecord(ctx context.Context, record reflect.Value, preloads []preloadConfig) error {
|
func (p *PgSQLSelectQuery) loadPreloadsForRecord(ctx context.Context, record reflect.Value, preloads []preloadConfig) error {
|
||||||
if record.Kind() == reflect.Ptr {
|
if record.Kind() == reflect.Pointer {
|
||||||
if record.IsNil() {
|
if record.IsNil() {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -1156,7 +1299,7 @@ func (p *PgSQLSelectQuery) executePreloadQuery(ctx context.Context, field reflec
|
|||||||
} else {
|
} else {
|
||||||
// Single struct - create a pointer if needed
|
// Single struct - create a pointer if needed
|
||||||
var target reflect.Value
|
var target reflect.Value
|
||||||
if field.Kind() == reflect.Ptr {
|
if field.Kind() == reflect.Pointer {
|
||||||
target = reflect.New(field.Type().Elem())
|
target = reflect.New(field.Type().Elem())
|
||||||
} else {
|
} else {
|
||||||
target = reflect.New(field.Type())
|
target = reflect.New(field.Type())
|
||||||
@@ -1169,7 +1312,7 @@ func (p *PgSQLSelectQuery) executePreloadQuery(ctx context.Context, field reflec
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Set the field
|
// Set the field
|
||||||
if field.Kind() == reflect.Ptr {
|
if field.Kind() == reflect.Pointer {
|
||||||
field.Set(target)
|
field.Set(target)
|
||||||
} else {
|
} else {
|
||||||
field.Set(target.Elem())
|
field.Set(target.Elem())
|
||||||
@@ -1186,7 +1329,7 @@ func (p *PgSQLSelectQuery) getRelationMetadata(fieldName string) *relationMetada
|
|||||||
}
|
}
|
||||||
|
|
||||||
modelType := reflect.TypeOf(p.model)
|
modelType := reflect.TypeOf(p.model)
|
||||||
if modelType.Kind() == reflect.Ptr {
|
if modelType.Kind() == reflect.Pointer {
|
||||||
modelType = modelType.Elem()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1235,7 +1378,7 @@ func (p *PgSQLSelectQuery) getRelationMetadataFromField(modelType reflect.Type,
|
|||||||
if fieldType.Kind() == reflect.Slice {
|
if fieldType.Kind() == reflect.Slice {
|
||||||
fieldType = fieldType.Elem()
|
fieldType = fieldType.Elem()
|
||||||
}
|
}
|
||||||
if fieldType.Kind() == reflect.Ptr {
|
if fieldType.Kind() == reflect.Pointer {
|
||||||
fieldType = fieldType.Elem()
|
fieldType = fieldType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1268,7 +1411,7 @@ func scanRows(rows *sql.Rows, dest interface{}) error {
|
|||||||
|
|
||||||
// Get destination type
|
// Get destination type
|
||||||
destValue := reflect.ValueOf(dest)
|
destValue := reflect.ValueOf(dest)
|
||||||
if destValue.Kind() != reflect.Ptr {
|
if destValue.Kind() != reflect.Pointer {
|
||||||
return fmt.Errorf("dest must be a pointer")
|
return fmt.Errorf("dest must be a pointer")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1323,7 +1466,7 @@ func scanRowsToMapSlice(rows *sql.Rows, columns []string, destValue reflect.Valu
|
|||||||
// scanRowsToStructSlice scans rows into a slice of structs
|
// scanRowsToStructSlice scans rows into a slice of structs
|
||||||
func scanRowsToStructSlice(rows *sql.Rows, columns []string, destValue reflect.Value) error {
|
func scanRowsToStructSlice(rows *sql.Rows, columns []string, destValue reflect.Value) error {
|
||||||
elemType := destValue.Type().Elem()
|
elemType := destValue.Type().Elem()
|
||||||
isPtr := elemType.Kind() == reflect.Ptr
|
isPtr := elemType.Kind() == reflect.Pointer
|
||||||
|
|
||||||
if isPtr {
|
if isPtr {
|
||||||
elemType = elemType.Elem()
|
elemType = elemType.Elem()
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ import (
|
|||||||
// Example demonstrates how to use the PgSQL adapter
|
// Example demonstrates how to use the PgSQL adapter
|
||||||
func ExamplePgSQLAdapter() error {
|
func ExamplePgSQLAdapter() error {
|
||||||
// Connect to PostgreSQL database
|
// 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)
|
db, err := sql.Open("pgx", dsn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to open database: %w", err)
|
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
|
// ExampleWithModel demonstrates using models with the PgSQL adapter
|
||||||
func ExampleWithModel() error {
|
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)
|
db, err := sql.Open("pgx", dsn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|||||||
@@ -51,7 +51,7 @@ func (c Comment) TableName() string {
|
|||||||
|
|
||||||
// ExamplePreload demonstrates the Preload functionality
|
// ExamplePreload demonstrates the Preload functionality
|
||||||
func ExamplePreload() error {
|
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)
|
db, err := sql.Open("pgx", dsn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -79,7 +79,7 @@ func ExamplePreload() error {
|
|||||||
|
|
||||||
// ExamplePreloadRelation demonstrates smart PreloadRelation with auto-detection
|
// ExamplePreloadRelation demonstrates smart PreloadRelation with auto-detection
|
||||||
func ExamplePreloadRelation() error {
|
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)
|
db, err := sql.Open("pgx", dsn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -148,7 +148,7 @@ func ExamplePreloadRelation() error {
|
|||||||
|
|
||||||
// ExampleJoinRelation demonstrates explicit JOIN loading
|
// ExampleJoinRelation demonstrates explicit JOIN loading
|
||||||
func ExampleJoinRelation() error {
|
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)
|
db, err := sql.Open("pgx", dsn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -185,7 +185,7 @@ func ExampleJoinRelation() error {
|
|||||||
|
|
||||||
// ExampleScanModel demonstrates ScanModel with struct destinations
|
// ExampleScanModel demonstrates ScanModel with struct destinations
|
||||||
func ExampleScanModel() error {
|
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)
|
db, err := sql.Open("pgx", dsn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -221,7 +221,7 @@ func ExampleScanModel() error {
|
|||||||
|
|
||||||
// ExampleCompleteWorkflow demonstrates a complete workflow with preloading
|
// ExampleCompleteWorkflow demonstrates a complete workflow with preloading
|
||||||
func ExampleCompleteWorkflow() error {
|
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)
|
db, err := sql.Open("pgx", dsn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|||||||
@@ -0,0 +1,335 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
)
|
||||||
|
|
||||||
|
const maxMetricFallbackEntityLength = 120
|
||||||
|
|
||||||
|
func recordQueryMetrics(enabled bool, operation, schema, entity, table string, startedAt time.Time, err error) {
|
||||||
|
if !enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
metrics.GetProvider().RecordDBQuery(
|
||||||
|
normalizeMetricOperation(operation),
|
||||||
|
normalizeMetricSchema(schema),
|
||||||
|
normalizeMetricEntity(entity, table),
|
||||||
|
normalizeMetricTable(table),
|
||||||
|
time.Since(startedAt),
|
||||||
|
err,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMetricOperation(operation string) string {
|
||||||
|
operation = strings.ToUpper(strings.TrimSpace(operation))
|
||||||
|
if operation == "" {
|
||||||
|
return "UNKNOWN"
|
||||||
|
}
|
||||||
|
return operation
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMetricSchema(schema string) string {
|
||||||
|
schema = cleanMetricIdentifier(schema)
|
||||||
|
if schema == "" {
|
||||||
|
return "default"
|
||||||
|
}
|
||||||
|
return schema
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMetricEntity(entity, table string) string {
|
||||||
|
entity = cleanMetricIdentifier(entity)
|
||||||
|
if entity != "" {
|
||||||
|
return entity
|
||||||
|
}
|
||||||
|
|
||||||
|
table = cleanMetricIdentifier(table)
|
||||||
|
if table != "" {
|
||||||
|
return table
|
||||||
|
}
|
||||||
|
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMetricTable(table string) string {
|
||||||
|
table = cleanMetricIdentifier(table)
|
||||||
|
if table == "" {
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
return table
|
||||||
|
}
|
||||||
|
|
||||||
|
func entityNameFromModel(model interface{}, table string) string {
|
||||||
|
if model == nil {
|
||||||
|
return cleanMetricIdentifier(table)
|
||||||
|
}
|
||||||
|
|
||||||
|
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 {
|
||||||
|
return cleanMetricIdentifier(table)
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelType.Kind() == reflect.Struct && modelType.Name() != "" {
|
||||||
|
return reflection.ToSnakeCase(modelType.Name())
|
||||||
|
}
|
||||||
|
|
||||||
|
return cleanMetricIdentifier(table)
|
||||||
|
}
|
||||||
|
|
||||||
|
func schemaAndTableFromModel(model interface{}, driverName string) (schema, table string) {
|
||||||
|
provider, ok := tableNameProviderFromModel(model)
|
||||||
|
if !ok {
|
||||||
|
return "", ""
|
||||||
|
}
|
||||||
|
|
||||||
|
return parseTableName(provider.TableName(), driverName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// tableNameProviderType is cached to avoid repeated reflection on every call.
|
||||||
|
var tableNameProviderType = reflect.TypeOf((*common.TableNameProvider)(nil)).Elem()
|
||||||
|
|
||||||
|
func tableNameProviderFromModel(model interface{}) (common.TableNameProvider, bool) {
|
||||||
|
if model == nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider, ok := model.(common.TableNameProvider); ok {
|
||||||
|
return provider, true
|
||||||
|
}
|
||||||
|
|
||||||
|
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 nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check whether *T implements TableNameProvider before allocating.
|
||||||
|
ptrType := reflect.PointerTo(modelType)
|
||||||
|
if !ptrType.Implements(tableNameProviderType) && !modelType.Implements(tableNameProviderType) {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
modelValue := reflect.New(modelType)
|
||||||
|
if provider, ok := modelValue.Interface().(common.TableNameProvider); ok {
|
||||||
|
return provider, true
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider, ok := modelValue.Elem().Interface().(common.TableNameProvider); ok {
|
||||||
|
return provider, true
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func metricTargetFromRawQuery(query, driverName string) (operation, schema, entity, table string) {
|
||||||
|
operation = normalizeMetricOperation(firstQueryKeyword(query))
|
||||||
|
tableRef := tableFromRawQuery(query, operation)
|
||||||
|
if tableRef == "" {
|
||||||
|
return operation, "", fallbackMetricEntityFromQuery(query), "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
schema, table = parseTableName(tableRef, driverName)
|
||||||
|
entity = cleanMetricIdentifier(table)
|
||||||
|
return operation, schema, entity, table
|
||||||
|
}
|
||||||
|
|
||||||
|
func fallbackMetricEntityFromQuery(query string) string {
|
||||||
|
query = sanitizeMetricQueryShape(query)
|
||||||
|
if query == "" {
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(query) > maxMetricFallbackEntityLength {
|
||||||
|
return query[:maxMetricFallbackEntityLength-3] + "..."
|
||||||
|
}
|
||||||
|
|
||||||
|
return query
|
||||||
|
}
|
||||||
|
|
||||||
|
func sanitizeMetricQueryShape(query string) string {
|
||||||
|
query = strings.TrimSpace(query)
|
||||||
|
if query == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var out strings.Builder
|
||||||
|
for i := 0; i < len(query); {
|
||||||
|
if query[i] == '\'' {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
for i < len(query) {
|
||||||
|
if query[i] == '\'' {
|
||||||
|
if i+1 < len(query) && query[i+1] == '\'' {
|
||||||
|
i += 2
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
i++
|
||||||
|
break
|
||||||
|
}
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if query[i] == '?' {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if query[i] == '$' && i+1 < len(query) && isASCIIDigit(query[i+1]) {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
for i < len(query) && isASCIIDigit(query[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if query[i] == ':' && (i == 0 || query[i-1] != ':') && i+1 < len(query) && isIdentifierStart(query[i+1]) {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
for i < len(query) && isIdentifierPart(query[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if query[i] == '@' && (i == 0 || query[i-1] != '@') && i+1 < len(query) && isIdentifierStart(query[i+1]) {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
for i < len(query) && isIdentifierPart(query[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if startsNumericLiteral(query, i) {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
for i < len(query) && (isASCIIDigit(query[i]) || query[i] == '.') {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
out.WriteByte(query[i])
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Join(strings.Fields(out.String()), " ")
|
||||||
|
}
|
||||||
|
|
||||||
|
func startsNumericLiteral(query string, idx int) bool {
|
||||||
|
if idx >= len(query) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
start := idx
|
||||||
|
if query[idx] == '-' {
|
||||||
|
if idx+1 >= len(query) || !isASCIIDigit(query[idx+1]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
start++
|
||||||
|
}
|
||||||
|
|
||||||
|
if !isASCIIDigit(query[start]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if idx > 0 && isIdentifierPart(query[idx-1]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if start+1 < len(query) && query[start] == '0' && (query[start+1] == 'x' || query[start+1] == 'X') {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func isASCIIDigit(ch byte) bool {
|
||||||
|
return ch >= '0' && ch <= '9'
|
||||||
|
}
|
||||||
|
|
||||||
|
func isIdentifierStart(ch byte) bool {
|
||||||
|
return (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') || ch == '_'
|
||||||
|
}
|
||||||
|
|
||||||
|
func isIdentifierPart(ch byte) bool {
|
||||||
|
return isIdentifierStart(ch) || isASCIIDigit(ch)
|
||||||
|
}
|
||||||
|
|
||||||
|
func firstQueryKeyword(query string) string {
|
||||||
|
query = strings.TrimSpace(query)
|
||||||
|
if query == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
fields := strings.Fields(query)
|
||||||
|
if len(fields) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
return fields[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
func tableFromRawQuery(query, operation string) string {
|
||||||
|
tokens := tokenizeQuery(query)
|
||||||
|
if len(tokens) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
switch operation {
|
||||||
|
case "SELECT":
|
||||||
|
return tokenAfter(tokens, "FROM")
|
||||||
|
case "INSERT":
|
||||||
|
return tokenAfter(tokens, "INTO")
|
||||||
|
case "UPDATE":
|
||||||
|
return tokenAfter(tokens, "UPDATE")
|
||||||
|
case "DELETE":
|
||||||
|
return tokenAfter(tokens, "FROM")
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func tokenAfter(tokens []string, keyword string) string {
|
||||||
|
for idx, token := range tokens {
|
||||||
|
if strings.EqualFold(token, keyword) && idx+1 < len(tokens) {
|
||||||
|
return cleanMetricIdentifier(tokens[idx+1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func tokenizeQuery(query string) []string {
|
||||||
|
replacer := strings.NewReplacer(
|
||||||
|
"\n", " ",
|
||||||
|
"\t", " ",
|
||||||
|
"(", " ",
|
||||||
|
")", " ",
|
||||||
|
",", " ",
|
||||||
|
)
|
||||||
|
return strings.Fields(replacer.Replace(query))
|
||||||
|
}
|
||||||
|
|
||||||
|
func cleanMetricIdentifier(value string) string {
|
||||||
|
value = strings.TrimSpace(value)
|
||||||
|
value = strings.Trim(value, "\"'`[]")
|
||||||
|
value = strings.TrimRight(value, ";")
|
||||||
|
return value
|
||||||
|
}
|
||||||
@@ -0,0 +1,394 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/DATA-DOG/go-sqlmock"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"github.com/uptrace/bun"
|
||||||
|
"github.com/uptrace/bun/dialect/sqlitedialect"
|
||||||
|
"github.com/uptrace/bun/driver/sqliteshim"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
)
|
||||||
|
|
||||||
|
type queryMetricCall struct {
|
||||||
|
operation string
|
||||||
|
schema string
|
||||||
|
entity string
|
||||||
|
table string
|
||||||
|
}
|
||||||
|
|
||||||
|
type capturingMetricsProvider struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
calls []queryMetricCall
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *capturingMetricsProvider) RecordHTTPRequest(method, path, status string, duration time.Duration) {
|
||||||
|
}
|
||||||
|
func (c *capturingMetricsProvider) IncRequestsInFlight() {}
|
||||||
|
func (c *capturingMetricsProvider) DecRequestsInFlight() {}
|
||||||
|
func (c *capturingMetricsProvider) RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error) {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
c.calls = append(c.calls, queryMetricCall{
|
||||||
|
operation: operation,
|
||||||
|
schema: schema,
|
||||||
|
entity: entity,
|
||||||
|
table: table,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
func (c *capturingMetricsProvider) RecordCacheHit(provider string) {}
|
||||||
|
func (c *capturingMetricsProvider) RecordCacheMiss(provider string) {}
|
||||||
|
func (c *capturingMetricsProvider) UpdateCacheSize(provider string, size int64) {
|
||||||
|
}
|
||||||
|
func (c *capturingMetricsProvider) RecordEventPublished(source, eventType string) {}
|
||||||
|
func (c *capturingMetricsProvider) RecordEventProcessed(source, eventType, status string, duration time.Duration) {
|
||||||
|
}
|
||||||
|
func (c *capturingMetricsProvider) UpdateEventQueueSize(size int64) {}
|
||||||
|
func (c *capturingMetricsProvider) RecordPanic(methodName string) {}
|
||||||
|
func (c *capturingMetricsProvider) Handler() http.Handler { return http.NewServeMux() }
|
||||||
|
|
||||||
|
func (c *capturingMetricsProvider) snapshot() []queryMetricCall {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
out := make([]queryMetricCall, len(c.calls))
|
||||||
|
copy(out, c.calls)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
type queryMetricsGormUser struct {
|
||||||
|
ID int `gorm:"primaryKey"`
|
||||||
|
Name string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (queryMetricsGormUser) TableName() string {
|
||||||
|
return "metrics_gorm_users"
|
||||||
|
}
|
||||||
|
|
||||||
|
type queryMetricsBunUser struct {
|
||||||
|
bun.BaseModel `bun:"table:metrics_bun_users"`
|
||||||
|
ID int64 `bun:"id,pk,autoincrement"`
|
||||||
|
Name string `bun:"name"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type queryMetricsBunParent struct {
|
||||||
|
bun.BaseModel `bun:"table:metrics_bun_parents"`
|
||||||
|
ID int64 `bun:"id,pk,autoincrement"`
|
||||||
|
Name string `bun:"name"`
|
||||||
|
Children []queryMetricsBunChild `bun:"rel:has-many,join:id=parent_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type queryMetricsBunChild struct {
|
||||||
|
bun.BaseModel `bun:"table:metrics_bun_children"`
|
||||||
|
ID int64 `bun:"id,pk,autoincrement"`
|
||||||
|
ParentID int64 `bun:"parent_id"`
|
||||||
|
Name string `bun:"name"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPgSQLAdapterRecordsSchemaEntityTableMetrics(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
provider := &capturingMetricsProvider{}
|
||||||
|
prev := metrics.GetProvider()
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
defer metrics.SetProvider(prev)
|
||||||
|
|
||||||
|
mock.ExpectExec(`UPDATE users SET name = \$1 WHERE id = \$2`).
|
||||||
|
WithArgs("Alice", 1).
|
||||||
|
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
_, err = adapter.NewUpdate().
|
||||||
|
Table("public.users").
|
||||||
|
Set("name", "Alice").
|
||||||
|
Where("id = ?", 1).
|
||||||
|
Exec(context.Background())
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, mock.ExpectationsWereMet())
|
||||||
|
|
||||||
|
calls := provider.snapshot()
|
||||||
|
require.Len(t, calls, 1)
|
||||||
|
assert.Equal(t, "UPDATE", calls[0].operation)
|
||||||
|
assert.Equal(t, "public", calls[0].schema)
|
||||||
|
assert.Equal(t, "users", calls[0].entity)
|
||||||
|
assert.Equal(t, "users", calls[0].table)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPgSQLAdapterDisableMetricsSuppressesEmission(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
provider := &capturingMetricsProvider{}
|
||||||
|
prev := metrics.GetProvider()
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
defer metrics.SetProvider(prev)
|
||||||
|
|
||||||
|
mock.ExpectExec(`DELETE FROM users WHERE id = \$1`).
|
||||||
|
WithArgs(1).
|
||||||
|
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db).SetMetricsEnabled(false)
|
||||||
|
_, err = adapter.NewDelete().
|
||||||
|
Table("users").
|
||||||
|
Where("id = ?", 1).
|
||||||
|
Exec(context.Background())
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, mock.ExpectationsWereMet())
|
||||||
|
assert.Empty(t, provider.snapshot())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGormAdapterRecordsEntityAndTableMetrics(t *testing.T) {
|
||||||
|
db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.NoError(t, db.AutoMigrate(&queryMetricsGormUser{}))
|
||||||
|
require.NoError(t, db.Create(&queryMetricsGormUser{Name: "Alice"}).Error)
|
||||||
|
|
||||||
|
provider := &capturingMetricsProvider{}
|
||||||
|
prev := metrics.GetProvider()
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
defer metrics.SetProvider(prev)
|
||||||
|
|
||||||
|
adapter := NewGormAdapter(db)
|
||||||
|
var users []queryMetricsGormUser
|
||||||
|
err = adapter.NewSelect().Model(&users).Scan(context.Background(), &users)
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotEmpty(t, users)
|
||||||
|
|
||||||
|
calls := provider.snapshot()
|
||||||
|
require.Len(t, calls, 1)
|
||||||
|
assert.Equal(t, "SELECT", calls[0].operation)
|
||||||
|
assert.Equal(t, "default", calls[0].schema)
|
||||||
|
assert.Equal(t, "query_metrics_gorm_user", calls[0].entity)
|
||||||
|
assert.Equal(t, "metrics_gorm_users", calls[0].table)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPgSQLAdapterRecordsErrorMetric(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
provider := &capturingMetricsProvider{}
|
||||||
|
prev := metrics.GetProvider()
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
defer metrics.SetProvider(prev)
|
||||||
|
|
||||||
|
mock.ExpectExec(`INSERT INTO users`).
|
||||||
|
WillReturnError(fmt.Errorf("unique constraint violation"))
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
_, err = adapter.NewInsert().
|
||||||
|
Table("users").
|
||||||
|
Value("name", "Alice").
|
||||||
|
Exec(context.Background())
|
||||||
|
|
||||||
|
require.Error(t, err)
|
||||||
|
|
||||||
|
calls := provider.snapshot()
|
||||||
|
require.Len(t, calls, 1)
|
||||||
|
assert.Equal(t, "INSERT", calls[0].operation)
|
||||||
|
assert.Equal(t, "users", calls[0].table)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPgSQLAdapterRecordsExistsMetric(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
provider := &capturingMetricsProvider{}
|
||||||
|
prev := metrics.GetProvider()
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
defer metrics.SetProvider(prev)
|
||||||
|
|
||||||
|
mock.ExpectQuery(`SELECT COUNT\(\*\) FROM users`).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(3))
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
exists, err := adapter.NewSelect().Table("users").Exists(context.Background())
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.True(t, exists)
|
||||||
|
|
||||||
|
calls := provider.snapshot()
|
||||||
|
require.Len(t, calls, 1)
|
||||||
|
assert.Equal(t, "EXISTS", calls[0].operation)
|
||||||
|
assert.Equal(t, "users", calls[0].table)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPgSQLAdapterRecordsCountMetric(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
provider := &capturingMetricsProvider{}
|
||||||
|
prev := metrics.GetProvider()
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
defer metrics.SetProvider(prev)
|
||||||
|
|
||||||
|
mock.ExpectQuery(`SELECT COUNT\(\*\) FROM users`).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(5))
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
count, err := adapter.NewSelect().Table("users").Count(context.Background())
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, 5, count)
|
||||||
|
|
||||||
|
calls := provider.snapshot()
|
||||||
|
require.Len(t, calls, 1)
|
||||||
|
assert.Equal(t, "COUNT", calls[0].operation)
|
||||||
|
assert.Equal(t, "users", calls[0].table)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPgSQLAdapterRawExecRecordsMetric(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
provider := &capturingMetricsProvider{}
|
||||||
|
prev := metrics.GetProvider()
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
defer metrics.SetProvider(prev)
|
||||||
|
|
||||||
|
mock.ExpectExec(`UPDATE public\.orders SET status = \$1`).
|
||||||
|
WithArgs("shipped").
|
||||||
|
WillReturnResult(sqlmock.NewResult(0, 2))
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
_, err = adapter.Exec(context.Background(), `UPDATE public.orders SET status = $1`, "shipped")
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
calls := provider.snapshot()
|
||||||
|
require.Len(t, calls, 1)
|
||||||
|
assert.Equal(t, "UPDATE", calls[0].operation)
|
||||||
|
assert.Equal(t, "public", calls[0].schema)
|
||||||
|
assert.Equal(t, "orders", calls[0].table)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPgSQLAdapterRawExecUsesSQLAsEntityWhenTargetUnknown(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
provider := &capturingMetricsProvider{}
|
||||||
|
prev := metrics.GetProvider()
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
defer metrics.SetProvider(prev)
|
||||||
|
|
||||||
|
query := `select core.c_setuserid($1)`
|
||||||
|
mock.ExpectExec(`select core\.c_setuserid\(\$1\)`).
|
||||||
|
WithArgs(42).
|
||||||
|
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
_, err = adapter.Exec(context.Background(), query, 42)
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
calls := provider.snapshot()
|
||||||
|
require.Len(t, calls, 1)
|
||||||
|
assert.Equal(t, "SELECT", calls[0].operation)
|
||||||
|
assert.Equal(t, "default", calls[0].schema)
|
||||||
|
assert.Equal(t, "select core.c_setuserid(?)", calls[0].entity)
|
||||||
|
assert.Equal(t, "unknown", calls[0].table)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallbackMetricEntityFromQuerySanitizesAndTruncates(t *testing.T) {
|
||||||
|
entity := fallbackMetricEntityFromQuery(" \n SELECT some_function(1, 'abc', $2, ?, :name, @p1, true, null) \t ")
|
||||||
|
assert.Equal(t, "SELECT some_function(?, ?, ?, ?, ?, ?, true, null)", entity)
|
||||||
|
|
||||||
|
entity = fallbackMetricEntityFromQuery("SELECT price::numeric, id FROM logs WHERE code = -42")
|
||||||
|
assert.Equal(t, "SELECT price::numeric, id FROM logs WHERE code = ?", entity)
|
||||||
|
|
||||||
|
longQuery := "SELECT " + strings.Repeat("x", maxMetricFallbackEntityLength)
|
||||||
|
entity = fallbackMetricEntityFromQuery(longQuery)
|
||||||
|
assert.Len(t, entity, maxMetricFallbackEntityLength)
|
||||||
|
assert.True(t, strings.HasSuffix(entity, "..."))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBunAdapterRecordsEntityAndTableMetrics(t *testing.T) {
|
||||||
|
sqldb, err := sql.Open(sqliteshim.ShimName, "file::memory:?cache=shared")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer sqldb.Close()
|
||||||
|
|
||||||
|
db := bun.NewDB(sqldb, sqlitedialect.New())
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
_, err = db.NewCreateTable().
|
||||||
|
Model((*queryMetricsBunUser)(nil)).
|
||||||
|
IfNotExists().
|
||||||
|
Exec(context.Background())
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
_, err = db.NewInsert().Model(&queryMetricsBunUser{Name: "Alice"}).Exec(context.Background())
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
provider := &capturingMetricsProvider{}
|
||||||
|
prev := metrics.GetProvider()
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
defer metrics.SetProvider(prev)
|
||||||
|
|
||||||
|
adapter := NewBunAdapter(db)
|
||||||
|
var users []queryMetricsBunUser
|
||||||
|
err = adapter.NewSelect().Model(&users).Scan(context.Background(), &users)
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotEmpty(t, users)
|
||||||
|
|
||||||
|
calls := provider.snapshot()
|
||||||
|
require.Len(t, calls, 1)
|
||||||
|
assert.Equal(t, "SELECT", calls[0].operation)
|
||||||
|
assert.Equal(t, "default", calls[0].schema)
|
||||||
|
assert.Equal(t, "query_metrics_bun_user", calls[0].entity)
|
||||||
|
assert.Equal(t, "metrics_bun_users", calls[0].table)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBunSelectQueryScanModelSupportsHasManyPreload(t *testing.T) {
|
||||||
|
sqldb, err := sql.Open(sqliteshim.ShimName, "file::memory:?cache=shared")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer sqldb.Close()
|
||||||
|
|
||||||
|
db := bun.NewDB(sqldb, sqlitedialect.New())
|
||||||
|
defer db.Close()
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
_, err = db.NewCreateTable().Model((*queryMetricsBunParent)(nil)).IfNotExists().Exec(ctx)
|
||||||
|
require.NoError(t, err)
|
||||||
|
_, err = db.NewCreateTable().Model((*queryMetricsBunChild)(nil)).IfNotExists().Exec(ctx)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
parent := &queryMetricsBunParent{Name: "parent"}
|
||||||
|
_, err = db.NewInsert().Model(parent).Exec(ctx)
|
||||||
|
require.NoError(t, err)
|
||||||
|
_, err = db.NewInsert().Model(&queryMetricsBunChild{ParentID: parent.ID, Name: "child"}).Exec(ctx)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
adapter := NewBunAdapter(db)
|
||||||
|
var parents []queryMetricsBunParent
|
||||||
|
err = adapter.NewSelect().Model(&parents).PreloadRelation("Children").Scan(ctx, &parents)
|
||||||
|
require.ErrorContains(t, err, "use Model instead of the dest parameter in Scan")
|
||||||
|
|
||||||
|
parents = nil
|
||||||
|
err = adapter.NewSelect().Model(&parents).PreloadRelation("Children").ScanModel(ctx)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, parents, 1)
|
||||||
|
require.Len(t, parents[0].Children, 1)
|
||||||
|
}
|
||||||
@@ -174,7 +174,9 @@ func (h *HTTPResponseWriter) Write(data []byte) (int, error) {
|
|||||||
|
|
||||||
func (h *HTTPResponseWriter) WriteJSON(data interface{}) error {
|
func (h *HTTPResponseWriter) WriteJSON(data interface{}) error {
|
||||||
h.SetHeader("Content-Type", "application/json")
|
h.SetHeader("Content-Type", "application/json")
|
||||||
return json.NewEncoder(h.resp).Encode(data)
|
enc := json.NewEncoder(h.resp)
|
||||||
|
enc.SetEscapeHTML(false)
|
||||||
|
return enc.Encode(data)
|
||||||
}
|
}
|
||||||
|
|
||||||
// UnderlyingResponseWriter returns the underlying http.ResponseWriter
|
// UnderlyingResponseWriter returns the underlying http.ResponseWriter
|
||||||
|
|||||||
+20
-13
@@ -115,32 +115,39 @@ func GetHeadSpecHeaders() []string {
|
|||||||
|
|
||||||
// SetCORSHeaders sets CORS headers on a response writer
|
// SetCORSHeaders sets CORS headers on a response writer
|
||||||
func SetCORSHeaders(w ResponseWriter, r Request, config CORSConfig) {
|
func SetCORSHeaders(w ResponseWriter, r Request, config CORSConfig) {
|
||||||
// Set allowed origins
|
// Reflect the request origin; fall back to wildcard only when no origin is present
|
||||||
// if len(config.AllowedOrigins) > 0 {
|
origin := r.Header("Origin")
|
||||||
// w.SetHeader("Access-Control-Allow-Origin", strings.Join(config.AllowedOrigins, ", "))
|
if origin == "" {
|
||||||
// }
|
origin = "*"
|
||||||
|
} else {
|
||||||
// Todo origin list parsing
|
// Vary must be set so caches don't serve one origin's response to another
|
||||||
w.SetHeader("Access-Control-Allow-Origin", "*")
|
httpW := w.UnderlyingResponseWriter()
|
||||||
|
httpW.Header().Set("Vary", "Origin")
|
||||||
|
}
|
||||||
|
w.SetHeader("Access-Control-Allow-Origin", origin)
|
||||||
|
|
||||||
// Set allowed methods
|
// Set allowed methods
|
||||||
if len(config.AllowedMethods) > 0 {
|
if len(config.AllowedMethods) > 0 {
|
||||||
w.SetHeader("Access-Control-Allow-Methods", strings.Join(config.AllowedMethods, ", "))
|
w.SetHeader("Access-Control-Allow-Methods", strings.Join(config.AllowedMethods, ", "))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Set allowed headers
|
// Reflect the preflight request headers when present; otherwise use the explicit config list
|
||||||
// if len(config.AllowedHeaders) > 0 {
|
requestedHeaders := r.Header("Access-Control-Request-Headers")
|
||||||
// w.SetHeader("Access-Control-Allow-Headers", strings.Join(config.AllowedHeaders, ", "))
|
if requestedHeaders != "" {
|
||||||
// }
|
w.SetHeader("Access-Control-Allow-Headers", requestedHeaders)
|
||||||
w.SetHeader("Access-Control-Allow-Headers", "*")
|
} else if len(config.AllowedHeaders) > 0 {
|
||||||
|
w.SetHeader("Access-Control-Allow-Headers", strings.Join(config.AllowedHeaders, ", "))
|
||||||
|
}
|
||||||
|
|
||||||
// Set max age
|
// Set max age
|
||||||
if config.MaxAge > 0 {
|
if config.MaxAge > 0 {
|
||||||
w.SetHeader("Access-Control-Max-Age", fmt.Sprintf("%d", config.MaxAge))
|
w.SetHeader("Access-Control-Max-Age", fmt.Sprintf("%d", config.MaxAge))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Allow credentials
|
// Allow credentials only when a specific origin is reflected (not wildcard)
|
||||||
|
if origin != "*" {
|
||||||
w.SetHeader("Access-Control-Allow-Credentials", "true")
|
w.SetHeader("Access-Control-Allow-Credentials", "true")
|
||||||
|
}
|
||||||
|
|
||||||
// Expose headers that clients can read
|
// Expose headers that clients can read
|
||||||
exposeHeaders := config.AllowedHeaders
|
exposeHeaders := config.AllowedHeaders
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package common
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
@@ -24,7 +25,7 @@ func ValidateAndUnwrapModel(model interface{}) (*ValidateAndUnwrapModelResult, e
|
|||||||
originalType := modelType
|
originalType := modelType
|
||||||
|
|
||||||
// Unwrap pointers, slices, and arrays to get to the base struct type
|
// Unwrap pointers, slices, and arrays to get to the base struct type
|
||||||
for modelType != nil && (modelType.Kind() == reflect.Ptr || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
modelType = modelType.Elem()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -125,15 +126,15 @@ func GetRelationshipInfo(modelType reflect.Type, relationName string) *Relations
|
|||||||
// Get related model type
|
// Get related model type
|
||||||
if field.Type.Kind() == reflect.Slice {
|
if field.Type.Kind() == reflect.Slice {
|
||||||
elemType := field.Type.Elem()
|
elemType := field.Type.Elem()
|
||||||
if elemType.Kind() == reflect.Ptr {
|
if elemType.Kind() == reflect.Pointer {
|
||||||
elemType = elemType.Elem()
|
elemType = elemType.Elem()
|
||||||
}
|
}
|
||||||
if elemType.Kind() == reflect.Struct {
|
if elemType.Kind() == reflect.Struct {
|
||||||
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||||
}
|
}
|
||||||
} else if field.Type.Kind() == reflect.Ptr || field.Type.Kind() == reflect.Struct {
|
} else if field.Type.Kind() == reflect.Pointer || field.Type.Kind() == reflect.Struct {
|
||||||
elemType := field.Type
|
elemType := field.Type
|
||||||
if elemType.Kind() == reflect.Ptr {
|
if elemType.Kind() == reflect.Pointer {
|
||||||
elemType = elemType.Elem()
|
elemType = elemType.Elem()
|
||||||
}
|
}
|
||||||
if elemType.Kind() == reflect.Struct {
|
if elemType.Kind() == reflect.Struct {
|
||||||
@@ -154,16 +155,16 @@ func GetRelationshipInfo(modelType reflect.Type, relationName string) *Relations
|
|||||||
info.RelationType = "hasMany"
|
info.RelationType = "hasMany"
|
||||||
// Get the element type for slice
|
// Get the element type for slice
|
||||||
elemType := field.Type.Elem()
|
elemType := field.Type.Elem()
|
||||||
if elemType.Kind() == reflect.Ptr {
|
if elemType.Kind() == reflect.Pointer {
|
||||||
elemType = elemType.Elem()
|
elemType = elemType.Elem()
|
||||||
}
|
}
|
||||||
if elemType.Kind() == reflect.Struct {
|
if elemType.Kind() == reflect.Struct {
|
||||||
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||||
}
|
}
|
||||||
} else if field.Type.Kind() == reflect.Ptr || field.Type.Kind() == reflect.Struct {
|
} else if field.Type.Kind() == reflect.Pointer || field.Type.Kind() == reflect.Struct {
|
||||||
info.RelationType = "belongsTo"
|
info.RelationType = "belongsTo"
|
||||||
elemType := field.Type
|
elemType := field.Type
|
||||||
if elemType.Kind() == reflect.Ptr {
|
if elemType.Kind() == reflect.Pointer {
|
||||||
elemType = elemType.Elem()
|
elemType = elemType.Elem()
|
||||||
}
|
}
|
||||||
if elemType.Kind() == reflect.Struct {
|
if elemType.Kind() == reflect.Struct {
|
||||||
@@ -176,7 +177,7 @@ func GetRelationshipInfo(modelType reflect.Type, relationName string) *Relations
|
|||||||
// Get the element type for many2many (always slice)
|
// Get the element type for many2many (always slice)
|
||||||
if field.Type.Kind() == reflect.Slice {
|
if field.Type.Kind() == reflect.Slice {
|
||||||
elemType := field.Type.Elem()
|
elemType := field.Type.Elem()
|
||||||
if elemType.Kind() == reflect.Ptr {
|
if elemType.Kind() == reflect.Pointer {
|
||||||
elemType = elemType.Elem()
|
elemType = elemType.Elem()
|
||||||
}
|
}
|
||||||
if elemType.Kind() == reflect.Struct {
|
if elemType.Kind() == reflect.Struct {
|
||||||
@@ -238,7 +239,7 @@ func GetTableNameFromModel(model interface{}) string {
|
|||||||
modelType := reflect.TypeOf(model)
|
modelType := reflect.TypeOf(model)
|
||||||
|
|
||||||
// Unwrap pointers
|
// Unwrap pointers
|
||||||
for modelType != nil && modelType.Kind() == reflect.Ptr {
|
for modelType != nil && modelType.Kind() == reflect.Pointer {
|
||||||
modelType = modelType.Elem()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -261,3 +262,48 @@ func GetTableNameFromModel(model interface{}) string {
|
|||||||
// This handles cases like "MasterTaskItem" -> "mastertaskitem"
|
// This handles cases like "MasterTaskItem" -> "mastertaskitem"
|
||||||
return strings.ToLower(modelType.Name())
|
return strings.ToLower(modelType.Name())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ConvertSliceForBun converts []interface{} values to PostgreSQL array literal strings.
|
||||||
|
// BUN's fallback appender for []interface{} is JSON encoding, which produces "[]" —
|
||||||
|
// invalid PostgreSQL array syntax. PostgreSQL expects "{}" for empty arrays and
|
||||||
|
// "{elem1,elem2}" for non-empty ones. All other value types are returned unchanged.
|
||||||
|
func ConvertSliceForBun(value interface{}) interface{} {
|
||||||
|
arr, ok := value.([]interface{})
|
||||||
|
if !ok {
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
if len(arr) == 0 {
|
||||||
|
return "{}"
|
||||||
|
}
|
||||||
|
parts := make([]string, len(arr))
|
||||||
|
for i, elem := range arr {
|
||||||
|
switch e := elem.(type) {
|
||||||
|
case string:
|
||||||
|
needsQuote := e == "" || strings.ContainsAny(e, `,"\\{}`+"\t\n\r ")
|
||||||
|
if needsQuote {
|
||||||
|
e = strings.ReplaceAll(e, `\`, `\\`)
|
||||||
|
e = strings.ReplaceAll(e, `"`, `""`)
|
||||||
|
parts[i] = `"` + e + `"`
|
||||||
|
} else {
|
||||||
|
parts[i] = e
|
||||||
|
}
|
||||||
|
case float64:
|
||||||
|
if e == float64(int64(e)) {
|
||||||
|
parts[i] = strconv.FormatInt(int64(e), 10)
|
||||||
|
} else {
|
||||||
|
parts[i] = strconv.FormatFloat(e, 'f', -1, 64)
|
||||||
|
}
|
||||||
|
case bool:
|
||||||
|
if e {
|
||||||
|
parts[i] = "t"
|
||||||
|
} else {
|
||||||
|
parts[i] = "f"
|
||||||
|
}
|
||||||
|
case nil:
|
||||||
|
parts[i] = "NULL"
|
||||||
|
default:
|
||||||
|
parts[i] = fmt.Sprintf("%v", e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return "{" + strings.Join(parts, ",") + "}"
|
||||||
|
}
|
||||||
|
|||||||
@@ -106,3 +106,66 @@ func TestExtractTagValue(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestConvertSliceForBun(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
expected interface{}
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "empty slice produces empty pg array",
|
||||||
|
input: []interface{}{},
|
||||||
|
expected: "{}",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "string elements",
|
||||||
|
input: []interface{}{"a", "b", "c"},
|
||||||
|
expected: "{a,b,c}",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "string element needing quotes",
|
||||||
|
input: []interface{}{"hello world", "ok"},
|
||||||
|
expected: `{"hello world",ok}`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "string with comma",
|
||||||
|
input: []interface{}{"a,b"},
|
||||||
|
expected: `{"a,b"}`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "integer elements (JSON float64)",
|
||||||
|
input: []interface{}{float64(1), float64(2), float64(3)},
|
||||||
|
expected: "{1,2,3}",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "bool elements",
|
||||||
|
input: []interface{}{true, false},
|
||||||
|
expected: "{t,f}",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "nil input passthrough",
|
||||||
|
input: nil,
|
||||||
|
expected: nil,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "string input passthrough",
|
||||||
|
input: "hello",
|
||||||
|
expected: "hello",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "int input passthrough",
|
||||||
|
input: 42,
|
||||||
|
expected: 42,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
result := ConvertSliceForBun(tt.input)
|
||||||
|
if result != tt.expected {
|
||||||
|
t.Errorf("ConvertSliceForBun(%v) = %v; want %v", tt.input, result, tt.expected)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -75,6 +75,7 @@ type InsertQuery interface {
|
|||||||
|
|
||||||
// Execution
|
// Execution
|
||||||
Exec(ctx context.Context) (Result, error)
|
Exec(ctx context.Context) (Result, error)
|
||||||
|
Scan(ctx context.Context, dest interface{}) error
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateQuery interface for building UPDATE queries
|
// UpdateQuery interface for building UPDATE queries
|
||||||
@@ -177,7 +178,9 @@ func (s *StandardResponseWriter) Write(data []byte) (int, error) {
|
|||||||
|
|
||||||
func (s *StandardResponseWriter) WriteJSON(data interface{}) error {
|
func (s *StandardResponseWriter) WriteJSON(data interface{}) error {
|
||||||
s.SetHeader("Content-Type", "application/json")
|
s.SetHeader("Content-Type", "application/json")
|
||||||
return json.NewEncoder(s.w).Encode(data)
|
enc := json.NewEncoder(s.w)
|
||||||
|
enc.SetEscapeHTML(false)
|
||||||
|
return enc.Encode(data)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *StandardResponseWriter) UnderlyingResponseWriter() http.ResponseWriter {
|
func (s *StandardResponseWriter) UnderlyingResponseWriter() http.ResponseWriter {
|
||||||
|
|||||||
@@ -0,0 +1,402 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// This file implements a single canonical parser + SQL builder for column
|
||||||
|
// references that traverse into JSON / JSONB values. It is used by SELECT,
|
||||||
|
// WHERE (filter) and ORDER BY handling so that all three treat JSON access
|
||||||
|
// consistently and safely.
|
||||||
|
//
|
||||||
|
// Supported input syntaxes (all PostgreSQL-oriented):
|
||||||
|
//
|
||||||
|
// data->>'city' arrow chain, text extraction
|
||||||
|
// data->'addr'->>'city' nested arrow chain
|
||||||
|
// data->2->>'name' arrow chain with array index
|
||||||
|
// data#>>'{addr,city}' hash-path, text extraction
|
||||||
|
// data#>'{addr,city}' hash-path, jsonb result
|
||||||
|
// data.addr.city dotted shorthand (Ambiguous: caller must
|
||||||
|
// confirm "data" is a JSON column)
|
||||||
|
// data->>'age'::int trailing cast (whitelisted targets only)
|
||||||
|
// (data->>'city') AS city parenthesised, with output alias
|
||||||
|
//
|
||||||
|
// JSON path segments are never interpolated into SQL: SQL() emits a `#>>` /
|
||||||
|
// `#>` operator with the path bound as a single `text[]` parameter.
|
||||||
|
|
||||||
|
// ColumnRef is a parsed reference to a (possibly JSON-traversing) column.
|
||||||
|
type ColumnRef struct {
|
||||||
|
// Base is the bare base column name, e.g. "data". Always a simple
|
||||||
|
// identifier ([A-Za-z_][A-Za-z0-9_]*); qualified names are rejected.
|
||||||
|
Base string
|
||||||
|
// Path is the JSON key / array-index path, e.g. ["address", "city"].
|
||||||
|
// Empty for a plain column reference.
|
||||||
|
Path []string
|
||||||
|
// AsText is true when the final extraction should yield text (->> / #>>)
|
||||||
|
// rather than jsonb (-> / #>).
|
||||||
|
AsText bool
|
||||||
|
// Cast is a normalised SQL type name to cast the whole expression to
|
||||||
|
// (e.g. "integer", "numeric", "timestamptz"), or "" for no cast.
|
||||||
|
Cast string
|
||||||
|
// Alias is a validated output identifier for `AS <alias>`, or "".
|
||||||
|
Alias string
|
||||||
|
// Ambiguous is true when Path was produced from the dotted "a.b.c"
|
||||||
|
// shorthand. The caller MUST verify that Base is a JSON column
|
||||||
|
// (reflection.IsJSONColumn) before treating this as a JSON expression,
|
||||||
|
// otherwise "a.b" is an ordinary table-qualified column.
|
||||||
|
Ambiguous bool
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
reSimpleIdent = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
|
||||||
|
reSimpleSegment = regexp.MustCompile(`^[A-Za-z0-9_]+$`)
|
||||||
|
reAliasSuffix = regexp.MustCompile(`(?i)\s+AS\s+("?[A-Za-z_][A-Za-z0-9_]*"?)\s*$`)
|
||||||
|
reArrowStep = regexp.MustCompile(`^\s*(->>|->)\s*(?:'((?:[^']|'')*)'|(\d+))\s*`)
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
maxJSONPathDepth = 32
|
||||||
|
maxJSONSegmentSize = 128
|
||||||
|
)
|
||||||
|
|
||||||
|
// castAliases maps accepted cast spellings to their canonical PostgreSQL type.
|
||||||
|
var castAliases = map[string]string{
|
||||||
|
"int": "integer",
|
||||||
|
"int4": "integer",
|
||||||
|
"integer": "integer",
|
||||||
|
"int2": "smallint",
|
||||||
|
"smallint": "smallint",
|
||||||
|
"int8": "bigint",
|
||||||
|
"bigint": "bigint",
|
||||||
|
"numeric": "numeric",
|
||||||
|
"decimal": "numeric",
|
||||||
|
"real": "real",
|
||||||
|
"float4": "real",
|
||||||
|
"float": "double precision",
|
||||||
|
"float8": "double precision",
|
||||||
|
"double precision": "double precision",
|
||||||
|
"bool": "boolean",
|
||||||
|
"boolean": "boolean",
|
||||||
|
"text": "text",
|
||||||
|
"varchar": "text",
|
||||||
|
"uuid": "uuid",
|
||||||
|
"date": "date",
|
||||||
|
"time": "time",
|
||||||
|
"timestamp": "timestamp",
|
||||||
|
"timestamptz": "timestamptz",
|
||||||
|
"json": "json",
|
||||||
|
"jsonb": "jsonb",
|
||||||
|
}
|
||||||
|
|
||||||
|
// NormalizeCastTarget returns the canonical PostgreSQL type name for a
|
||||||
|
// user-supplied cast spelling, and whether it is on the allowlist.
|
||||||
|
func NormalizeCastTarget(s string) (string, bool) {
|
||||||
|
c, ok := castAliases[strings.ToLower(strings.TrimSpace(s))]
|
||||||
|
return c, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseColumnRef parses a column token that traverses into a JSON value.
|
||||||
|
//
|
||||||
|
// ok is true only when the token carries JSON traversal syntax (arrow chain,
|
||||||
|
// hash-path, or dotted shorthand with at least one sub-key). For a plain
|
||||||
|
// column name — with or without an alias/cast — ok is false and the caller
|
||||||
|
// should handle the token the way it did before.
|
||||||
|
//
|
||||||
|
// When ok is true and ref.Ambiguous is true, the caller must confirm that
|
||||||
|
// ref.Base is a JSON column before using ref.SQL; otherwise the dotted token
|
||||||
|
// is an ordinary "table.column" reference.
|
||||||
|
func ParseColumnRef(raw string) (ColumnRef, bool) {
|
||||||
|
expr := strings.TrimSpace(raw)
|
||||||
|
if expr == "" {
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
var ref ColumnRef
|
||||||
|
|
||||||
|
// 1. Trailing `AS <alias>`.
|
||||||
|
if m := reAliasSuffix.FindStringSubmatch(expr); m != nil {
|
||||||
|
ref.Alias = strings.Trim(m[1], `"`)
|
||||||
|
expr = strings.TrimSpace(expr[:len(expr)-len(m[0])])
|
||||||
|
if expr == "" {
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2. Trailing `::<type>` cast (take the last `::` in the string).
|
||||||
|
if idx := strings.LastIndex(expr, "::"); idx != -1 {
|
||||||
|
candidate := strings.TrimSpace(expr[idx+2:])
|
||||||
|
if canonical, allowed := NormalizeCastTarget(candidate); allowed {
|
||||||
|
ref.Cast = canonical
|
||||||
|
expr = strings.TrimSpace(expr[:idx])
|
||||||
|
} else if candidate != "" && looksLikeCastTail(candidate) {
|
||||||
|
// An explicit but unsupported cast target — reject rather than
|
||||||
|
// silently dropping it.
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. One layer of wrapping parentheses: "(expr)" -> "expr".
|
||||||
|
if wrapped, ok := stripWrappingParens(expr); ok {
|
||||||
|
expr = strings.TrimSpace(wrapped)
|
||||||
|
if expr == "" {
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 4. Parse the core expression.
|
||||||
|
switch {
|
||||||
|
case strings.Contains(expr, "#>>") || strings.Contains(expr, "#>"):
|
||||||
|
if !parseHashPath(expr, &ref) {
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
case strings.Contains(expr, "->"):
|
||||||
|
if !parseArrowChain(expr, &ref) {
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
case strings.Contains(expr, "."):
|
||||||
|
if !parseDottedPath(expr, &ref) {
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
// Plain column — nothing JSON about it.
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
if !validateRef(&ref) {
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
return ref, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// SQL renders the reference as a parameterised SQL expression plus its args.
|
||||||
|
// tableAlias, when non-empty, qualifies the base column (each dot-separated
|
||||||
|
// part is quoted independently, so "public.users" -> `"public"."users"`).
|
||||||
|
func (r ColumnRef) SQL(tableAlias string) (expr string, args []interface{}) {
|
||||||
|
base := quoteQualifiedIdent(r.Base)
|
||||||
|
if tableAlias != "" {
|
||||||
|
base = quoteQualifiedIdent(tableAlias) + "." + QuoteIdent(r.Base)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(r.Path) == 0 {
|
||||||
|
if r.Cast != "" {
|
||||||
|
return fmt.Sprintf("(%s)::%s", base, r.Cast), nil
|
||||||
|
}
|
||||||
|
return base, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
op := "#>"
|
||||||
|
if r.AsText {
|
||||||
|
op = "#>>"
|
||||||
|
}
|
||||||
|
expr = fmt.Sprintf("(%s %s ?::text[])", base, op)
|
||||||
|
args = []interface{}{pgTextArrayLiteral(r.Path)}
|
||||||
|
|
||||||
|
if r.Cast != "" {
|
||||||
|
expr = fmt.Sprintf("(%s)::%s", expr, r.Cast)
|
||||||
|
}
|
||||||
|
return expr, args
|
||||||
|
}
|
||||||
|
|
||||||
|
// OutputAlias returns the alias to use for this reference in a SELECT list:
|
||||||
|
// the explicit alias when given, otherwise a deterministic name derived from
|
||||||
|
// the base column and path (e.g. "data_address_city").
|
||||||
|
func (r ColumnRef) OutputAlias() string {
|
||||||
|
if r.Alias != "" {
|
||||||
|
return r.Alias
|
||||||
|
}
|
||||||
|
if len(r.Path) == 0 {
|
||||||
|
return r.Base
|
||||||
|
}
|
||||||
|
parts := make([]string, 0, len(r.Path)+1)
|
||||||
|
parts = append(parts, r.Base)
|
||||||
|
for _, p := range r.Path {
|
||||||
|
parts = append(parts, sanitizeAliasPart(p))
|
||||||
|
}
|
||||||
|
return strings.Join(parts, "_")
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── parsing helpers ─────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
func parseHashPath(expr string, ref *ColumnRef) bool {
|
||||||
|
op := "#>>"
|
||||||
|
ref.AsText = true
|
||||||
|
if !strings.Contains(expr, "#>>") {
|
||||||
|
op = "#>"
|
||||||
|
ref.AsText = false
|
||||||
|
}
|
||||||
|
parts := strings.SplitN(expr, op, 2)
|
||||||
|
if len(parts) != 2 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
ref.Base = strings.TrimSpace(parts[0])
|
||||||
|
|
||||||
|
rhs := strings.TrimSpace(parts[1])
|
||||||
|
// Expect a single-quoted array literal: '{a,b,c}'
|
||||||
|
if len(rhs) < 2 || rhs[0] != '\'' || rhs[len(rhs)-1] != '\'' {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
rhs = rhs[1 : len(rhs)-1]
|
||||||
|
rhs = strings.TrimSpace(rhs)
|
||||||
|
rhs = strings.TrimPrefix(rhs, "{")
|
||||||
|
rhs = strings.TrimSuffix(rhs, "}")
|
||||||
|
if strings.TrimSpace(rhs) == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, seg := range strings.Split(rhs, ",") {
|
||||||
|
seg = strings.TrimSpace(seg)
|
||||||
|
seg = strings.Trim(seg, `"`)
|
||||||
|
if seg == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
ref.Path = append(ref.Path, seg)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseArrowChain(expr string, ref *ColumnRef) bool {
|
||||||
|
arrowIdx := strings.Index(expr, "->")
|
||||||
|
if arrowIdx <= 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
ref.Base = strings.TrimSpace(expr[:arrowIdx])
|
||||||
|
|
||||||
|
rest := expr[arrowIdx:]
|
||||||
|
for strings.TrimSpace(rest) != "" {
|
||||||
|
m := reArrowStep.FindStringSubmatch(rest)
|
||||||
|
if m == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
ref.AsText = m[1] == "->>"
|
||||||
|
if m[3] != "" {
|
||||||
|
// unquoted array index
|
||||||
|
ref.Path = append(ref.Path, m[3])
|
||||||
|
} else {
|
||||||
|
// quoted key; unescape doubled single quotes
|
||||||
|
ref.Path = append(ref.Path, strings.ReplaceAll(m[2], "''", "'"))
|
||||||
|
}
|
||||||
|
rest = rest[len(m[0]):]
|
||||||
|
}
|
||||||
|
return len(ref.Path) > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseDottedPath(expr string, ref *ColumnRef) bool {
|
||||||
|
segs := strings.Split(expr, ".")
|
||||||
|
if len(segs) < 2 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for i, s := range segs {
|
||||||
|
s = strings.TrimSpace(s)
|
||||||
|
if !reSimpleSegment.MatchString(s) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if i == 0 {
|
||||||
|
ref.Base = s
|
||||||
|
} else {
|
||||||
|
ref.Path = append(ref.Path, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
ref.AsText = true
|
||||||
|
ref.Ambiguous = true
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func validateRef(ref *ColumnRef) bool {
|
||||||
|
if !reSimpleIdent.MatchString(ref.Base) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if len(ref.Path) == 0 || len(ref.Path) > maxJSONPathDepth {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, seg := range ref.Path {
|
||||||
|
if seg == "" || len(seg) > maxJSONSegmentSize || strings.ContainsRune(seg, 0) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if ref.Alias != "" && !reSimpleIdent.MatchString(ref.Alias) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// looksLikeCastTail reports whether s is plausibly meant as a `::type` target
|
||||||
|
// (letters/digits/spaces only) rather than, say, part of a JSON operator.
|
||||||
|
func looksLikeCastTail(s string) bool {
|
||||||
|
for _, r := range s {
|
||||||
|
isCastChar := r == ' ' || r == '_' ||
|
||||||
|
(r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9')
|
||||||
|
if !isCastChar {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// stripWrappingParens removes one layer of parentheses when they wrap the whole
|
||||||
|
// expression, e.g. "(a->>'b')" -> "a->>'b'". It respects single-quoted strings.
|
||||||
|
func stripWrappingParens(s string) (string, bool) {
|
||||||
|
s = strings.TrimSpace(s)
|
||||||
|
if len(s) < 2 || s[0] != '(' || s[len(s)-1] != ')' {
|
||||||
|
return s, false
|
||||||
|
}
|
||||||
|
depth := 0
|
||||||
|
inQuote := false
|
||||||
|
for i := 0; i < len(s); i++ {
|
||||||
|
c := s[i]
|
||||||
|
switch {
|
||||||
|
case c == '\'':
|
||||||
|
inQuote = !inQuote
|
||||||
|
case inQuote:
|
||||||
|
// skip
|
||||||
|
case c == '(':
|
||||||
|
depth++
|
||||||
|
case c == ')':
|
||||||
|
depth--
|
||||||
|
if depth == 0 && i != len(s)-1 {
|
||||||
|
// closing paren is not the last char -> not a full wrap
|
||||||
|
return s, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if depth != 0 {
|
||||||
|
return s, false
|
||||||
|
}
|
||||||
|
return s[1 : len(s)-1], true
|
||||||
|
}
|
||||||
|
|
||||||
|
// pgTextArrayLiteral builds a PostgreSQL text[] array literal ("{a,b,c}") from
|
||||||
|
// path segments, quoting and escaping any segment that is not a bare word.
|
||||||
|
func pgTextArrayLiteral(segs []string) string {
|
||||||
|
escaper := strings.NewReplacer(`\`, `\\`, `"`, `\"`)
|
||||||
|
parts := make([]string, len(segs))
|
||||||
|
for i, s := range segs {
|
||||||
|
if reSimpleSegment.MatchString(s) {
|
||||||
|
parts[i] = s
|
||||||
|
} else {
|
||||||
|
parts[i] = `"` + escaper.Replace(s) + `"`
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return "{" + strings.Join(parts, ",") + "}"
|
||||||
|
}
|
||||||
|
|
||||||
|
// quoteQualifiedIdent quotes each dot-separated part of an identifier.
|
||||||
|
func quoteQualifiedIdent(ident string) string {
|
||||||
|
parts := strings.Split(ident, ".")
|
||||||
|
for i, p := range parts {
|
||||||
|
parts[i] = QuoteIdent(p)
|
||||||
|
}
|
||||||
|
return strings.Join(parts, ".")
|
||||||
|
}
|
||||||
|
|
||||||
|
func sanitizeAliasPart(s string) string {
|
||||||
|
var b strings.Builder
|
||||||
|
for _, r := range s {
|
||||||
|
if (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9') || r == '_' {
|
||||||
|
b.WriteRune(r)
|
||||||
|
} else {
|
||||||
|
b.WriteRune('_')
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return b.String()
|
||||||
|
}
|
||||||
@@ -0,0 +1,274 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestParseColumnRef_Valid(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
base string
|
||||||
|
path []string
|
||||||
|
asText bool
|
||||||
|
cast string
|
||||||
|
alias string
|
||||||
|
ambiguous bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "arrow text extraction",
|
||||||
|
input: "data->>'city'",
|
||||||
|
base: "data", path: []string{"city"}, asText: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "arrow whitespace tolerant",
|
||||||
|
input: "data ->> 'city'",
|
||||||
|
base: "data", path: []string{"city"}, asText: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "nested arrow chain",
|
||||||
|
input: "data->'address'->>'city'",
|
||||||
|
base: "data", path: []string{"address", "city"}, asText: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "arrow jsonb result",
|
||||||
|
input: "data->'address'",
|
||||||
|
base: "data", path: []string{"address"}, asText: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "arrow array index",
|
||||||
|
input: "items->0->>'name'",
|
||||||
|
base: "items", path: []string{"0", "name"}, asText: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "hash path text",
|
||||||
|
input: "data#>>'{address,city}'",
|
||||||
|
base: "data", path: []string{"address", "city"}, asText: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "hash path jsonb",
|
||||||
|
input: "data#>'{address,city}'",
|
||||||
|
base: "data", path: []string{"address", "city"}, asText: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "dotted shorthand",
|
||||||
|
input: "data.address.city",
|
||||||
|
base: "data", path: []string{"address", "city"}, asText: true, ambiguous: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "trailing cast",
|
||||||
|
input: "data->>'age'::int",
|
||||||
|
base: "data", path: []string{"age"}, asText: true, cast: "integer",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "cast normalises",
|
||||||
|
input: "data->>'ts'::timestamptz",
|
||||||
|
base: "data", path: []string{"ts"}, asText: true, cast: "timestamptz",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "parenthesised with alias",
|
||||||
|
input: "(data->>'city') AS city_name",
|
||||||
|
base: "data", path: []string{"city"}, asText: true, alias: "city_name",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "paren wrap and cast",
|
||||||
|
input: "(data->>'age')::numeric",
|
||||||
|
base: "data", path: []string{"age"}, asText: true, cast: "numeric",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "quoted key with spaces",
|
||||||
|
input: "data->>'key with space'",
|
||||||
|
base: "data", path: []string{"key with space"}, asText: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "quoted key with escaped quote",
|
||||||
|
input: "data->>'o''brien'",
|
||||||
|
base: "data", path: []string{"o'brien"}, asText: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "relation column is ambiguous json",
|
||||||
|
input: "orders.total",
|
||||||
|
base: "orders", path: []string{"total"}, asText: true, ambiguous: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
ref, ok := ParseColumnRef(tc.input)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("ParseColumnRef(%q) returned ok=false", tc.input)
|
||||||
|
}
|
||||||
|
if ref.Base != tc.base {
|
||||||
|
t.Errorf("Base = %q, want %q", ref.Base, tc.base)
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(ref.Path, tc.path) {
|
||||||
|
t.Errorf("Path = %#v, want %#v", ref.Path, tc.path)
|
||||||
|
}
|
||||||
|
if ref.AsText != tc.asText {
|
||||||
|
t.Errorf("AsText = %v, want %v", ref.AsText, tc.asText)
|
||||||
|
}
|
||||||
|
if ref.Cast != tc.cast {
|
||||||
|
t.Errorf("Cast = %q, want %q", ref.Cast, tc.cast)
|
||||||
|
}
|
||||||
|
if ref.Alias != tc.alias {
|
||||||
|
t.Errorf("Alias = %q, want %q", ref.Alias, tc.alias)
|
||||||
|
}
|
||||||
|
if ref.Ambiguous != tc.ambiguous {
|
||||||
|
t.Errorf("Ambiguous = %v, want %v", ref.Ambiguous, tc.ambiguous)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseColumnRef_NotJSON(t *testing.T) {
|
||||||
|
// These must return ok=false so callers fall back to their normal handling.
|
||||||
|
inputs := []string{
|
||||||
|
"",
|
||||||
|
" ",
|
||||||
|
"name",
|
||||||
|
"data",
|
||||||
|
"created_at",
|
||||||
|
"(id)",
|
||||||
|
}
|
||||||
|
for _, in := range inputs {
|
||||||
|
if ref, ok := ParseColumnRef(in); ok {
|
||||||
|
t.Errorf("ParseColumnRef(%q) = %+v, ok=true; want ok=false", in, ref)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseColumnRef_Rejected(t *testing.T) {
|
||||||
|
// Malformed or unsafe tokens must be rejected outright.
|
||||||
|
inputs := []string{
|
||||||
|
"data->>'x'::bogus", // cast not on allowlist
|
||||||
|
"data->>'x' AS 1bad", // invalid alias
|
||||||
|
"(data->>'a') OR (x->>'b')", // not a single wrapped expr
|
||||||
|
"data->>'x'); DROP TABLE users; --", // injection attempt
|
||||||
|
"data->b", // unquoted non-numeric key
|
||||||
|
"data->>''", // empty key
|
||||||
|
"data#>>'{}'", // empty hash path
|
||||||
|
"data#>>address", // hash path not a quoted literal
|
||||||
|
"weird col->>'x'", // base not an identifier
|
||||||
|
"data.address.city.but.way.too...deep.", // trailing dot -> empty segment
|
||||||
|
}
|
||||||
|
for _, in := range inputs {
|
||||||
|
if ref, ok := ParseColumnRef(in); ok {
|
||||||
|
t.Errorf("ParseColumnRef(%q) = %+v, ok=true; want rejected", in, ref)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestColumnRef_SQL(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
ref ColumnRef
|
||||||
|
alias string
|
||||||
|
wantExpr string
|
||||||
|
wantArgs []interface{}
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "text extraction qualified",
|
||||||
|
ref: ColumnRef{Base: "data", Path: []string{"address", "city"}, AsText: true},
|
||||||
|
alias: "u",
|
||||||
|
wantExpr: `("u"."data" #>> ?::text[])`,
|
||||||
|
wantArgs: []interface{}{"{address,city}"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "jsonb extraction unqualified",
|
||||||
|
ref: ColumnRef{Base: "data", Path: []string{"a"}, AsText: false},
|
||||||
|
alias: "",
|
||||||
|
wantExpr: `("data" #> ?::text[])`,
|
||||||
|
wantArgs: []interface{}{"{a}"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "with cast",
|
||||||
|
ref: ColumnRef{Base: "data", Path: []string{"age"}, AsText: true, Cast: "integer"},
|
||||||
|
alias: "t",
|
||||||
|
wantExpr: `(("t"."data" #>> ?::text[]))::integer`,
|
||||||
|
wantArgs: []interface{}{"{age}"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "schema qualified alias",
|
||||||
|
ref: ColumnRef{Base: "data", Path: []string{"k"}, AsText: true},
|
||||||
|
alias: "public.users",
|
||||||
|
wantExpr: `("public"."users"."data" #>> ?::text[])`,
|
||||||
|
wantArgs: []interface{}{"{k}"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "key needing quoting",
|
||||||
|
ref: ColumnRef{Base: "data", Path: []string{"key with space", `ev"il`}, AsText: true},
|
||||||
|
alias: "",
|
||||||
|
wantExpr: `("data" #>> ?::text[])`,
|
||||||
|
wantArgs: []interface{}{`{"key with space","ev\"il"}`},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
expr, args := tc.ref.SQL(tc.alias)
|
||||||
|
if expr != tc.wantExpr {
|
||||||
|
t.Errorf("expr = %q, want %q", expr, tc.wantExpr)
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(args, tc.wantArgs) {
|
||||||
|
t.Errorf("args = %#v, want %#v", args, tc.wantArgs)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestColumnRef_SQL_RoundTrip(t *testing.T) {
|
||||||
|
ref, ok := ParseColumnRef("profile->'contact'->>'email'")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("parse failed")
|
||||||
|
}
|
||||||
|
expr, args := ref.SQL("customers")
|
||||||
|
wantExpr := `("customers"."profile" #>> ?::text[])`
|
||||||
|
if expr != wantExpr {
|
||||||
|
t.Errorf("expr = %q, want %q", expr, wantExpr)
|
||||||
|
}
|
||||||
|
if len(args) != 1 || args[0] != "{contact,email}" {
|
||||||
|
t.Errorf("args = %#v, want [{contact,email}]", args)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestColumnRef_OutputAlias(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
ref ColumnRef
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{ColumnRef{Base: "data", Path: []string{"address", "city"}, AsText: true}, "data_address_city"},
|
||||||
|
{ColumnRef{Base: "data", Path: []string{"city"}, Alias: "city"}, "city"},
|
||||||
|
{ColumnRef{Base: "data", Path: []string{"weird key"}}, "data_weird_key"},
|
||||||
|
{ColumnRef{Base: "data"}, "data"},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if got := c.ref.OutputAlias(); got != c.want {
|
||||||
|
t.Errorf("OutputAlias(%+v) = %q, want %q", c.ref, got, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeCastTarget(t *testing.T) {
|
||||||
|
ok := map[string]string{
|
||||||
|
"int": "integer",
|
||||||
|
"INT": "integer",
|
||||||
|
" bigint ": "bigint",
|
||||||
|
"decimal": "numeric",
|
||||||
|
"float8": "double precision",
|
||||||
|
"bool": "boolean",
|
||||||
|
"timestamptz": "timestamptz",
|
||||||
|
"uuid": "uuid",
|
||||||
|
}
|
||||||
|
for in, want := range ok {
|
||||||
|
got, allowed := NormalizeCastTarget(in)
|
||||||
|
if !allowed || got != want {
|
||||||
|
t.Errorf("NormalizeCastTarget(%q) = %q, %v; want %q, true", in, got, allowed, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, in := range []string{"", "regclass", "int; drop", "text[]"} {
|
||||||
|
if got, allowed := NormalizeCastTarget(in); allowed {
|
||||||
|
t.Errorf("NormalizeCastTarget(%q) = %q, true; want not allowed", in, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,219 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
)
|
||||||
|
|
||||||
|
// This file wires the canonical JSON column parser (json_column.go) into the
|
||||||
|
// three query-building paths that every spec handler shares: SELECT column
|
||||||
|
// lists, WHERE filters and ORDER BY. The helpers here are the single place
|
||||||
|
// those paths call so that JSON access is resolved (and made injection-safe)
|
||||||
|
// identically everywhere. They mirror the style of BuildSpatialCondition /
|
||||||
|
// BuildVectorCondition: a boolean ok result tells the caller whether the token
|
||||||
|
// was a JSON reference it should take over, otherwise the caller keeps its
|
||||||
|
// existing (non-JSON) behaviour.
|
||||||
|
|
||||||
|
// jsonComparisonOps are the operators for which a JSON text extraction should be
|
||||||
|
// cast to a concrete type when the value looks numeric — otherwise "10" < "9".
|
||||||
|
var jsonComparisonOps = map[string]bool{
|
||||||
|
"gt": true, "greater_than": true, ">": true,
|
||||||
|
"gte": true, "greater_than_equals": true, "ge": true, ">=": true,
|
||||||
|
"lt": true, "less_than": true, "<": true,
|
||||||
|
"lte": true, "less_than_equals": true, "le": true, "<=": true,
|
||||||
|
"between": true, "between_inclusive": true,
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolveJSONColumnRef parses token and, when it is a usable JSON reference for
|
||||||
|
// model, returns the parsed ColumnRef. For the dotted "a.b" shorthand (which is
|
||||||
|
// otherwise indistinguishable from a table-qualified column) ok is true only
|
||||||
|
// when model confirms the base is a JSON column.
|
||||||
|
func ResolveJSONColumnRef(model interface{}, token string) (ColumnRef, bool) {
|
||||||
|
ref, ok := ParseColumnRef(token)
|
||||||
|
if !ok {
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
if ref.Ambiguous && !reflection.IsJSONColumn(model, ref.Base) {
|
||||||
|
return ColumnRef{}, false
|
||||||
|
}
|
||||||
|
return ref, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsJSONColumnToken reports whether token is a JSON reference this package can
|
||||||
|
// resolve for model (arrow/hash syntax always; dotted shorthand only when the
|
||||||
|
// base is a JSON column).
|
||||||
|
func IsJSONColumnToken(model interface{}, token string) bool {
|
||||||
|
_, ok := ResolveJSONColumnRef(model, token)
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolveJSONColumnExpr resolves a raw column token that traverses into a JSON
|
||||||
|
// value into a parameterised SQL expression plus its args and a deterministic
|
||||||
|
// output alias. ok is false when the token is not a JSON reference, in which
|
||||||
|
// case the caller should handle it the way it did before.
|
||||||
|
//
|
||||||
|
// tableAlias, when non-empty, qualifies the base column.
|
||||||
|
func ResolveJSONColumnExpr(model interface{}, tableAlias, token string) (expr string, args []interface{}, alias string, ok bool) {
|
||||||
|
ref, ok := ResolveJSONColumnRef(model, token)
|
||||||
|
if !ok {
|
||||||
|
return "", nil, "", false
|
||||||
|
}
|
||||||
|
expr, args = ref.SQL(tableAlias)
|
||||||
|
return expr, args, ref.OutputAlias(), true
|
||||||
|
}
|
||||||
|
|
||||||
|
// ApplySelectColumns adds the requested columns to query, resolving any that are
|
||||||
|
// JSON sub-field references (data->>'x', data#>>'{a,b}', or the dotted data.x
|
||||||
|
// shorthand for a JSON column) into safe parameterised expressions with a
|
||||||
|
// deterministic alias. Plain columns are passed through reflection.ExtractSourceColumn
|
||||||
|
// exactly as before. tableAlias, when non-empty, qualifies JSON base columns.
|
||||||
|
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
|
||||||
|
}
|
||||||
|
query = query.Column(reflection.ExtractSourceColumn(col))
|
||||||
|
}
|
||||||
|
return query
|
||||||
|
}
|
||||||
|
|
||||||
|
// BuildJSONFilterCondition builds a complete WHERE condition for a JSON column
|
||||||
|
// token. ok is false when the token is not a JSON reference or the operator is
|
||||||
|
// not one this builder handles (the caller then keeps its existing behaviour).
|
||||||
|
//
|
||||||
|
// The JSON path is always bound as a parameter, never interpolated. When the
|
||||||
|
// reference carries no explicit ::cast and the operator is an ordered
|
||||||
|
// comparison against a numeric value, the extracted text is cast to numeric so
|
||||||
|
// the comparison is numeric rather than lexical.
|
||||||
|
func BuildJSONFilterCondition(model interface{}, tableAlias, token, operator string, value interface{}) (condition string, args []interface{}, ok bool) {
|
||||||
|
ref, ok := ResolveJSONColumnRef(model, token)
|
||||||
|
if !ok {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
op := strings.ToLower(strings.TrimSpace(operator))
|
||||||
|
|
||||||
|
// Infer a cast for ordered comparisons on numeric values so "10" > "9".
|
||||||
|
if ref.Cast == "" && jsonComparisonOps[op] && jsonValueIsNumeric(value) {
|
||||||
|
ref.Cast = "numeric"
|
||||||
|
}
|
||||||
|
|
||||||
|
colExpr, colArgs := ref.SQL(tableAlias)
|
||||||
|
|
||||||
|
// prepend copies the column-expression args (the bound JSON path, and any
|
||||||
|
// others) ahead of the value args so placeholder order matches the SQL.
|
||||||
|
prepend := func(valueArgs ...interface{}) []interface{} {
|
||||||
|
out := make([]interface{}, 0, len(colArgs)+len(valueArgs))
|
||||||
|
out = append(out, colArgs...)
|
||||||
|
out = append(out, valueArgs...)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
switch op {
|
||||||
|
case "eq", "equals", "=":
|
||||||
|
return fmt.Sprintf("%s = ?", colExpr), prepend(value), true
|
||||||
|
case "neq", "not_equals", "ne", "!=", "<>":
|
||||||
|
return fmt.Sprintf("%s != ?", colExpr), prepend(value), true
|
||||||
|
case "gt", "greater_than", ">":
|
||||||
|
return fmt.Sprintf("%s > ?", colExpr), prepend(value), true
|
||||||
|
case "gte", "greater_than_equals", "ge", ">=":
|
||||||
|
return fmt.Sprintf("%s >= ?", colExpr), prepend(value), true
|
||||||
|
case "lt", "less_than", "<":
|
||||||
|
return fmt.Sprintf("%s < ?", colExpr), prepend(value), true
|
||||||
|
case "lte", "less_than_equals", "le", "<=":
|
||||||
|
return fmt.Sprintf("%s <= ?", colExpr), prepend(value), true
|
||||||
|
case "like":
|
||||||
|
return fmt.Sprintf("%s LIKE ?", colExpr), prepend(value), true
|
||||||
|
case "ilike":
|
||||||
|
return fmt.Sprintf("%s ILIKE ?", colExpr), prepend(value), true
|
||||||
|
case "in":
|
||||||
|
inCond, inArgs := BuildInCondition(colExpr, value)
|
||||||
|
if inCond == "" {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
return inCond, prepend(inArgs...), true
|
||||||
|
case "between", "between_inclusive":
|
||||||
|
lo, hi, bok := twoBoundValues(value)
|
||||||
|
if !bok {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
loOp, hiOp := ">", "<"
|
||||||
|
if op == "between_inclusive" {
|
||||||
|
loOp, hiOp = ">=", "<="
|
||||||
|
}
|
||||||
|
// colExpr appears twice, so its bound args (the JSON path) appear twice.
|
||||||
|
betweenArgs := make([]interface{}, 0, 2*len(colArgs)+2)
|
||||||
|
betweenArgs = append(betweenArgs, colArgs...)
|
||||||
|
betweenArgs = append(betweenArgs, lo)
|
||||||
|
betweenArgs = append(betweenArgs, colArgs...)
|
||||||
|
betweenArgs = append(betweenArgs, hi)
|
||||||
|
return fmt.Sprintf("(%s %s ? AND %s %s ?)", colExpr, loOp, colExpr, hiOp), betweenArgs, true
|
||||||
|
case "is_null", "isnull":
|
||||||
|
return fmt.Sprintf("%s IS NULL", colExpr), prepend(), true
|
||||||
|
case "is_not_null", "isnotnull":
|
||||||
|
return fmt.Sprintf("%s IS NOT NULL", colExpr), prepend(), true
|
||||||
|
default:
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// jsonValueIsNumeric reports whether value (or every element of a 2-slice) is a
|
||||||
|
// number or a numeric-looking string.
|
||||||
|
func jsonValueIsNumeric(value interface{}) bool {
|
||||||
|
switch v := value.(type) {
|
||||||
|
case []interface{}:
|
||||||
|
if len(v) == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, e := range v {
|
||||||
|
if !jsonValueIsNumeric(e) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
case []string:
|
||||||
|
if len(v) == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, e := range v {
|
||||||
|
if _, ok := toFloat(e); !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
case string:
|
||||||
|
_, ok := toFloat(v)
|
||||||
|
return ok
|
||||||
|
default:
|
||||||
|
_, ok := toFloat(value)
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// twoBoundValues extracts the low/high bounds from a BETWEEN filter value.
|
||||||
|
func twoBoundValues(value interface{}) (lo, hi interface{}, ok bool) {
|
||||||
|
switch v := value.(type) {
|
||||||
|
case []interface{}:
|
||||||
|
if len(v) == 2 {
|
||||||
|
return v[0], v[1], true
|
||||||
|
}
|
||||||
|
case []string:
|
||||||
|
if len(v) == 2 {
|
||||||
|
return v[0], v[1], true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil, nil, false
|
||||||
|
}
|
||||||
@@ -0,0 +1,205 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||||
|
)
|
||||||
|
|
||||||
|
type jsonCondModel struct {
|
||||||
|
ID int64 `json:"id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Data spectypes.SqlJSONB `json:"data"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveJSONColumnRef_Gate(t *testing.T) {
|
||||||
|
m := jsonCondModel{}
|
||||||
|
|
||||||
|
// Explicit operator syntax needs no model confirmation.
|
||||||
|
if _, ok := ResolveJSONColumnRef(m, "data->>'city'"); !ok {
|
||||||
|
t.Error("arrow syntax should resolve")
|
||||||
|
}
|
||||||
|
// Dotted shorthand on a real JSON column resolves.
|
||||||
|
if ref, ok := ResolveJSONColumnRef(m, "data.city"); !ok || !reflect.DeepEqual(ref.Path, []string{"city"}) {
|
||||||
|
t.Errorf("dotted shorthand on JSON column should resolve, got ok=%v ref=%+v", ok, ref)
|
||||||
|
}
|
||||||
|
// Dotted shorthand on a non-JSON column must NOT be treated as JSON.
|
||||||
|
if _, ok := ResolveJSONColumnRef(m, "name.first"); ok {
|
||||||
|
t.Error("dotted shorthand on non-JSON column must not resolve as JSON")
|
||||||
|
}
|
||||||
|
// Plain columns never resolve.
|
||||||
|
if _, ok := ResolveJSONColumnRef(m, "name"); ok {
|
||||||
|
t.Error("plain column must not resolve")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveJSONColumnExpr(t *testing.T) {
|
||||||
|
m := jsonCondModel{}
|
||||||
|
|
||||||
|
expr, args, alias, ok := ResolveJSONColumnExpr(m, "t", "data->'addr'->>'city'")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected ok")
|
||||||
|
}
|
||||||
|
if expr != `("t"."data" #>> ?::text[])` {
|
||||||
|
t.Errorf("expr = %q", expr)
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(args, []interface{}{"{addr,city}"}) {
|
||||||
|
t.Errorf("args = %#v", args)
|
||||||
|
}
|
||||||
|
if alias != "data_addr_city" {
|
||||||
|
t.Errorf("alias = %q", alias)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, _, _, ok := ResolveJSONColumnExpr(m, "t", "name"); ok {
|
||||||
|
t.Error("plain column must not resolve")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildJSONFilterCondition(t *testing.T) {
|
||||||
|
m := jsonCondModel{}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
token string
|
||||||
|
operator string
|
||||||
|
value interface{}
|
||||||
|
wantCond string
|
||||||
|
wantArgs []interface{}
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "eq stays text", token: "data->>'city'", operator: "eq", value: "LA",
|
||||||
|
wantCond: `("data" #>> ?::text[]) = ?`,
|
||||||
|
wantArgs: []interface{}{"{city}", "LA"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "gt numeric value infers numeric cast", token: "data->>'age'", operator: "gt", value: 18,
|
||||||
|
wantCond: `(("data" #>> ?::text[]))::numeric > ?`,
|
||||||
|
wantArgs: []interface{}{"{age}", 18},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "gt non-numeric value stays text", token: "data->>'name'", operator: "gt", value: "m",
|
||||||
|
wantCond: `("data" #>> ?::text[]) > ?`,
|
||||||
|
wantArgs: []interface{}{"{name}", "m"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "explicit cast is respected for lt", token: "data->>'ts'::timestamptz", operator: "lt", value: "2020-01-01",
|
||||||
|
wantCond: `(("data" #>> ?::text[]))::timestamptz < ?`,
|
||||||
|
wantArgs: []interface{}{"{ts}", "2020-01-01"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "ilike", token: "data->>'city'", operator: "ilike", value: "%la%",
|
||||||
|
wantCond: `("data" #>> ?::text[]) ILIKE ?`,
|
||||||
|
wantArgs: []interface{}{"{city}", "%la%"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "in", token: "data->>'tier'", operator: "in", value: []string{"a", "b"},
|
||||||
|
wantCond: `("data" #>> ?::text[]) IN (?,?)`,
|
||||||
|
wantArgs: []interface{}{"{tier}", "a", "b"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "between numeric", token: "data->>'age'", operator: "between", value: []interface{}{10, 20},
|
||||||
|
wantCond: `((("data" #>> ?::text[]))::numeric > ? AND (("data" #>> ?::text[]))::numeric < ?)`,
|
||||||
|
wantArgs: []interface{}{"{age}", 10, "{age}", 20},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "is_null", token: "data->>'city'", operator: "is_null", value: nil,
|
||||||
|
wantCond: `("data" #>> ?::text[]) IS NULL`,
|
||||||
|
wantArgs: []interface{}{"{city}"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "hash path", token: "data#>>'{a,b}'", operator: "eq", value: "x",
|
||||||
|
wantCond: `("data" #>> ?::text[]) = ?`,
|
||||||
|
wantArgs: []interface{}{"{a,b}", "x"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "dotted shorthand on json column", token: "data.city", operator: "eq", value: "x",
|
||||||
|
wantCond: `("data" #>> ?::text[]) = ?`,
|
||||||
|
wantArgs: []interface{}{"{city}", "x"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
cond, args, ok := BuildJSONFilterCondition(m, "", tc.token, tc.operator, tc.value)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("ok=false for %q", tc.token)
|
||||||
|
}
|
||||||
|
if cond != tc.wantCond {
|
||||||
|
t.Errorf("cond = %q, want %q", cond, tc.wantCond)
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(args, tc.wantArgs) {
|
||||||
|
t.Errorf("args = %#v, want %#v", args, tc.wantArgs)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildJSONFilterCondition_NotJSON(t *testing.T) {
|
||||||
|
m := jsonCondModel{}
|
||||||
|
for _, tok := range []string{"name", "id", "name.first"} {
|
||||||
|
if _, _, ok := BuildJSONFilterCondition(m, "", tok, "eq", "x"); ok {
|
||||||
|
t.Errorf("BuildJSONFilterCondition(%q) ok=true, want false", tok)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Unknown operator on a real JSON ref -> caller keeps its own handling.
|
||||||
|
if _, _, ok := BuildJSONFilterCondition(m, "", "data->>'x'", "st_intersects", "y"); ok {
|
||||||
|
t.Error("unknown operator must yield ok=false")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildJSONFilterCondition_QualifiedAndInjectionSafe(t *testing.T) {
|
||||||
|
m := jsonCondModel{}
|
||||||
|
// A hostile key never reaches the SQL string — it is bound in the text[] arg.
|
||||||
|
cond, args, ok := BuildJSONFilterCondition(m, "pub.tbl", "data->>'ev\"il'", "eq", "x")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("ok=false")
|
||||||
|
}
|
||||||
|
if cond != `("pub"."tbl"."data" #>> ?::text[]) = ?` {
|
||||||
|
t.Errorf("cond = %q", cond)
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(args, []interface{}{`{"ev\"il"}`, "x"}) {
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -69,7 +69,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
|||||||
|
|
||||||
// Get model type for reflection
|
// Get model type for reflection
|
||||||
modelType := reflect.TypeOf(model)
|
modelType := reflect.TypeOf(model)
|
||||||
for modelType != nil && (modelType.Kind() == reflect.Ptr || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
modelType = modelType.Elem()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -113,7 +113,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
|||||||
|
|
||||||
// Process based on operation
|
// Process based on operation
|
||||||
switch strings.ToLower(operation) {
|
switch strings.ToLower(operation) {
|
||||||
case "insert", "create":
|
case "insert", "create", "add":
|
||||||
// Only perform insert if we have data to insert
|
// Only perform insert if we have data to insert
|
||||||
if hasData {
|
if hasData {
|
||||||
id, err := p.processInsert(ctx, regularData, tableName)
|
id, err := p.processInsert(ctx, regularData, tableName)
|
||||||
@@ -125,6 +125,13 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
|||||||
result.AffectedRows = 1
|
result.AffectedRows = 1
|
||||||
result.Data = regularData
|
result.Data = regularData
|
||||||
|
|
||||||
|
// Re-select the inserted row so result.Data reflects DB-generated defaults.
|
||||||
|
if row, err := p.processSelect(ctx, tableName, id); err != nil {
|
||||||
|
logger.Warn("Select after insert failed: table=%s, id=%v, error=%v", tableName, id, err)
|
||||||
|
} else if len(row) > 0 {
|
||||||
|
result.Data = row
|
||||||
|
}
|
||||||
|
|
||||||
// Process child relations after parent insert (to get parent ID)
|
// Process child relations after parent insert (to get parent ID)
|
||||||
if err := p.processChildRelations(ctx, "insert", id, relationFields, result.RelationData, modelType, parentIDs); err != nil {
|
if err := p.processChildRelations(ctx, "insert", id, relationFields, result.RelationData, modelType, parentIDs); err != nil {
|
||||||
logger.Error("Failed to process child relations after insert: table=%s, parentID=%v, relations=%+v, error=%v", tableName, id, relationFields, err)
|
logger.Error("Failed to process child relations after insert: table=%s, parentID=%v, relations=%+v, error=%v", tableName, id, relationFields, err)
|
||||||
@@ -134,8 +141,12 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
|||||||
logger.Debug("Skipping insert for %s - no data columns besides _request", tableName)
|
logger.Debug("Skipping insert for %s - no data columns besides _request", tableName)
|
||||||
}
|
}
|
||||||
|
|
||||||
case "update":
|
case "update", "change", "modify":
|
||||||
// Only perform update if we have data to update
|
// Only perform update if we have data to update
|
||||||
|
if reflection.IsEmptyValue(data[pkName]) {
|
||||||
|
logger.Warn("Skipping update for %s - no primary key", tableName)
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
if hasData {
|
if hasData {
|
||||||
rows, err := p.processUpdate(ctx, regularData, tableName, data[pkName])
|
rows, err := p.processUpdate(ctx, regularData, tableName, data[pkName])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -146,9 +157,16 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
|||||||
result.AffectedRows = rows
|
result.AffectedRows = rows
|
||||||
result.Data = regularData
|
result.Data = regularData
|
||||||
|
|
||||||
|
// Re-select the updated row so result.Data reflects current DB state.
|
||||||
|
if row, err := p.processSelect(ctx, tableName, result.ID); err != nil {
|
||||||
|
logger.Warn("Select after update failed: table=%s, id=%v, error=%v", tableName, result.ID, err)
|
||||||
|
} else if len(row) > 0 {
|
||||||
|
result.Data = row
|
||||||
|
}
|
||||||
|
|
||||||
// Process child relations for update
|
// Process child relations for update
|
||||||
if err := p.processChildRelations(ctx, "update", data[pkName], relationFields, result.RelationData, modelType, parentIDs); err != nil {
|
if err := p.processChildRelations(ctx, "update", data[pkName], relationFields, result.RelationData, modelType, parentIDs); err != nil {
|
||||||
logger.Error("Failed to process child relations after update: table=%s, parentID=%v, relations=%+v, error=%v", tableName, data[pkName], relationFields, err)
|
logger.Error("Failed to process child relations after update: table=%s, parentID=%v, relations=%+v, error=%v", tableName, data[pkName], regularData, err)
|
||||||
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
@@ -156,11 +174,16 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
|||||||
result.ID = data[pkName]
|
result.ID = data[pkName]
|
||||||
}
|
}
|
||||||
|
|
||||||
case "delete":
|
case "delete", "remove":
|
||||||
|
if reflection.IsEmptyValue(data[pkName]) {
|
||||||
|
logger.Warn("Skipping delete for %s - no primary key", tableName)
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
// Process child relations first (for referential integrity)
|
// Process child relations first (for referential integrity)
|
||||||
if err := p.processChildRelations(ctx, "delete", data[pkName], relationFields, result.RelationData, modelType, parentIDs); err != nil {
|
if err := p.processChildRelations(ctx, "delete", data[pkName], relationFields, result.RelationData, modelType, parentIDs); err != nil {
|
||||||
logger.Error("Failed to process child relations before delete: table=%s, id=%v, relations=%+v, error=%v", tableName, data[pkName], relationFields, err)
|
logger.Error("Failed to process child relations before delete: table=%s, id=%v, relations=%+v, error=%v", tableName, data[pkName], relationFields, err)
|
||||||
return nil, fmt.Errorf("failed to process child relations before delete: %w", err)
|
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
rows, err := p.processDelete(ctx, tableName, data[pkName])
|
rows, err := p.processDelete(ctx, tableName, data[pkName])
|
||||||
@@ -201,7 +224,7 @@ func (p *NestedCUDProcessor) filterValidFields(data map[string]interface{}, mode
|
|||||||
}
|
}
|
||||||
|
|
||||||
modelType := reflect.TypeOf(model)
|
modelType := reflect.TypeOf(model)
|
||||||
for modelType != nil && (modelType.Kind() == reflect.Ptr || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
modelType = modelType.Elem()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -234,28 +257,38 @@ func (p *NestedCUDProcessor) injectForeignKeys(data map[string]interface{}, mode
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Iterate through model fields to find foreign key fields
|
pkCol := reflection.GetPrimaryKeyName(reflect.New(modelType).Interface())
|
||||||
|
|
||||||
|
for parentKey, parentID := range parentIDs {
|
||||||
|
dbColNames := reflection.GetForeignKeyColumn(modelType, parentKey)
|
||||||
|
|
||||||
|
if len(dbColNames) == 0 {
|
||||||
|
// No explicit tag found — fall back to naming convention by scanning scalar fields.
|
||||||
for i := 0; i < modelType.NumField(); i++ {
|
for i := 0; i < modelType.NumField(); i++ {
|
||||||
field := modelType.Field(i)
|
field := modelType.Field(i)
|
||||||
jsonTag := field.Tag.Get("json")
|
jsonName := strings.Split(field.Tag.Get("json"), ",")[0]
|
||||||
jsonName := strings.Split(jsonTag, ",")[0]
|
if strings.EqualFold(jsonName, "rid"+parentKey) ||
|
||||||
|
strings.EqualFold(jsonName, "rid_"+parentKey) ||
|
||||||
// Check if this field is a foreign key and we have a parent ID for it
|
strings.EqualFold(jsonName, "id_"+parentKey) ||
|
||||||
// Common patterns: DepartmentID, ManagerID, ProjectID, etc.
|
strings.EqualFold(jsonName, parentKey+"_id") ||
|
||||||
for parentKey, parentID := range parentIDs {
|
|
||||||
// Match field name patterns like "department_id" with parent key "department"
|
|
||||||
if strings.EqualFold(jsonName, parentKey+"_id") ||
|
|
||||||
strings.EqualFold(jsonName, parentKey+"id") ||
|
strings.EqualFold(jsonName, parentKey+"id") ||
|
||||||
strings.EqualFold(field.Name, parentKey+"ID") {
|
strings.EqualFold(field.Name, parentKey+"ID") {
|
||||||
// Use the DB column name as the key, since data is keyed by DB column names
|
dbColNames = []string{reflection.GetColumnName(field)}
|
||||||
dbColName := reflection.GetColumnName(field)
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, dbColName := range dbColNames {
|
||||||
|
if pkCol != "" && strings.EqualFold(dbColName, pkCol) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
if _, exists := data[dbColName]; !exists {
|
if _, exists := data[dbColName]; !exists {
|
||||||
logger.Debug("Injecting foreign key: %s = %v", dbColName, parentID)
|
logger.Debug("Injecting foreign key: %s = %v", dbColName, parentID)
|
||||||
data[dbColName] = parentID
|
data[dbColName] = parentID
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// processInsert handles insert operation
|
// processInsert handles insert operation
|
||||||
@@ -269,30 +302,35 @@ func (p *NestedCUDProcessor) processInsert(
|
|||||||
query := p.db.NewInsert().Table(tableName)
|
query := p.db.NewInsert().Table(tableName)
|
||||||
|
|
||||||
for key, value := range data {
|
for key, value := range data {
|
||||||
query = query.Value(key, value)
|
query = query.Value(key, ConvertSliceForBun(value))
|
||||||
}
|
}
|
||||||
pkName := reflection.GetPrimaryKeyName(tableName)
|
pkName := reflection.GetPrimaryKeyName(tableName)
|
||||||
// Add RETURNING clause to get the inserted ID
|
|
||||||
query = query.Returning(pkName)
|
query = query.Returning(pkName)
|
||||||
|
|
||||||
result, err := query.Exec(ctx)
|
var id interface{}
|
||||||
if err != nil {
|
if err := query.Scan(ctx, &id); err != nil {
|
||||||
logger.Error("Insert execution failed: table=%s, data=%+v, error=%v", tableName, data, err)
|
logger.Error("Insert execution failed: table=%s, data=%+v, error=%v", tableName, data, err)
|
||||||
return nil, fmt.Errorf("insert exec failed: %w", err)
|
return nil, fmt.Errorf("insert exec failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Try to get the ID
|
logger.Debug("Insert successful, ID: %v", id)
|
||||||
var id interface{}
|
|
||||||
if lastID, err := result.LastInsertId(); err == nil && lastID > 0 {
|
|
||||||
id = lastID
|
|
||||||
} else if data[pkName] != nil {
|
|
||||||
id = data[pkName]
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.Debug("Insert successful, ID: %v, rows affected: %d", id, result.RowsAffected())
|
|
||||||
return id, nil
|
return id, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// processSelect fetches the row identified by id from tableName into a flat map.
|
||||||
|
// Used to populate result.Data with the actual DB state after insert/update.
|
||||||
|
func (p *NestedCUDProcessor) processSelect(ctx context.Context, tableName string, id interface{}) (map[string]interface{}, error) {
|
||||||
|
pkName := reflection.GetPrimaryKeyName(tableName)
|
||||||
|
var row map[string]interface{}
|
||||||
|
if err := p.db.NewSelect().
|
||||||
|
Table(tableName).
|
||||||
|
Where(fmt.Sprintf("%s = ?", QuoteIdent(pkName)), id).
|
||||||
|
Scan(ctx, &row); err != nil {
|
||||||
|
return nil, fmt.Errorf("select after write failed: %w", err)
|
||||||
|
}
|
||||||
|
return row, nil
|
||||||
|
}
|
||||||
|
|
||||||
// processUpdate handles update operation
|
// processUpdate handles update operation
|
||||||
func (p *NestedCUDProcessor) processUpdate(
|
func (p *NestedCUDProcessor) processUpdate(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
@@ -372,7 +410,7 @@ func (p *NestedCUDProcessor) processChildRelations(
|
|||||||
if relatedModelType.Kind() == reflect.Slice {
|
if relatedModelType.Kind() == reflect.Slice {
|
||||||
relatedModelType = relatedModelType.Elem()
|
relatedModelType = relatedModelType.Elem()
|
||||||
}
|
}
|
||||||
if relatedModelType.Kind() == reflect.Ptr {
|
if relatedModelType.Kind() == reflect.Pointer {
|
||||||
relatedModelType = relatedModelType.Elem()
|
relatedModelType = relatedModelType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -433,13 +471,17 @@ func (p *NestedCUDProcessor) processChildRelations(
|
|||||||
// Priority: Use foreign key field name if specified
|
// Priority: Use foreign key field name if specified
|
||||||
var foreignKeyFieldName string
|
var foreignKeyFieldName string
|
||||||
if relInfo.ForeignKey != "" {
|
if relInfo.ForeignKey != "" {
|
||||||
// Get the JSON name for the foreign key field in the child model
|
// For has-many/has-one: join:parentCol=childCol
|
||||||
foreignKeyFieldName = reflection.GetJSONNameForField(relatedModelType, relInfo.ForeignKey)
|
// ForeignKey = parent side, References = child side (where we actually set the value)
|
||||||
if foreignKeyFieldName == "" {
|
childField := relInfo.ForeignKey
|
||||||
// Fallback to lowercase field name
|
if (relInfo.RelationType == "hasMany" || relInfo.RelationType == "hasOne") && relInfo.References != "" {
|
||||||
foreignKeyFieldName = strings.ToLower(relInfo.ForeignKey)
|
childField = relInfo.References
|
||||||
}
|
}
|
||||||
logger.Debug("Using foreign key field for direct assignment: %s (from FK %s)", foreignKeyFieldName, relInfo.ForeignKey)
|
foreignKeyFieldName = reflection.GetJSONNameForField(relatedModelType, childField)
|
||||||
|
if foreignKeyFieldName == "" {
|
||||||
|
foreignKeyFieldName = strings.ToLower(childField)
|
||||||
|
}
|
||||||
|
logger.Debug("Using foreign key field for direct assignment: %s (from FK %s -> child %s)", foreignKeyFieldName, relInfo.ForeignKey, childField)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get the primary key name for the child model to avoid overwriting it in recursive relationships
|
// Get the primary key name for the child model to avoid overwriting it in recursive relationships
|
||||||
@@ -548,7 +590,7 @@ func shouldUseNestedProcessorDepth(data map[string]interface{}, model interface{
|
|||||||
|
|
||||||
// Get model type
|
// Get model type
|
||||||
modelType := reflect.TypeOf(model)
|
modelType := reflect.TypeOf(model)
|
||||||
for modelType != nil && (modelType.Kind() == reflect.Ptr || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
modelType = modelType.Elem()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -101,12 +101,18 @@ func (m *mockInsertQuery) Value(column string, value interface{}) InsertQuery {
|
|||||||
func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m }
|
func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m }
|
||||||
func (m *mockInsertQuery) Returning(columns ...string) InsertQuery { return m }
|
func (m *mockInsertQuery) Returning(columns ...string) InsertQuery { return m }
|
||||||
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
|
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
|
||||||
// Record the insert call
|
|
||||||
m.db.insertCalls = append(m.db.insertCalls, m.values)
|
m.db.insertCalls = append(m.db.insertCalls, m.values)
|
||||||
m.db.lastID++
|
m.db.lastID++
|
||||||
return &mockResult{lastID: m.db.lastID, rowsAffected: 1}, nil
|
return &mockResult{lastID: m.db.lastID, rowsAffected: 1}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (m *mockInsertQuery) Scan(ctx context.Context, dest interface{}) error {
|
||||||
|
m.db.insertCalls = append(m.db.insertCalls, m.values)
|
||||||
|
m.db.lastID++
|
||||||
|
reflect.ValueOf(dest).Elem().Set(reflect.ValueOf(m.db.lastID))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Mock UpdateQuery
|
// Mock UpdateQuery
|
||||||
type mockUpdateQuery struct {
|
type mockUpdateQuery struct {
|
||||||
db *mockDatabase
|
db *mockDatabase
|
||||||
@@ -707,6 +713,220 @@ func TestInjectForeignKeys(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Models for asymmetric join column tests (mirrors the bun has-many join:parentCol=childCol pattern).
|
||||||
|
// ActionOption has-many ActionOptionLinks via join:rid_actionoption=rid_actionoption_child.
|
||||||
|
// The child column ("rid_actionoption_child") differs from the parent column ("rid_actionoption").
|
||||||
|
type ActionOption struct {
|
||||||
|
RidActionoption int64 `json:"rid_actionoption" bun:"rid_actionoption,pk"`
|
||||||
|
Label string `json:"label"`
|
||||||
|
Links []*ActionOptionLink `json:"aol_rid_actionoption_child,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a ActionOption) TableName() string { return "action_options" }
|
||||||
|
func (a ActionOption) GetIDName() string { return "RidActionoption" }
|
||||||
|
|
||||||
|
type ActionOptionLink struct {
|
||||||
|
RidActionoptionlink int64 `json:"rid_actionoptionlink" bun:"rid_actionoptionlink,pk"`
|
||||||
|
RidActionoptionChild int64 `json:"rid_actionoption_child" bun:"rid_actionoption_child"`
|
||||||
|
Label string `json:"label"`
|
||||||
|
// Note: no field named "rid_actionoption" — that is the parent's column.
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a ActionOptionLink) TableName() string { return "action_option_links" }
|
||||||
|
func (a ActionOptionLink) GetIDName() string { return "RidActionoptionlink" }
|
||||||
|
|
||||||
|
// TestProcessNestedCUD_AsymmetricJoinColumns verifies that for a has-many relation with
|
||||||
|
// join:parentCol=childCol, the child rows are stamped with the child-side column (References),
|
||||||
|
// not the parent-side column (ForeignKey).
|
||||||
|
func TestProcessNestedCUD_AsymmetricJoinColumns(t *testing.T) {
|
||||||
|
db := newMockDatabase()
|
||||||
|
registry := &mockModelRegistry{}
|
||||||
|
relProvider := newMockRelationshipProvider()
|
||||||
|
|
||||||
|
// Mirrors: bun:"rel:has-many,join:rid_actionoption=rid_actionoption_child"
|
||||||
|
relProvider.RegisterRelation("ActionOption", "aol_rid_actionoption_child", &RelationshipInfo{
|
||||||
|
FieldName: "Links",
|
||||||
|
JSONName: "aol_rid_actionoption_child",
|
||||||
|
RelationType: "hasMany",
|
||||||
|
ForeignKey: "rid_actionoption", // parent-side column (left of join:)
|
||||||
|
References: "rid_actionoption_child", // child-side column (right of join:)
|
||||||
|
RelatedModel: ActionOptionLink{},
|
||||||
|
})
|
||||||
|
|
||||||
|
processor := NewNestedCUDProcessor(db, registry, relProvider)
|
||||||
|
|
||||||
|
data := map[string]interface{}{
|
||||||
|
"label": "option-a",
|
||||||
|
"aol_rid_actionoption_child": []interface{}{
|
||||||
|
map[string]interface{}{"label": "link-1"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := processor.ProcessNestedCUD(
|
||||||
|
context.Background(),
|
||||||
|
"insert",
|
||||||
|
data,
|
||||||
|
ActionOption{},
|
||||||
|
nil,
|
||||||
|
"action_options",
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ProcessNestedCUD failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(db.insertCalls) < 2 {
|
||||||
|
t.Fatalf("Expected at least 2 insert calls (parent + child), got %d", len(db.insertCalls))
|
||||||
|
}
|
||||||
|
|
||||||
|
childInsert := db.insertCalls[1]
|
||||||
|
|
||||||
|
// The fix: child must receive "rid_actionoption_child", NOT "rid_actionoption".
|
||||||
|
if childInsert["rid_actionoption_child"] == nil {
|
||||||
|
t.Error("Expected child to have rid_actionoption_child set (child-side FK column)")
|
||||||
|
}
|
||||||
|
if childInsert["rid_actionoption"] != nil {
|
||||||
|
t.Errorf("Child must not receive parent-side column rid_actionoption, got %v", childInsert["rid_actionoption"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestProcessNestedCUD_BelongsToUnchanged verifies that the fix does not regress belongsTo
|
||||||
|
// relations, where ForeignKey is already the local (child) column.
|
||||||
|
func TestProcessNestedCUD_BelongsToUnchanged(t *testing.T) {
|
||||||
|
db := newMockDatabase()
|
||||||
|
registry := &mockModelRegistry{}
|
||||||
|
relProvider := newMockRelationshipProvider()
|
||||||
|
|
||||||
|
// For belongsTo, ForeignKey is the column on the child; References is on the parent.
|
||||||
|
// The old and new code must behave identically here.
|
||||||
|
relProvider.RegisterRelation("Employee", "department", &RelationshipInfo{
|
||||||
|
FieldName: "Department",
|
||||||
|
JSONName: "department",
|
||||||
|
RelationType: "belongsTo",
|
||||||
|
ForeignKey: "DepartmentID", // child's own column
|
||||||
|
References: "ID", // parent's PK
|
||||||
|
RelatedModel: Department{},
|
||||||
|
})
|
||||||
|
relProvider.RegisterRelation("Department", "employees", &RelationshipInfo{
|
||||||
|
FieldName: "Employees",
|
||||||
|
JSONName: "employees",
|
||||||
|
RelationType: "has_many",
|
||||||
|
ForeignKey: "DepartmentID",
|
||||||
|
RelatedModel: Employee{},
|
||||||
|
})
|
||||||
|
|
||||||
|
processor := NewNestedCUDProcessor(db, registry, relProvider)
|
||||||
|
|
||||||
|
data := map[string]interface{}{
|
||||||
|
"name": "Engineering",
|
||||||
|
"employees": []interface{}{
|
||||||
|
map[string]interface{}{"name": "Alice"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := processor.ProcessNestedCUD(
|
||||||
|
context.Background(),
|
||||||
|
"insert",
|
||||||
|
data,
|
||||||
|
Department{},
|
||||||
|
nil,
|
||||||
|
"departments",
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ProcessNestedCUD failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(db.insertCalls) < 2 {
|
||||||
|
t.Fatalf("Expected at least 2 inserts, got %d", len(db.insertCalls))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Employees relation uses has_many (old-style) so it goes through the parentIDs injection path,
|
||||||
|
// not the foreignKeyFieldName path. Just confirm no panic and employee is inserted.
|
||||||
|
if db.insertCalls[0]["name"] != "Engineering" {
|
||||||
|
t.Errorf("Expected department name 'Engineering', got %v", db.insertCalls[0]["name"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcessNestedCUD_AddAlias(t *testing.T) {
|
||||||
|
db := newMockDatabase()
|
||||||
|
registry := &mockModelRegistry{}
|
||||||
|
relProvider := newMockRelationshipProvider()
|
||||||
|
|
||||||
|
processor := NewNestedCUDProcessor(db, registry, relProvider)
|
||||||
|
|
||||||
|
data := map[string]interface{}{
|
||||||
|
"_request": "add",
|
||||||
|
"name": "New Department",
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := processor.ProcessNestedCUD(context.Background(), "insert", data, Department{}, nil, "departments")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ProcessNestedCUD with _request=add failed: %v", err)
|
||||||
|
}
|
||||||
|
if result.ID == nil {
|
||||||
|
t.Error("Expected result.ID to be set after add")
|
||||||
|
}
|
||||||
|
if len(db.insertCalls) != 1 {
|
||||||
|
t.Errorf("Expected 1 insert call, got %d", len(db.insertCalls))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcessNestedCUD_RemoveAlias(t *testing.T) {
|
||||||
|
db := newMockDatabase()
|
||||||
|
registry := &mockModelRegistry{}
|
||||||
|
relProvider := newMockRelationshipProvider()
|
||||||
|
|
||||||
|
processor := NewNestedCUDProcessor(db, registry, relProvider)
|
||||||
|
|
||||||
|
data := map[string]interface{}{
|
||||||
|
"_request": "remove",
|
||||||
|
"ID": int64(42),
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := processor.ProcessNestedCUD(context.Background(), "delete", data, Department{}, nil, "departments")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ProcessNestedCUD with _request=remove failed: %v", err)
|
||||||
|
}
|
||||||
|
if len(db.deleteCalls) != 1 {
|
||||||
|
t.Errorf("Expected 1 delete call, got %d", len(db.deleteCalls))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcessNestedCUD_NestedAddRemoveAliases(t *testing.T) {
|
||||||
|
db := newMockDatabase()
|
||||||
|
registry := &mockModelRegistry{}
|
||||||
|
relProvider := newMockRelationshipProvider()
|
||||||
|
|
||||||
|
relProvider.RegisterRelation("Department", "employees", &RelationshipInfo{
|
||||||
|
FieldName: "Employees",
|
||||||
|
JSONName: "employees",
|
||||||
|
RelationType: "has_many",
|
||||||
|
ForeignKey: "DepartmentID",
|
||||||
|
RelatedModel: Employee{},
|
||||||
|
})
|
||||||
|
|
||||||
|
processor := NewNestedCUDProcessor(db, registry, relProvider)
|
||||||
|
|
||||||
|
data := map[string]interface{}{
|
||||||
|
"ID": int64(1),
|
||||||
|
"name": "Engineering",
|
||||||
|
"employees": []interface{}{
|
||||||
|
map[string]interface{}{"_request": "add", "name": "Alice"},
|
||||||
|
map[string]interface{}{"_request": "remove", "ID": int64(5)},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := processor.ProcessNestedCUD(context.Background(), "update", data, Department{}, nil, "departments")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ProcessNestedCUD with nested add/remove failed: %v", err)
|
||||||
|
}
|
||||||
|
if len(db.insertCalls) != 1 {
|
||||||
|
t.Errorf("Expected 1 insert (add alias) for employee, got %d", len(db.insertCalls))
|
||||||
|
}
|
||||||
|
if len(db.deleteCalls) != 1 {
|
||||||
|
t.Errorf("Expected 1 delete (remove alias) for employee, got %d", len(db.deleteCalls))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestGetPrimaryKeyName(t *testing.T) {
|
func TestGetPrimaryKeyName(t *testing.T) {
|
||||||
dept := Department{}
|
dept := Department{}
|
||||||
pkName := reflection.GetPrimaryKeyName(dept)
|
pkName := reflection.GetPrimaryKeyName(dept)
|
||||||
|
|||||||
@@ -0,0 +1,317 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// This file implements PostGIS spatial and pgvector similarity filter operators.
|
||||||
|
// The builders return parameterised SQL fragments (with `?` placeholders) plus
|
||||||
|
// their args, matching the style of BuildInCondition / BuildArrayOverlapCondition.
|
||||||
|
// PostgreSQL only — on other databases these operators simply will not resolve.
|
||||||
|
|
||||||
|
// ── vector similarity ───────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
// VectorOperator maps a metric name to its pgvector distance operator.
|
||||||
|
//
|
||||||
|
// "l2" / "euclidean" / "" -> <->
|
||||||
|
// "cosine" -> <=>
|
||||||
|
// "ip" / "inner" / "dot" -> <#>
|
||||||
|
func VectorOperator(metric string) string {
|
||||||
|
switch strings.ToLower(strings.TrimSpace(metric)) {
|
||||||
|
case "cosine", "cos":
|
||||||
|
return "<=>"
|
||||||
|
case "ip", "inner", "dot", "innerproduct", "inner_product":
|
||||||
|
return "<#>"
|
||||||
|
default:
|
||||||
|
return "<->"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// VectorLiteral converts a vector value into a pgvector literal string
|
||||||
|
// "[1,2,3]". Accepts []float32, []float64, []int, []any (of numbers), or an
|
||||||
|
// already-formatted string.
|
||||||
|
func VectorLiteral(value any) (string, error) {
|
||||||
|
switch v := value.(type) {
|
||||||
|
case string:
|
||||||
|
s := strings.TrimSpace(v)
|
||||||
|
if strings.HasPrefix(s, "[") && strings.HasSuffix(s, "]") {
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
return "", fmt.Errorf("vector literal: malformed string %q", v)
|
||||||
|
case []float32:
|
||||||
|
return floatsToVectorLiteral(len(v), func(i int) float64 { return float64(v[i]) }), nil
|
||||||
|
case []float64:
|
||||||
|
return floatsToVectorLiteral(len(v), func(i int) float64 { return v[i] }), nil
|
||||||
|
case []int:
|
||||||
|
return floatsToVectorLiteral(len(v), func(i int) float64 { return float64(v[i]) }), nil
|
||||||
|
case []any:
|
||||||
|
nums := make([]float64, len(v))
|
||||||
|
for i, e := range v {
|
||||||
|
f, ok := toFloat(e)
|
||||||
|
if !ok {
|
||||||
|
return "", fmt.Errorf("vector literal: element %d is not a number (%T)", i, e)
|
||||||
|
}
|
||||||
|
nums[i] = f
|
||||||
|
}
|
||||||
|
return floatsToVectorLiteral(len(nums), func(i int) float64 { return nums[i] }), nil
|
||||||
|
default:
|
||||||
|
return "", fmt.Errorf("vector literal: unsupported type %T", value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func floatsToVectorLiteral(n int, at func(int) float64) string {
|
||||||
|
parts := make([]string, n)
|
||||||
|
for i := 0; i < n; i++ {
|
||||||
|
parts[i] = strconv.FormatFloat(at(i), 'f', -1, 32)
|
||||||
|
}
|
||||||
|
return "[" + strings.Join(parts, ",") + "]"
|
||||||
|
}
|
||||||
|
|
||||||
|
// BuildVectorCondition builds a pgvector distance-threshold filter.
|
||||||
|
//
|
||||||
|
// operator: "l2_within" | "cosine_within" | "ip_within"
|
||||||
|
// value: {"vector": [...], "distance": <n>}
|
||||||
|
// {"vector": [...], "lt"|"lte"|"gt"|"gte": <n>}
|
||||||
|
//
|
||||||
|
// Produces e.g. `embedding <=> ? < ?` with args [vectorLiteral, threshold].
|
||||||
|
func BuildVectorCondition(column, operator string, value any) (query string, args []interface{}, ok bool) {
|
||||||
|
var op string
|
||||||
|
switch strings.ToLower(operator) {
|
||||||
|
case "l2_within", "l2distance_within", "euclidean_within":
|
||||||
|
op = "<->"
|
||||||
|
case "cosine_within", "cosinedistance_within":
|
||||||
|
op = "<=>"
|
||||||
|
case "ip_within", "inner_within", "negativeinnerproduct_within":
|
||||||
|
op = "<#>"
|
||||||
|
default:
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
m, mok := value.(map[string]any)
|
||||||
|
if !mok {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
lit, err := VectorLiteral(m["vector"])
|
||||||
|
if err != nil {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
cmp := "<"
|
||||||
|
var threshold any
|
||||||
|
if t, ok := m["distance"]; ok {
|
||||||
|
threshold = t
|
||||||
|
} else {
|
||||||
|
for _, k := range []string{"lt", "lte", "gt", "gte"} {
|
||||||
|
if t, ok := m[k]; ok {
|
||||||
|
threshold = t
|
||||||
|
switch k {
|
||||||
|
case "lt":
|
||||||
|
cmp = "<"
|
||||||
|
case "lte":
|
||||||
|
cmp = "<="
|
||||||
|
case "gt":
|
||||||
|
cmp = ">"
|
||||||
|
case "gte":
|
||||||
|
cmp = ">="
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if threshold == nil {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
f, fok := toFloat(threshold)
|
||||||
|
if !fok {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Sprintf("%s %s ? %s ?", column, op, cmp), []interface{}{lit, f}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── PostGIS spatial ─────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
var spatialPredicates = map[string]string{
|
||||||
|
"st_intersects": "ST_Intersects",
|
||||||
|
"st_contains": "ST_Contains",
|
||||||
|
"st_within": "ST_Within",
|
||||||
|
"st_covers": "ST_Covers",
|
||||||
|
"st_coveredby": "ST_CoveredBy",
|
||||||
|
"st_overlaps": "ST_Overlaps",
|
||||||
|
"st_touches": "ST_Touches",
|
||||||
|
"st_crosses": "ST_Crosses",
|
||||||
|
"st_equals": "ST_Equals",
|
||||||
|
"st_disjoint": "ST_Disjoint",
|
||||||
|
}
|
||||||
|
|
||||||
|
// BuildSpatialCondition builds a PostGIS spatial filter.
|
||||||
|
//
|
||||||
|
// "st_dwithin" value: {"geom": <geojson|ewkt|hex>, "distance": <n>}
|
||||||
|
// "st_intersects" / "st_contains" / "st_within" / "st_covers" /
|
||||||
|
// "st_coveredby" / "st_overlaps" / "st_touches" / "st_crosses" /
|
||||||
|
// "st_equals" / "st_disjoint" value: <geojson|ewkt|hex>
|
||||||
|
// "bbox" (alias "&&") value: <geom> or {"bbox":[minx,miny,maxx,maxy],"srid":4326}
|
||||||
|
func BuildSpatialCondition(column, operator string, value any) (query string, args []interface{}, ok bool) {
|
||||||
|
operator = strings.ToLower(strings.TrimSpace(operator))
|
||||||
|
|
||||||
|
if fn, isPred := spatialPredicates[operator]; isPred {
|
||||||
|
expr, arg, err := geomArgExpr(value)
|
||||||
|
if err != nil {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%s(%s, %s)", fn, column, expr), []interface{}{arg}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
switch operator {
|
||||||
|
case "st_dwithin":
|
||||||
|
m, mok := value.(map[string]any)
|
||||||
|
if !mok {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
expr, arg, err := geomArgExpr(m["geom"])
|
||||||
|
if err != nil {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
dist, dok := toFloat(m["distance"])
|
||||||
|
if !dok {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("ST_DWithin(%s, %s, ?)", column, expr), []interface{}{arg, dist}, true
|
||||||
|
|
||||||
|
case "bbox", "&&":
|
||||||
|
if m, mok := value.(map[string]any); mok {
|
||||||
|
if bboxRaw, has := m["bbox"]; has {
|
||||||
|
coords, cok := toFloatSlice(bboxRaw)
|
||||||
|
if !cok || len(coords) != 4 {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
srid := 4326
|
||||||
|
if s, sok := toFloat(m["srid"]); sok {
|
||||||
|
srid = int(s)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%s && ST_MakeEnvelope(?, ?, ?, ?, ?)", column),
|
||||||
|
[]interface{}{coords[0], coords[1], coords[2], coords[3], srid}, true
|
||||||
|
}
|
||||||
|
// fall through: treat the map as a GeoJSON geometry
|
||||||
|
}
|
||||||
|
expr, arg, err := geomArgExpr(value)
|
||||||
|
if err != nil {
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%s && %s", column, expr), []interface{}{arg}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
return "", nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// geomArgExpr inspects a geometry value and returns the SQL placeholder
|
||||||
|
// expression that turns a bound argument into a geometry, plus that argument.
|
||||||
|
//
|
||||||
|
// GeoJSON object -> "ST_GeomFromGeoJSON(?)", <json string>
|
||||||
|
// hex EWKB -> "?::geometry", <hex string>
|
||||||
|
// WKT / EWKT -> "ST_GeomFromEWKT(?)", <ewkt string>
|
||||||
|
func geomArgExpr(value any) (expr string, arg any, err error) {
|
||||||
|
switch v := value.(type) {
|
||||||
|
case nil:
|
||||||
|
return "", nil, fmt.Errorf("geometry: nil value")
|
||||||
|
case map[string]any:
|
||||||
|
b, mErr := json.Marshal(v)
|
||||||
|
if mErr != nil {
|
||||||
|
return "", nil, mErr
|
||||||
|
}
|
||||||
|
return "ST_GeomFromGeoJSON(?)", string(b), nil
|
||||||
|
case json.RawMessage:
|
||||||
|
return "ST_GeomFromGeoJSON(?)", string(v), nil
|
||||||
|
case []byte:
|
||||||
|
return geomArgExpr(string(v))
|
||||||
|
case string:
|
||||||
|
s := strings.TrimSpace(v)
|
||||||
|
if s == "" {
|
||||||
|
return "", nil, fmt.Errorf("geometry: empty value")
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(s, "{") {
|
||||||
|
return "ST_GeomFromGeoJSON(?)", s, nil
|
||||||
|
}
|
||||||
|
if isHexString(s) {
|
||||||
|
return "?::geometry", s, nil
|
||||||
|
}
|
||||||
|
return "ST_GeomFromEWKT(?)", s, nil
|
||||||
|
default:
|
||||||
|
return "", nil, fmt.Errorf("geometry: unsupported type %T", value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func isHexString(s string) bool {
|
||||||
|
if len(s) < 10 || len(s)%2 != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
_, err := hex.DecodeString(s)
|
||||||
|
return err == nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func toFloat(v any) (float64, bool) {
|
||||||
|
switch n := v.(type) {
|
||||||
|
case float64:
|
||||||
|
return n, true
|
||||||
|
case float32:
|
||||||
|
return float64(n), true
|
||||||
|
case int:
|
||||||
|
return float64(n), true
|
||||||
|
case int64:
|
||||||
|
return float64(n), true
|
||||||
|
case json.Number:
|
||||||
|
f, err := n.Float64()
|
||||||
|
return f, err == nil
|
||||||
|
case string:
|
||||||
|
f, err := strconv.ParseFloat(strings.TrimSpace(n), 64)
|
||||||
|
return f, err == nil
|
||||||
|
default:
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func toFloatSlice(v any) ([]float64, bool) {
|
||||||
|
switch s := v.(type) {
|
||||||
|
case []float64:
|
||||||
|
return s, true
|
||||||
|
case []any:
|
||||||
|
out := make([]float64, len(s))
|
||||||
|
for i, e := range s {
|
||||||
|
f, ok := toFloat(e)
|
||||||
|
if !ok {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
out[i] = f
|
||||||
|
}
|
||||||
|
return out, true
|
||||||
|
default:
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsSpatialOperator reports whether op is a spatial filter operator handled by
|
||||||
|
// BuildSpatialCondition.
|
||||||
|
func IsSpatialOperator(op string) bool {
|
||||||
|
op = strings.ToLower(strings.TrimSpace(op))
|
||||||
|
if _, ok := spatialPredicates[op]; ok {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return op == "st_dwithin" || op == "bbox" || op == "&&"
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsVectorOperator reports whether op is a vector similarity filter operator
|
||||||
|
// handled by BuildVectorCondition.
|
||||||
|
func IsVectorOperator(op string) bool {
|
||||||
|
switch strings.ToLower(strings.TrimSpace(op)) {
|
||||||
|
case "l2_within", "l2distance_within", "euclidean_within",
|
||||||
|
"cosine_within", "cosinedistance_within",
|
||||||
|
"ip_within", "inner_within", "negativeinnerproduct_within":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,145 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestVectorOperator(t *testing.T) {
|
||||||
|
cases := map[string]string{
|
||||||
|
"": "<->", "l2": "<->", "euclidean": "<->",
|
||||||
|
"cosine": "<=>", "cos": "<=>",
|
||||||
|
"ip": "<#>", "inner": "<#>", "dot": "<#>",
|
||||||
|
}
|
||||||
|
for in, want := range cases {
|
||||||
|
if got := VectorOperator(in); got != want {
|
||||||
|
t.Errorf("VectorOperator(%q) = %q, want %q", in, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestVectorLiteral(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
in any
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{[]float32{1, 2, 3}, "[1,2,3]"},
|
||||||
|
{[]float64{1.5, -2}, "[1.5,-2]"},
|
||||||
|
{[]int{1, 2}, "[1,2]"},
|
||||||
|
{[]any{1.0, 2.0}, "[1,2]"},
|
||||||
|
{"[4,5,6]", "[4,5,6]"},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
got, err := VectorLiteral(c.in)
|
||||||
|
if err != nil || got != c.want {
|
||||||
|
t.Errorf("VectorLiteral(%v) = %q, %v; want %q", c.in, got, err, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, err := VectorLiteral("not-a-vector"); err == nil {
|
||||||
|
t.Error("expected error for malformed string")
|
||||||
|
}
|
||||||
|
if _, err := VectorLiteral(42); err == nil {
|
||||||
|
t.Error("expected error for unsupported type")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildVectorCondition(t *testing.T) {
|
||||||
|
q, args, ok := BuildVectorCondition("embedding", "cosine_within", map[string]any{
|
||||||
|
"vector": []any{1.0, 2.0, 3.0}, "distance": 0.5,
|
||||||
|
})
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected ok")
|
||||||
|
}
|
||||||
|
if q != "embedding <=> ? < ?" {
|
||||||
|
t.Errorf("query = %q", q)
|
||||||
|
}
|
||||||
|
if len(args) != 2 || args[0] != "[1,2,3]" || args[1] != 0.5 {
|
||||||
|
t.Errorf("args = %v", args)
|
||||||
|
}
|
||||||
|
|
||||||
|
// explicit comparator
|
||||||
|
q, _, ok = BuildVectorCondition("v", "l2_within", map[string]any{
|
||||||
|
"vector": []float32{1}, "lte": 2.0,
|
||||||
|
})
|
||||||
|
if !ok || q != "v <-> ? <= ?" {
|
||||||
|
t.Errorf("lte: q=%q ok=%v", q, ok)
|
||||||
|
}
|
||||||
|
|
||||||
|
// unknown operator
|
||||||
|
if _, _, ok := BuildVectorCondition("v", "bogus", map[string]any{}); ok {
|
||||||
|
t.Error("expected not ok for unknown operator")
|
||||||
|
}
|
||||||
|
// missing threshold
|
||||||
|
if _, _, ok := BuildVectorCondition("v", "l2_within", map[string]any{"vector": []float32{1}}); ok {
|
||||||
|
t.Error("expected not ok without threshold")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildSpatialCondition_Predicates(t *testing.T) {
|
||||||
|
q, args, ok := BuildSpatialCondition("geom", "st_intersects", "SRID=4326;POINT(0 0)")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected ok")
|
||||||
|
}
|
||||||
|
if q != "ST_Intersects(geom, ST_GeomFromEWKT(?))" {
|
||||||
|
t.Errorf("query = %q", q)
|
||||||
|
}
|
||||||
|
if len(args) != 1 || args[0] != "SRID=4326;POINT(0 0)" {
|
||||||
|
t.Errorf("args = %v", args)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GeoJSON value
|
||||||
|
q, args, ok = BuildSpatialCondition("geom", "st_contains", map[string]any{
|
||||||
|
"type": "Point", "coordinates": []any{1.0, 2.0},
|
||||||
|
})
|
||||||
|
if !ok || q != "ST_Contains(geom, ST_GeomFromGeoJSON(?))" {
|
||||||
|
t.Errorf("geojson: q=%q ok=%v", q, ok)
|
||||||
|
}
|
||||||
|
if len(args) != 1 {
|
||||||
|
t.Errorf("args = %v", args)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildSpatialCondition_DWithin(t *testing.T) {
|
||||||
|
q, args, ok := BuildSpatialCondition("geom", "st_dwithin", map[string]any{
|
||||||
|
"geom": "SRID=4326;POINT(0 0)", "distance": 1000.0,
|
||||||
|
})
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected ok")
|
||||||
|
}
|
||||||
|
if q != "ST_DWithin(geom, ST_GeomFromEWKT(?), ?)" {
|
||||||
|
t.Errorf("query = %q", q)
|
||||||
|
}
|
||||||
|
if len(args) != 2 || args[1] != 1000.0 {
|
||||||
|
t.Errorf("args = %v", args)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildSpatialCondition_BBox(t *testing.T) {
|
||||||
|
q, args, ok := BuildSpatialCondition("geom", "bbox", map[string]any{
|
||||||
|
"bbox": []any{0.0, 0.0, 10.0, 10.0}, "srid": 4326.0,
|
||||||
|
})
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected ok")
|
||||||
|
}
|
||||||
|
if q != "geom && ST_MakeEnvelope(?, ?, ?, ?, ?)" {
|
||||||
|
t.Errorf("query = %q", q)
|
||||||
|
}
|
||||||
|
if len(args) != 5 || args[4] != 4326 {
|
||||||
|
t.Errorf("args = %v", args)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsSpatialAndVectorOperator(t *testing.T) {
|
||||||
|
for _, op := range []string{"st_dwithin", "st_intersects", "bbox", "&&"} {
|
||||||
|
if !IsSpatialOperator(op) {
|
||||||
|
t.Errorf("%q should be spatial", op)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, op := range []string{"l2_within", "cosine_within", "ip_within"} {
|
||||||
|
if !IsVectorOperator(op) {
|
||||||
|
t.Errorf("%q should be vector", op)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if IsSpatialOperator("eq") || IsVectorOperator("eq") {
|
||||||
|
t.Error("eq is neither spatial nor vector")
|
||||||
|
}
|
||||||
|
}
|
||||||
+115
-17
@@ -59,6 +59,38 @@ func IsSQLExpression(cond string) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// reEmptyCompMid matches a simple column comparison with an empty RHS that is immediately
|
||||||
|
// followed by AND/OR (only whitespace between the operator and the next keyword).
|
||||||
|
// Removing the match leaves the preceding AND/OR connector intact.
|
||||||
|
// Example: "cond1 and col = \n and cond2" → "cond1 and cond2"
|
||||||
|
var reEmptyCompMid = regexp.MustCompile(`(?i)[\w.]+\s*(?:=|<>|!=|>=|<=|>|<)\s+(?:and|or)\s+`)
|
||||||
|
|
||||||
|
// reEmptyCompEnd matches AND/OR + a simple column comparison with an empty RHS at the end
|
||||||
|
// of the string (or sub-clause).
|
||||||
|
// Example: "cond1 and col = " → "cond1"
|
||||||
|
var reEmptyCompEnd = regexp.MustCompile(`(?i)\s+(?:and|or)\s+[\w.]+\s*(?:=|<>|!=|>=|<=|>|<)\s*$`)
|
||||||
|
|
||||||
|
// stripEmptyComparisonClauses removes comparison conditions that have no right-hand side
|
||||||
|
// value from a raw SQL string. Operates on the whole string so it also cleans up conditions
|
||||||
|
// inside subqueries, not just top-level AND splits.
|
||||||
|
func stripEmptyComparisonClauses(sql string) string {
|
||||||
|
sql = reEmptyCompMid.ReplaceAllLiteralString(sql, "")
|
||||||
|
sql = reEmptyCompEnd.ReplaceAllLiteralString(sql, "")
|
||||||
|
return sql
|
||||||
|
}
|
||||||
|
|
||||||
|
// hasEmptyRHS returns true when a condition ends with a comparison operator and has no
|
||||||
|
// right-hand side value — e.g., "col = ", "com.rid_parent = ", "col >= ".
|
||||||
|
func hasEmptyRHS(cond string) bool {
|
||||||
|
cond = strings.TrimSpace(cond)
|
||||||
|
for _, op := range []string{"<>", "!=", ">=", "<=", "=", ">", "<"} {
|
||||||
|
if strings.HasSuffix(cond, op) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
// IsTrivialCondition checks if a condition is trivial and always evaluates to true
|
// IsTrivialCondition checks if a condition is trivial and always evaluates to true
|
||||||
// These conditions should be removed from WHERE clauses as they have no filtering effect
|
// These conditions should be removed from WHERE clauses as they have no filtering effect
|
||||||
func IsTrivialCondition(cond string) bool {
|
func IsTrivialCondition(cond string) bool {
|
||||||
@@ -147,6 +179,14 @@ func SanitizeWhereClause(where string, tableName string, options ...*RequestOpti
|
|||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Strip comparison conditions with empty RHS throughout the SQL string (including
|
||||||
|
// inside subqueries), before condition splitting.
|
||||||
|
where = stripEmptyComparisonClauses(where)
|
||||||
|
if where == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
where = strings.TrimSpace(where)
|
||||||
|
|
||||||
// Check if the original clause has outer parentheses and contains OR operators
|
// Check if the original clause has outer parentheses and contains OR operators
|
||||||
// If so, we need to preserve the outer parentheses to prevent OR logic from escaping
|
// If so, we need to preserve the outer parentheses to prevent OR logic from escaping
|
||||||
hasOuterParens := false
|
hasOuterParens := false
|
||||||
@@ -212,6 +252,12 @@ func SanitizeWhereClause(where string, tableName string, options ...*RequestOpti
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Skip conditions with no right-hand side value (e.g. "col = " with empty value)
|
||||||
|
if hasEmptyRHS(condToCheck) {
|
||||||
|
logger.Debug("Removing condition with empty value: '%s'", cond)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
// If tableName is provided and the condition HAS a table prefix, check if it's correct
|
// If tableName is provided and the condition HAS a table prefix, check if it's correct
|
||||||
if tableName != "" && hasTablePrefix(condToCheck) {
|
if tableName != "" && hasTablePrefix(condToCheck) {
|
||||||
// Extract the current prefix and column name
|
// Extract the current prefix and column name
|
||||||
@@ -400,18 +446,36 @@ func containsTopLevelOR(clause string) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// splitByAND splits a WHERE clause by AND operators (case-insensitive)
|
// splitByAND splits a WHERE clause by AND operators (case-insensitive).
|
||||||
// This is parenthesis-aware and won't split on AND operators inside subqueries
|
// It is parenthesis-aware (won't split inside subqueries), quote-aware
|
||||||
|
// (won't split on AND inside single-quoted strings), and BETWEEN-aware
|
||||||
|
// (won't split on the AND that separates the two operands of BETWEEN x AND y).
|
||||||
func splitByAND(where string) []string {
|
func splitByAND(where string) []string {
|
||||||
conditions := []string{}
|
conditions := []string{}
|
||||||
currentCondition := strings.Builder{}
|
currentCondition := strings.Builder{}
|
||||||
depth := 0 // Track parenthesis depth
|
depth := 0 // parenthesis nesting depth
|
||||||
|
inSingleQuote := false
|
||||||
|
afterBetween := false // true after seeing BETWEEN at depth 0; next AND belongs to it
|
||||||
i := 0
|
i := 0
|
||||||
|
|
||||||
for i < len(where) {
|
for i < len(where) {
|
||||||
ch := where[i]
|
ch := where[i]
|
||||||
|
|
||||||
// Track parenthesis depth
|
// Track single-quote state so we never split on AND inside string literals.
|
||||||
|
if ch == '\'' {
|
||||||
|
inSingleQuote = !inSingleQuote
|
||||||
|
currentCondition.WriteByte(ch)
|
||||||
|
i++
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if inSingleQuote {
|
||||||
|
currentCondition.WriteByte(ch)
|
||||||
|
i++
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Track parenthesis depth (outside quotes only).
|
||||||
if ch == '(' {
|
if ch == '(' {
|
||||||
depth++
|
depth++
|
||||||
currentCondition.WriteByte(ch)
|
currentCondition.WriteByte(ch)
|
||||||
@@ -424,32 +488,39 @@ func splitByAND(where string) []string {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Only look for AND operators at depth 0 (not inside parentheses)
|
// All keyword checks only apply at depth 0 (not inside subqueries).
|
||||||
if depth == 0 {
|
if depth == 0 {
|
||||||
// Check if we're at an AND operator (case-insensitive)
|
// Detect " BETWEEN " (9 chars, case-insensitive) so the very next
|
||||||
// We need at least " AND " (5 chars) or " and " (5 chars)
|
// top-level AND is recognised as part of the BETWEEN syntax.
|
||||||
if i+5 <= len(where) {
|
if i+9 <= len(where) && strings.ToLower(where[i:i+9]) == " between " {
|
||||||
substring := where[i : i+5]
|
afterBetween = true
|
||||||
lowerSubstring := strings.ToLower(substring)
|
currentCondition.WriteString(where[i : i+9])
|
||||||
|
i += 9
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
if lowerSubstring == " and " {
|
// Detect " AND " (5 chars, case-insensitive).
|
||||||
// Found an AND operator at the top level
|
if i+5 <= len(where) && strings.ToLower(where[i:i+5]) == " and " {
|
||||||
// Add the current condition to the list
|
if afterBetween {
|
||||||
|
// This AND closes a BETWEEN expression — do NOT split.
|
||||||
|
afterBetween = false
|
||||||
|
currentCondition.WriteString(where[i : i+5])
|
||||||
|
i += 5
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Regular conjunction — split here.
|
||||||
conditions = append(conditions, currentCondition.String())
|
conditions = append(conditions, currentCondition.String())
|
||||||
currentCondition.Reset()
|
currentCondition.Reset()
|
||||||
// Skip past the AND operator
|
|
||||||
i += 5
|
i += 5
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
// Not an AND operator or we're inside parentheses, just add the character
|
|
||||||
currentCondition.WriteByte(ch)
|
currentCondition.WriteByte(ch)
|
||||||
i++
|
i++
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add the last condition
|
// Add the last condition.
|
||||||
if currentCondition.Len() > 0 {
|
if currentCondition.Len() > 0 {
|
||||||
conditions = append(conditions, currentCondition.String())
|
conditions = append(conditions, currentCondition.String())
|
||||||
}
|
}
|
||||||
@@ -568,6 +639,15 @@ func extractTableAndColumn(cond string) (table string, column string) {
|
|||||||
// Remove any quotes
|
// Remove any quotes
|
||||||
columnRef = strings.Trim(columnRef, "`\"'")
|
columnRef = strings.Trim(columnRef, "`\"'")
|
||||||
|
|
||||||
|
// If the left side is a parenthesized subquery (starts with '(' and contains SQL keywords),
|
||||||
|
// don't attempt prefix extraction from inside it.
|
||||||
|
if len(columnRef) > 0 && columnRef[0] == '(' {
|
||||||
|
lowerRef := strings.ToLower(columnRef)
|
||||||
|
if strings.Contains(lowerRef, "select ") || strings.Contains(lowerRef, " from ") || strings.Contains(lowerRef, " where ") {
|
||||||
|
return "", ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Check if there's a function call (contains opening parenthesis)
|
// Check if there's a function call (contains opening parenthesis)
|
||||||
openParenIdx := strings.Index(columnRef, "(")
|
openParenIdx := strings.Index(columnRef, "(")
|
||||||
|
|
||||||
@@ -960,3 +1040,21 @@ func BuildInCondition(column string, v interface{}) (query string, args []interf
|
|||||||
}
|
}
|
||||||
return fmt.Sprintf("%s IN (%s)", column, strings.Join(placeholders, ",")), values
|
return fmt.Sprintf("%s IN (%s)", column, strings.Join(placeholders, ",")), values
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// BuildArrayOverlapCondition builds a parameterized condition testing whether an
|
||||||
|
// array column has at least one element in common with the given value(s), using
|
||||||
|
// PostgreSQL's array overlap operator (&&). Unlike a text-cast ILIKE, this performs
|
||||||
|
// real element-wise containment (no substring false positives) and can use a GIN
|
||||||
|
// index on the column. A single value is treated as a one-element array.
|
||||||
|
// Returns ("", nil) if the value is empty.
|
||||||
|
func BuildArrayOverlapCondition(column string, v interface{}) (query string, args []interface{}) {
|
||||||
|
values := FilterValueToSlice(v)
|
||||||
|
if len(values) == 0 {
|
||||||
|
return "", nil
|
||||||
|
}
|
||||||
|
placeholders := make([]string, len(values))
|
||||||
|
for i := range values {
|
||||||
|
placeholders[i] = "?"
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%s && ARRAY[%s]", column, strings.Join(placeholders, ",")), values
|
||||||
|
}
|
||||||
|
|||||||
+201
-68
@@ -134,6 +134,30 @@ func TestSanitizeWhereClause(t *testing.T) {
|
|||||||
tableName: "apiprovider",
|
tableName: "apiprovider",
|
||||||
expected: "apiprovider.type in ('softphone') AND (apiprovider.rid_apiprovider in (select l.rid_apiprovider from core.apiproviderlink l where l.rid_hub = 2576))",
|
expected: "apiprovider.type in ('softphone') AND (apiprovider.rid_apiprovider in (select l.rid_apiprovider from core.apiproviderlink l where l.rid_hub = 2576))",
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
name: "empty RHS stripped mid-clause",
|
||||||
|
where: "com.tableprefix = 'tcli' and com.rid_parent = \n and com.status = 'Active'",
|
||||||
|
tableName: "",
|
||||||
|
expected: "com.tableprefix = 'tcli' AND com.status = 'Active'",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty RHS stripped at end of clause",
|
||||||
|
where: "com.tableprefix = 'tcli' and com.rid_parent =",
|
||||||
|
tableName: "",
|
||||||
|
expected: "com.tableprefix = 'tcli'",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "non-empty value not stripped",
|
||||||
|
where: "com.tableprefix = 'tcli' and com.rid_parent = 123 and com.status = 'Active'",
|
||||||
|
tableName: "",
|
||||||
|
expected: "com.tableprefix = 'tcli' AND com.rid_parent = 123 AND com.status = 'Active'",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty RHS inside subquery stripped",
|
||||||
|
where: "a = 1 and b in (select x from t where c.rid = \n and d = 2)",
|
||||||
|
tableName: "",
|
||||||
|
expected: "a = 1 AND b in (select x from t where d = 2)",
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
@@ -496,6 +520,38 @@ func TestSplitByAND(t *testing.T) {
|
|||||||
input: "a = 1 AND b = 2 AND c = 3 and (select s from generate_series(1,10) s where s < 10 and s > 0 offset 2 limit 1) = 3",
|
input: "a = 1 AND b = 2 AND c = 3 and (select s from generate_series(1,10) s where s < 10 and s > 0 offset 2 limit 1) = 3",
|
||||||
expected: []string{"a = 1", "b = 2", "c = 3", "(select s from generate_series(1,10) s where s < 10 and s > 0 offset 2 limit 1) = 3"},
|
expected: []string{"a = 1", "b = 2", "c = 3", "(select s from generate_series(1,10) s where s < 10 and s > 0 offset 2 limit 1) = 3"},
|
||||||
},
|
},
|
||||||
|
// BETWEEN-aware cases: the AND inside BETWEEN x AND y must not cause a split.
|
||||||
|
{
|
||||||
|
name: "BETWEEN does not split on its AND",
|
||||||
|
input: "col between '2025-08-31' and '1970-01-01'",
|
||||||
|
expected: []string{"col between '2025-08-31' and '1970-01-01'"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "BETWEEN uppercase AND",
|
||||||
|
input: "col BETWEEN '2025-08-31' AND '1970-01-01'",
|
||||||
|
expected: []string{"col BETWEEN '2025-08-31' AND '1970-01-01'"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "BETWEEN followed by a regular AND conjunction",
|
||||||
|
input: "col between 1 and 5 and other = 'x'",
|
||||||
|
expected: []string{"col between 1 and 5", "other = 'x'"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "two BETWEEN conditions joined by AND",
|
||||||
|
input: "col1 between 1 and 5 and col2 between 10 and 20",
|
||||||
|
expected: []string{"col1 between 1 and 5", "col2 between 10 and 20"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "complex OR block with multiple BETWEENs (real-world case)",
|
||||||
|
input: "tbl.applicationdate between '2025-08-31' and '1970-01-01'\n or tbl.capturedate between '2025-08-31' and '1970-01-01'\n or tbl.startdate between '2025-08-31' AND '1970-01-01'",
|
||||||
|
expected: []string{"tbl.applicationdate between '2025-08-31' and '1970-01-01'\n or tbl.capturedate between '2025-08-31' and '1970-01-01'\n or tbl.startdate between '2025-08-31' AND '1970-01-01'"},
|
||||||
|
},
|
||||||
|
// Quote-aware cases: AND inside a string literal must not split.
|
||||||
|
{
|
||||||
|
name: "AND inside single-quoted string is not a split point",
|
||||||
|
input: "comment = 'this and that' and status = 'active'",
|
||||||
|
expected: []string{"comment = 'this and that'", "status = 'active'"},
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
@@ -833,74 +889,151 @@ func TestSanitizeWhereClause_PreservesParenthesesWithOR(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestAddTablePrefixToColumns_ComplexConditions(t *testing.T) {
|
func TestAddTablePrefixToColumns_ComplexConditions(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
where string
|
where string
|
||||||
tableName string
|
tableName string
|
||||||
expected string
|
expected string
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "Parentheses with true AND condition - should not prefix true",
|
name: "Parentheses with true AND condition - should not prefix true",
|
||||||
where: "(true AND status = 'active')",
|
where: "(true AND status = 'active')",
|
||||||
tableName: "mastertask",
|
tableName: "mastertask",
|
||||||
expected: "(true AND mastertask.status = 'active')",
|
expected: "(true AND mastertask.status = 'active')",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Parentheses with multiple conditions including true",
|
name: "Parentheses with multiple conditions including true",
|
||||||
where: "(true AND status = 'active' AND id > 5)",
|
where: "(true AND status = 'active' AND id > 5)",
|
||||||
tableName: "mastertask",
|
tableName: "mastertask",
|
||||||
expected: "(true AND mastertask.status = 'active' AND mastertask.id > 5)",
|
expected: "(true AND mastertask.status = 'active' AND mastertask.id > 5)",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Nested parentheses with true",
|
name: "Nested parentheses with true",
|
||||||
where: "((true AND status = 'active'))",
|
where: "((true AND status = 'active'))",
|
||||||
tableName: "mastertask",
|
tableName: "mastertask",
|
||||||
expected: "((true AND mastertask.status = 'active'))",
|
expected: "((true AND mastertask.status = 'active'))",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Mixed: false AND valid conditions",
|
name: "Mixed: false AND valid conditions",
|
||||||
where: "(false AND name = 'test')",
|
where: "(false AND name = 'test')",
|
||||||
tableName: "mastertask",
|
tableName: "mastertask",
|
||||||
expected: "(false AND mastertask.name = 'test')",
|
expected: "(false AND mastertask.name = 'test')",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Mixed: null AND valid conditions",
|
name: "Mixed: null AND valid conditions",
|
||||||
where: "(null AND status = 'active')",
|
where: "(null AND status = 'active')",
|
||||||
tableName: "mastertask",
|
tableName: "mastertask",
|
||||||
expected: "(null AND mastertask.status = 'active')",
|
expected: "(null AND mastertask.status = 'active')",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Multiple true conditions in parentheses",
|
name: "Multiple true conditions in parentheses",
|
||||||
where: "(true AND true AND status = 'active')",
|
where: "(true AND true AND status = 'active')",
|
||||||
tableName: "mastertask",
|
tableName: "mastertask",
|
||||||
expected: "(true AND true AND mastertask.status = 'active')",
|
expected: "(true AND true AND mastertask.status = 'active')",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Simple true without parens - should not prefix",
|
name: "Simple true without parens - should not prefix",
|
||||||
where: "true",
|
where: "true",
|
||||||
tableName: "mastertask",
|
tableName: "mastertask",
|
||||||
expected: "true",
|
expected: "true",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Simple condition without parens - should prefix",
|
name: "Simple condition without parens - should prefix",
|
||||||
where: "status = 'active'",
|
where: "status = 'active'",
|
||||||
tableName: "mastertask",
|
tableName: "mastertask",
|
||||||
expected: "mastertask.status = 'active'",
|
expected: "mastertask.status = 'active'",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Unregistered table with true - should not prefix true",
|
name: "Unregistered table with true - should not prefix true",
|
||||||
where: "(true AND status = 'active')",
|
where: "(true AND status = 'active')",
|
||||||
tableName: "unregistered_table",
|
tableName: "unregistered_table",
|
||||||
expected: "(true AND unregistered_table.status = 'active')",
|
expected: "(true AND unregistered_table.status = 'active')",
|
||||||
},
|
},
|
||||||
|
// BETWEEN regression: date literals inside BETWEEN must not be prefixed as columns.
|
||||||
|
{
|
||||||
|
name: "BETWEEN date range - second date must not be prefixed",
|
||||||
|
where: "applicationdate between '2025-08-31' and '1970-01-01'",
|
||||||
|
tableName: "unregistered_table",
|
||||||
|
expected: "unregistered_table.applicationdate between '2025-08-31' and '1970-01-01'",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Already-prefixed BETWEEN column - unchanged",
|
||||||
|
where: `"v_webui_clients".applicationdate between '2025-08-31' and '1970-01-01'`,
|
||||||
|
tableName: "v_webui_clients",
|
||||||
|
expected: `"v_webui_clients".applicationdate between '2025-08-31' and '1970-01-01'`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Complex OR block with multiple BETWEENs - date values must not be prefixed",
|
||||||
|
where: `("v_webui_clients".applicationdate between '2025-08-31' and '1970-01-01' or "v_webui_clients".clientcapturedate between '2025-08-31' and '1970-01-01' or "v_webui_clients".startdate between '2025-08-31' AND '1970-01-01')`,
|
||||||
|
tableName: "v_webui_clients",
|
||||||
|
expected: `("v_webui_clients".applicationdate between '2025-08-31' and '1970-01-01' or "v_webui_clients".clientcapturedate between '2025-08-31' and '1970-01-01' or "v_webui_clients".startdate between '2025-08-31' AND '1970-01-01')`,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
result := AddTablePrefixToColumns(tt.where, tt.tableName)
|
||||||
|
if result != tt.expected {
|
||||||
|
t.Errorf("AddTablePrefixToColumns(%q, %q) = %q; want %q", tt.where, tt.tableName, result, tt.expected)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
func TestBuildArrayOverlapCondition(t *testing.T) {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
tests := []struct {
|
||||||
result := AddTablePrefixToColumns(tt.where, tt.tableName)
|
name string
|
||||||
if result != tt.expected {
|
column string
|
||||||
t.Errorf("AddTablePrefixToColumns(%q, %q) = %q; want %q", tt.where, tt.tableName, result, tt.expected)
|
value interface{}
|
||||||
}
|
expectedCond string
|
||||||
})
|
expectedArgs int
|
||||||
}
|
}{
|
||||||
|
{
|
||||||
|
name: "single scalar value",
|
||||||
|
column: "tags",
|
||||||
|
value: "urgent",
|
||||||
|
expectedCond: "tags && ARRAY[?]",
|
||||||
|
expectedArgs: 1,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "multiple values",
|
||||||
|
column: "tags",
|
||||||
|
value: []string{"urgent", "billing", "vip"},
|
||||||
|
expectedCond: "tags && ARRAY[?,?,?]",
|
||||||
|
expectedArgs: 3,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "JSON-decoded []interface{} value",
|
||||||
|
column: "tags",
|
||||||
|
value: []interface{}{"urgent", "billing"},
|
||||||
|
expectedCond: "tags && ARRAY[?,?]",
|
||||||
|
expectedArgs: 2,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "nil value",
|
||||||
|
column: "tags",
|
||||||
|
value: nil,
|
||||||
|
expectedCond: "",
|
||||||
|
expectedArgs: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty slice value",
|
||||||
|
column: "tags",
|
||||||
|
value: []string{},
|
||||||
|
expectedCond: "",
|
||||||
|
expectedArgs: 0,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
cond, args := BuildArrayOverlapCondition(tt.column, tt.value)
|
||||||
|
if cond != tt.expectedCond {
|
||||||
|
t.Errorf("BuildArrayOverlapCondition(%q, %v) condition = %q; want %q", tt.column, tt.value, cond, tt.expectedCond)
|
||||||
|
}
|
||||||
|
if len(args) != tt.expectedArgs {
|
||||||
|
t.Errorf("BuildArrayOverlapCondition(%q, %v) args = %d; want %d", tt.column, tt.value, len(args), tt.expectedArgs)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,5 +1,23 @@
|
|||||||
package common
|
package common
|
||||||
|
|
||||||
|
// SQLError wraps a database error together with the SQL that caused it,
|
||||||
|
// so callers can surface the query in API error responses for easier debugging.
|
||||||
|
type SQLError struct {
|
||||||
|
Err error
|
||||||
|
SQL string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *SQLError) Error() string { return e.Err.Error() }
|
||||||
|
func (e *SQLError) Unwrap() error { return e.Err }
|
||||||
|
|
||||||
|
// WrapSQLError wraps err with the given SQL. If err is nil it returns nil.
|
||||||
|
func WrapSQLError(err error, sql string) error {
|
||||||
|
if err == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return &SQLError{Err: err, SQL: sql}
|
||||||
|
}
|
||||||
|
|
||||||
type RequestBody struct {
|
type RequestBody struct {
|
||||||
Operation string `json:"operation"`
|
Operation string `json:"operation"`
|
||||||
Data interface{} `json:"data"`
|
Data interface{} `json:"data"`
|
||||||
@@ -24,6 +42,10 @@ type RequestOptions struct {
|
|||||||
CursorBackward string `json:"cursor_backward"`
|
CursorBackward string `json:"cursor_backward"`
|
||||||
FetchRowNumber *string `json:"fetch_row_number"`
|
FetchRowNumber *string `json:"fetch_row_number"`
|
||||||
|
|
||||||
|
// VectorSearch performs a pgvector nearest-neighbour ordering (KNN) and
|
||||||
|
// optionally returns the computed distance as an extra column.
|
||||||
|
VectorSearch *VectorSearchOption `json:"vector_search"`
|
||||||
|
|
||||||
// Join table aliases (used for validation of prefixed columns in filters/sorts)
|
// Join table aliases (used for validation of prefixed columns in filters/sorts)
|
||||||
// Not serialized to JSON as it's internal validation state
|
// Not serialized to JSON as it's internal validation state
|
||||||
JoinAliases []string `json:"-"`
|
JoinAliases []string `json:"-"`
|
||||||
@@ -72,6 +94,41 @@ type SortOption struct {
|
|||||||
Direction string `json:"direction"`
|
Direction string `json:"direction"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// PrimaryKeySortColumn is a sentinel SortOption.Column value that gets
|
||||||
|
// resolved to a model's actual primary key column name at query time.
|
||||||
|
// This lets a single global default sort (e.g. set once via
|
||||||
|
// Handler.SetDefaultSort) work across models with different primary keys,
|
||||||
|
// e.g. common.SortOption{Column: common.PrimaryKeySortColumn, Direction: "asc"}.
|
||||||
|
const PrimaryKeySortColumn = "$pk"
|
||||||
|
|
||||||
|
// ResolveSortColumns returns a copy of sort with any PrimaryKeySortColumn
|
||||||
|
// entries replaced by pkName. If pkName is empty, matching entries are
|
||||||
|
// dropped since there is no column to sort by.
|
||||||
|
func ResolveSortColumns(sort []SortOption, pkName string) []SortOption {
|
||||||
|
resolved := make([]SortOption, 0, len(sort))
|
||||||
|
for _, s := range sort {
|
||||||
|
if s.Column == PrimaryKeySortColumn {
|
||||||
|
if pkName == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
s.Column = pkName
|
||||||
|
}
|
||||||
|
resolved = append(resolved, s)
|
||||||
|
}
|
||||||
|
return resolved
|
||||||
|
}
|
||||||
|
|
||||||
|
// VectorSearchOption describes a pgvector KNN search: order rows by the distance
|
||||||
|
// between Column and Vector using Metric, and (when As is set) select that
|
||||||
|
// distance as an additional result column.
|
||||||
|
type VectorSearchOption struct {
|
||||||
|
Column string `json:"column"`
|
||||||
|
Vector []float32 `json:"vector"`
|
||||||
|
Metric string `json:"metric"` // "l2" (default) | "cosine" | "ip"
|
||||||
|
As string `json:"as"` // distance column alias; default "_distance"
|
||||||
|
Direction string `json:"direction"` // "asc" (default) | "desc"
|
||||||
|
}
|
||||||
|
|
||||||
type CustomOperator struct {
|
type CustomOperator struct {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
SQL string `json:"sql"`
|
SQL string `json:"sql"`
|
||||||
@@ -104,6 +161,7 @@ type APIError struct {
|
|||||||
Message string `json:"message"`
|
Message string `json:"message"`
|
||||||
Details interface{} `json:"details,omitempty"`
|
Details interface{} `json:"details,omitempty"`
|
||||||
Detail string `json:"detail,omitempty"`
|
Detail string `json:"detail,omitempty"`
|
||||||
|
SQL string `json:"sql,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type Column struct {
|
type Column struct {
|
||||||
|
|||||||
@@ -0,0 +1,58 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestResolveSortColumns(t *testing.T) {
|
||||||
|
sort := []SortOption{
|
||||||
|
{Column: PrimaryKeySortColumn, Direction: "asc"},
|
||||||
|
{Column: "name", Direction: "desc"},
|
||||||
|
}
|
||||||
|
|
||||||
|
got := ResolveSortColumns(sort, "user_id")
|
||||||
|
want := []SortOption{
|
||||||
|
{Column: "user_id", Direction: "asc"},
|
||||||
|
{Column: "name", Direction: "desc"},
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Errorf("ResolveSortColumns() = %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Original slice must not be mutated
|
||||||
|
if sort[0].Column != PrimaryKeySortColumn {
|
||||||
|
t.Errorf("ResolveSortColumns() mutated input slice: %v", sort)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSortColumnsDesc(t *testing.T) {
|
||||||
|
sort := []SortOption{
|
||||||
|
{Column: PrimaryKeySortColumn, Direction: "desc"},
|
||||||
|
}
|
||||||
|
|
||||||
|
got := ResolveSortColumns(sort, "id")
|
||||||
|
want := []SortOption{{Column: "id", Direction: "desc"}}
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Errorf("ResolveSortColumns() = %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSortColumnsEmptyPK(t *testing.T) {
|
||||||
|
sort := []SortOption{
|
||||||
|
{Column: PrimaryKeySortColumn, Direction: "asc"},
|
||||||
|
{Column: "name", Direction: "desc"},
|
||||||
|
}
|
||||||
|
|
||||||
|
got := ResolveSortColumns(sort, "")
|
||||||
|
want := []SortOption{{Column: "name", Direction: "desc"}}
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Errorf("ResolveSortColumns() with empty pk = %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSortColumnsNil(t *testing.T) {
|
||||||
|
if got := ResolveSortColumns(nil, "id"); len(got) != 0 {
|
||||||
|
t.Errorf("ResolveSortColumns(nil) = %v, want empty", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -3,6 +3,7 @@ package common
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
@@ -30,7 +31,7 @@ func (v *ColumnValidator) buildValidColumns() {
|
|||||||
modelType := reflect.TypeOf(v.model)
|
modelType := reflect.TypeOf(v.model)
|
||||||
|
|
||||||
// Unwrap pointers, slices, and arrays to get to the base struct type
|
// Unwrap pointers, slices, and arrays to get to the base struct type
|
||||||
for modelType != nil && (modelType.Kind() == reflect.Ptr || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
modelType = modelType.Elem()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -43,7 +44,7 @@ func (v *ColumnValidator) buildValidColumns() {
|
|||||||
for i := 0; i < modelType.NumField(); i++ {
|
for i := 0; i < modelType.NumField(); i++ {
|
||||||
field := modelType.Field(i)
|
field := modelType.Field(i)
|
||||||
|
|
||||||
if !field.IsExported() {
|
if !field.IsExported() || field.Anonymous {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -108,6 +109,19 @@ func (v *ColumnValidator) ValidateColumn(column string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// JSON-traversing references (data->>'x', data#>>'{a,b}', or the dotted
|
||||||
|
// data.x shorthand): validate the base column, and for the ambiguous
|
||||||
|
// dotted form require that the base is actually a JSON column.
|
||||||
|
if ref, isJSON := ParseColumnRef(column); isJSON {
|
||||||
|
if ref.Ambiguous && !reflection.IsJSONColumn(v.model, ref.Base) {
|
||||||
|
return fmt.Errorf("invalid column '%s': '%s' is not a JSON column", column, ref.Base)
|
||||||
|
}
|
||||||
|
if _, exists := v.validColumns[strings.ToLower(ref.Base)]; !exists {
|
||||||
|
return fmt.Errorf("invalid column '%s': column does not exist in model", column)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Extract source column name (remove JSON operators like ->> or ->)
|
// Extract source column name (remove JSON operators like ->> or ->)
|
||||||
sourceColumn := reflection.ExtractSourceColumn(column)
|
sourceColumn := reflection.ExtractSourceColumn(column)
|
||||||
|
|
||||||
@@ -125,6 +139,16 @@ func (v *ColumnValidator) IsValidColumn(column string) bool {
|
|||||||
return v.ValidateColumn(column) == nil
|
return v.ValidateColumn(column) == nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Columns returns all valid column names known to this validator
|
||||||
|
func (v *ColumnValidator) Columns() []string {
|
||||||
|
cols := make([]string, 0, len(v.validColumns))
|
||||||
|
for col := range v.validColumns {
|
||||||
|
cols = append(cols, col)
|
||||||
|
}
|
||||||
|
sort.Strings(cols)
|
||||||
|
return cols
|
||||||
|
}
|
||||||
|
|
||||||
// FilterValidColumns filters a list of columns, returning only valid ones
|
// FilterValidColumns filters a list of columns, returning only valid ones
|
||||||
// Logs warnings for any invalid columns
|
// Logs warnings for any invalid columns
|
||||||
func (v *ColumnValidator) FilterValidColumns(columns []string) []string {
|
func (v *ColumnValidator) FilterValidColumns(columns []string) []string {
|
||||||
@@ -224,7 +248,19 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
|||||||
// Filter Filter columns
|
// Filter Filter columns
|
||||||
validFilters := make([]FilterOption, 0, len(options.Filters))
|
validFilters := make([]FilterOption, 0, len(options.Filters))
|
||||||
for _, filter := range options.Filters {
|
for _, filter := range options.Filters {
|
||||||
if v.IsValidColumn(filter.Column) {
|
if strings.EqualFold(filter.Column, "all") {
|
||||||
|
allCols := v.Columns()
|
||||||
|
if len(filtered.Columns) > 0 {
|
||||||
|
allCols = filtered.Columns
|
||||||
|
}
|
||||||
|
for _, col := range allCols {
|
||||||
|
expanded := filter
|
||||||
|
expanded.Column = col
|
||||||
|
expanded.LogicOperator = "OR"
|
||||||
|
|
||||||
|
validFilters = append(validFilters, expanded)
|
||||||
|
}
|
||||||
|
} else if v.IsValidColumn(filter.Column) {
|
||||||
validFilters = append(validFilters, filter)
|
validFilters = append(validFilters, filter)
|
||||||
} else {
|
} else {
|
||||||
logger.Warn("Invalid column in filter '%s' removed", filter.Column)
|
logger.Warn("Invalid column in filter '%s' removed", filter.Column)
|
||||||
@@ -266,11 +302,24 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
|||||||
|
|
||||||
// Filter Preload columns
|
// Filter Preload columns
|
||||||
validPreloads := make([]PreloadOption, 0, len(options.Preload))
|
validPreloads := make([]PreloadOption, 0, len(options.Preload))
|
||||||
|
modelType := reflect.TypeOf(v.model)
|
||||||
|
if modelType != nil && modelType.Kind() == reflect.Pointer {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
for idx := range options.Preload {
|
for idx := range options.Preload {
|
||||||
preload := options.Preload[idx]
|
preload := options.Preload[idx]
|
||||||
filteredPreload := preload
|
filteredPreload := preload
|
||||||
filteredPreload.Columns = v.FilterValidColumns(preload.Columns)
|
|
||||||
filteredPreload.OmitColumns = v.FilterValidColumns(preload.OmitColumns)
|
// Use the related model's validator for preload columns/filters/sorts
|
||||||
|
preloadValidator := v
|
||||||
|
if modelType != nil {
|
||||||
|
if relInfo := GetRelationshipInfo(modelType, preload.Relation); relInfo != nil && relInfo.RelatedModel != nil {
|
||||||
|
preloadValidator = NewColumnValidator(relInfo.RelatedModel)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
filteredPreload.Columns = preloadValidator.FilterValidColumns(preload.Columns)
|
||||||
|
filteredPreload.OmitColumns = preloadValidator.FilterValidColumns(preload.OmitColumns)
|
||||||
|
|
||||||
// Preserve SqlJoins and JoinAliases for preloads with custom joins
|
// Preserve SqlJoins and JoinAliases for preloads with custom joins
|
||||||
filteredPreload.SqlJoins = preload.SqlJoins
|
filteredPreload.SqlJoins = preload.SqlJoins
|
||||||
@@ -279,7 +328,7 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
|||||||
// Filter preload filters
|
// Filter preload filters
|
||||||
validPreloadFilters := make([]FilterOption, 0, len(preload.Filters))
|
validPreloadFilters := make([]FilterOption, 0, len(preload.Filters))
|
||||||
for _, filter := range preload.Filters {
|
for _, filter := range preload.Filters {
|
||||||
if v.IsValidColumn(filter.Column) {
|
if preloadValidator.IsValidColumn(filter.Column) {
|
||||||
validPreloadFilters = append(validPreloadFilters, filter)
|
validPreloadFilters = append(validPreloadFilters, filter)
|
||||||
} else {
|
} else {
|
||||||
// Check if the filter column references a joined table alias
|
// Check if the filter column references a joined table alias
|
||||||
@@ -302,7 +351,7 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
|||||||
// Filter preload sort columns
|
// Filter preload sort columns
|
||||||
validPreloadSorts := make([]SortOption, 0, len(preload.Sort))
|
validPreloadSorts := make([]SortOption, 0, len(preload.Sort))
|
||||||
for _, sort := range preload.Sort {
|
for _, sort := range preload.Sort {
|
||||||
if v.IsValidColumn(sort.Column) {
|
if preloadValidator.IsValidColumn(sort.Column) {
|
||||||
validPreloadSorts = append(validPreloadSorts, sort)
|
validPreloadSorts = append(validPreloadSorts, sort)
|
||||||
} else if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
} else if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||||
// Allow sort by expression/subquery, but validate for security
|
// Allow sort by expression/subquery, but validate for security
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestExtractSourceColumn(t *testing.T) {
|
func TestExtractSourceColumn(t *testing.T) {
|
||||||
@@ -124,3 +125,35 @@ func TestValidateColumnWithJSONOperators(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestValidateColumn_JSONPathsAndDottedShorthand(t *testing.T) {
|
||||||
|
type Model struct {
|
||||||
|
ID int64 `json:"id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Data spectypes.SqlJSONB `json:"data"`
|
||||||
|
}
|
||||||
|
v := NewColumnValidator(Model{})
|
||||||
|
|
||||||
|
valid := []string{
|
||||||
|
"data->>'city'",
|
||||||
|
"data->'addr'->>'city'",
|
||||||
|
"data#>>'{addr,city}'",
|
||||||
|
"data.addr.city", // dotted shorthand, base is JSON -> allowed
|
||||||
|
"(data->>'age')::int", // cast + paren
|
||||||
|
}
|
||||||
|
for _, c := range valid {
|
||||||
|
if err := v.ValidateColumn(c); err != nil {
|
||||||
|
t.Errorf("ValidateColumn(%q) = %v, want nil", c, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
invalid := []string{
|
||||||
|
"nope->>'city'", // base column does not exist
|
||||||
|
"name.first", // dotted shorthand but 'name' is not a JSON column
|
||||||
|
}
|
||||||
|
for _, c := range invalid {
|
||||||
|
if err := v.ValidateColumn(c); err == nil {
|
||||||
|
t.Errorf("ValidateColumn(%q) = nil, want error", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -464,3 +464,84 @@ func TestFilterRequestOptions_WithSortExpressions(t *testing.T) {
|
|||||||
t.Errorf("Expected third sort to be 'name', got '%s'", filtered.Sort[2].Column)
|
t.Errorf("Expected third sort to be 'name', got '%s'", filtered.Sort[2].Column)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RelatedModel is used by PreloadParentModel to test preload column validation.
|
||||||
|
type RelatedModel struct {
|
||||||
|
RelatedID int64 `bun:"related_id,pk"`
|
||||||
|
Functionname string `bun:"functionname"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PreloadParentModel has a has-one relation to RelatedModel. The json tag on
|
||||||
|
// the relation field is the name used in x-preload headers.
|
||||||
|
type PreloadParentModel struct {
|
||||||
|
ID int64 `bun:"id,pk"`
|
||||||
|
Name string `bun:"name"`
|
||||||
|
RELATED *RelatedModel `json:"RELATED" bun:"rel:has-one,join:id=related_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFilterRequestOptions_PreloadColumnsValidatedAgainstRelatedModel verifies
|
||||||
|
// that preload columns are validated against the related model's fields, not the
|
||||||
|
// parent model's fields. This is the fix for the bug where specifying a column
|
||||||
|
// that exists only on the relation (e.g. "functionname") was incorrectly filtered
|
||||||
|
// out because it doesn't exist on the parent model.
|
||||||
|
func TestFilterRequestOptions_PreloadColumnsValidatedAgainstRelatedModel(t *testing.T) {
|
||||||
|
validator := NewColumnValidator(PreloadParentModel{})
|
||||||
|
|
||||||
|
options := RequestOptions{
|
||||||
|
Preload: []PreloadOption{
|
||||||
|
{
|
||||||
|
Relation: "RELATED",
|
||||||
|
// "functionname" exists on RelatedModel but NOT on PreloadParentModel.
|
||||||
|
// "name" exists on PreloadParentModel but NOT on RelatedModel.
|
||||||
|
// "nonexistent" exists on neither.
|
||||||
|
Columns: []string{"functionname", "name", "nonexistent"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
filtered := validator.FilterRequestOptions(options)
|
||||||
|
|
||||||
|
if len(filtered.Preload) != 1 {
|
||||||
|
t.Fatalf("Expected 1 preload, got %d", len(filtered.Preload))
|
||||||
|
}
|
||||||
|
|
||||||
|
cols := filtered.Preload[0].Columns
|
||||||
|
// Only "functionname" should survive: it belongs to RelatedModel.
|
||||||
|
if len(cols) != 1 {
|
||||||
|
t.Errorf("Expected 1 preload column, got %d: %v", len(cols), cols)
|
||||||
|
}
|
||||||
|
if len(cols) > 0 && cols[0] != "functionname" {
|
||||||
|
t.Errorf("Expected preload column 'functionname', got '%s'", cols[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFilterRequestOptions_PreloadColumnsParentModelFallback verifies that when
|
||||||
|
// a preload relation is not found on the parent model, column validation falls
|
||||||
|
// back to the parent model's validator (no panic, no silent pass-through).
|
||||||
|
func TestFilterRequestOptions_PreloadColumnsParentModelFallback(t *testing.T) {
|
||||||
|
validator := NewColumnValidator(PreloadParentModel{})
|
||||||
|
|
||||||
|
options := RequestOptions{
|
||||||
|
Preload: []PreloadOption{
|
||||||
|
{
|
||||||
|
Relation: "UNKNOWN_RELATION",
|
||||||
|
Columns: []string{"id", "functionname"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
filtered := validator.FilterRequestOptions(options)
|
||||||
|
|
||||||
|
if len(filtered.Preload) != 1 {
|
||||||
|
t.Fatalf("Expected 1 preload, got %d", len(filtered.Preload))
|
||||||
|
}
|
||||||
|
|
||||||
|
cols := filtered.Preload[0].Columns
|
||||||
|
// Falls back to parent model: only "id" is valid on PreloadParentModel.
|
||||||
|
if len(cols) != 1 {
|
||||||
|
t.Errorf("Expected 1 preload column (fallback to parent), got %d: %v", len(cols), cols)
|
||||||
|
}
|
||||||
|
if len(cols) > 0 && cols[0] != "id" {
|
||||||
|
t.Errorf("Expected preload column 'id', got '%s'", cols[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+34
-1
@@ -1,6 +1,9 @@
|
|||||||
package config
|
package config
|
||||||
|
|
||||||
import "time"
|
import (
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
// Config represents the complete application configuration
|
// Config represents the complete application configuration
|
||||||
type Config struct {
|
type Config struct {
|
||||||
@@ -50,6 +53,10 @@ type ServerInstanceConfig struct {
|
|||||||
// GZIP enables GZIP compression middleware
|
// GZIP enables GZIP compression middleware
|
||||||
GZIP bool `mapstructure:"gzip"`
|
GZIP bool `mapstructure:"gzip"`
|
||||||
|
|
||||||
|
// HTTP2 enables HTTP/2 with the Extended CONNECT protocol (RFC 8441) for WebSocket support.
|
||||||
|
// Requires TLS; pair with SSLCert/SSLKey, SelfSignedSSL, or AutoTLS.
|
||||||
|
HTTP2 bool `mapstructure:"http2"`
|
||||||
|
|
||||||
// TLS/HTTPS configuration options (mutually exclusive)
|
// TLS/HTTPS configuration options (mutually exclusive)
|
||||||
// Option 1: Provide certificate and key files directly
|
// Option 1: Provide certificate and key files directly
|
||||||
SSLCert string `mapstructure:"ssl_cert"`
|
SSLCert string `mapstructure:"ssl_cert"`
|
||||||
@@ -84,6 +91,12 @@ type TracingConfig struct {
|
|||||||
ServiceName string `mapstructure:"service_name"`
|
ServiceName string `mapstructure:"service_name"`
|
||||||
ServiceVersion string `mapstructure:"service_version"`
|
ServiceVersion string `mapstructure:"service_version"`
|
||||||
Endpoint string `mapstructure:"endpoint"`
|
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
|
// CacheConfig holds cache provider configuration
|
||||||
@@ -194,3 +207,23 @@ type EventBrokerRetryPolicyConfig struct {
|
|||||||
// This is a map of path name to file system path
|
// This is a map of path name to file system path
|
||||||
// Example: "data_dir": "/var/lib/myapp/data"
|
// Example: "data_dir": "/var/lib/myapp/data"
|
||||||
type PathsConfig map[string]string
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
+80
-18
@@ -2,37 +2,60 @@ package config
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
"github.com/spf13/viper"
|
"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 {
|
type Manager struct {
|
||||||
|
mu sync.RWMutex
|
||||||
v *viper.Viper
|
v *viper.Viper
|
||||||
}
|
}
|
||||||
|
|
||||||
var configInstance *Manager
|
var (
|
||||||
|
configInstance *Manager
|
||||||
|
configMu sync.Mutex
|
||||||
|
)
|
||||||
|
|
||||||
// GetConfigManager returns a singleton configuration manager instance
|
// GetConfigManager returns a singleton configuration manager instance
|
||||||
func GetConfigManager() *Manager {
|
func GetConfigManager() *Manager {
|
||||||
|
configMu.Lock()
|
||||||
|
defer configMu.Unlock()
|
||||||
if configInstance == nil {
|
if configInstance == nil {
|
||||||
configInstance = NewManager()
|
configInstance = NewManager()
|
||||||
}
|
}
|
||||||
return configInstance
|
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 {
|
func NewManager() *Manager {
|
||||||
v := viper.New()
|
v := viper.New()
|
||||||
|
|
||||||
// Set configuration file settings
|
// Set configuration file settings
|
||||||
v.SetConfigName("config")
|
v.SetConfigName("config")
|
||||||
v.SetConfigType("yaml")
|
v.SetConfigType("yaml")
|
||||||
v.AddConfigPath(".")
|
// Most trusted location first; the working directory is the least trustworthy
|
||||||
v.AddConfigPath("./config")
|
// and is searched last (viper takes the first match).
|
||||||
v.AddConfigPath("/etc/resolvespec")
|
v.AddConfigPath("/etc/resolvespec")
|
||||||
v.AddConfigPath("$HOME/.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
|
// Enable environment variable support
|
||||||
v.SetEnvPrefix("RESOLVESPEC")
|
v.SetEnvPrefix("RESOLVESPEC")
|
||||||
@@ -42,8 +65,7 @@ func NewManager() *Manager {
|
|||||||
// Set default values
|
// Set default values
|
||||||
setDefaults(v)
|
setDefaults(v)
|
||||||
|
|
||||||
configInstance = &Manager{v: v}
|
return &Manager{v: v}
|
||||||
return configInstance
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewManagerWithOptions creates a new configuration manager with custom options
|
// NewManagerWithOptions creates a new configuration manager with custom options
|
||||||
@@ -61,6 +83,8 @@ type Option func(*Manager)
|
|||||||
// WithConfigFile sets a specific config file path
|
// WithConfigFile sets a specific config file path
|
||||||
func WithConfigFile(path string) Option {
|
func WithConfigFile(path string) Option {
|
||||||
return func(m *Manager) {
|
return func(m *Manager) {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
m.v.SetConfigFile(path)
|
m.v.SetConfigFile(path)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -68,6 +92,8 @@ func WithConfigFile(path string) Option {
|
|||||||
// WithConfigName sets the config file name (without extension)
|
// WithConfigName sets the config file name (without extension)
|
||||||
func WithConfigName(name string) Option {
|
func WithConfigName(name string) Option {
|
||||||
return func(m *Manager) {
|
return func(m *Manager) {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
m.v.SetConfigName(name)
|
m.v.SetConfigName(name)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -75,6 +101,8 @@ func WithConfigName(name string) Option {
|
|||||||
// WithConfigPath adds a path to search for config files
|
// WithConfigPath adds a path to search for config files
|
||||||
func WithConfigPath(path string) Option {
|
func WithConfigPath(path string) Option {
|
||||||
return func(m *Manager) {
|
return func(m *Manager) {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
m.v.AddConfigPath(path)
|
m.v.AddConfigPath(path)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -82,13 +110,19 @@ func WithConfigPath(path string) Option {
|
|||||||
// WithEnvPrefix sets the environment variable prefix
|
// WithEnvPrefix sets the environment variable prefix
|
||||||
func WithEnvPrefix(prefix string) Option {
|
func WithEnvPrefix(prefix string) Option {
|
||||||
return func(m *Manager) {
|
return func(m *Manager) {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
m.v.SetEnvPrefix(prefix)
|
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 {
|
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 err := m.v.ReadInConfig(); err != nil {
|
||||||
if _, ok := err.(viper.ConfigFileNotFoundError); !ok {
|
if _, ok := err.(viper.ConfigFileNotFoundError); !ok {
|
||||||
return fmt.Errorf("error reading config file: %w", err)
|
return fmt.Errorf("error reading config file: %w", err)
|
||||||
@@ -99,8 +133,19 @@ func (m *Manager) Load() error {
|
|||||||
return nil
|
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
|
// GetConfig returns the complete configuration
|
||||||
func (m *Manager) GetConfig() (*Config, error) {
|
func (m *Manager) GetConfig() (*Config, error) {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
|
|
||||||
var cfg Config
|
var cfg Config
|
||||||
if err := m.v.Unmarshal(&cfg); err != nil {
|
if err := m.v.Unmarshal(&cfg); err != nil {
|
||||||
return nil, fmt.Errorf("failed to unmarshal config: %w", err)
|
return nil, fmt.Errorf("failed to unmarshal config: %w", err)
|
||||||
@@ -108,16 +153,11 @@ func (m *Manager) GetConfig() (*Config, error) {
|
|||||||
return &cfg, nil
|
return &cfg, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetConfig sets the complete configuration
|
// SetConfig sets the complete configuration atomically
|
||||||
func (m *Manager) SetConfig(cfg *Config) error {
|
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("servers", cfg.Servers)
|
||||||
m.v.Set("tracing", cfg.Tracing)
|
m.v.Set("tracing", cfg.Tracing)
|
||||||
m.v.Set("cache", cfg.Cache)
|
m.v.Set("cache", cfg.Cache)
|
||||||
@@ -135,34 +175,54 @@ func (m *Manager) SetConfig(cfg *Config) error {
|
|||||||
|
|
||||||
// Get returns a configuration value by key
|
// Get returns a configuration value by key
|
||||||
func (m *Manager) Get(key string) interface{} {
|
func (m *Manager) Get(key string) interface{} {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
return m.v.Get(key)
|
return m.v.Get(key)
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetString returns a string configuration value
|
// GetString returns a string configuration value
|
||||||
func (m *Manager) GetString(key string) string {
|
func (m *Manager) GetString(key string) string {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
return m.v.GetString(key)
|
return m.v.GetString(key)
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetInt returns an int configuration value
|
// GetInt returns an int configuration value
|
||||||
func (m *Manager) GetInt(key string) int {
|
func (m *Manager) GetInt(key string) int {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
return m.v.GetInt(key)
|
return m.v.GetInt(key)
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetBool returns a bool configuration value
|
// GetBool returns a bool configuration value
|
||||||
func (m *Manager) GetBool(key string) bool {
|
func (m *Manager) GetBool(key string) bool {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
return m.v.GetBool(key)
|
return m.v.GetBool(key)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Set sets a configuration value
|
// Set sets a configuration value
|
||||||
func (m *Manager) Set(key string, value interface{}) {
|
func (m *Manager) Set(key string, value interface{}) {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
m.v.Set(key, value)
|
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 {
|
func (m *Manager) SaveConfig(path string) error {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
|
|
||||||
if err := m.v.WriteConfigAs(path); err != nil {
|
if err := m.v.WriteConfigAs(path); err != nil {
|
||||||
return fmt.Errorf("failed to save config to %s: %w", path, err)
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -190,6 +250,8 @@ func setDefaults(v *viper.Viper) {
|
|||||||
v.SetDefault("tracing.service_name", "resolvespec")
|
v.SetDefault("tracing.service_name", "resolvespec")
|
||||||
v.SetDefault("tracing.service_version", "1.0.0")
|
v.SetDefault("tracing.service_version", "1.0.0")
|
||||||
v.SetDefault("tracing.endpoint", "")
|
v.SetDefault("tracing.endpoint", "")
|
||||||
|
v.SetDefault("tracing.insecure", false)
|
||||||
|
v.SetDefault("tracing.sample_rate", 0.1)
|
||||||
|
|
||||||
// Cache defaults
|
// Cache defaults
|
||||||
v.SetDefault("cache.provider", "memory")
|
v.SetDefault("cache.provider", "memory")
|
||||||
|
|||||||
+18
-5
@@ -4,6 +4,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Get retrieves a path by name
|
// Get retrieves a path by name
|
||||||
@@ -34,9 +35,13 @@ func (pc PathsConfig) GetOrDefault(name, defaultPath string) string {
|
|||||||
return path
|
return path
|
||||||
}
|
}
|
||||||
|
|
||||||
// Set sets a path by name
|
// Set sets a path by name. It takes a pointer so a nil map can be allocated.
|
||||||
func (pc PathsConfig) Set(name, path string) {
|
// PathsConfig is not safe for concurrent mutation; populate it before sharing.
|
||||||
pc[name] = path
|
func (pc *PathsConfig) Set(name, path string) {
|
||||||
|
if *pc == nil {
|
||||||
|
*pc = make(PathsConfig)
|
||||||
|
}
|
||||||
|
(*pc)[name] = path
|
||||||
}
|
}
|
||||||
|
|
||||||
// Has checks if a path exists by name
|
// Has checks if a path exists by name
|
||||||
@@ -92,7 +97,8 @@ func (pc PathsConfig) AbsPath(name string) (string, error) {
|
|||||||
return absPath, nil
|
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) {
|
func (pc PathsConfig) Join(name string, elem ...string) (string, error) {
|
||||||
base, err := pc.Get(name)
|
base, err := pc.Get(name)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -100,7 +106,14 @@ func (pc PathsConfig) Join(name string, elem ...string) (string, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
parts := append([]string{base}, elem...)
|
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
|
// List returns all configured path names
|
||||||
|
|||||||
+26
-27
@@ -1,10 +1,12 @@
|
|||||||
package config
|
package config
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ApplyGlobalDefaults applies global server defaults to this instance
|
// ApplyGlobalDefaults applies global server defaults to this instance
|
||||||
@@ -95,7 +97,8 @@ func (sc *ServersConfig) Validate() error {
|
|||||||
return nil
|
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) {
|
func (sc *ServersConfig) GetDefault() (*ServerInstanceConfig, error) {
|
||||||
if sc.DefaultServer == "" {
|
if sc.DefaultServer == "" {
|
||||||
return nil, fmt.Errorf("no default server configured")
|
return nil, fmt.Errorf("no default server configured")
|
||||||
@@ -109,41 +112,37 @@ func (sc *ServersConfig) GetDefault() (*ServerInstanceConfig, error) {
|
|||||||
return &instance, nil
|
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) {
|
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()
|
hostname, _ = os.Hostname()
|
||||||
ipaddrlist := make([]net.IP, 0)
|
ipNetList = make([]net.IP, 0)
|
||||||
iplist := ""
|
|
||||||
addrs, err := net.LookupIP(hostname)
|
|
||||||
if err != nil {
|
|
||||||
return hostname, iplist, ipaddrlist
|
|
||||||
}
|
|
||||||
|
|
||||||
|
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 {
|
for _, a := range addrs {
|
||||||
// cfg.LogInfo("\nFound IP Host Address: %s", a)
|
if a.IP.IsLoopback() {
|
||||||
if strings.Contains(a.String(), "127.0.0.1") {
|
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
iplist = fmt.Sprintf("%s,%s", iplist, a)
|
ips = append(ips, a.IP.String())
|
||||||
ipaddrlist = append(ipaddrlist, a)
|
ipNetList = append(ipNetList, a.IP)
|
||||||
}
|
}
|
||||||
if iplist == "" {
|
}
|
||||||
iff, _ := net.InterfaceAddrs()
|
|
||||||
for _, a := range iff {
|
if len(ips) == 0 {
|
||||||
// cfg.LogInfo("\nFound IP Address: %s", a)
|
ifaceAddrs, _ := net.InterfaceAddrs()
|
||||||
if strings.Contains(a.String(), "127.0.0.1") {
|
for _, a := range ifaceAddrs {
|
||||||
|
ipn, ok := a.(*net.IPNet)
|
||||||
|
if !ok || ipn.IP.IsLoopback() {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
iplist = fmt.Sprintf("%s,%s", iplist, a)
|
ips = append(ips, ipn.IP.String())
|
||||||
|
ipNetList = append(ipNetList, ipn.IP)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
return hostname, strings.Join(ips, ","), ipNetList
|
||||||
iplist = strings.TrimLeft(iplist, ",")
|
|
||||||
return hostname, iplist, ipaddrlist
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -50,7 +50,6 @@ dbmanager:
|
|||||||
|
|
||||||
# Health checks
|
# Health checks
|
||||||
health_check_interval: 30s
|
health_check_interval: 30s
|
||||||
enable_auto_reconnect: true
|
|
||||||
|
|
||||||
connections:
|
connections:
|
||||||
# Primary PostgreSQL connection
|
# Primary PostgreSQL connection
|
||||||
@@ -256,7 +255,7 @@ db, _ := mgr.GetDefaultDatabase()
|
|||||||
| `retry_delay` | duration | 1s | Initial retry delay |
|
| `retry_delay` | duration | 1s | Initial retry delay |
|
||||||
| `retry_max_delay` | duration | 10s | Maximum retry delay |
|
| `retry_max_delay` | duration | 10s | Maximum retry delay |
|
||||||
| `health_check_interval` | duration | 30s | Interval between health checks |
|
| `health_check_interval` | duration | 30s | Interval between health checks |
|
||||||
| `enable_auto_reconnect` | bool | true | Auto-reconnect on health check failure |
|
| `enable_auto_reconnect` | bool | - | Deprecated and ignored: the manager never closes the pool to recover from errors |
|
||||||
|
|
||||||
### Connection Configuration
|
### Connection Configuration
|
||||||
|
|
||||||
@@ -451,7 +450,6 @@ db.NewSelect().Model(&User{}).Scan(ctx)
|
|||||||
3. **Enable Health Checks**: Catch connection issues early
|
3. **Enable Health Checks**: Catch connection issues early
|
||||||
```yaml
|
```yaml
|
||||||
health_check_interval: 30s
|
health_check_interval: 30s
|
||||||
enable_auto_reconnect: true
|
|
||||||
```
|
```
|
||||||
|
|
||||||
4. **Use Appropriate ORM**: Choose based on your needs
|
4. **Use Appropriate ORM**: Choose based on your needs
|
||||||
|
|||||||
+105
-70
@@ -2,6 +2,10 @@ package dbmanager
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/url"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/config"
|
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||||
@@ -57,8 +61,14 @@ type ManagerConfig struct {
|
|||||||
RetryDelay time.Duration `mapstructure:"retry_delay"`
|
RetryDelay time.Duration `mapstructure:"retry_delay"`
|
||||||
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
|
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
|
||||||
|
|
||||||
// Health checks
|
// Health checks. A zero HealthCheckInterval selects the default (15s); a
|
||||||
|
// negative value disables the background health checker.
|
||||||
HealthCheckInterval time.Duration `mapstructure:"health_check_interval"`
|
HealthCheckInterval time.Duration `mapstructure:"health_check_interval"`
|
||||||
|
|
||||||
|
// 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"`
|
EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -103,6 +113,11 @@ type ConnectionConfig struct {
|
|||||||
ConnectTimeout time.Duration `mapstructure:"connect_timeout"`
|
ConnectTimeout time.Duration `mapstructure:"connect_timeout"`
|
||||||
QueryTimeout time.Duration `mapstructure:"query_timeout"`
|
QueryTimeout time.Duration `mapstructure:"query_timeout"`
|
||||||
|
|
||||||
|
// Retry policy for the initial connect (inherited from the manager config)
|
||||||
|
RetryAttempts int `mapstructure:"retry_attempts"`
|
||||||
|
RetryDelay time.Duration `mapstructure:"retry_delay"`
|
||||||
|
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
|
||||||
|
|
||||||
// Features
|
// Features
|
||||||
EnableTracing bool `mapstructure:"enable_tracing"`
|
EnableTracing bool `mapstructure:"enable_tracing"`
|
||||||
EnableMetrics bool `mapstructure:"enable_metrics"`
|
EnableMetrics bool `mapstructure:"enable_metrics"`
|
||||||
@@ -129,7 +144,6 @@ func DefaultManagerConfig() ManagerConfig {
|
|||||||
RetryDelay: 1 * time.Second,
|
RetryDelay: 1 * time.Second,
|
||||||
RetryMaxDelay: 10 * time.Second,
|
RetryMaxDelay: 10 * time.Second,
|
||||||
HealthCheckInterval: 15 * time.Second,
|
HealthCheckInterval: 15 * time.Second,
|
||||||
EnableAutoReconnect: true,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -161,11 +175,6 @@ func (c *ManagerConfig) ApplyDefaults() {
|
|||||||
if c.HealthCheckInterval == 0 {
|
if c.HealthCheckInterval == 0 {
|
||||||
c.HealthCheckInterval = defaults.HealthCheckInterval
|
c.HealthCheckInterval = defaults.HealthCheckInterval
|
||||||
}
|
}
|
||||||
// EnableAutoReconnect defaults to true - apply if not explicitly set
|
|
||||||
// Since this is a boolean, we apply the default unconditionally when it's false
|
|
||||||
if !c.EnableAutoReconnect {
|
|
||||||
c.EnableAutoReconnect = defaults.EnableAutoReconnect
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate validates the manager configuration
|
// Validate validates the manager configuration
|
||||||
@@ -222,9 +231,18 @@ func (cc *ConnectionConfig) ApplyDefaults(global *ManagerConfig) {
|
|||||||
}
|
}
|
||||||
if cc.QueryTimeout == 0 {
|
if cc.QueryTimeout == 0 {
|
||||||
cc.QueryTimeout = 2 * time.Minute // Default to 2 minutes
|
cc.QueryTimeout = 2 * time.Minute // Default to 2 minutes
|
||||||
} else if cc.QueryTimeout < 2*time.Minute {
|
}
|
||||||
// Enforce minimum of 2 minutes
|
|
||||||
cc.QueryTimeout = 2 * time.Minute
|
if global != nil {
|
||||||
|
if cc.RetryAttempts == 0 {
|
||||||
|
cc.RetryAttempts = global.RetryAttempts
|
||||||
|
}
|
||||||
|
if cc.RetryDelay == 0 {
|
||||||
|
cc.RetryDelay = global.RetryDelay
|
||||||
|
}
|
||||||
|
if cc.RetryMaxDelay == 0 {
|
||||||
|
cc.RetryMaxDelay = global.RetryMaxDelay
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Default ORM
|
// Default ORM
|
||||||
@@ -314,108 +332,122 @@ func (cc *ConnectionConfig) BuildDSN() (string, error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// buildPostgresDSN builds a postgres:// URL so credentials and other values are
|
||||||
|
// escaped rather than spliced into a key=value string. statement_timeout is
|
||||||
|
// applied by the provider as a runtime parameter.
|
||||||
func (cc *ConnectionConfig) buildPostgresDSN() string {
|
func (cc *ConnectionConfig) buildPostgresDSN() string {
|
||||||
dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s",
|
q := url.Values{}
|
||||||
cc.Host, cc.Port, cc.User, cc.Password, cc.Database)
|
|
||||||
|
|
||||||
if cc.SSLMode != "" {
|
if cc.SSLMode != "" {
|
||||||
dsn += fmt.Sprintf(" sslmode=%s", cc.SSLMode)
|
q.Set("sslmode", cc.SSLMode)
|
||||||
} else {
|
} else {
|
||||||
dsn += " sslmode=disable"
|
// prefer: use TLS when the server offers it, without failing on
|
||||||
|
// servers that do not.
|
||||||
|
q.Set("sslmode", "prefer")
|
||||||
}
|
}
|
||||||
|
|
||||||
if cc.Schema != "" {
|
if cc.Schema != "" {
|
||||||
dsn += fmt.Sprintf(" search_path=%s", cc.Schema)
|
q.Set("search_path", cc.Schema)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add statement_timeout for query execution timeout (in milliseconds)
|
u := url.URL{
|
||||||
if cc.QueryTimeout > 0 {
|
Scheme: "postgres",
|
||||||
timeoutMs := int(cc.QueryTimeout.Milliseconds())
|
Host: hostPort(cc.Host, cc.Port),
|
||||||
dsn += fmt.Sprintf(" statement_timeout=%d", timeoutMs)
|
Path: "/" + cc.Database,
|
||||||
|
RawQuery: q.Encode(),
|
||||||
}
|
}
|
||||||
|
if cc.User != "" || cc.Password != "" {
|
||||||
return dsn
|
u.User = url.UserPassword(cc.User, cc.Password)
|
||||||
|
}
|
||||||
|
return u.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func hostPort(host string, port int) string {
|
||||||
|
if port == 0 {
|
||||||
|
return host
|
||||||
|
}
|
||||||
|
// JoinHostPort brackets IPv6 literals.
|
||||||
|
return net.JoinHostPort(host, strconv.Itoa(port))
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildSQLiteDSN puts per-connection settings in the DSN as _pragma parameters
|
||||||
|
// so every pooled connection gets them, not just the one that ran an Exec.
|
||||||
func (cc *ConnectionConfig) buildSQLiteDSN() string {
|
func (cc *ConnectionConfig) buildSQLiteDSN() string {
|
||||||
filepath := cc.FilePath
|
filepath := cc.FilePath
|
||||||
if filepath == "" {
|
if filepath == "" {
|
||||||
filepath = ":memory:"
|
filepath = ":memory:"
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add query parameters for timeouts
|
var pragmas []string
|
||||||
// Note: SQLite driver supports _timeout parameter (in milliseconds)
|
|
||||||
if cc.QueryTimeout > 0 {
|
if cc.QueryTimeout > 0 {
|
||||||
timeoutMs := int(cc.QueryTimeout.Milliseconds())
|
pragmas = append(pragmas, fmt.Sprintf("busy_timeout(%d)", cc.QueryTimeout.Milliseconds()))
|
||||||
filepath += fmt.Sprintf("?_timeout=%d", timeoutMs)
|
}
|
||||||
|
if filepath != ":memory:" {
|
||||||
|
pragmas = append(pragmas, "journal_mode(WAL)")
|
||||||
|
}
|
||||||
|
if len(pragmas) == 0 {
|
||||||
|
return filepath
|
||||||
}
|
}
|
||||||
|
|
||||||
return filepath
|
q := url.Values{}
|
||||||
|
for _, p := range pragmas {
|
||||||
|
q.Add("_pragma", p)
|
||||||
|
}
|
||||||
|
sep := "?"
|
||||||
|
if strings.Contains(filepath, "?") {
|
||||||
|
sep = "&"
|
||||||
|
}
|
||||||
|
return filepath + sep + q.Encode()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cc *ConnectionConfig) buildMSSQLDSN() string {
|
func (cc *ConnectionConfig) buildMSSQLDSN() string {
|
||||||
// Format: sqlserver://username:password@host:port?database=dbname
|
// Format: sqlserver://username:password@host:port?database=dbname
|
||||||
dsn := fmt.Sprintf("sqlserver://%s:%s@%s:%d?database=%s",
|
q := url.Values{}
|
||||||
cc.User, cc.Password, cc.Host, cc.Port, cc.Database)
|
q.Set("database", cc.Database)
|
||||||
|
|
||||||
if cc.Schema != "" {
|
if cc.Schema != "" {
|
||||||
dsn += fmt.Sprintf("&schema=%s", cc.Schema)
|
q.Set("schema", cc.Schema)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add connection timeout (in seconds)
|
|
||||||
if cc.ConnectTimeout > 0 {
|
if cc.ConnectTimeout > 0 {
|
||||||
timeoutSec := int(cc.ConnectTimeout.Seconds())
|
sec := strconv.Itoa(int(cc.ConnectTimeout.Seconds()))
|
||||||
dsn += fmt.Sprintf("&connection timeout=%d", timeoutSec)
|
q.Set("connection timeout", sec)
|
||||||
|
q.Set("dial timeout", sec)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add dial timeout for TCP connection (in seconds)
|
|
||||||
if cc.ConnectTimeout > 0 {
|
|
||||||
dialTimeoutSec := int(cc.ConnectTimeout.Seconds())
|
|
||||||
dsn += fmt.Sprintf("&dial timeout=%d", dialTimeoutSec)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Add read timeout (in seconds) - enforces timeout for reading data
|
|
||||||
if cc.QueryTimeout > 0 {
|
if cc.QueryTimeout > 0 {
|
||||||
readTimeoutSec := int(cc.QueryTimeout.Seconds())
|
q.Set("read timeout", strconv.Itoa(int(cc.QueryTimeout.Seconds())))
|
||||||
dsn += fmt.Sprintf("&read timeout=%d", readTimeoutSec)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return dsn
|
u := url.URL{
|
||||||
|
Scheme: "sqlserver",
|
||||||
|
Host: hostPort(cc.Host, cc.Port),
|
||||||
|
RawQuery: q.Encode(),
|
||||||
|
}
|
||||||
|
if cc.User != "" || cc.Password != "" {
|
||||||
|
u.User = url.UserPassword(cc.User, cc.Password)
|
||||||
|
}
|
||||||
|
return u.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cc *ConnectionConfig) buildMongoDSN() string {
|
func (cc *ConnectionConfig) buildMongoDSN() string {
|
||||||
// Format: mongodb://username:password@host:port/database?authSource=admin
|
// Format: mongodb://username:password@host:port/database?authSource=admin
|
||||||
var dsn string
|
q := url.Values{}
|
||||||
|
|
||||||
if cc.User != "" && cc.Password != "" {
|
|
||||||
dsn = fmt.Sprintf("mongodb://%s:%s@%s:%d/%s",
|
|
||||||
cc.User, cc.Password, cc.Host, cc.Port, cc.Database)
|
|
||||||
} else {
|
|
||||||
dsn = fmt.Sprintf("mongodb://%s:%d/%s", cc.Host, cc.Port, cc.Database)
|
|
||||||
}
|
|
||||||
|
|
||||||
params := ""
|
|
||||||
if cc.AuthSource != "" {
|
if cc.AuthSource != "" {
|
||||||
params += fmt.Sprintf("authSource=%s", cc.AuthSource)
|
q.Set("authSource", cc.AuthSource)
|
||||||
}
|
}
|
||||||
if cc.ReplicaSet != "" {
|
if cc.ReplicaSet != "" {
|
||||||
if params != "" {
|
q.Set("replicaSet", cc.ReplicaSet)
|
||||||
params += "&"
|
|
||||||
}
|
|
||||||
params += fmt.Sprintf("replicaSet=%s", cc.ReplicaSet)
|
|
||||||
}
|
}
|
||||||
if cc.ReadPreference != "" {
|
if cc.ReadPreference != "" {
|
||||||
if params != "" {
|
q.Set("readPreference", cc.ReadPreference)
|
||||||
params += "&"
|
|
||||||
}
|
|
||||||
params += fmt.Sprintf("readPreference=%s", cc.ReadPreference)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if params != "" {
|
u := url.URL{
|
||||||
dsn += "?" + params
|
Scheme: "mongodb",
|
||||||
|
Host: hostPort(cc.Host, cc.Port),
|
||||||
|
Path: "/" + cc.Database,
|
||||||
|
RawQuery: q.Encode(),
|
||||||
}
|
}
|
||||||
|
if cc.User != "" && cc.Password != "" {
|
||||||
return dsn
|
u.User = url.UserPassword(cc.User, cc.Password)
|
||||||
|
}
|
||||||
|
return u.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
// FromConfig converts config.DBManagerConfig to internal ManagerConfig
|
// FromConfig converts config.DBManagerConfig to internal ManagerConfig
|
||||||
@@ -487,3 +519,6 @@ func (cc *ConnectionConfig) GetConnMaxIdleTime() *time.Duration { return cc.Conn
|
|||||||
func (cc *ConnectionConfig) GetQueryTimeout() time.Duration { return cc.QueryTimeout }
|
func (cc *ConnectionConfig) GetQueryTimeout() time.Duration { return cc.QueryTimeout }
|
||||||
func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics }
|
func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics }
|
||||||
func (cc *ConnectionConfig) GetReadPreference() string { return cc.ReadPreference }
|
func (cc *ConnectionConfig) GetReadPreference() string { return cc.ReadPreference }
|
||||||
|
func (cc *ConnectionConfig) GetRetryAttempts() int { return cc.RetryAttempts }
|
||||||
|
func (cc *ConnectionConfig) GetRetryDelay() time.Duration { return cc.RetryDelay }
|
||||||
|
func (cc *ConnectionConfig) GetRetryMaxDelay() time.Duration { return cc.RetryMaxDelay }
|
||||||
|
|||||||
@@ -0,0 +1,95 @@
|
|||||||
|
package dbmanager
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPostgresDSNEscapesCredentials(t *testing.T) {
|
||||||
|
cc := ConnectionConfig{
|
||||||
|
Type: DatabaseTypePostgreSQL, Host: "db", Port: 5432, Database: "app",
|
||||||
|
User: "u@x", Password: "p w'd sslmode=disable&x=y/?#",
|
||||||
|
}
|
||||||
|
dsn := cc.buildPostgresDSN()
|
||||||
|
u, err := url.Parse(dsn)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("DSN is not a valid URL: %v", err)
|
||||||
|
}
|
||||||
|
if pw, _ := u.User.Password(); pw != cc.Password {
|
||||||
|
t.Errorf("password did not round-trip: %q", pw)
|
||||||
|
}
|
||||||
|
if u.User.Username() != cc.User {
|
||||||
|
t.Errorf("user did not round-trip: %q", u.User.Username())
|
||||||
|
}
|
||||||
|
if got := u.Query().Get("sslmode"); got != "prefer" {
|
||||||
|
t.Errorf("sslmode = %q, want prefer (password must not inject parameters)", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMSSQLAndMongoDSNEscapeCredentials(t *testing.T) {
|
||||||
|
cc := ConnectionConfig{Host: "h", Port: 1, Database: "d", User: "u", Password: "a@b:c/d?e&f"}
|
||||||
|
for name, dsn := range map[string]string{"mssql": cc.buildMSSQLDSN(), "mongo": cc.buildMongoDSN()} {
|
||||||
|
u, err := url.Parse(dsn)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("%s: %v", name, err)
|
||||||
|
}
|
||||||
|
if pw, _ := u.User.Password(); pw != cc.Password {
|
||||||
|
t.Errorf("%s: password did not round-trip: %q", name, pw)
|
||||||
|
}
|
||||||
|
if u.Host != "h:1" {
|
||||||
|
t.Errorf("%s: host = %q", name, u.Host)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSQLiteDSNUsesPragmas(t *testing.T) {
|
||||||
|
cc := ConnectionConfig{FilePath: "/tmp/x.db", QueryTimeout: 3 * time.Second}
|
||||||
|
dsn := cc.buildSQLiteDSN()
|
||||||
|
if strings.Contains(dsn, "?_timeout=") {
|
||||||
|
t.Errorf("unsupported _timeout parameter present: %s", dsn)
|
||||||
|
}
|
||||||
|
if !strings.Contains(dsn, "busy_timeout%283000%29") || !strings.Contains(dsn, "journal_mode%28WAL%29") {
|
||||||
|
t.Errorf("expected busy_timeout and WAL pragmas in DSN: %s", dsn)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQueryTimeoutHonoredWithoutFloor(t *testing.T) {
|
||||||
|
cc := ConnectionConfig{QueryTimeout: 30 * time.Second}
|
||||||
|
cc.ApplyDefaults(&ManagerConfig{})
|
||||||
|
if cc.QueryTimeout != 30*time.Second {
|
||||||
|
t.Errorf("QueryTimeout = %v, want 30s", cc.QueryTimeout)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRetryPolicyInherited(t *testing.T) {
|
||||||
|
g := ManagerConfig{RetryAttempts: 5, RetryDelay: time.Second, RetryMaxDelay: time.Minute}
|
||||||
|
cc := ConnectionConfig{}
|
||||||
|
cc.ApplyDefaults(&g)
|
||||||
|
if cc.GetRetryAttempts() != 5 || cc.GetRetryMaxDelay() != time.Minute {
|
||||||
|
t.Errorf("retry policy not inherited: %+v", cc)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSQLiteMemoryPoolPinned(t *testing.T) {
|
||||||
|
mgr, _ := NewManager(sqliteManagerConfig())
|
||||||
|
if err := mgr.Connect(t.Context()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer mgr.Close()
|
||||||
|
conn, _ := mgr.GetDefault()
|
||||||
|
db, _ := conn.Native()
|
||||||
|
if got := db.Stats().MaxOpenConnections; got != 1 {
|
||||||
|
t.Fatalf("MaxOpenConnections = %d, want 1 for :memory:", got)
|
||||||
|
}
|
||||||
|
if _, err := db.Exec("CREATE TABLE t(a int)"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for i := 0; i < 5; i++ {
|
||||||
|
var n int
|
||||||
|
if err := db.QueryRow("SELECT count(*) FROM t").Scan(&n); err != nil {
|
||||||
|
t.Fatalf("table missing on later use: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+150
-65
@@ -3,6 +3,7 @@ package dbmanager
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -14,6 +15,7 @@ import (
|
|||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Connection represents a single named database connection
|
// Connection represents a single named database connection
|
||||||
@@ -82,6 +84,9 @@ type sqlConnection struct {
|
|||||||
// State
|
// State
|
||||||
connected bool
|
connected bool
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
|
// lifecycleMu serialises Connect/Close/Reconnect against health-check pings.
|
||||||
|
// Lock order: lifecycleMu before mu.
|
||||||
|
lifecycleMu sync.RWMutex
|
||||||
|
|
||||||
// Health check
|
// Health check
|
||||||
lastHealthCheck time.Time
|
lastHealthCheck time.Time
|
||||||
@@ -110,9 +115,16 @@ func (c *sqlConnection) Type() DatabaseType {
|
|||||||
|
|
||||||
// Connect establishes the database connection
|
// Connect establishes the database connection
|
||||||
func (c *sqlConnection) Connect(ctx context.Context) error {
|
func (c *sqlConnection) Connect(ctx context.Context) error {
|
||||||
|
c.lifecycleMu.Lock()
|
||||||
|
defer c.lifecycleMu.Unlock()
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
return c.connectLocked(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// connectLocked requires lifecycleMu and mu held for writing.
|
||||||
|
func (c *sqlConnection) connectLocked(ctx context.Context) error {
|
||||||
if c.connected {
|
if c.connected {
|
||||||
return ErrAlreadyConnected
|
return ErrAlreadyConnected
|
||||||
}
|
}
|
||||||
@@ -127,17 +139,29 @@ func (c *sqlConnection) Connect(ctx context.Context) error {
|
|||||||
|
|
||||||
// Close closes the database connection and all ORM instances
|
// Close closes the database connection and all ORM instances
|
||||||
func (c *sqlConnection) Close() error {
|
func (c *sqlConnection) Close() error {
|
||||||
|
c.lifecycleMu.Lock()
|
||||||
|
defer c.lifecycleMu.Unlock()
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
return c.closeLocked()
|
||||||
|
}
|
||||||
|
|
||||||
|
// closeLocked requires lifecycleMu and mu held for writing. The connection is
|
||||||
|
// always marked disconnected and its cached handles dropped, even when closing
|
||||||
|
// fails, so accessors never hand out handles over a half-closed pool.
|
||||||
|
func (c *sqlConnection) closeLocked() error {
|
||||||
if !c.connected {
|
if !c.connected {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close Bun if initialized
|
var errs []error
|
||||||
if c.bunDB != nil {
|
|
||||||
|
// Close Bun if initialized. bun.DB.Close closes the underlying *sql.DB, so
|
||||||
|
// skip it when the pool belongs to the caller.
|
||||||
|
if o, ok := c.provider.(interface{ OwnsDB() bool }); c.bunDB != nil && (!ok || o.OwnsDB()) {
|
||||||
if err := c.bunDB.Close(); err != nil {
|
if err := c.bunDB.Close(); err != nil {
|
||||||
return NewConnectionError(c.name, "close bun", err)
|
errs = append(errs, NewConnectionError(c.name, "close bun", err))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -145,7 +169,7 @@ func (c *sqlConnection) Close() error {
|
|||||||
|
|
||||||
// Close the provider (which closes the underlying sql.DB)
|
// Close the provider (which closes the underlying sql.DB)
|
||||||
if err := c.provider.Close(); err != nil {
|
if err := c.provider.Close(); err != nil {
|
||||||
return NewConnectionError(c.name, "close", err)
|
errs = append(errs, NewConnectionError(c.name, "close", err))
|
||||||
}
|
}
|
||||||
|
|
||||||
c.connected = false
|
c.connected = false
|
||||||
@@ -156,39 +180,75 @@ func (c *sqlConnection) Close() error {
|
|||||||
c.gormAdapter = nil
|
c.gormAdapter = nil
|
||||||
c.nativeAdapter = nil
|
c.nativeAdapter = nil
|
||||||
|
|
||||||
return nil
|
return errors.Join(errs...)
|
||||||
}
|
}
|
||||||
|
|
||||||
// HealthCheck verifies the connection is alive
|
// HealthCheck verifies the connection is alive. The network ping runs without
|
||||||
|
// holding mu, so handle accessors are never blocked behind a slow ping.
|
||||||
func (c *sqlConnection) HealthCheck(ctx context.Context) error {
|
func (c *sqlConnection) HealthCheck(ctx context.Context) error {
|
||||||
if c == nil {
|
if c == nil {
|
||||||
return fmt.Errorf("connection is nil")
|
return fmt.Errorf("connection is nil")
|
||||||
}
|
}
|
||||||
c.mu.Lock()
|
|
||||||
defer c.mu.Unlock()
|
|
||||||
|
|
||||||
c.lastHealthCheck = time.Now()
|
// lifecycleMu (read) keeps Close/Reconnect from tearing the provider down
|
||||||
|
// mid-ping without blocking the accessors that only need mu.
|
||||||
|
c.lifecycleMu.RLock()
|
||||||
|
defer c.lifecycleMu.RUnlock()
|
||||||
|
|
||||||
if !c.connected {
|
c.mu.RLock()
|
||||||
c.healthCheckStatus = "disconnected"
|
connected := c.connected
|
||||||
|
provider := c.provider
|
||||||
|
c.mu.RUnlock()
|
||||||
|
|
||||||
|
if !connected {
|
||||||
|
c.setHealth("disconnected")
|
||||||
return ErrConnectionClosed
|
return ErrConnectionClosed
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.provider.HealthCheck(ctx); err != nil {
|
if err := provider.HealthCheck(ctx); err != nil {
|
||||||
c.healthCheckStatus = "unhealthy: " + err.Error()
|
c.setHealth("unhealthy: " + err.Error())
|
||||||
return NewConnectionError(c.name, "health check", err)
|
return NewConnectionError(c.name, "health check", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
c.healthCheckStatus = "healthy"
|
c.setHealth("healthy")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reconnect closes and re-establishes the connection
|
func (c *sqlConnection) setHealth(status string) {
|
||||||
func (c *sqlConnection) Reconnect(ctx context.Context) error {
|
c.mu.Lock()
|
||||||
if err := c.Close(); err != nil {
|
c.lastHealthCheck = time.Now()
|
||||||
|
c.healthCheckStatus = status
|
||||||
|
c.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reconnect refreshes the connection as a single critical section.
|
||||||
|
//
|
||||||
|
// Providers that support it (PostgreSQL) retire their pooled connections and
|
||||||
|
// dial fresh ones without closing the *sql.DB, so handles handed out earlier
|
||||||
|
// keep working. Other providers fall back to Close+Connect, which invalidates
|
||||||
|
// earlier handles; that is meant for explicit operator use only, since
|
||||||
|
// *sql.DB already replaces broken connections by itself.
|
||||||
|
func (c *sqlConnection) Reconnect(ctx context.Context) (err error) {
|
||||||
|
c.lifecycleMu.Lock()
|
||||||
|
defer c.lifecycleMu.Unlock()
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
defer func() { RecordReconnectAttempt(c.name, c.dbType, err == nil) }()
|
||||||
|
|
||||||
|
if c.connected {
|
||||||
|
if r, ok := c.provider.(providers.Refresher); ok {
|
||||||
|
if err := r.Refresh(ctx); err != nil {
|
||||||
|
return NewConnectionError(c.name, "reconnect", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := c.closeLocked(); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return c.Connect(ctx)
|
return c.connectLocked(ctx)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Native returns the native *sql.DB connection
|
// Native returns the native *sql.DB connection
|
||||||
@@ -250,6 +310,10 @@ func (c *sqlConnection) Bun() (*bun.DB, error) {
|
|||||||
return c.bunDB, nil
|
return c.bunDB, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !c.connected {
|
||||||
|
return nil, ErrConnectionClosed
|
||||||
|
}
|
||||||
|
|
||||||
// Get native connection first
|
// Get native connection first
|
||||||
native, err := c.provider.GetNative()
|
native, err := c.provider.GetNative()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -283,6 +347,10 @@ func (c *sqlConnection) GORM() (*gorm.DB, error) {
|
|||||||
return c.gormDB, nil
|
return c.gormDB, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !c.connected {
|
||||||
|
return nil, ErrConnectionClosed
|
||||||
|
}
|
||||||
|
|
||||||
// Get native connection first
|
// Get native connection first
|
||||||
native, err := c.provider.GetNative()
|
native, err := c.provider.GetNative()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -359,39 +427,18 @@ func (c *sqlConnection) Stats() *ConnectionStats {
|
|||||||
return stats
|
return stats
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *sqlConnection) reconnectForAdapter() error {
|
// The adapter factories only re-fetch the current handle. They must not close
|
||||||
timeout := c.config.ConnectTimeout
|
// the shared pool: *sql.DB discards bad connections on its own, and closing it
|
||||||
if timeout <= 0 {
|
// here would break every other holder of the pool.
|
||||||
timeout = 10 * time.Second
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
return c.Reconnect(ctx)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *sqlConnection) reopenNativeForAdapter() (*sql.DB, error) {
|
func (c *sqlConnection) reopenNativeForAdapter() (*sql.DB, error) {
|
||||||
if err := c.reconnectForAdapter(); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return c.Native()
|
return c.Native()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *sqlConnection) reopenBunForAdapter() (*bun.DB, error) {
|
func (c *sqlConnection) reopenBunForAdapter() (*bun.DB, error) {
|
||||||
if err := c.reconnectForAdapter(); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return c.Bun()
|
return c.Bun()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *sqlConnection) reopenGORMForAdapter() (*gorm.DB, error) {
|
func (c *sqlConnection) reopenGORMForAdapter() (*gorm.DB, error) {
|
||||||
if err := c.reconnectForAdapter(); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return c.GORM()
|
return c.GORM()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -427,7 +474,9 @@ func (c *sqlConnection) getBunAdapter() (common.Database, error) {
|
|||||||
c.bunDB = bun.NewDB(native, dialect)
|
c.bunDB = bun.NewDB(native, dialect)
|
||||||
}
|
}
|
||||||
|
|
||||||
c.bunAdapter = database.NewBunAdapter(c.bunDB).WithDBFactory(c.reopenBunForAdapter)
|
c.bunAdapter = database.NewBunAdapter(c.bunDB).
|
||||||
|
WithDBFactory(c.reopenBunForAdapter).
|
||||||
|
SetMetricsEnabled(c.config.EnableMetrics)
|
||||||
return c.bunAdapter, nil
|
return c.bunAdapter, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -468,7 +517,9 @@ func (c *sqlConnection) getGORMAdapter() (common.Database, error) {
|
|||||||
c.gormDB = db
|
c.gormDB = db
|
||||||
}
|
}
|
||||||
|
|
||||||
c.gormAdapter = database.NewGormAdapter(c.gormDB).WithDBFactory(c.reopenGORMForAdapter)
|
c.gormAdapter = database.NewGormAdapter(c.gormDB).
|
||||||
|
WithDBFactory(c.reopenGORMForAdapter).
|
||||||
|
SetMetricsEnabled(c.config.EnableMetrics)
|
||||||
return c.gormAdapter, nil
|
return c.gormAdapter, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -508,12 +559,11 @@ func (c *sqlConnection) getNativeAdapter() (common.Database, error) {
|
|||||||
|
|
||||||
// Create a native adapter based on database type
|
// Create a native adapter based on database type
|
||||||
switch c.dbType {
|
switch c.dbType {
|
||||||
case DatabaseTypePostgreSQL:
|
case DatabaseTypePostgreSQL, DatabaseTypeSQLite, DatabaseTypeMSSQL:
|
||||||
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).WithDBFactory(c.reopenNativeForAdapter)
|
// The adapter takes the driver name so it can adjust its dialect.
|
||||||
case DatabaseTypeSQLite:
|
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
|
||||||
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).WithDBFactory(c.reopenNativeForAdapter)
|
WithDBFactory(c.reopenNativeForAdapter).
|
||||||
case DatabaseTypeMSSQL:
|
SetMetricsEnabled(c.config.EnableMetrics)
|
||||||
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).WithDBFactory(c.reopenNativeForAdapter)
|
|
||||||
default:
|
default:
|
||||||
return nil, ErrUnsupportedDatabase
|
return nil, ErrUnsupportedDatabase
|
||||||
}
|
}
|
||||||
@@ -564,6 +614,7 @@ type mongoConnection struct {
|
|||||||
// State
|
// State
|
||||||
connected bool
|
connected bool
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
|
lifecycleMu sync.RWMutex // see sqlConnection.lifecycleMu
|
||||||
|
|
||||||
// Health check
|
// Health check
|
||||||
lastHealthCheck time.Time
|
lastHealthCheck time.Time
|
||||||
@@ -591,9 +642,16 @@ func (c *mongoConnection) Type() DatabaseType {
|
|||||||
|
|
||||||
// Connect establishes the MongoDB connection
|
// Connect establishes the MongoDB connection
|
||||||
func (c *mongoConnection) Connect(ctx context.Context) error {
|
func (c *mongoConnection) Connect(ctx context.Context) error {
|
||||||
|
c.lifecycleMu.Lock()
|
||||||
|
defer c.lifecycleMu.Unlock()
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
return c.connectLocked(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// connectLocked requires lifecycleMu and mu held for writing.
|
||||||
|
func (c *mongoConnection) connectLocked(ctx context.Context) error {
|
||||||
if c.connected {
|
if c.connected {
|
||||||
return ErrAlreadyConnected
|
return ErrAlreadyConnected
|
||||||
}
|
}
|
||||||
@@ -605,6 +663,7 @@ func (c *mongoConnection) Connect(ctx context.Context) error {
|
|||||||
// Get the mongo client
|
// Get the mongo client
|
||||||
client, err := c.provider.GetMongo()
|
client, err := c.provider.GetMongo()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
_ = c.provider.Close()
|
||||||
return NewConnectionError(c.name, "get mongo client", err)
|
return NewConnectionError(c.name, "get mongo client", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -615,49 +674,75 @@ func (c *mongoConnection) Connect(ctx context.Context) error {
|
|||||||
|
|
||||||
// Close closes the MongoDB connection
|
// Close closes the MongoDB connection
|
||||||
func (c *mongoConnection) Close() error {
|
func (c *mongoConnection) Close() error {
|
||||||
|
c.lifecycleMu.Lock()
|
||||||
|
defer c.lifecycleMu.Unlock()
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
return c.closeLocked()
|
||||||
|
}
|
||||||
|
|
||||||
|
// closeLocked requires lifecycleMu and mu held for writing. The connection is
|
||||||
|
// marked disconnected even when the provider fails to close cleanly.
|
||||||
|
func (c *mongoConnection) closeLocked() error {
|
||||||
if !c.connected {
|
if !c.connected {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.provider.Close(); err != nil {
|
err := c.provider.Close()
|
||||||
return NewConnectionError(c.name, "close", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
c.connected = false
|
c.connected = false
|
||||||
c.client = nil
|
c.client = nil
|
||||||
|
if err != nil {
|
||||||
|
return NewConnectionError(c.name, "close", err)
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// HealthCheck verifies the MongoDB connection is alive
|
// HealthCheck verifies the MongoDB connection is alive. The ping runs without
|
||||||
|
// holding mu so handle accessors are never blocked behind it.
|
||||||
func (c *mongoConnection) HealthCheck(ctx context.Context) error {
|
func (c *mongoConnection) HealthCheck(ctx context.Context) error {
|
||||||
c.mu.Lock()
|
c.lifecycleMu.RLock()
|
||||||
defer c.mu.Unlock()
|
defer c.lifecycleMu.RUnlock()
|
||||||
|
|
||||||
c.lastHealthCheck = time.Now()
|
c.mu.RLock()
|
||||||
|
connected := c.connected
|
||||||
|
c.mu.RUnlock()
|
||||||
|
|
||||||
if !c.connected {
|
if !connected {
|
||||||
c.healthCheckStatus = "disconnected"
|
c.setHealth("disconnected")
|
||||||
return ErrConnectionClosed
|
return ErrConnectionClosed
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.provider.HealthCheck(ctx); err != nil {
|
if err := c.provider.HealthCheck(ctx); err != nil {
|
||||||
c.healthCheckStatus = "unhealthy: " + err.Error()
|
c.setHealth("unhealthy: " + err.Error())
|
||||||
return NewConnectionError(c.name, "health check", err)
|
return NewConnectionError(c.name, "health check", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
c.healthCheckStatus = "healthy"
|
c.setHealth("healthy")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reconnect closes and re-establishes the MongoDB connection
|
func (c *mongoConnection) setHealth(status string) {
|
||||||
func (c *mongoConnection) Reconnect(ctx context.Context) error {
|
c.mu.Lock()
|
||||||
if err := c.Close(); err != nil {
|
c.lastHealthCheck = time.Now()
|
||||||
|
c.healthCheckStatus = status
|
||||||
|
c.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reconnect closes and re-establishes the MongoDB connection atomically.
|
||||||
|
func (c *mongoConnection) Reconnect(ctx context.Context) (err error) {
|
||||||
|
c.lifecycleMu.Lock()
|
||||||
|
defer c.lifecycleMu.Unlock()
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
defer func() { RecordReconnectAttempt(c.name, DatabaseTypeMongoDB, err == nil) }()
|
||||||
|
|
||||||
|
if err := c.closeLocked(); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return c.Connect(ctx)
|
return c.connectLocked(ctx)
|
||||||
}
|
}
|
||||||
|
|
||||||
// MongoDB returns the MongoDB client
|
// MongoDB returns the MongoDB client
|
||||||
|
|||||||
@@ -4,13 +4,8 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
|
||||||
|
|
||||||
_ "github.com/mattn/go-sqlite3"
|
_ "github.com/mattn/go-sqlite3"
|
||||||
"gorm.io/gorm"
|
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestNewConnectionFromDB(t *testing.T) {
|
func TestNewConnectionFromDB(t *testing.T) {
|
||||||
@@ -213,157 +208,3 @@ func TestNewConnectionFromDB_PostgreSQL(t *testing.T) {
|
|||||||
t.Errorf("Expected type DatabaseTypePostgreSQL, got '%s'", conn.Type())
|
t.Errorf("Expected type DatabaseTypePostgreSQL, got '%s'", conn.Type())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDatabaseNativeAdapterReconnectFactory(t *testing.T) {
|
|
||||||
conn := newSQLConnection("test-native", DatabaseTypeSQLite, ConnectionConfig{
|
|
||||||
Name: "test-native",
|
|
||||||
Type: DatabaseTypeSQLite,
|
|
||||||
FilePath: ":memory:",
|
|
||||||
DefaultORM: string(ORMTypeNative),
|
|
||||||
ConnectTimeout: 2 * time.Second,
|
|
||||||
}, providers.NewSQLiteProvider())
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
if err := conn.Connect(ctx); err != nil {
|
|
||||||
t.Fatalf("Failed to connect: %v", err)
|
|
||||||
}
|
|
||||||
defer conn.Close()
|
|
||||||
|
|
||||||
db, err := conn.Database()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to get database adapter: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
adapter, ok := db.(*database.PgSQLAdapter)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("Expected PgSQLAdapter, got %T", db)
|
|
||||||
}
|
|
||||||
|
|
||||||
underlyingBefore, ok := adapter.GetUnderlyingDB().(*sql.DB)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("Expected underlying *sql.DB, got %T", adapter.GetUnderlyingDB())
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := underlyingBefore.Close(); err != nil {
|
|
||||||
t.Fatalf("Failed to close underlying database: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := db.Exec(ctx, "SELECT 1"); err != nil {
|
|
||||||
t.Fatalf("Expected native adapter to reconnect, got error: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
underlyingAfter, ok := adapter.GetUnderlyingDB().(*sql.DB)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("Expected reconnected *sql.DB, got %T", adapter.GetUnderlyingDB())
|
|
||||||
}
|
|
||||||
|
|
||||||
if underlyingAfter == underlyingBefore {
|
|
||||||
t.Fatal("Expected adapter to swap to a fresh *sql.DB after reconnect")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDatabaseBunAdapterReconnectFactory(t *testing.T) {
|
|
||||||
conn := newSQLConnection("test-bun", DatabaseTypeSQLite, ConnectionConfig{
|
|
||||||
Name: "test-bun",
|
|
||||||
Type: DatabaseTypeSQLite,
|
|
||||||
FilePath: ":memory:",
|
|
||||||
DefaultORM: string(ORMTypeBun),
|
|
||||||
ConnectTimeout: 2 * time.Second,
|
|
||||||
}, providers.NewSQLiteProvider())
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
if err := conn.Connect(ctx); err != nil {
|
|
||||||
t.Fatalf("Failed to connect: %v", err)
|
|
||||||
}
|
|
||||||
defer conn.Close()
|
|
||||||
|
|
||||||
db, err := conn.Database()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to get database adapter: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
adapter, ok := db.(*database.BunAdapter)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("Expected BunAdapter, got %T", db)
|
|
||||||
}
|
|
||||||
|
|
||||||
underlyingBefore, ok := adapter.GetUnderlyingDB().(interface{ Close() error })
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("Expected underlying Bun DB with Close method, got %T", adapter.GetUnderlyingDB())
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := underlyingBefore.Close(); err != nil {
|
|
||||||
t.Fatalf("Failed to close underlying Bun database: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := db.Exec(ctx, "SELECT 1"); err != nil {
|
|
||||||
t.Fatalf("Expected Bun adapter to reconnect, got error: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
underlyingAfter := adapter.GetUnderlyingDB()
|
|
||||||
if underlyingAfter == underlyingBefore {
|
|
||||||
t.Fatal("Expected adapter to swap to a fresh Bun DB after reconnect")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDatabaseGormAdapterReconnectFactory(t *testing.T) {
|
|
||||||
conn := newSQLConnection("test-gorm", DatabaseTypeSQLite, ConnectionConfig{
|
|
||||||
Name: "test-gorm",
|
|
||||||
Type: DatabaseTypeSQLite,
|
|
||||||
FilePath: ":memory:",
|
|
||||||
DefaultORM: string(ORMTypeGORM),
|
|
||||||
ConnectTimeout: 2 * time.Second,
|
|
||||||
}, providers.NewSQLiteProvider())
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
if err := conn.Connect(ctx); err != nil {
|
|
||||||
t.Fatalf("Failed to connect: %v", err)
|
|
||||||
}
|
|
||||||
defer conn.Close()
|
|
||||||
|
|
||||||
db, err := conn.Database()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to get database adapter: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
adapter, ok := db.(*database.GormAdapter)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("Expected GormAdapter, got %T", db)
|
|
||||||
}
|
|
||||||
|
|
||||||
gormBefore, ok := adapter.GetUnderlyingDB().(*gorm.DB)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("Expected underlying *gorm.DB, got %T", adapter.GetUnderlyingDB())
|
|
||||||
}
|
|
||||||
|
|
||||||
sqlBefore, err := gormBefore.DB()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to get underlying *sql.DB: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := sqlBefore.Close(); err != nil {
|
|
||||||
t.Fatalf("Failed to close underlying database: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
count, err := db.NewSelect().Table("sqlite_master").Count(ctx)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Expected GORM query builder to reconnect, got error: %v", err)
|
|
||||||
}
|
|
||||||
if count < 0 {
|
|
||||||
t.Fatalf("Expected non-negative count, got %d", count)
|
|
||||||
}
|
|
||||||
|
|
||||||
gormAfter, ok := adapter.GetUnderlyingDB().(*gorm.DB)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("Expected reconnected *gorm.DB, got %T", adapter.GetUnderlyingDB())
|
|
||||||
}
|
|
||||||
|
|
||||||
sqlAfter, err := gormAfter.DB()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to get reconnected *sql.DB: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if sqlAfter == sqlBefore {
|
|
||||||
t.Fatal("Expected GORM adapter to use a fresh *sql.DB after reconnect")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -0,0 +1,199 @@
|
|||||||
|
package dbmanager
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func sqliteManagerConfig() ManagerConfig {
|
||||||
|
return ManagerConfig{
|
||||||
|
DefaultConnection: "test",
|
||||||
|
Connections: map[string]ConnectionConfig{
|
||||||
|
"test": {Name: "test", Type: DatabaseTypeSQLite, FilePath: ":memory:"},
|
||||||
|
},
|
||||||
|
HealthCheckInterval: time.Hour,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestManagerConnectCloseCycleTwice(t *testing.T) {
|
||||||
|
mgr, err := NewManager(sqliteManagerConfig())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ctx := context.Background()
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
if err := mgr.Connect(ctx); err != nil {
|
||||||
|
t.Fatalf("cycle %d connect: %v", i, err)
|
||||||
|
}
|
||||||
|
cm := mgr.(*connectionManager)
|
||||||
|
cm.healthMu.Lock()
|
||||||
|
running := cm.healthTicker != nil
|
||||||
|
cm.healthMu.Unlock()
|
||||||
|
if !running {
|
||||||
|
t.Fatalf("cycle %d: health checker not running", i)
|
||||||
|
}
|
||||||
|
if err := mgr.Close(); err != nil {
|
||||||
|
t.Fatalf("cycle %d close: %v", i, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// A further Close must not panic.
|
||||||
|
if err := mgr.Close(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestManagerConnectIsIdempotent(t *testing.T) {
|
||||||
|
mgr, _ := NewManager(sqliteManagerConfig())
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := mgr.Connect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer mgr.Close()
|
||||||
|
if err := mgr.Connect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := mgr.Stats().TotalConnections; got != 1 {
|
||||||
|
t.Fatalf("expected 1 connection, got %d", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConcurrentReconnectIsAtomic(t *testing.T) {
|
||||||
|
mgr, _ := NewManager(sqliteManagerConfig())
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := mgr.Connect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer mgr.Close()
|
||||||
|
conn, _ := mgr.GetDefault()
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
if err := conn.Reconnect(ctx); err != nil {
|
||||||
|
t.Errorf("reconnect: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
|
db, err := conn.Native()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := db.PingContext(ctx); err != nil {
|
||||||
|
t.Fatalf("pool unusable after concurrent reconnects: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAdapterFactoryDoesNotClosePool(t *testing.T) {
|
||||||
|
mgr, _ := NewManager(sqliteManagerConfig())
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := mgr.Connect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer mgr.Close()
|
||||||
|
conn, _ := mgr.GetDefault()
|
||||||
|
sc := conn.(*sqlConnection)
|
||||||
|
|
||||||
|
held, _ := sc.Native()
|
||||||
|
if _, err := sc.reopenNativeForAdapter(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := held.PingContext(ctx); err != nil {
|
||||||
|
t.Fatalf("existing handle broken by adapter factory: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHealthCheckDoesNotBlockAccessors(t *testing.T) {
|
||||||
|
mgr, _ := NewManager(sqliteManagerConfig())
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := mgr.Connect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer mgr.Close()
|
||||||
|
conn, _ := mgr.GetDefault()
|
||||||
|
sc := conn.(*sqlConnection)
|
||||||
|
|
||||||
|
// Simulate a health check in flight: it holds lifecycleMu (read) only.
|
||||||
|
sc.lifecycleMu.RLock()
|
||||||
|
defer sc.lifecycleMu.RUnlock()
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
_, _ = sc.Bun()
|
||||||
|
_, _ = sc.GORM()
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("accessors blocked while health check in flight")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCloseAlwaysMarksDisconnected(t *testing.T) {
|
||||||
|
mgr, _ := NewManager(sqliteManagerConfig())
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := mgr.Connect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
conn, _ := mgr.GetDefault()
|
||||||
|
sc := conn.(*sqlConnection)
|
||||||
|
_, _ = sc.Bun()
|
||||||
|
_ = mgr.Close()
|
||||||
|
|
||||||
|
if _, err := sc.Bun(); err == nil {
|
||||||
|
t.Fatal("Bun() should fail after Close")
|
||||||
|
}
|
||||||
|
if _, err := sc.GORM(); err == nil {
|
||||||
|
t.Fatal("GORM() should fail after Close")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReconnectOnExistingDBKeepsCallersPool(t *testing.T) {
|
||||||
|
db, err := sql.Open("sqlite", ":memory:")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
conn := NewConnectionFromDB("existing", DatabaseTypeSQLite, db)
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := conn.Connect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := conn.Reconnect(ctx); err != nil {
|
||||||
|
t.Fatalf("reconnect: %v", err)
|
||||||
|
}
|
||||||
|
if err := db.PingContext(ctx); err != nil {
|
||||||
|
t.Fatalf("caller's pool was closed by Reconnect: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCloseOnExistingDBLeavesCallersPoolOpen(t *testing.T) {
|
||||||
|
db, err := sql.Open("sqlite", ":memory:")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
conn := NewConnectionFromDB("existing", DatabaseTypeSQLite, db)
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := conn.Connect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := conn.Bun(); err != nil { // bun.DB.Close would close the pool
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := conn.Close(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := db.PingContext(ctx); err != nil {
|
||||||
|
t.Fatalf("caller's pool was closed: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
+63
-49
@@ -2,9 +2,7 @@ package dbmanager
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -49,6 +47,7 @@ type connectionManager struct {
|
|||||||
// Background health check
|
// Background health check
|
||||||
healthTicker *time.Ticker
|
healthTicker *time.Ticker
|
||||||
stopChan chan struct{}
|
stopChan chan struct{}
|
||||||
|
healthMu sync.Mutex // guards healthTicker and stopChan
|
||||||
wg sync.WaitGroup
|
wg sync.WaitGroup
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -100,7 +99,9 @@ func ResetInstance() {
|
|||||||
defer instanceMu.Unlock()
|
defer instanceMu.Unlock()
|
||||||
|
|
||||||
if instance != nil {
|
if instance != nil {
|
||||||
_ = instance.Close()
|
if err := instance.Close(); err != nil {
|
||||||
|
logger.Error("Failed to close manager during reset: %v", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
instance = nil
|
instance = nil
|
||||||
}
|
}
|
||||||
@@ -116,7 +117,6 @@ func NewManager(cfg ManagerConfig) (Manager, error) {
|
|||||||
mgr := &connectionManager{
|
mgr := &connectionManager{
|
||||||
connections: make(map[string]Connection),
|
connections: make(map[string]Connection),
|
||||||
config: cfg,
|
config: cfg,
|
||||||
stopChan: make(chan struct{}),
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return mgr, nil
|
return mgr, nil
|
||||||
@@ -195,11 +195,26 @@ func (m *connectionManager) SetDefaultDatabase(name string) error {
|
|||||||
|
|
||||||
// Connect establishes all configured database connections
|
// Connect establishes all configured database connections
|
||||||
func (m *connectionManager) Connect(ctx context.Context) error {
|
func (m *connectionManager) Connect(ctx context.Context) error {
|
||||||
m.mu.Lock()
|
// Dial outside m.mu so a slow connect never blocks Get/Stats/HealthCheck.
|
||||||
defer m.mu.Unlock()
|
m.mu.RLock()
|
||||||
|
names := make([]string, 0, len(m.config.Connections))
|
||||||
// Create connections from configuration
|
|
||||||
for name := range m.config.Connections {
|
for name := range m.config.Connections {
|
||||||
|
if _, exists := m.connections[name]; !exists {
|
||||||
|
names = append(names, name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
m.mu.RUnlock()
|
||||||
|
|
||||||
|
opened := make(map[string]Connection, len(names))
|
||||||
|
closeOpened := func() {
|
||||||
|
for name, conn := range opened {
|
||||||
|
if err := conn.Close(); err != nil {
|
||||||
|
logger.Error("Failed to close connection after failed Connect: name=%s, error=%v", name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, name := range names {
|
||||||
// Get a copy of the connection config
|
// Get a copy of the connection config
|
||||||
connCfg := m.config.Connections[name]
|
connCfg := m.config.Connections[name]
|
||||||
// Apply global defaults to connection config
|
// Apply global defaults to connection config
|
||||||
@@ -209,25 +224,39 @@ func (m *connectionManager) Connect(ctx context.Context) error {
|
|||||||
// Create connection using factory
|
// Create connection using factory
|
||||||
conn, err := createConnection(connCfg)
|
conn, err := createConnection(connCfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
closeOpened()
|
||||||
return fmt.Errorf("failed to create connection '%s': %w", name, err)
|
return fmt.Errorf("failed to create connection '%s': %w", name, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Connect
|
// Connect
|
||||||
if err := conn.Connect(ctx); err != nil {
|
if err := conn.Connect(ctx); err != nil {
|
||||||
|
closeOpened()
|
||||||
return fmt.Errorf("failed to connect '%s': %w", name, err)
|
return fmt.Errorf("failed to connect '%s': %w", name, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
m.connections[name] = conn
|
opened[name] = conn
|
||||||
logger.Info("Database connection established: name=%s, type=%s", name, connCfg.Type)
|
logger.Info("Database connection established: name=%s, type=%s", name, connCfg.Type)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
m.mu.Lock()
|
||||||
|
for name, conn := range opened {
|
||||||
|
if _, exists := m.connections[name]; exists {
|
||||||
|
// Lost a race with a concurrent Connect; drop our duplicate.
|
||||||
|
_ = conn.Close()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
m.connections[name] = conn
|
||||||
|
}
|
||||||
|
total := len(m.connections)
|
||||||
|
m.mu.Unlock()
|
||||||
|
|
||||||
// Always start background health checks
|
// Always start background health checks
|
||||||
if m.config.HealthCheckInterval > 0 {
|
if m.config.HealthCheckInterval > 0 {
|
||||||
m.startHealthChecker()
|
m.startHealthChecker()
|
||||||
logger.Info("Background health checker started: interval=%v", m.config.HealthCheckInterval)
|
logger.Info("Background health checker started: interval=%v", m.config.HealthCheckInterval)
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.Info("Database manager initialized: connections=%d", len(m.connections))
|
logger.Info("Database manager initialized: connections=%d", total)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -246,7 +275,7 @@ func (m *connectionManager) Close() error {
|
|||||||
for name, conn := range m.connections {
|
for name, conn := range m.connections {
|
||||||
if err := conn.Close(); err != nil {
|
if err := conn.Close(); err != nil {
|
||||||
errors = append(errors, fmt.Errorf("failed to close connection '%s': %w", name, err))
|
errors = append(errors, fmt.Errorf("failed to close connection '%s': %w", name, err))
|
||||||
logger.Error("Failed to close connection", "name", name, "error", err)
|
logger.Error("Failed to close connection: name=%s, error=%v", name, err)
|
||||||
} else {
|
} else {
|
||||||
logger.Info("Connection closed: name=%s", name)
|
logger.Info("Connection closed: name=%s", name)
|
||||||
}
|
}
|
||||||
@@ -311,11 +340,17 @@ func (m *connectionManager) Stats() *ManagerStats {
|
|||||||
|
|
||||||
// startHealthChecker starts background health checking
|
// startHealthChecker starts background health checking
|
||||||
func (m *connectionManager) startHealthChecker() {
|
func (m *connectionManager) startHealthChecker() {
|
||||||
|
m.healthMu.Lock()
|
||||||
|
defer m.healthMu.Unlock()
|
||||||
|
|
||||||
if m.healthTicker != nil {
|
if m.healthTicker != nil {
|
||||||
return // Already running
|
return // Already running
|
||||||
}
|
}
|
||||||
|
|
||||||
m.healthTicker = time.NewTicker(m.config.HealthCheckInterval)
|
ticker := time.NewTicker(m.config.HealthCheckInterval)
|
||||||
|
stop := make(chan struct{})
|
||||||
|
m.healthTicker = ticker
|
||||||
|
m.stopChan = stop
|
||||||
|
|
||||||
m.wg.Add(1)
|
m.wg.Add(1)
|
||||||
go func() {
|
go func() {
|
||||||
@@ -324,9 +359,9 @@ func (m *connectionManager) startHealthChecker() {
|
|||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-m.healthTicker.C:
|
case <-ticker.C:
|
||||||
m.performHealthCheck()
|
m.performHealthCheck()
|
||||||
case <-m.stopChan:
|
case <-stop:
|
||||||
logger.Info("Health checker stopped")
|
logger.Info("Health checker stopped")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -334,14 +369,19 @@ func (m *connectionManager) startHealthChecker() {
|
|||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
|
|
||||||
// stopHealthChecker stops background health checking
|
// stopHealthChecker stops background health checking. Safe to call repeatedly.
|
||||||
func (m *connectionManager) stopHealthChecker() {
|
func (m *connectionManager) stopHealthChecker() {
|
||||||
if m.healthTicker != nil {
|
m.healthMu.Lock()
|
||||||
|
defer m.healthMu.Unlock()
|
||||||
|
|
||||||
|
if m.healthTicker == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
m.healthTicker.Stop()
|
m.healthTicker.Stop()
|
||||||
close(m.stopChan)
|
close(m.stopChan)
|
||||||
m.wg.Wait()
|
m.wg.Wait()
|
||||||
m.healthTicker = nil
|
m.healthTicker = nil
|
||||||
}
|
m.stopChan = nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// performHealthCheck performs a health check on all connections
|
// performHealthCheck performs a health check on all connections
|
||||||
@@ -362,40 +402,14 @@ func (m *connectionManager) performHealthCheck() {
|
|||||||
}
|
}
|
||||||
m.mu.RUnlock()
|
m.mu.RUnlock()
|
||||||
|
|
||||||
|
defer m.PublishMetrics()
|
||||||
|
|
||||||
for _, item := range connections {
|
for _, item := range connections {
|
||||||
if err := item.conn.HealthCheck(ctx); err != nil {
|
if err := item.conn.HealthCheck(ctx); err != nil {
|
||||||
logger.Warn("Health check failed",
|
// Do not reconnect here: *sql.DB discards bad connections and dials
|
||||||
"connection", item.name,
|
// new ones by itself, while Reconnect closes the pool and breaks
|
||||||
"error", err)
|
// every handle already handed out. Reconnect is operator-only.
|
||||||
|
logger.Warn("Health check failed: connection=%s, error=%v", item.name, err)
|
||||||
// Only reconnect when the client handle itself is closed/disconnected.
|
|
||||||
// For transient database restarts or network blips, *sql.DB can recover
|
|
||||||
// on its own; forcing Close()+Connect() here invalidates any cached ORM
|
|
||||||
// wrappers and callers that still hold the old handle.
|
|
||||||
if m.config.EnableAutoReconnect && shouldReconnectAfterHealthCheck(err) {
|
|
||||||
logger.Info("Attempting reconnection: connection=%s", item.name)
|
|
||||||
if err := item.conn.Reconnect(ctx); err != nil {
|
|
||||||
logger.Error("Reconnection failed",
|
|
||||||
"connection", item.name,
|
|
||||||
"error", err)
|
|
||||||
} else {
|
|
||||||
logger.Info("Reconnection successful: connection=%s", item.name)
|
|
||||||
}
|
|
||||||
} else if m.config.EnableAutoReconnect {
|
|
||||||
logger.Info("Skipping reconnect for transient health check failure: connection=%s", item.name)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func shouldReconnectAfterHealthCheck(err error) bool {
|
|
||||||
if err == nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
if errors.Is(err, ErrConnectionClosed) {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
return strings.Contains(err.Error(), "sql: database is closed")
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -24,15 +24,26 @@ type healthCheckStubConnection struct {
|
|||||||
func (c *healthCheckStubConnection) Name() string { return "stub" }
|
func (c *healthCheckStubConnection) Name() string { return "stub" }
|
||||||
func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL }
|
func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL }
|
||||||
func (c *healthCheckStubConnection) Bun() (*bun.DB, error) { return nil, fmt.Errorf("not implemented") }
|
func (c *healthCheckStubConnection) Bun() (*bun.DB, error) { return nil, fmt.Errorf("not implemented") }
|
||||||
func (c *healthCheckStubConnection) GORM() (*gorm.DB, error) { return nil, fmt.Errorf("not implemented") }
|
func (c *healthCheckStubConnection) GORM() (*gorm.DB, error) {
|
||||||
func (c *healthCheckStubConnection) Native() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") }
|
return nil, fmt.Errorf("not implemented")
|
||||||
|
}
|
||||||
|
func (c *healthCheckStubConnection) 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) 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) Database() (common.Database, error) {
|
||||||
func (c *healthCheckStubConnection) MongoDB() (*mongo.Client, error) { return nil, fmt.Errorf("not implemented") }
|
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) Connect(ctx context.Context) error { return nil }
|
||||||
func (c *healthCheckStubConnection) Close() error { return nil }
|
func (c *healthCheckStubConnection) Close() error { return nil }
|
||||||
func (c *healthCheckStubConnection) HealthCheck(ctx context.Context) error { return c.healthErr }
|
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) Reconnect(ctx context.Context) error {
|
||||||
|
c.reconnectCalls++
|
||||||
|
return nil
|
||||||
|
}
|
||||||
func (c *healthCheckStubConnection) Stats() *ConnectionStats { return &ConnectionStats{} }
|
func (c *healthCheckStubConnection) Stats() *ConnectionStats { return &ConnectionStats{} }
|
||||||
|
|
||||||
func TestBackgroundHealthChecker(t *testing.T) {
|
func TestBackgroundHealthChecker(t *testing.T) {
|
||||||
@@ -117,41 +128,21 @@ func TestDefaultHealthCheckInterval(t *testing.T) {
|
|||||||
t.Errorf("Expected default health check interval to be %v, got %v",
|
t.Errorf("Expected default health check interval to be %v, got %v",
|
||||||
expectedInterval, defaults.HealthCheckInterval)
|
expectedInterval, defaults.HealthCheckInterval)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !defaults.EnableAutoReconnect {
|
|
||||||
t.Error("Expected EnableAutoReconnect to be true by default")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestApplyDefaultsEnablesAutoReconnect(t *testing.T) {
|
func TestApplyDefaultsHealthCheckInterval(t *testing.T) {
|
||||||
// Create a config without setting EnableAutoReconnect
|
cfg := ManagerConfig{}
|
||||||
cfg := ManagerConfig{
|
|
||||||
Connections: map[string]ConnectionConfig{
|
|
||||||
"test": {
|
|
||||||
Name: "test",
|
|
||||||
Type: DatabaseTypeSQLite,
|
|
||||||
FilePath: ":memory:",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify it's false initially (Go's zero value for bool)
|
|
||||||
if cfg.EnableAutoReconnect {
|
|
||||||
t.Error("Expected EnableAutoReconnect to be false before ApplyDefaults")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Apply defaults
|
|
||||||
cfg.ApplyDefaults()
|
cfg.ApplyDefaults()
|
||||||
|
|
||||||
// Verify it's now true
|
|
||||||
if !cfg.EnableAutoReconnect {
|
|
||||||
t.Error("Expected EnableAutoReconnect to be true after ApplyDefaults")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify health check interval is also set
|
|
||||||
if cfg.HealthCheckInterval != 15*time.Second {
|
if cfg.HealthCheckInterval != 15*time.Second {
|
||||||
t.Errorf("Expected health check interval to be 15s, got %v", cfg.HealthCheckInterval)
|
t.Errorf("Expected health check interval to be 15s, got %v", cfg.HealthCheckInterval)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// A negative interval disables the background checker and is preserved.
|
||||||
|
cfg = ManagerConfig{HealthCheckInterval: -1}
|
||||||
|
cfg.ApplyDefaults()
|
||||||
|
if cfg.HealthCheckInterval >= 0 {
|
||||||
|
t.Errorf("Expected negative interval to be preserved, got %v", cfg.HealthCheckInterval)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestManagerHealthCheck(t *testing.T) {
|
func TestManagerHealthCheck(t *testing.T) {
|
||||||
@@ -270,7 +261,7 @@ func TestPerformHealthCheckSkipsReconnectForTransientFailures(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestPerformHealthCheckReconnectsClosedConnections(t *testing.T) {
|
func TestPerformHealthCheckNeverReconnects(t *testing.T) {
|
||||||
conn := &healthCheckStubConnection{
|
conn := &healthCheckStubConnection{
|
||||||
healthErr: NewConnectionError("primary", "health check", fmt.Errorf("sql: database is closed")),
|
healthErr: NewConnectionError("primary", "health check", fmt.Errorf("sql: database is closed")),
|
||||||
}
|
}
|
||||||
@@ -284,7 +275,7 @@ func TestPerformHealthCheckReconnectsClosedConnections(t *testing.T) {
|
|||||||
|
|
||||||
mgr.performHealthCheck()
|
mgr.performHealthCheck()
|
||||||
|
|
||||||
if conn.reconnectCalls != 1 {
|
if conn.reconnectCalls != 0 {
|
||||||
t.Fatalf("expected reconnect attempt for closed database handle, got %d", conn.reconnectCalls)
|
t.Fatalf("health check must not close the shared pool via Reconnect, got %d", conn.reconnectCalls)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+39
-15
@@ -1,6 +1,8 @@
|
|||||||
package dbmanager
|
package dbmanager
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"sync"
|
||||||
|
|
||||||
"github.com/prometheus/client_golang/prometheus"
|
"github.com/prometheus/client_golang/prometheus"
|
||||||
"github.com/prometheus/client_golang/prometheus/promauto"
|
"github.com/prometheus/client_golang/prometheus/promauto"
|
||||||
)
|
)
|
||||||
@@ -34,8 +36,8 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// connectionWaitCount tracks how many times connections had to wait for availability
|
// connectionWaitCount tracks how many times connections had to wait for availability
|
||||||
connectionWaitCount = promauto.NewGaugeVec(
|
connectionWaitCount = promauto.NewCounterVec(
|
||||||
prometheus.GaugeOpts{
|
prometheus.CounterOpts{
|
||||||
Name: "dbmanager_connection_wait_count",
|
Name: "dbmanager_connection_wait_count",
|
||||||
Help: "Number of times connections had to wait for availability",
|
Help: "Number of times connections had to wait for availability",
|
||||||
},
|
},
|
||||||
@@ -43,8 +45,8 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// connectionWaitDuration tracks total time connections spent waiting
|
// connectionWaitDuration tracks total time connections spent waiting
|
||||||
connectionWaitDuration = promauto.NewGaugeVec(
|
connectionWaitDuration = promauto.NewCounterVec(
|
||||||
prometheus.GaugeOpts{
|
prometheus.CounterOpts{
|
||||||
Name: "dbmanager_connection_wait_duration_seconds",
|
Name: "dbmanager_connection_wait_duration_seconds",
|
||||||
Help: "Total time connections spent waiting for availability",
|
Help: "Total time connections spent waiting for availability",
|
||||||
},
|
},
|
||||||
@@ -61,8 +63,8 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// connectionLifetimeClosed tracks connections closed due to max lifetime
|
// connectionLifetimeClosed tracks connections closed due to max lifetime
|
||||||
connectionLifetimeClosed = promauto.NewGaugeVec(
|
connectionLifetimeClosed = promauto.NewCounterVec(
|
||||||
prometheus.GaugeOpts{
|
prometheus.CounterOpts{
|
||||||
Name: "dbmanager_connection_lifetime_closed_total",
|
Name: "dbmanager_connection_lifetime_closed_total",
|
||||||
Help: "Total connections closed due to exceeding max lifetime",
|
Help: "Total connections closed due to exceeding max lifetime",
|
||||||
},
|
},
|
||||||
@@ -70,8 +72,8 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// connectionIdleClosed tracks connections closed due to max idle time
|
// connectionIdleClosed tracks connections closed due to max idle time
|
||||||
connectionIdleClosed = promauto.NewGaugeVec(
|
connectionIdleClosed = promauto.NewCounterVec(
|
||||||
prometheus.GaugeOpts{
|
prometheus.CounterOpts{
|
||||||
Name: "dbmanager_connection_idle_closed_total",
|
Name: "dbmanager_connection_idle_closed_total",
|
||||||
Help: "Total connections closed due to exceeding max idle time",
|
Help: "Total connections closed due to exceeding max idle time",
|
||||||
},
|
},
|
||||||
@@ -114,13 +116,13 @@ func (m *connectionManager) PublishMetrics() {
|
|||||||
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "idle").Set(float64(connStats.Idle))
|
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "idle").Set(float64(connStats.Idle))
|
||||||
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "in_use").Set(float64(connStats.InUse))
|
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "in_use").Set(float64(connStats.InUse))
|
||||||
|
|
||||||
// Wait stats
|
// sql.DBStats values are cumulative, so add only the growth since
|
||||||
connectionWaitCount.With(labels).Set(float64(connStats.WaitCount))
|
// the last publish to keep these true counters.
|
||||||
connectionWaitDuration.With(labels).Set(connStats.WaitDuration.Seconds())
|
prev := lastPublished.swap(name, connStats)
|
||||||
|
connectionWaitCount.With(labels).Add(float64(connStats.WaitCount - prev.WaitCount))
|
||||||
// Lifetime/idle closed
|
connectionWaitDuration.With(labels).Add((connStats.WaitDuration - prev.WaitDuration).Seconds())
|
||||||
connectionLifetimeClosed.With(labels).Set(float64(connStats.MaxLifetimeClosed))
|
connectionLifetimeClosed.With(labels).Add(float64(connStats.MaxLifetimeClosed - prev.MaxLifetimeClosed))
|
||||||
connectionIdleClosed.With(labels).Set(float64(connStats.MaxIdleClosed))
|
connectionIdleClosed.With(labels).Add(float64(connStats.MaxIdleClosed - prev.MaxIdleClosed))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -134,3 +136,25 @@ func RecordReconnectAttempt(name string, dbType DatabaseType, success bool) {
|
|||||||
|
|
||||||
reconnectAttempts.WithLabelValues(name, string(dbType), result).Inc()
|
reconnectAttempts.WithLabelValues(name, string(dbType), result).Inc()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// publishedStats remembers the cumulative pool stats last exported per
|
||||||
|
// connection so counters can be advanced by the delta.
|
||||||
|
type publishedStats struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
last map[string]ConnectionStats
|
||||||
|
}
|
||||||
|
|
||||||
|
var lastPublished = &publishedStats{last: make(map[string]ConnectionStats)}
|
||||||
|
|
||||||
|
// swap stores cur and returns the previous value. A counter reset (a new pool
|
||||||
|
// after Close+Connect) is treated as starting from zero.
|
||||||
|
func (p *publishedStats) swap(name string, cur *ConnectionStats) ConnectionStats {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
prev := p.last[name]
|
||||||
|
if cur.WaitCount < prev.WaitCount || cur.MaxIdleClosed < prev.MaxIdleClosed || cur.MaxLifetimeClosed < prev.MaxLifetimeClosed {
|
||||||
|
prev = ConnectionStats{}
|
||||||
|
}
|
||||||
|
p.last[name] = *cur
|
||||||
|
return prev
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,93 @@
|
|||||||
|
package dbmanager
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestLivePostgresRefreshKeepsHandles(t *testing.T) {
|
||||||
|
if os.Getenv("PG_LIVE") == "" {
|
||||||
|
t.Skip("PG_LIVE not set")
|
||||||
|
}
|
||||||
|
mgr, err := NewManager(ManagerConfig{
|
||||||
|
DefaultConnection: "pg",
|
||||||
|
Connections: map[string]ConnectionConfig{"pg": {
|
||||||
|
Name: "pg", Type: DatabaseTypePostgreSQL, Host: "127.0.0.1", Port: 54329,
|
||||||
|
User: "postgres", Database: "postgres", QueryTimeout: 30 * time.Second,
|
||||||
|
}},
|
||||||
|
HealthCheckInterval: -1,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := mgr.Connect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer mgr.Close()
|
||||||
|
conn, _ := mgr.GetDefault()
|
||||||
|
held, _ := conn.Bun()
|
||||||
|
gormDB, _ := conn.GORM()
|
||||||
|
|
||||||
|
var pid1, pid2 int
|
||||||
|
var st string
|
||||||
|
if err := held.DB.QueryRow("select pg_backend_pid(), current_setting('statement_timeout')").Scan(&pid1, &st); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if st != "30s" {
|
||||||
|
t.Errorf("statement_timeout = %q", st)
|
||||||
|
}
|
||||||
|
if err := conn.Reconnect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := held.DB.QueryRow("select pg_backend_pid()").Scan(&pid2); err != nil {
|
||||||
|
t.Fatalf("held bun handle broken after reconnect: %v", err)
|
||||||
|
}
|
||||||
|
if pid1 == pid2 {
|
||||||
|
t.Error("expected a new backend after reconnect")
|
||||||
|
}
|
||||||
|
var n int
|
||||||
|
if err := gormDB.Raw("select 1").Scan(&n).Error; err != nil || n != 1 {
|
||||||
|
t.Fatalf("held gorm handle broken: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLiveListenerListenNotify(t *testing.T) {
|
||||||
|
if os.Getenv("PG_LIVE") == "" {
|
||||||
|
t.Skip("PG_LIVE not set")
|
||||||
|
}
|
||||||
|
cc := ConnectionConfig{Name: "pg", Type: DatabaseTypePostgreSQL, Host: "127.0.0.1", Port: 54329,
|
||||||
|
User: "postgres", Database: "postgres", ConnectTimeout: 5 * time.Second}
|
||||||
|
p := providers.NewPostgresProvider()
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := p.Connect(ctx, &cc); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer p.Close()
|
||||||
|
l, err := p.GetListener(ctx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got := make(chan string, 4)
|
||||||
|
for _, ch := range []string{"a", "b"} {
|
||||||
|
if err := l.Listen(ch, func(c, payload string) { got <- c + ":" + payload }); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
if err := l.Notify(ctx, "a", "x"); err != nil {
|
||||||
|
t.Fatalf("notify: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case v := <-got:
|
||||||
|
if v != "a:x" {
|
||||||
|
t.Fatalf("got %q", v)
|
||||||
|
}
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
t.Fatal("no notification")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -7,6 +7,8 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"go.mongodb.org/mongo-driver/mongo"
|
"go.mongodb.org/mongo-driver/mongo"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ExistingDBProvider wraps an existing *sql.DB connection
|
// ExistingDBProvider wraps an existing *sql.DB connection
|
||||||
@@ -44,16 +46,27 @@ func (p *ExistingDBProvider) Connect(ctx context.Context, cfg ConnectionConfig)
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close closes the underlying database connection
|
// Refresh verifies the wrapped database is still reachable. The pool belongs to
|
||||||
func (p *ExistingDBProvider) Close() error {
|
// the caller and cannot be re-dialed here, so it is never closed to "reconnect".
|
||||||
p.mu.Lock()
|
func (p *ExistingDBProvider) Refresh(ctx context.Context) error {
|
||||||
defer p.mu.Unlock()
|
p.mu.RLock()
|
||||||
|
defer p.mu.RUnlock()
|
||||||
|
|
||||||
if p.db == nil {
|
if p.db == nil {
|
||||||
return nil
|
return fmt.Errorf("database connection is nil")
|
||||||
}
|
}
|
||||||
|
return p.db.PingContext(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
return p.db.Close()
|
// OwnsDB reports whether Close releases the wrapped database. It never does:
|
||||||
|
// the *sql.DB was opened by the caller, who is responsible for closing it.
|
||||||
|
func (p *ExistingDBProvider) OwnsDB() bool { return false }
|
||||||
|
|
||||||
|
// Close is a no-op for the wrapped database. The pool belongs to the caller, so
|
||||||
|
// closing it here would break the caller's other users of it.
|
||||||
|
func (p *ExistingDBProvider) Close() error {
|
||||||
|
logger.Warn("Not closing externally provided database: name=%s; the caller owns this *sql.DB and must close it", p.name)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// HealthCheck verifies the connection is alive
|
// HealthCheck verifies the connection is alive
|
||||||
|
|||||||
@@ -164,7 +164,7 @@ func TestExistingDBProvider_Stats(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestExistingDBProvider_Close(t *testing.T) {
|
func TestExistingDBProvider_Close_LeavesDBOpen(t *testing.T) {
|
||||||
db, err := sql.Open("sqlite3", ":memory:")
|
db, err := sql.Open("sqlite3", ":memory:")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Failed to open database: %v", err)
|
t.Fatalf("Failed to open database: %v", err)
|
||||||
@@ -177,10 +177,10 @@ func TestExistingDBProvider_Close(t *testing.T) {
|
|||||||
t.Errorf("Expected Close to succeed, got error: %v", err)
|
t.Errorf("Expected Close to succeed, got error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Verify the database is closed
|
// The caller owns the database, so Close must leave it open
|
||||||
err = db.Ping()
|
defer db.Close()
|
||||||
if err == nil {
|
if err := db.Ping(); err != nil {
|
||||||
t.Error("Expected database to be closed")
|
t.Errorf("Expected caller's database to stay open, got: %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -37,13 +37,14 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
|
|
||||||
// Set connection pool size
|
// Set connection pool size
|
||||||
if cfg.GetMaxOpenConns() != nil {
|
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)
|
clientOpts.SetMaxPoolSize(maxPoolSize)
|
||||||
}
|
}
|
||||||
|
|
||||||
if cfg.GetMaxIdleConns() != nil {
|
// MaxIdleConns is a ceiling on idle connections, not a pre-warmed minimum
|
||||||
minPoolSize := uint64(*cfg.GetMaxIdleConns())
|
// (MinPoolSize), so only the idle-time limit maps onto the Mongo pool.
|
||||||
clientOpts.SetMinPoolSize(minPoolSize)
|
if cfg.GetConnMaxIdleTime() != nil {
|
||||||
|
clientOpts.SetMaxConnIdleTime(*cfg.GetConnMaxIdleTime())
|
||||||
}
|
}
|
||||||
|
|
||||||
// Set timeouts
|
// Set timeouts
|
||||||
@@ -65,12 +66,11 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
var client *mongo.Client
|
var client *mongo.Client
|
||||||
var lastErr error
|
var lastErr error
|
||||||
|
|
||||||
retryAttempts := 3
|
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
|
||||||
retryDelay := 1 * time.Second
|
|
||||||
|
|
||||||
for attempt := 0; attempt < retryAttempts; attempt++ {
|
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||||
if attempt > 0 {
|
if attempt > 0 {
|
||||||
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
|
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Info("Retrying MongoDB connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
logger.Info("Retrying MongoDB connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
||||||
}
|
}
|
||||||
@@ -87,7 +87,7 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
lastErr = err
|
lastErr = err
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Warn("Failed to connect to MongoDB", "error", err)
|
logger.Warn("Failed to connect to MongoDB: %v", err)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -101,7 +101,7 @@ func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
lastErr = err
|
lastErr = err
|
||||||
_ = client.Disconnect(ctx)
|
_ = client.Disconnect(ctx)
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Warn("Failed to ping MongoDB", "error", err)
|
logger.Warn("Failed to ping MongoDB: %v", err)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -35,12 +35,11 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
var db *sql.DB
|
var db *sql.DB
|
||||||
var lastErr error
|
var lastErr error
|
||||||
|
|
||||||
retryAttempts := 3 // Default retry attempts
|
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
|
||||||
retryDelay := 1 * time.Second
|
|
||||||
|
|
||||||
for attempt := 0; attempt < retryAttempts; attempt++ {
|
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||||
if attempt > 0 {
|
if attempt > 0 {
|
||||||
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
|
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Info("Retrying MSSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
logger.Info("Retrying MSSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
||||||
}
|
}
|
||||||
@@ -57,7 +56,7 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
lastErr = err
|
lastErr = err
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Warn("Failed to open MSSQL connection", "error", err)
|
logger.Warn("Failed to open MSSQL connection: %v", err)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -69,9 +68,9 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
|||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
lastErr = err
|
lastErr = err
|
||||||
db.Close()
|
db.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Warn("Failed to ping MSSQL database", "error", err)
|
logger.Warn("Failed to ping MSSQL database: %v", err)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,132 @@
|
|||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql/driver"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/stdlib"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// tcpKeepAlive is how often keepalive probes are sent on idle connections.
|
||||||
|
tcpKeepAlive = 30 * time.Second
|
||||||
|
// tcpUserTimeout bounds how long written data may stay unacknowledged before
|
||||||
|
// the kernel drops the socket. Without it a query on a silently dead peer
|
||||||
|
// waits for tcp_retries2 (about 15 minutes).
|
||||||
|
tcpUserTimeout = 30 * time.Second
|
||||||
|
// resetSessionTimeout bounds the liveness ping database/sql triggers when a
|
||||||
|
// pooled connection is reused, which otherwise runs on the request context.
|
||||||
|
resetSessionTimeout = 5 * time.Second
|
||||||
|
)
|
||||||
|
|
||||||
|
// pgConnector is a driver.Connector whose connections can be retired without
|
||||||
|
// closing the *sql.DB. Reconnecting bumps a generation; connections created
|
||||||
|
// under an older generation report themselves invalid and database/sql
|
||||||
|
// discards them and dials new ones. Every handle wrapping the *sql.DB keeps
|
||||||
|
// working across a reconnect.
|
||||||
|
type pgConnector struct {
|
||||||
|
inner atomic.Pointer[connectorState]
|
||||||
|
generation atomic.Uint64
|
||||||
|
}
|
||||||
|
|
||||||
|
type connectorState struct {
|
||||||
|
connector driver.Connector
|
||||||
|
gen uint64
|
||||||
|
}
|
||||||
|
|
||||||
|
func newPGConnector(cfg *pgx.ConnConfig) *pgConnector {
|
||||||
|
c := &pgConnector{}
|
||||||
|
c.swap(cfg)
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// swap installs a new connection config under a fresh generation.
|
||||||
|
func (c *pgConnector) swap(cfg *pgx.ConnConfig) {
|
||||||
|
gen := c.generation.Add(1)
|
||||||
|
c.inner.Store(&connectorState{connector: stdlib.GetConnector(*cfg), gen: gen})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pgConnector) Connect(ctx context.Context) (driver.Conn, error) {
|
||||||
|
st := c.inner.Load()
|
||||||
|
conn, err := st.connector.Connect(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
sc, ok := conn.(*stdlib.Conn)
|
||||||
|
if !ok {
|
||||||
|
conn.Close() //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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -3,12 +3,12 @@ package providers
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"math"
|
"math"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
_ "github.com/jackc/pgx/v5/stdlib" // PostgreSQL driver
|
|
||||||
"go.mongodb.org/mongo-driver/mongo"
|
"go.mongodb.org/mongo-driver/mongo"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
@@ -17,6 +17,7 @@ import (
|
|||||||
// PostgresProvider implements Provider for PostgreSQL databases
|
// PostgresProvider implements Provider for PostgreSQL databases
|
||||||
type PostgresProvider struct {
|
type PostgresProvider struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
|
connector *pgConnector
|
||||||
config ConnectionConfig
|
config ConnectionConfig
|
||||||
listener *PostgresListener
|
listener *PostgresListener
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
@@ -29,22 +30,24 @@ func NewPostgresProvider() *PostgresProvider {
|
|||||||
|
|
||||||
// Connect establishes a PostgreSQL connection
|
// Connect establishes a PostgreSQL connection
|
||||||
func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
||||||
// Build DSN
|
connCfg, err := buildPGXConfig(cfg)
|
||||||
dsn, err := cfg.BuildDSN()
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to build DSN: %w", err)
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// The connector and *sql.DB are created once; the pool is never closed to
|
||||||
|
// recover from errors (see Refresh).
|
||||||
|
connector := newPGConnector(connCfg)
|
||||||
|
db := sql.OpenDB(connector)
|
||||||
|
|
||||||
// Connect with retry logic
|
// Connect with retry logic
|
||||||
var db *sql.DB
|
|
||||||
var lastErr error
|
var lastErr error
|
||||||
|
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
|
||||||
|
|
||||||
retryAttempts := 3 // Default retry attempts
|
connected := false
|
||||||
retryDelay := 1 * time.Second
|
|
||||||
|
|
||||||
for attempt := 0; attempt < retryAttempts; attempt++ {
|
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||||
if attempt > 0 {
|
if attempt > 0 {
|
||||||
delay := calculateBackoff(attempt, retryDelay, 10*time.Second)
|
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Info("Retrying PostgreSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
logger.Info("Retrying PostgreSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
||||||
}
|
}
|
||||||
@@ -52,20 +55,11 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
|
|||||||
select {
|
select {
|
||||||
case <-time.After(delay):
|
case <-time.After(delay):
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
|
db.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
|
||||||
return ctx.Err()
|
return ctx.Err()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Open database connection
|
|
||||||
db, err = sql.Open("pgx", dsn)
|
|
||||||
if err != nil {
|
|
||||||
lastErr = err
|
|
||||||
if cfg.GetEnableLogging() {
|
|
||||||
logger.Warn("Failed to open PostgreSQL connection", "error", err)
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
// Test the connection with context timeout
|
// Test the connection with context timeout
|
||||||
connectCtx, cancel := context.WithTimeout(ctx, cfg.GetConnectTimeout())
|
connectCtx, cancel := context.WithTimeout(ctx, cfg.GetConnectTimeout())
|
||||||
err = db.PingContext(connectCtx)
|
err = db.PingContext(connectCtx)
|
||||||
@@ -73,18 +67,18 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
|
|||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
lastErr = err
|
lastErr = err
|
||||||
db.Close()
|
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
logger.Warn("Failed to ping PostgreSQL database", "error", err)
|
logger.Warn("Failed to ping PostgreSQL database: %v", err)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Connection successful
|
connected = true
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if !connected {
|
||||||
|
db.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
|
||||||
return fmt.Errorf("failed to connect after %d attempts: %w", retryAttempts, lastErr)
|
return fmt.Errorf("failed to connect after %d attempts: %w", retryAttempts, lastErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -103,6 +97,7 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
|
|||||||
}
|
}
|
||||||
|
|
||||||
p.db = db
|
p.db = db
|
||||||
|
p.connector = connector
|
||||||
p.config = cfg
|
p.config = cfg
|
||||||
|
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
@@ -112,34 +107,55 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close closes the PostgreSQL connection
|
// Refresh retires every pooled connection and dials fresh ones on demand,
|
||||||
func (p *PostgresProvider) Close() error {
|
// without closing the *sql.DB. Handles already handed out keep working:
|
||||||
// Close listener if it exists
|
// connections in use finish their current query and are then discarded.
|
||||||
p.mu.Lock()
|
func (p *PostgresProvider) Refresh(ctx context.Context) error {
|
||||||
if p.listener != nil {
|
if p.db == nil || p.connector == nil {
|
||||||
if err := p.listener.Close(); err != nil {
|
return fmt.Errorf("database connection is not initialized")
|
||||||
p.mu.Unlock()
|
|
||||||
return fmt.Errorf("failed to close listener: %w", err)
|
|
||||||
}
|
|
||||||
p.listener = nil
|
|
||||||
}
|
|
||||||
p.mu.Unlock()
|
|
||||||
|
|
||||||
if p.db == nil {
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
err := p.db.Close()
|
connCfg, err := buildPGXConfig(p.config)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to close PostgreSQL connection: %w", err)
|
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 listener != nil {
|
||||||
|
if err := listener.Close(); err != nil {
|
||||||
|
errs = append(errs, fmt.Errorf("failed to close listener: %w", err))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if p.config.GetEnableLogging() {
|
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())
|
logger.Info("PostgreSQL connection closed: name=%s", p.config.GetName())
|
||||||
}
|
}
|
||||||
|
|
||||||
p.db = nil
|
p.db = nil
|
||||||
return nil
|
p.connector = nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return errors.Join(errs...)
|
||||||
}
|
}
|
||||||
|
|
||||||
// HealthCheck verifies the PostgreSQL connection is alive
|
// HealthCheck verifies the PostgreSQL connection is alive
|
||||||
|
|||||||
@@ -23,6 +23,10 @@ type PostgresListener struct {
|
|||||||
// Channel subscriptions
|
// Channel subscriptions
|
||||||
channels map[string]NotificationHandler
|
channels map[string]NotificationHandler
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
|
// connMu serialises use of the single pgx.Conn: it is not safe for
|
||||||
|
// concurrent use, so the notification wait and LISTEN/UNLISTEN/NOTIFY take
|
||||||
|
// turns. Lock order: connMu before mu.
|
||||||
|
connMu sync.Mutex
|
||||||
|
|
||||||
// Lifecycle management
|
// Lifecycle management
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
@@ -30,6 +34,7 @@ type PostgresListener struct {
|
|||||||
closed bool
|
closed bool
|
||||||
closeMu sync.Mutex
|
closeMu sync.Mutex
|
||||||
reconnectC chan struct{}
|
reconnectC chan struct{}
|
||||||
|
startOnce sync.Once // background goroutines start exactly once
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewPostgresListener creates a new PostgreSQL listener
|
// NewPostgresListener creates a new PostgreSQL listener
|
||||||
@@ -44,76 +49,20 @@ func NewPostgresListener(cfg ConnectionConfig) *PostgresListener {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Connect establishes a dedicated connection for listening
|
// Connect establishes a dedicated connection for listening and starts the
|
||||||
|
// background loops (once per listener).
|
||||||
func (l *PostgresListener) Connect(ctx context.Context) error {
|
func (l *PostgresListener) Connect(ctx context.Context) error {
|
||||||
dsn, err := l.config.BuildDSN()
|
conn, err := l.dial(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to build DSN: %w", err)
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse connection config
|
l.swapConn(conn)
|
||||||
connConfig, err := pgx.ParseConfig(dsn)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to parse connection config: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Connect with retry logic
|
l.startOnce.Do(func() {
|
||||||
var conn *pgx.Conn
|
|
||||||
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()
|
go l.handleNotifications()
|
||||||
|
|
||||||
// Start reconnection handler
|
|
||||||
go l.handleReconnection()
|
go l.handleReconnection()
|
||||||
|
})
|
||||||
|
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Info("PostgreSQL listener connected: name=%s", l.config.GetName())
|
logger.Info("PostgreSQL listener connected: name=%s", l.config.GetName())
|
||||||
@@ -122,30 +71,105 @@ func (l *PostgresListener) Connect(ctx context.Context) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// dial opens and verifies a new dedicated connection, with retries.
|
||||||
|
func (l *PostgresListener) dial(ctx context.Context) (*pgx.Conn, error) {
|
||||||
|
connConfig, err := buildPGXConfig(l.config)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var lastErr error
|
||||||
|
|
||||||
|
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(l.config)
|
||||||
|
|
||||||
|
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||||
|
if attempt > 0 {
|
||||||
|
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||||
|
if l.config.GetEnableLogging() {
|
||||||
|
logger.Info("Retrying PostgreSQL listener connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-time.After(delay):
|
||||||
|
case <-ctx.Done():
|
||||||
|
return nil, ctx.Err()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
conn, err := pgx.ConnectConfig(ctx, connConfig)
|
||||||
|
if err != nil {
|
||||||
|
lastErr = err
|
||||||
|
if l.config.GetEnableLogging() {
|
||||||
|
logger.Warn("Failed to connect PostgreSQL listener: %v", err)
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test the connection
|
||||||
|
if err = conn.Ping(ctx); err != nil {
|
||||||
|
lastErr = err
|
||||||
|
_ = closeConnBounded(conn)
|
||||||
|
if l.config.GetEnableLogging() {
|
||||||
|
logger.Warn("Failed to ping PostgreSQL listener: %v", err)
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
return conn, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("failed to connect listener after %d attempts: %w", retryAttempts, lastErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// closeConnBounded closes a pgx connection without ever waiting on a dead socket.
|
||||||
|
func closeConnBounded(conn *pgx.Conn) error {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), listenerCloseTimeout)
|
||||||
|
defer cancel()
|
||||||
|
return conn.Close(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
listenerCloseTimeout = 2 * time.Second
|
||||||
|
notificationPollInterval = 500 * time.Millisecond
|
||||||
|
)
|
||||||
|
|
||||||
|
// currentConn returns the live connection, or an error if the listener is
|
||||||
|
// closed or not yet connected.
|
||||||
|
func (l *PostgresListener) currentConn() (*pgx.Conn, error) {
|
||||||
|
l.closeMu.Lock()
|
||||||
|
closed := l.closed
|
||||||
|
l.closeMu.Unlock()
|
||||||
|
if closed {
|
||||||
|
return nil, fmt.Errorf("listener is closed")
|
||||||
|
}
|
||||||
|
|
||||||
|
l.mu.RLock()
|
||||||
|
conn := l.conn
|
||||||
|
l.mu.RUnlock()
|
||||||
|
if conn == nil {
|
||||||
|
return nil, fmt.Errorf("listener connection is not initialized")
|
||||||
|
}
|
||||||
|
return conn, nil
|
||||||
|
}
|
||||||
|
|
||||||
// Listen subscribes to a PostgreSQL notification channel
|
// Listen subscribes to a PostgreSQL notification channel
|
||||||
func (l *PostgresListener) Listen(channel string, handler NotificationHandler) error {
|
func (l *PostgresListener) Listen(channel string, handler NotificationHandler) error {
|
||||||
l.closeMu.Lock()
|
// Take the connection between notification waits (each wait is short).
|
||||||
if l.closed {
|
l.connMu.Lock()
|
||||||
l.closeMu.Unlock()
|
defer l.connMu.Unlock()
|
||||||
return fmt.Errorf("listener is closed")
|
|
||||||
}
|
|
||||||
l.closeMu.Unlock()
|
|
||||||
|
|
||||||
l.mu.Lock()
|
conn, err := l.currentConn()
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
if l.conn == nil {
|
|
||||||
return fmt.Errorf("listener connection is not initialized")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Execute LISTEN command
|
|
||||||
_, err := l.conn.Exec(l.ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{channel}.Sanitize()))
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := conn.Exec(l.ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{channel}.Sanitize())); err != nil {
|
||||||
return fmt.Errorf("failed to listen on channel %s: %w", channel, err)
|
return fmt.Errorf("failed to listen on channel %s: %w", channel, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Store the handler
|
l.mu.Lock()
|
||||||
l.channels[channel] = handler
|
l.channels[channel] = handler
|
||||||
|
l.mu.Unlock()
|
||||||
|
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Info("Listening on channel: name=%s, channel=%s", l.config.GetName(), channel)
|
logger.Info("Listening on channel: name=%s, channel=%s", l.config.GetName(), channel)
|
||||||
@@ -156,28 +180,21 @@ func (l *PostgresListener) Listen(channel string, handler NotificationHandler) e
|
|||||||
|
|
||||||
// Unlisten unsubscribes from a PostgreSQL notification channel
|
// Unlisten unsubscribes from a PostgreSQL notification channel
|
||||||
func (l *PostgresListener) Unlisten(channel string) error {
|
func (l *PostgresListener) Unlisten(channel string) error {
|
||||||
l.closeMu.Lock()
|
l.connMu.Lock()
|
||||||
if l.closed {
|
defer l.connMu.Unlock()
|
||||||
l.closeMu.Unlock()
|
|
||||||
return fmt.Errorf("listener is closed")
|
|
||||||
}
|
|
||||||
l.closeMu.Unlock()
|
|
||||||
|
|
||||||
l.mu.Lock()
|
conn, err := l.currentConn()
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
if l.conn == nil {
|
|
||||||
return fmt.Errorf("listener connection is not initialized")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Execute UNLISTEN command
|
|
||||||
_, err := l.conn.Exec(l.ctx, fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize()))
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := conn.Exec(l.ctx, fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize())); err != nil {
|
||||||
return fmt.Errorf("failed to unlisten from channel %s: %w", channel, err)
|
return fmt.Errorf("failed to unlisten from channel %s: %w", channel, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Remove the handler
|
l.mu.Lock()
|
||||||
delete(l.channels, channel)
|
delete(l.channels, channel)
|
||||||
|
l.mu.Unlock()
|
||||||
|
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Info("Unlistened from channel: name=%s, channel=%s", l.config.GetName(), channel)
|
logger.Info("Unlistened from channel: name=%s, channel=%s", l.config.GetName(), channel)
|
||||||
@@ -188,31 +205,24 @@ func (l *PostgresListener) Unlisten(channel string) error {
|
|||||||
|
|
||||||
// Notify sends a notification to a PostgreSQL channel
|
// Notify sends a notification to a PostgreSQL channel
|
||||||
func (l *PostgresListener) Notify(ctx context.Context, channel string, payload string) error {
|
func (l *PostgresListener) Notify(ctx context.Context, channel string, payload string) error {
|
||||||
l.closeMu.Lock()
|
l.connMu.Lock()
|
||||||
if l.closed {
|
defer l.connMu.Unlock()
|
||||||
l.closeMu.Unlock()
|
|
||||||
return fmt.Errorf("listener is closed")
|
|
||||||
}
|
|
||||||
l.closeMu.Unlock()
|
|
||||||
|
|
||||||
l.mu.RLock()
|
conn, err := l.currentConn()
|
||||||
conn := l.conn
|
|
||||||
l.mu.RUnlock()
|
|
||||||
|
|
||||||
if conn == nil {
|
|
||||||
return fmt.Errorf("listener connection is not initialized")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Execute NOTIFY command
|
|
||||||
_, err := conn.Exec(ctx, "SELECT pg_notify($1, $2)", channel, payload)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := conn.Exec(ctx, "SELECT pg_notify($1, $2)", channel, payload); err != nil {
|
||||||
return fmt.Errorf("failed to notify channel %s: %w", channel, err)
|
return fmt.Errorf("failed to notify channel %s: %w", channel, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close closes the listener and all subscriptions
|
// Close closes the listener and all subscriptions. Closing the connection drops
|
||||||
|
// every subscription server-side, so no UNLISTEN round trips are needed, and
|
||||||
|
// the close itself is bounded so a dead socket cannot hang the caller.
|
||||||
func (l *PostgresListener) Close() error {
|
func (l *PostgresListener) Close() error {
|
||||||
l.closeMu.Lock()
|
l.closeMu.Lock()
|
||||||
if l.closed {
|
if l.closed {
|
||||||
@@ -225,27 +235,26 @@ func (l *PostgresListener) Close() error {
|
|||||||
// Cancel context to stop background goroutines
|
// Cancel context to stop background goroutines
|
||||||
l.cancel()
|
l.cancel()
|
||||||
|
|
||||||
|
// The cancelled ctx makes the notification wait return promptly, releasing
|
||||||
|
// connMu; closing the conn while it is being read would race inside pgx.
|
||||||
|
l.connMu.Lock()
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
conn := l.conn
|
||||||
|
l.conn = nil
|
||||||
|
l.channels = make(map[string]NotificationHandler)
|
||||||
|
l.mu.Unlock()
|
||||||
|
|
||||||
if l.conn == nil {
|
if conn == nil {
|
||||||
|
l.connMu.Unlock()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Unlisten from all channels
|
err := closeConnBounded(conn)
|
||||||
for channel := range l.channels {
|
l.connMu.Unlock()
|
||||||
_, _ = l.conn.Exec(context.Background(), fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize()))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Close connection
|
|
||||||
err := l.conn.Close(context.Background())
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to close listener connection: %w", err)
|
return fmt.Errorf("failed to close listener connection: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
l.conn = nil
|
|
||||||
l.channels = make(map[string]NotificationHandler)
|
|
||||||
|
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Info("PostgreSQL listener closed: name=%s", l.config.GetName())
|
logger.Info("PostgreSQL listener closed: name=%s", l.config.GetName())
|
||||||
}
|
}
|
||||||
@@ -262,20 +271,26 @@ func (l *PostgresListener) handleNotifications() {
|
|||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
|
|
||||||
|
l.connMu.Lock()
|
||||||
l.mu.RLock()
|
l.mu.RLock()
|
||||||
conn := l.conn
|
conn := l.conn
|
||||||
l.mu.RUnlock()
|
l.mu.RUnlock()
|
||||||
|
|
||||||
if conn == nil {
|
if conn == nil {
|
||||||
|
l.connMu.Unlock()
|
||||||
// Connection not available, wait for reconnection
|
// Connection not available, wait for reconnection
|
||||||
time.Sleep(100 * time.Millisecond)
|
if !l.sleep(100 * time.Millisecond) {
|
||||||
|
return
|
||||||
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Wait for notification with timeout
|
// Wait for a notification with a short timeout, so Listen/Unlisten/Notify
|
||||||
ctx, cancel := context.WithTimeout(l.ctx, 5*time.Second)
|
// waiting on connMu are served promptly.
|
||||||
|
ctx, cancel := context.WithTimeout(l.ctx, notificationPollInterval)
|
||||||
notification, err := conn.WaitForNotification(ctx)
|
notification, err := conn.WaitForNotification(ctx)
|
||||||
cancel()
|
cancel()
|
||||||
|
l.connMu.Unlock()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Check if context was cancelled
|
// Check if context was cancelled
|
||||||
@@ -291,13 +306,15 @@ func (l *PostgresListener) handleNotifications() {
|
|||||||
|
|
||||||
// Connection error, trigger reconnection
|
// Connection error, trigger reconnection
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Warn("Notification error, triggering reconnection", "error", err)
|
logger.Warn("Notification error, triggering reconnection: %v", err)
|
||||||
}
|
}
|
||||||
select {
|
select {
|
||||||
case l.reconnectC <- struct{}{}:
|
case l.reconnectC <- struct{}{}:
|
||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
time.Sleep(1 * time.Second)
|
if !l.sleep(1 * time.Second) {
|
||||||
|
return
|
||||||
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -322,7 +339,22 @@ func (l *PostgresListener) handleNotifications() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleReconnection manages automatic reconnection
|
// sleep waits for d or until the listener is closed; it reports whether the
|
||||||
|
// listener is still running.
|
||||||
|
func (l *PostgresListener) sleep(d time.Duration) bool {
|
||||||
|
t := time.NewTimer(d)
|
||||||
|
defer t.Stop()
|
||||||
|
select {
|
||||||
|
case <-t.C:
|
||||||
|
return true
|
||||||
|
case <-l.ctx.Done():
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// handleReconnection manages automatic reconnection. It runs as a single
|
||||||
|
// goroutine and dials replacement connections directly rather than through the
|
||||||
|
// public Connect, so no extra loops are started.
|
||||||
func (l *PostgresListener) handleReconnection() {
|
func (l *PostgresListener) handleReconnection() {
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
@@ -333,31 +365,21 @@ func (l *PostgresListener) handleReconnection() {
|
|||||||
logger.Info("Attempting to reconnect listener: name=%s", l.config.GetName())
|
logger.Info("Attempting to reconnect listener: name=%s", l.config.GetName())
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close existing connection
|
ctx, cancel := context.WithTimeout(l.ctx, 30*time.Second)
|
||||||
l.mu.Lock()
|
err := l.reconnect(ctx)
|
||||||
if l.conn != nil {
|
|
||||||
l.conn.Close(context.Background())
|
|
||||||
l.conn = nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Save current subscriptions
|
|
||||||
channels := make(map[string]NotificationHandler)
|
|
||||||
for ch, handler := range l.channels {
|
|
||||||
channels[ch] = handler
|
|
||||||
}
|
|
||||||
l.mu.Unlock()
|
|
||||||
|
|
||||||
// Attempt reconnection
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
|
||||||
err := l.Connect(ctx)
|
|
||||||
cancel()
|
cancel()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
if l.ctx.Err() != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Error("Failed to reconnect listener: name=%s, error=%v", l.config.GetName(), err)
|
logger.Error("Failed to reconnect listener: name=%s, error=%v", l.config.GetName(), err)
|
||||||
}
|
}
|
||||||
// Retry after delay
|
// Retry after delay
|
||||||
time.Sleep(5 * time.Second)
|
if !l.sleep(5 * time.Second) {
|
||||||
|
return
|
||||||
|
}
|
||||||
select {
|
select {
|
||||||
case l.reconnectC <- struct{}{}:
|
case l.reconnectC <- struct{}{}:
|
||||||
default:
|
default:
|
||||||
@@ -365,15 +387,6 @@ func (l *PostgresListener) handleReconnection() {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Resubscribe to all channels
|
|
||||||
for channel, handler := range channels {
|
|
||||||
if err := l.Listen(channel, handler); err != nil {
|
|
||||||
if l.config.GetEnableLogging() {
|
|
||||||
logger.Error("Failed to resubscribe to channel: name=%s, channel=%s, error=%v", l.config.GetName(), channel, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if l.config.GetEnableLogging() {
|
if l.config.GetEnableLogging() {
|
||||||
logger.Info("Listener reconnected successfully: name=%s", l.config.GetName())
|
logger.Info("Listener reconnected successfully: name=%s", l.config.GetName())
|
||||||
}
|
}
|
||||||
@@ -381,6 +394,51 @@ func (l *PostgresListener) handleReconnection() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// reconnect replaces the connection and resubscribes every channel on the new
|
||||||
|
// connection before publishing it, so the notification loop never touches a
|
||||||
|
// half-initialised conn.
|
||||||
|
func (l *PostgresListener) reconnect(ctx context.Context) error {
|
||||||
|
conn, err := l.dial(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
l.mu.RLock()
|
||||||
|
channels := make([]string, 0, len(l.channels))
|
||||||
|
for ch := range l.channels {
|
||||||
|
channels = append(channels, ch)
|
||||||
|
}
|
||||||
|
l.mu.RUnlock()
|
||||||
|
|
||||||
|
for _, ch := range channels {
|
||||||
|
if _, err := conn.Exec(ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{ch}.Sanitize())); err != nil {
|
||||||
|
_ = closeConnBounded(conn)
|
||||||
|
return fmt.Errorf("failed to resubscribe to channel %s: %w", ch, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if l.ctx.Err() != nil {
|
||||||
|
_ = closeConnBounded(conn)
|
||||||
|
return l.ctx.Err()
|
||||||
|
}
|
||||||
|
l.swapConn(conn)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// swapConn installs conn and closes the previous one. The old connection is
|
||||||
|
// closed under connMu so it is never closed while another goroutine is using it.
|
||||||
|
func (l *PostgresListener) swapConn(conn *pgx.Conn) {
|
||||||
|
l.connMu.Lock()
|
||||||
|
l.mu.Lock()
|
||||||
|
old := l.conn
|
||||||
|
l.conn = conn
|
||||||
|
l.mu.Unlock()
|
||||||
|
if old != nil {
|
||||||
|
_ = closeConnBounded(old)
|
||||||
|
}
|
||||||
|
l.connMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
// IsConnected returns true if the listener is connected
|
// IsConnected returns true if the listener is connected
|
||||||
func (l *PostgresListener) IsConnected() bool {
|
func (l *PostgresListener) IsConnected() bool {
|
||||||
l.mu.RLock()
|
l.mu.RLock()
|
||||||
|
|||||||
@@ -4,17 +4,11 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"errors"
|
"errors"
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"go.mongodb.org/mongo-driver/mongo"
|
"go.mongodb.org/mongo-driver/mongo"
|
||||||
)
|
)
|
||||||
|
|
||||||
// isDBClosed reports whether err indicates the *sql.DB has been closed.
|
|
||||||
func isDBClosed(err error) bool {
|
|
||||||
return err != nil && strings.Contains(err.Error(), "sql: database is closed")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Common errors
|
// Common errors
|
||||||
var (
|
var (
|
||||||
// ErrNotSQLDatabase is returned when attempting SQL operations on a non-SQL database
|
// ErrNotSQLDatabase is returned when attempting SQL operations on a non-SQL database
|
||||||
@@ -63,6 +57,31 @@ type ConnectionConfig interface {
|
|||||||
GetConnMaxLifetime() *time.Duration
|
GetConnMaxLifetime() *time.Duration
|
||||||
GetConnMaxIdleTime() *time.Duration
|
GetConnMaxIdleTime() *time.Duration
|
||||||
GetReadPreference() string
|
GetReadPreference() string
|
||||||
|
GetRetryAttempts() int
|
||||||
|
GetRetryDelay() time.Duration
|
||||||
|
GetRetryMaxDelay() time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
// retryPolicy returns the configured retry settings, falling back to defaults.
|
||||||
|
func retryPolicy(cfg ConnectionConfig) (attempts int, delay, maxDelay time.Duration) {
|
||||||
|
attempts, delay, maxDelay = cfg.GetRetryAttempts(), cfg.GetRetryDelay(), cfg.GetRetryMaxDelay()
|
||||||
|
if attempts <= 0 {
|
||||||
|
attempts = 3
|
||||||
|
}
|
||||||
|
if delay <= 0 {
|
||||||
|
delay = time.Second
|
||||||
|
}
|
||||||
|
if maxDelay <= 0 {
|
||||||
|
maxDelay = 10 * time.Second
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Refresher is implemented by providers that can retire their pooled
|
||||||
|
// connections and dial fresh ones without closing the shared *sql.DB, so
|
||||||
|
// handles already handed out keep working.
|
||||||
|
type Refresher interface {
|
||||||
|
Refresh(ctx context.Context) error
|
||||||
}
|
}
|
||||||
|
|
||||||
// Provider creates and manages the underlying database connection
|
// Provider creates and manages the underlying database connection
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -17,7 +18,6 @@ import (
|
|||||||
type SQLiteProvider struct {
|
type SQLiteProvider struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
dbMu sync.RWMutex
|
dbMu sync.RWMutex
|
||||||
dbFactory func() (*sql.DB, error)
|
|
||||||
config ConnectionConfig
|
config ConnectionConfig
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -26,6 +26,22 @@ func NewSQLiteProvider() *SQLiteProvider {
|
|||||||
return &SQLiteProvider{}
|
return &SQLiteProvider{}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isMemoryDSN reports whether the SQLite DSN refers to a private in-memory
|
||||||
|
// database (each pooled connection would get its own empty database).
|
||||||
|
func isMemoryDSN(dsn string) bool {
|
||||||
|
path := dsn
|
||||||
|
if i := strings.IndexByte(path, '?'); i >= 0 {
|
||||||
|
path = path[:i]
|
||||||
|
}
|
||||||
|
if path == ":memory:" || path == "" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if strings.Contains(dsn, "mode=memory") && !strings.Contains(dsn, "cache=shared") {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return path == "file::memory:" && !strings.Contains(dsn, "cache=shared")
|
||||||
|
}
|
||||||
|
|
||||||
// Connect establishes a SQLite connection
|
// Connect establishes a SQLite connection
|
||||||
func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
||||||
// Build DSN
|
// Build DSN
|
||||||
@@ -46,20 +62,25 @@ func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) erro
|
|||||||
cancel()
|
cancel()
|
||||||
|
|
||||||
if err != nil {
|
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)
|
return fmt.Errorf("failed to ping SQLite database: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Configure connection pool
|
if isMemoryDSN(dsn) {
|
||||||
// Note: SQLite works best with MaxOpenConns=1 for write operations
|
// A private in-memory database exists per connection and disappears when
|
||||||
// but can handle multiple readers
|
// that connection closes, so pin the pool to one connection that is
|
||||||
|
// never recycled.
|
||||||
|
db.SetMaxOpenConns(1)
|
||||||
|
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 {
|
if cfg.GetMaxOpenConns() != nil {
|
||||||
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
|
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
|
||||||
} else {
|
} else {
|
||||||
// Default to 1 for SQLite to avoid "database is locked" errors
|
|
||||||
db.SetMaxOpenConns(1)
|
db.SetMaxOpenConns(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
if cfg.GetMaxIdleConns() != nil {
|
if cfg.GetMaxIdleConns() != nil {
|
||||||
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
|
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
|
||||||
}
|
}
|
||||||
@@ -69,29 +90,11 @@ func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) erro
|
|||||||
if cfg.GetConnMaxIdleTime() != nil {
|
if cfg.GetConnMaxIdleTime() != nil {
|
||||||
db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime())
|
db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime())
|
||||||
}
|
}
|
||||||
|
|
||||||
// Enable WAL mode for better concurrent access
|
|
||||||
_, err = db.ExecContext(ctx, "PRAGMA journal_mode=WAL")
|
|
||||||
if err != nil {
|
|
||||||
if cfg.GetEnableLogging() {
|
|
||||||
logger.Warn("Failed to enable WAL mode for SQLite", "error", err)
|
|
||||||
}
|
|
||||||
// Don't fail connection if WAL mode cannot be enabled
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
p.dbMu.Lock()
|
||||||
p.db = db
|
p.db = db
|
||||||
|
p.dbMu.Unlock()
|
||||||
p.config = cfg
|
p.config = cfg
|
||||||
|
|
||||||
if cfg.GetEnableLogging() {
|
if cfg.GetEnableLogging() {
|
||||||
@@ -132,14 +135,7 @@ func (p *SQLiteProvider) HealthCheck(ctx context.Context) error {
|
|||||||
|
|
||||||
// Execute a simple query to verify the database is accessible
|
// Execute a simple query to verify the database is accessible
|
||||||
var result int
|
var result int
|
||||||
run := func() error { return p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result) }
|
if err := p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result); err != nil {
|
||||||
err := run()
|
|
||||||
if isDBClosed(err) {
|
|
||||||
if reconnErr := p.reconnectDB(); reconnErr == nil {
|
|
||||||
err = run()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("health check failed: %w", err)
|
return fmt.Errorf("health check failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -150,32 +146,12 @@ func (p *SQLiteProvider) HealthCheck(ctx context.Context) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// WithDBFactory configures a factory used to reopen the database connection if it is closed.
|
|
||||||
func (p *SQLiteProvider) WithDBFactory(factory func() (*sql.DB, error)) *SQLiteProvider {
|
|
||||||
p.dbFactory = factory
|
|
||||||
return p
|
|
||||||
}
|
|
||||||
|
|
||||||
func (p *SQLiteProvider) getDB() *sql.DB {
|
func (p *SQLiteProvider) getDB() *sql.DB {
|
||||||
p.dbMu.RLock()
|
p.dbMu.RLock()
|
||||||
defer p.dbMu.RUnlock()
|
defer p.dbMu.RUnlock()
|
||||||
return p.db
|
return p.db
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *SQLiteProvider) reconnectDB() error {
|
|
||||||
if p.dbFactory == nil {
|
|
||||||
return fmt.Errorf("no db factory configured for reconnect")
|
|
||||||
}
|
|
||||||
newDB, err := p.dbFactory()
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
p.dbMu.Lock()
|
|
||||||
p.db = newDB
|
|
||||||
p.dbMu.Unlock()
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetNative returns the native *sql.DB connection
|
// GetNative returns the native *sql.DB connection
|
||||||
func (p *SQLiteProvider) GetNative() (*sql.DB, error) {
|
func (p *SQLiteProvider) GetNative() (*sql.DB, error) {
|
||||||
if p.db == nil {
|
if p.db == nil {
|
||||||
|
|||||||
@@ -0,0 +1,24 @@
|
|||||||
|
//go:build linux
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"syscall"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
// setTCPUserTimeout returns a net.Dialer Control func setting TCP_USER_TIMEOUT.
|
||||||
|
func setTCPUserTimeout(d time.Duration) func(network, address string, c syscall.RawConn) error {
|
||||||
|
return func(network, address string, c syscall.RawConn) error {
|
||||||
|
var sockErr error
|
||||||
|
err := c.Control(func(fd uintptr) {
|
||||||
|
sockErr = unix.SetsockoptInt(int(fd), unix.IPPROTO_TCP, unix.TCP_USER_TIMEOUT, int(d.Milliseconds()))
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return sockErr
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,14 @@
|
|||||||
|
//go:build !linux
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"syscall"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// setTCPUserTimeout is a no-op where TCP_USER_TIMEOUT is unavailable; TCP
|
||||||
|
// keepalive still applies.
|
||||||
|
func setTCPUserTimeout(time.Duration) func(network, address string, c syscall.RawConn) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,69 @@
|
|||||||
|
package dbmanager
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestLiveServerRestart(t *testing.T) {
|
||||||
|
dir := os.Getenv("PG_RESTART_DIR")
|
||||||
|
if dir == "" {
|
||||||
|
t.Skip("PG_RESTART_DIR not set")
|
||||||
|
}
|
||||||
|
mgr, _ := NewManager(ManagerConfig{
|
||||||
|
DefaultConnection: "pg",
|
||||||
|
Connections: map[string]ConnectionConfig{"pg": {
|
||||||
|
Name: "pg", Type: DatabaseTypePostgreSQL, Host: "127.0.0.1", Port: 54329,
|
||||||
|
User: "postgres", Database: "postgres", ConnectTimeout: 2 * time.Second,
|
||||||
|
}},
|
||||||
|
HealthCheckInterval: 500 * time.Millisecond,
|
||||||
|
})
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := mgr.Connect(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer mgr.Close()
|
||||||
|
conn, _ := mgr.GetDefault()
|
||||||
|
db, _ := conn.Bun()
|
||||||
|
gdb, _ := conn.GORM()
|
||||||
|
query := func() error { var n int; return db.DB.QueryRow("select 1").Scan(&n) }
|
||||||
|
if err := query(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
run := func(args ...string) {
|
||||||
|
// Output must not be piped: the daemonised server would hold the pipe open.
|
||||||
|
if err := exec.Command("pg_ctl", append([]string{"-D", dir}, args...)...).Run(); err != nil {
|
||||||
|
t.Fatalf("pg_ctl %v: %v", args, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
run("-m", "immediate", "-w", "stop") // crash-style shutdown
|
||||||
|
time.Sleep(1500 * time.Millisecond) // health checks fail meanwhile
|
||||||
|
if err := query(); err == nil {
|
||||||
|
t.Fatal("expected failure while server is down")
|
||||||
|
}
|
||||||
|
run("-l", dir+"/restart.log", "-o", "-p 54329 -k "+dir+" -c listen_addresses=127.0.0.1", "-w", "start")
|
||||||
|
|
||||||
|
var last error
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
if last = query(); last == nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
t.Logf("attempt %d after restart: %v", i, last)
|
||||||
|
time.Sleep(200 * time.Millisecond)
|
||||||
|
}
|
||||||
|
if last != nil {
|
||||||
|
t.Fatalf("held bun handle never recovered: %v", last)
|
||||||
|
}
|
||||||
|
var n int
|
||||||
|
if err := gdb.Raw("select 1").Scan(&n).Error; err != nil {
|
||||||
|
t.Fatalf("held gorm handle: %v", err)
|
||||||
|
}
|
||||||
|
time.Sleep(time.Second)
|
||||||
|
if err := conn.HealthCheck(ctx); err != nil {
|
||||||
|
t.Fatalf("health check after restart: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -66,7 +66,7 @@ func (s *SentryProvider) CaptureError(ctx context.Context, err error, severity S
|
|||||||
}
|
}
|
||||||
|
|
||||||
if extra != nil {
|
if extra != nil {
|
||||||
event.Extra = extra
|
event.Contexts["extra"] = sentry.Context(extra)
|
||||||
}
|
}
|
||||||
|
|
||||||
hub.CaptureEvent(event)
|
hub.CaptureEvent(event)
|
||||||
@@ -88,7 +88,7 @@ func (s *SentryProvider) CaptureMessage(ctx context.Context, message string, sev
|
|||||||
event.Message = message
|
event.Message = message
|
||||||
|
|
||||||
if extra != nil {
|
if extra != nil {
|
||||||
event.Extra = extra
|
event.Contexts["extra"] = sentry.Context(extra)
|
||||||
}
|
}
|
||||||
|
|
||||||
hub.CaptureEvent(event)
|
hub.CaptureEvent(event)
|
||||||
@@ -115,12 +115,15 @@ func (s *SentryProvider) CapturePanic(ctx context.Context, recovered interface{}
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
if extra != nil {
|
extraCtx := sentry.Context{}
|
||||||
event.Extra = extra
|
for k, v := range extra {
|
||||||
|
extraCtx[k] = v
|
||||||
}
|
}
|
||||||
|
|
||||||
if stackTrace != nil {
|
if stackTrace != nil {
|
||||||
event.Extra["stack_trace"] = string(stackTrace)
|
extraCtx["stack_trace"] = string(stackTrace)
|
||||||
|
}
|
||||||
|
if len(extraCtx) > 0 {
|
||||||
|
event.Contexts["extra"] = extraCtx
|
||||||
}
|
}
|
||||||
|
|
||||||
hub.CaptureEvent(event)
|
hub.CaptureEvent(event)
|
||||||
|
|||||||
@@ -502,9 +502,9 @@ func TestBrokerProcessingModes(t *testing.T) {
|
|||||||
broker.Start(context.Background())
|
broker.Start(context.Background())
|
||||||
defer broker.Stop(context.Background())
|
defer broker.Stop(context.Background())
|
||||||
|
|
||||||
called := false
|
var called atomic.Bool
|
||||||
broker.Subscribe("test.*", EventHandlerFunc(func(ctx context.Context, event *Event) error {
|
broker.Subscribe("test.*", EventHandlerFunc(func(ctx context.Context, event *Event) error {
|
||||||
called = true
|
called.Store(true)
|
||||||
return nil
|
return nil
|
||||||
}))
|
}))
|
||||||
|
|
||||||
@@ -516,7 +516,7 @@ func TestBrokerProcessingModes(t *testing.T) {
|
|||||||
time.Sleep(50 * time.Millisecond)
|
time.Sleep(50 * time.Millisecond)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !called {
|
if !called.Load() {
|
||||||
t.Error("Expected handler to be called")
|
t.Error("Expected handler to be called")
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -584,7 +584,7 @@ func (dp *DatabaseProvider) pollEvents() {
|
|||||||
dp.stats.EventsConsumed.Add(1)
|
dp.stats.EventsConsumed.Add(1)
|
||||||
sub.lastSeenID = event.ID
|
sub.lastSeenID = event.ID
|
||||||
case <-sub.ctx.Done():
|
case <-sub.ctx.Done():
|
||||||
rows.Close()
|
rows.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
|
||||||
return
|
return
|
||||||
default:
|
default:
|
||||||
// Channel full, skip
|
// Channel full, skip
|
||||||
@@ -595,7 +595,7 @@ func (dp *DatabaseProvider) pollEvents() {
|
|||||||
sub.lastSeenID = event.ID
|
sub.lastSeenID = event.ID
|
||||||
}
|
}
|
||||||
|
|
||||||
rows.Close()
|
rows.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+146
-20
@@ -3,6 +3,7 @@ package funcspec
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -16,6 +17,7 @@ import (
|
|||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/restheadspec"
|
"github.com/bitechdev/ResolveSpec/pkg/restheadspec"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
)
|
)
|
||||||
@@ -168,9 +170,16 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
|||||||
// Replace meta variables in SQL
|
// Replace meta variables in SQL
|
||||||
sqlquery = h.replaceMetaVariables(sqlquery, r, userCtx, metainfo, variables)
|
sqlquery = h.replaceMetaVariables(sqlquery, r, userCtx, metainfo, variables)
|
||||||
|
|
||||||
// Remove unused input variables
|
// Replace variables from provided values, then blank any remaining unused ones
|
||||||
if options.BlankParams {
|
|
||||||
for _, kw := range inputvars {
|
for _, kw := range inputvars {
|
||||||
|
varName := kw[1 : len(kw)-1] // strip [ and ]
|
||||||
|
if val, ok := variables[varName]; ok {
|
||||||
|
if strVal := fmt.Sprintf("%v", val); strVal != "" {
|
||||||
|
sqlquery = strings.ReplaceAll(sqlquery, kw, safeSubstituteVar(sqlquery, kw, strVal))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if options.BlankParams {
|
||||||
replacement := getReplacementForBlankParam(sqlquery, kw)
|
replacement := getReplacementForBlankParam(sqlquery, kw)
|
||||||
sqlquery = strings.ReplaceAll(sqlquery, kw, replacement)
|
sqlquery = strings.ReplaceAll(sqlquery, kw, replacement)
|
||||||
logger.Debug("Replaced unused variable %s with: %s", kw, replacement)
|
logger.Debug("Replaced unused variable %s with: %s", kw, replacement)
|
||||||
@@ -188,7 +197,7 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
|||||||
hookCtx.Tx = tx
|
hookCtx.Tx = tx
|
||||||
|
|
||||||
// Execute BeforeQueryList hook (inside transaction)
|
// Execute BeforeQueryList hook (inside transaction)
|
||||||
if err := h.hooks.Execute(BeforeQueryList, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeQueryList, hookCtx); err != nil {
|
||||||
logger.Error("BeforeQueryList hook failed: %v", err)
|
logger.Error("BeforeQueryList hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
return err
|
||||||
@@ -252,7 +261,7 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
|||||||
|
|
||||||
// Execute BeforeSQLExec hook
|
// Execute BeforeSQLExec hook
|
||||||
hookCtx.SQLQuery = sqlquery
|
hookCtx.SQLQuery = sqlquery
|
||||||
if err := h.hooks.Execute(BeforeSQLExec, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
|
||||||
logger.Error("BeforeSQLExec hook failed: %v", err)
|
logger.Error("BeforeSQLExec hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
return err
|
||||||
@@ -322,7 +331,10 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
|||||||
w.Header().Set("Content-Range", fmt.Sprintf("items %d-%d/%d", respOffset, respOffset+len(dbobjlist), total))
|
w.Header().Set("Content-Range", fmt.Sprintf("items %d-%d/%d", respOffset, respOffset+len(dbobjlist), total))
|
||||||
logger.Info("Serving: Records %d of %d", len(dbobjlist), total)
|
logger.Info("Serving: Records %d of %d", len(dbobjlist), total)
|
||||||
|
|
||||||
// Execute BeforeResponse hook
|
// Execute BeforeResponse hook. The transaction has already committed by
|
||||||
|
// this point, so hooks must use the pooled connection rather than the
|
||||||
|
// now-dead tx.
|
||||||
|
hookCtx.Tx = h.db
|
||||||
hookCtx.Result = dbobjlist
|
hookCtx.Result = dbobjlist
|
||||||
hookCtx.Total = total
|
hookCtx.Total = total
|
||||||
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
||||||
@@ -359,13 +371,17 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
|||||||
}
|
}
|
||||||
|
|
||||||
case "detail":
|
case "detail":
|
||||||
// Detail format: complex API with metadata
|
// Detail format: { count, fields, items, tablename, tableprefix, total }
|
||||||
|
tableName := r.URL.Path
|
||||||
|
tablePrefix := reflection.ExtractTableNameOnly(tableName)
|
||||||
|
fields := buildDetailFieldsFromRows(dbobjlist)
|
||||||
metaobj := map[string]interface{}{
|
metaobj := map[string]interface{}{
|
||||||
"items": dbobjlist,
|
|
||||||
"count": fmt.Sprintf("%d", len(dbobjlist)),
|
"count": fmt.Sprintf("%d", len(dbobjlist)),
|
||||||
|
"fields": fields,
|
||||||
|
"items": dbobjlist,
|
||||||
|
"tablename": tableName,
|
||||||
|
"tableprefix": tablePrefix,
|
||||||
"total": fmt.Sprintf("%d", total),
|
"total": fmt.Sprintf("%d", total),
|
||||||
"tablename": r.URL.Path,
|
|
||||||
"tableprefix": "gsql",
|
|
||||||
}
|
}
|
||||||
data, err := json.Marshal(metaobj)
|
data, err := json.Marshal(metaobj)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -520,9 +536,16 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Remove unused input variables
|
// Replace variables from provided values, then blank any remaining unused ones
|
||||||
if options.BlankParams {
|
|
||||||
for _, kw := range inputvars {
|
for _, kw := range inputvars {
|
||||||
|
varName := kw[1 : len(kw)-1] // strip [ and ]
|
||||||
|
if val, ok := variables[varName]; ok {
|
||||||
|
if strVal := fmt.Sprintf("%v", val); strVal != "" {
|
||||||
|
sqlquery = strings.ReplaceAll(sqlquery, kw, safeSubstituteVar(sqlquery, kw, strVal))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if options.BlankParams {
|
||||||
replacement := getReplacementForBlankParam(sqlquery, kw)
|
replacement := getReplacementForBlankParam(sqlquery, kw)
|
||||||
sqlquery = strings.ReplaceAll(sqlquery, kw, replacement)
|
sqlquery = strings.ReplaceAll(sqlquery, kw, replacement)
|
||||||
logger.Debug("Replaced unused variable %s with: %s", kw, replacement)
|
logger.Debug("Replaced unused variable %s with: %s", kw, replacement)
|
||||||
@@ -540,7 +563,7 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
|||||||
hookCtx.Tx = tx
|
hookCtx.Tx = tx
|
||||||
|
|
||||||
// Execute BeforeQuery hook (inside transaction)
|
// Execute BeforeQuery hook (inside transaction)
|
||||||
if err := h.hooks.Execute(BeforeQuery, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeQuery, hookCtx); err != nil {
|
||||||
logger.Error("BeforeQuery hook failed: %v", err)
|
logger.Error("BeforeQuery hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
return err
|
||||||
@@ -559,7 +582,7 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
|||||||
sqlquery = hookCtx.SQLQuery
|
sqlquery = hookCtx.SQLQuery
|
||||||
|
|
||||||
// Execute BeforeSQLExec hook
|
// Execute BeforeSQLExec hook
|
||||||
if err := h.hooks.Execute(BeforeSQLExec, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
|
||||||
logger.Error("BeforeSQLExec hook failed: %v", err)
|
logger.Error("BeforeSQLExec hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
return err
|
||||||
@@ -611,7 +634,10 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute BeforeResponse hook
|
// Execute BeforeResponse hook. The transaction has already committed by
|
||||||
|
// this point, so hooks must use the pooled connection rather than the
|
||||||
|
// now-dead tx.
|
||||||
|
hookCtx.Tx = h.db
|
||||||
hookCtx.Result = dbobj
|
hookCtx.Result = dbobj
|
||||||
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
||||||
logger.Error("BeforeResponse hook failed: %v", err)
|
logger.Error("BeforeResponse hook failed: %v", err)
|
||||||
@@ -715,8 +741,10 @@ func (h *Handler) mergeQueryParams(r *http.Request, sqlquery string, variables m
|
|||||||
propQry[parmk] = val
|
propQry[parmk] = val
|
||||||
}
|
}
|
||||||
|
|
||||||
// Apply filters if allowed
|
// Apply filters if allowed — check only the SELECT list to avoid matching function
|
||||||
if allowFilter && len(parmk) > 1 && strings.Contains(strings.ToLower(sqlquery), strings.ToLower(parmk)) {
|
// parameters in the FROM clause (e.g. [p_rid_doctype] in a set-returning function call)
|
||||||
|
// or names inside quoted string arguments.
|
||||||
|
if allowFilter && len(parmk) > 1 && strings.Contains(sqlSelectList(sqlStripStringLiterals(sqlquery)), strings.ToLower(parmk)) {
|
||||||
if len(parmv) > 1 {
|
if len(parmv) > 1 {
|
||||||
// Sanitize each value in the IN clause with appropriate quoting
|
// Sanitize each value in the IN clause with appropriate quoting
|
||||||
sanitizedValues := make([]string, len(parmv))
|
sanitizedValues := make([]string, len(parmv))
|
||||||
@@ -739,7 +767,7 @@ func (h *Handler) mergeQueryParams(r *http.Request, sqlquery string, variables m
|
|||||||
colval = strings.ReplaceAll(colval, "\\", "\\\\")
|
colval = strings.ReplaceAll(colval, "\\", "\\\\")
|
||||||
colval = strings.ReplaceAll(colval, "'", "''")
|
colval = strings.ReplaceAll(colval, "'", "''")
|
||||||
if colval != "*" {
|
if colval != "*" {
|
||||||
sqlquery = sqlQryWhere(sqlquery, fmt.Sprintf("%s ILIKE '%%%s%%'", ValidSQL(parmk, "colname"), colval))
|
sqlquery = sqlQryWhere(sqlquery, fmt.Sprintf("CAST(%s AS TEXT) ILIKE '%%%s%%'", ValidSQL(parmk, "colname"), colval))
|
||||||
}
|
}
|
||||||
} else if val == "" || val == "0" {
|
} else if val == "" || val == "0" {
|
||||||
// For empty/zero values, treat as literal 0 or empty string with quotes
|
// For empty/zero values, treat as literal 0 or empty string with quotes
|
||||||
@@ -806,7 +834,7 @@ func (h *Handler) mergeHeaderParams(r *http.Request, sqlquery string, variables
|
|||||||
colname := strings.ReplaceAll(k, "x-searchfilter-", "")
|
colname := strings.ReplaceAll(k, "x-searchfilter-", "")
|
||||||
sval := strings.ReplaceAll(val, "'", "")
|
sval := strings.ReplaceAll(val, "'", "")
|
||||||
if sval != "" {
|
if sval != "" {
|
||||||
sqlquery = sqlQryWhere(sqlquery, fmt.Sprintf("%s ILIKE '%%%s%%'", ValidSQL(colname, "colname"), ValidSQL(sval, "colvalue")))
|
sqlquery = sqlQryWhere(sqlquery, fmt.Sprintf("CAST(%s AS TEXT) ILIKE '%%%s%%'", ValidSQL(colname, "colname"), ValidSQL(sval, "colvalue")))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -824,6 +852,26 @@ func (h *Handler) mergeHeaderParams(r *http.Request, sqlquery string, variables
|
|||||||
return sqlquery
|
return sqlquery
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// sqlStripStringLiterals removes the contents of single-quoted string literals from SQL,
|
||||||
|
// leaving the structural identifiers (column names, table names) intact.
|
||||||
|
// Used to check column presence without matching inside string arguments.
|
||||||
|
func sqlStripStringLiterals(sql string) string {
|
||||||
|
re := regexp.MustCompile(`'(?:[^']|'')*'`)
|
||||||
|
return re.ReplaceAllString(sql, "''")
|
||||||
|
}
|
||||||
|
|
||||||
|
// sqlSelectList returns the column list portion of a SELECT query (between SELECT and FROM).
|
||||||
|
// Returns the full query lowercased if no clear SELECT…FROM boundary is found.
|
||||||
|
func sqlSelectList(sql string) string {
|
||||||
|
lower := strings.ToLower(sql)
|
||||||
|
selectPos := strings.Index(lower, "select ")
|
||||||
|
fromPos := strings.Index(lower, " from ")
|
||||||
|
if selectPos < 0 || fromPos <= selectPos {
|
||||||
|
return lower
|
||||||
|
}
|
||||||
|
return lower[selectPos+7 : fromPos]
|
||||||
|
}
|
||||||
|
|
||||||
// replaceMetaVariables replaces meta variables like [rid_user], [user], etc. in the SQL query
|
// replaceMetaVariables replaces meta variables like [rid_user], [user], etc. in the SQL query
|
||||||
func (h *Handler) replaceMetaVariables(sqlquery string, r *http.Request, userCtx *security.UserContext, metainfo map[string]interface{}, variables map[string]interface{}) string {
|
func (h *Handler) replaceMetaVariables(sqlquery string, r *http.Request, userCtx *security.UserContext, metainfo map[string]interface{}, variables map[string]interface{}) string {
|
||||||
if strings.Contains(sqlquery, "[p_meta_default]") {
|
if strings.Contains(sqlquery, "[p_meta_default]") {
|
||||||
@@ -969,6 +1017,37 @@ func IsNumeric(s string) bool {
|
|||||||
return err == nil
|
return err == nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isInsideDollarQuote reports whether the first occurrence of placeholder in sqlquery
|
||||||
|
// is immediately surrounded by dollar-sign characters (i.e. inside a $...$-quoted string).
|
||||||
|
// Dollar-quoted strings pass content through literally — no backslash processing — so
|
||||||
|
// values placed there must NOT have their backslashes escaped.
|
||||||
|
func isInsideDollarQuote(sqlquery, placeholder string) bool {
|
||||||
|
idx := strings.Index(sqlquery, placeholder)
|
||||||
|
if idx < 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
endIdx := idx + len(placeholder)
|
||||||
|
charBefore := byte(0)
|
||||||
|
charAfter := byte(0)
|
||||||
|
if idx > 0 {
|
||||||
|
charBefore = sqlquery[idx-1]
|
||||||
|
}
|
||||||
|
if endIdx < len(sqlquery) {
|
||||||
|
charAfter = sqlquery[endIdx]
|
||||||
|
}
|
||||||
|
return charBefore == '$' || charAfter == '$'
|
||||||
|
}
|
||||||
|
|
||||||
|
// safeSubstituteVar returns value sanitised for the quoting context that surrounds
|
||||||
|
// placeholder in sqlquery: raw (no backslash escaping) for dollar-quoted contexts,
|
||||||
|
// ValidSQL("colvalue") escaping for everything else.
|
||||||
|
func safeSubstituteVar(sqlquery, placeholder, value string) string {
|
||||||
|
if isInsideDollarQuote(sqlquery, placeholder) {
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
return ValidSQL(value, "colvalue")
|
||||||
|
}
|
||||||
|
|
||||||
// getReplacementForBlankParam determines the replacement value for an unused parameter
|
// getReplacementForBlankParam determines the replacement value for an unused parameter
|
||||||
// based on whether it appears within quotes in the SQL query.
|
// based on whether it appears within quotes in the SQL query.
|
||||||
// It checks for PostgreSQL quotes: single quotes (”) and dollar quotes ($...$)
|
// It checks for PostgreSQL quotes: single quotes (”) and dollar quotes ($...$)
|
||||||
@@ -991,8 +1070,8 @@ func getReplacementForBlankParam(sqlquery, param string) string {
|
|||||||
charAfter = sqlquery[endIdx]
|
charAfter = sqlquery[endIdx]
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if parameter is surrounded by quotes (single quote or dollar sign for PostgreSQL dollar-quoted strings)
|
// Check if parameter is surrounded by quotes (single quote, dollar sign for PostgreSQL dollar-quoted strings, or double quote for JSON string values)
|
||||||
if (charBefore == '\'' || charBefore == '$') && (charAfter == '\'' || charAfter == '$') {
|
if (charBefore == '\'' || charBefore == '$' || charBefore == '"') && (charAfter == '\'' || charAfter == '$' || charAfter == '"') {
|
||||||
// Parameter is in quotes, return empty string
|
// Parameter is in quotes, return empty string
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
@@ -1011,6 +1090,49 @@ func getReplacementForBlankParam(sqlquery, param string) string {
|
|||||||
// return result
|
// return result
|
||||||
// }
|
// }
|
||||||
|
|
||||||
|
// buildDetailFieldsFromRows builds a field metadata list from the column names and value types
|
||||||
|
// of a raw SQL result set. Used when no model struct is available (funcspec raw queries).
|
||||||
|
func buildDetailFieldsFromRows(rows []map[string]interface{}) []reflection.ModelFieldDetail {
|
||||||
|
if len(rows) == 0 {
|
||||||
|
return []reflection.ModelFieldDetail{}
|
||||||
|
}
|
||||||
|
first := rows[0]
|
||||||
|
fields := make([]reflection.ModelFieldDetail, 0, len(first))
|
||||||
|
for colName, val := range first {
|
||||||
|
dataType := inferGoType(val)
|
||||||
|
fields = append(fields, reflection.ModelFieldDetail{
|
||||||
|
Name: colName,
|
||||||
|
DataType: dataType,
|
||||||
|
SQLName: colName,
|
||||||
|
SQLDataType: "",
|
||||||
|
SQLKey: "",
|
||||||
|
Nullable: val == nil,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return fields
|
||||||
|
}
|
||||||
|
|
||||||
|
// inferGoType returns a simple type name for a value, used for detail field metadata.
|
||||||
|
func inferGoType(val interface{}) string {
|
||||||
|
if val == nil {
|
||||||
|
return "interface{}"
|
||||||
|
}
|
||||||
|
switch val.(type) {
|
||||||
|
case bool:
|
||||||
|
return "bool"
|
||||||
|
case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64:
|
||||||
|
return "int64"
|
||||||
|
case float32, float64:
|
||||||
|
return "float64"
|
||||||
|
case string:
|
||||||
|
return "string"
|
||||||
|
case []byte:
|
||||||
|
return "[]byte"
|
||||||
|
default:
|
||||||
|
return "interface{}"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// getIPAddress extracts the real IP address from the request
|
// getIPAddress extracts the real IP address from the request
|
||||||
func getIPAddress(r *http.Request) string {
|
func getIPAddress(r *http.Request) string {
|
||||||
if forwarded := r.Header.Get("X-Forwarded-For"); forwarded != "" {
|
if forwarded := r.Header.Get("X-Forwarded-For"); forwarded != "" {
|
||||||
@@ -1035,6 +1157,10 @@ func sendError(w http.ResponseWriter, status int, code, message string, err erro
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
errObj.Detail = err.Error()
|
errObj.Detail = err.Error()
|
||||||
|
var sqlErr *common.SQLError
|
||||||
|
if errors.As(err, &sqlErr) {
|
||||||
|
errObj.SQL = sqlErr.SQL
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
data, _ := json.Marshal(map[string]interface{}{
|
data, _ := json.Marshal(map[string]interface{}{
|
||||||
|
|||||||
@@ -617,6 +617,91 @@ func TestSqlQueryList(t *testing.T) {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
name: "x-detailapi header returns detail format",
|
||||||
|
sqlQuery: "SELECT * FROM myschema.myentity",
|
||||||
|
noCount: false,
|
||||||
|
blankParams: false,
|
||||||
|
allowFilter: false,
|
||||||
|
headers: map[string]string{"x-detailapi": "true"},
|
||||||
|
setupDB: func() *MockDatabase {
|
||||||
|
return &MockDatabase{
|
||||||
|
RunInTransactionFunc: func(ctx context.Context, fn func(common.Database) error) error {
|
||||||
|
db := &MockDatabase{
|
||||||
|
QueryFunc: func(ctx context.Context, dest interface{}, query string, args ...interface{}) error {
|
||||||
|
if strings.Contains(query, "COUNT") {
|
||||||
|
dest.(*struct{ Count int64 }).Count = 3
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
*dest.(*[]map[string]interface{}) = []map[string]interface{}{
|
||||||
|
{"id": float64(1), "name": "Alice"},
|
||||||
|
{"id": float64(2), "name": "Bob"},
|
||||||
|
{"id": float64(3), "name": "Carol"},
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
return fn(db)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
},
|
||||||
|
expectedStatus: 200,
|
||||||
|
validateResp: func(t *testing.T, w *httptest.ResponseRecorder) {
|
||||||
|
var resp map[string]json.RawMessage
|
||||||
|
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("expected JSON object, got: %s", w.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, key := range []string{"count", "fields", "items", "tablename", "tableprefix", "total"} {
|
||||||
|
if _, ok := resp[key]; !ok {
|
||||||
|
t.Errorf("missing key %q in detail response", key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var count, total string
|
||||||
|
json.Unmarshal(resp["count"], &count)
|
||||||
|
json.Unmarshal(resp["total"], &total)
|
||||||
|
if count != "3" {
|
||||||
|
t.Errorf("expected count %q, got %q", "3", count)
|
||||||
|
}
|
||||||
|
if total != "3" {
|
||||||
|
t.Errorf("expected total %q, got %q", "3", total)
|
||||||
|
}
|
||||||
|
|
||||||
|
var items []map[string]interface{}
|
||||||
|
if err := json.Unmarshal(resp["items"], &items); err != nil {
|
||||||
|
t.Fatalf("items is not an array: %v", err)
|
||||||
|
}
|
||||||
|
if len(items) != 3 {
|
||||||
|
t.Errorf("expected 3 items, got %d", len(items))
|
||||||
|
}
|
||||||
|
|
||||||
|
var fields []map[string]interface{}
|
||||||
|
if err := json.Unmarshal(resp["fields"], &fields); err != nil {
|
||||||
|
t.Fatalf("fields is not an array: %v", err)
|
||||||
|
}
|
||||||
|
if len(fields) == 0 {
|
||||||
|
t.Error("expected non-empty fields list")
|
||||||
|
}
|
||||||
|
for _, f := range fields {
|
||||||
|
for _, key := range []string{"name", "datatype", "sqlname", "sqldatatype", "sqlkey", "nullable"} {
|
||||||
|
if _, ok := f[key]; !ok {
|
||||||
|
t.Errorf("field %v missing key %q", f, key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var tablename, tableprefix string
|
||||||
|
json.Unmarshal(resp["tablename"], &tablename)
|
||||||
|
json.Unmarshal(resp["tableprefix"], &tableprefix)
|
||||||
|
if tablename == "" {
|
||||||
|
t.Error("expected non-empty tablename")
|
||||||
|
}
|
||||||
|
if tableprefix == "" {
|
||||||
|
t.Error("expected non-empty tableprefix")
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
{
|
{
|
||||||
name: "List query with noCount",
|
name: "List query with noCount",
|
||||||
sqlQuery: "SELECT * FROM users",
|
sqlQuery: "SELECT * FROM users",
|
||||||
@@ -821,7 +906,7 @@ func TestReplaceMetaVariables(t *testing.T) {
|
|||||||
name: "Replace [user]",
|
name: "Replace [user]",
|
||||||
sqlQuery: "SELECT * FROM audit WHERE username = [user]",
|
sqlQuery: "SELECT * FROM audit WHERE username = [user]",
|
||||||
expectedCheck: func(result string) bool {
|
expectedCheck: func(result string) bool {
|
||||||
return strings.Contains(result, "'testuser'")
|
return strings.Contains(result, "$USR$testuser$USR$")
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -851,6 +936,285 @@ func TestReplaceMetaVariables(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestSqlStripStringLiterals tests that single-quoted string literals are removed
|
||||||
|
func TestSqlStripStringLiterals(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
expected string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "No string literals",
|
||||||
|
input: "SELECT rid, rid_parent FROM users",
|
||||||
|
expected: "SELECT rid, rid_parent FROM users",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Simple string literal",
|
||||||
|
input: "SELECT * FROM users WHERE mode = 'admin'",
|
||||||
|
expected: "SELECT * FROM users WHERE mode = ''",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "JSON argument containing column names",
|
||||||
|
input: `SELECT rid, rid_parent FROM crm_get_menu(1,'mode', '{"rid_parent":"[rid_parent]","CF:STARTDATE":"[cf_startdate]"}')`,
|
||||||
|
expected: `SELECT rid, rid_parent FROM crm_get_menu(1,'', '')`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Escaped single quotes inside literal",
|
||||||
|
input: "SELECT * FROM t WHERE name = 'O''Brien'",
|
||||||
|
expected: "SELECT * FROM t WHERE name = ''",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
result := sqlStripStringLiterals(tt.input)
|
||||||
|
if result != tt.expected {
|
||||||
|
t.Errorf("sqlStripStringLiterals() =\n %q\nwant\n %q", result, tt.expected)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestAllowFilterDoesNotMatchInsideJsonArgument verifies that AllowFilter will add WHERE
|
||||||
|
// clauses for real output columns (rid, rid_parent) but not for names that only appear
|
||||||
|
// inside a JSON string argument (cf_startdate, cf_rid_branch).
|
||||||
|
func TestAllowFilterDoesNotMatchInsideJsonArgument(t *testing.T) {
|
||||||
|
handler := NewHandler(&MockDatabase{})
|
||||||
|
|
||||||
|
sqlQuery := `select rid, rid_parent, description
|
||||||
|
from crm_get_menu([rid_user],'[p_mode]', 0, '', '{"rid_parent":"[rid_parent]", "CF:STARTDATE": "[cf_startdate]", "CF:RID_BRANCH": "[cf_rid_branch]"}')`
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
queryParams map[string]string
|
||||||
|
checkResult func(t *testing.T, result string)
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "rid_parent=0 is a real column — filter applied",
|
||||||
|
queryParams: map[string]string{"rid_parent": "0"},
|
||||||
|
checkResult: func(t *testing.T, result string) {
|
||||||
|
if !strings.Contains(strings.ToLower(result), "where") {
|
||||||
|
t.Error("Expected WHERE clause to be added for rid_parent")
|
||||||
|
}
|
||||||
|
if !strings.Contains(result, "rid_parent = 0 OR") && !strings.Contains(result, "rid_parent IS NULL") {
|
||||||
|
t.Errorf("Expected null-safe filter for rid_parent=0, got:\n%s", result)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "cf_startdate only appears in JSON string — no filter applied",
|
||||||
|
queryParams: map[string]string{"cf_startdate": "2024-01-01"},
|
||||||
|
checkResult: func(t *testing.T, result string) {
|
||||||
|
if strings.Contains(strings.ToLower(result), "where") {
|
||||||
|
t.Errorf("Expected no WHERE clause for cf_startdate (only in JSON arg), got:\n%s", result)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "cf_rid_branch only appears in JSON string — no filter applied",
|
||||||
|
queryParams: map[string]string{"cf_rid_branch": "5"},
|
||||||
|
checkResult: func(t *testing.T, result string) {
|
||||||
|
if strings.Contains(strings.ToLower(result), "where") {
|
||||||
|
t.Errorf("Expected no WHERE clause for cf_rid_branch (only in JSON arg), got:\n%s", result)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "description is a real column — filter applied",
|
||||||
|
queryParams: map[string]string{"description": "test"},
|
||||||
|
checkResult: func(t *testing.T, result string) {
|
||||||
|
if !strings.Contains(strings.ToLower(result), "where") {
|
||||||
|
t.Error("Expected WHERE clause for description")
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
req := createTestRequest("GET", "/test", tt.queryParams, nil, nil)
|
||||||
|
variables := make(map[string]interface{})
|
||||||
|
propQry := make(map[string]string)
|
||||||
|
|
||||||
|
result := handler.mergeQueryParams(req, sqlQuery, variables, true, propQry)
|
||||||
|
tt.checkResult(t, result)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestAllowFilterDoesNotMatchFunctionParams verifies that query params that appear only
|
||||||
|
// as function call arguments in the FROM clause (e.g. [p_rid_doctype]) are not treated
|
||||||
|
// as column filters, since they are not in the SELECT list.
|
||||||
|
func TestAllowFilterDoesNotMatchFunctionParams(t *testing.T) {
|
||||||
|
handler := NewHandler(&MockDatabase{})
|
||||||
|
|
||||||
|
sqlQuery := `select rid, rid_parent, description, row_cnt, filterstring, tableprefix, rid_table, tooltip, additionalfilter, haschildren
|
||||||
|
from crm_get_doc_menu($JQ$[p_tableprefix]$JQ$,[p_rid_parent],[p_rid_doctype],[p_removedup],[p_showall]) r`
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
queryParams map[string]string
|
||||||
|
checkResult func(t *testing.T, result string)
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "p_rid_doctype is a function param, not a column — no filter applied",
|
||||||
|
queryParams: map[string]string{"p_rid_doctype": "0"},
|
||||||
|
checkResult: func(t *testing.T, result string) {
|
||||||
|
if strings.Contains(strings.ToLower(result), "where") {
|
||||||
|
t.Errorf("Expected no WHERE clause for p_rid_doctype (function arg, not SELECT column), got:\n%s", result)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "p_showall is a function param, not a column — no filter applied",
|
||||||
|
queryParams: map[string]string{"p_showall": "1"},
|
||||||
|
checkResult: func(t *testing.T, result string) {
|
||||||
|
if strings.Contains(strings.ToLower(result), "where") {
|
||||||
|
t.Errorf("Expected no WHERE clause for p_showall (function arg, not SELECT column), got:\n%s", result)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "rid is a SELECT column — filter applied",
|
||||||
|
queryParams: map[string]string{"rid": "42"},
|
||||||
|
checkResult: func(t *testing.T, result string) {
|
||||||
|
if !strings.Contains(strings.ToLower(result), "where") {
|
||||||
|
t.Error("Expected WHERE clause for rid (real SELECT column)")
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
req := createTestRequest("GET", "/test", tt.queryParams, nil, nil)
|
||||||
|
variables := make(map[string]interface{})
|
||||||
|
propQry := make(map[string]string)
|
||||||
|
result := handler.mergeQueryParams(req, sqlQuery, variables, true, propQry)
|
||||||
|
tt.checkResult(t, result)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestGetReplacementForBlankParamDoubleQuote verifies that placeholders surrounded by
|
||||||
|
// double quotes (as in JSON string values) are blanked to "" not NULL.
|
||||||
|
func TestGetReplacementForBlankParamDoubleQuote(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
sqlQuery string
|
||||||
|
param string
|
||||||
|
expected string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "Parameter in double quotes (JSON value)",
|
||||||
|
sqlQuery: `SELECT * FROM f(1, '{"key":"[myparam]"}')`,
|
||||||
|
param: "[myparam]",
|
||||||
|
expected: "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Parameter not in any quotes",
|
||||||
|
sqlQuery: `SELECT * FROM f([myparam])`,
|
||||||
|
param: "[myparam]",
|
||||||
|
expected: "NULL",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Parameter in single quotes",
|
||||||
|
sqlQuery: `SELECT * FROM f('[myparam]')`,
|
||||||
|
param: "[myparam]",
|
||||||
|
expected: "",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
result := getReplacementForBlankParam(tt.sqlQuery, tt.param)
|
||||||
|
if result != tt.expected {
|
||||||
|
t.Errorf("getReplacementForBlankParam() = %q, want %q\nquery: %s", result, tt.expected, tt.sqlQuery)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestVariableReplacementFromQueryParams verifies that query params matching [placeholder]
|
||||||
|
// tokens are substituted even when they don't have the p- prefix.
|
||||||
|
func TestVariableReplacementFromQueryParams(t *testing.T) {
|
||||||
|
handler := NewHandler(&MockDatabase{})
|
||||||
|
|
||||||
|
sqlQuery := `select rid, rid_parent from crm_get_menu([rid_user],'[p_mode]', 0, '', '{"rid_parent":"[rid_parent]","CF:STARTDATE":"[cf_startdate]"}')`
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
queryParams map[string]string
|
||||||
|
checkResult func(t *testing.T, result string)
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "rid_parent replaced from query param",
|
||||||
|
queryParams: map[string]string{"rid_parent": "42"},
|
||||||
|
checkResult: func(t *testing.T, result string) {
|
||||||
|
if strings.Contains(result, "[rid_parent]") {
|
||||||
|
t.Errorf("Expected [rid_parent] to be replaced, still present in:\n%s", result)
|
||||||
|
}
|
||||||
|
if !strings.Contains(result, "42") {
|
||||||
|
t.Errorf("Expected value 42 in query, got:\n%s", result)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "cf_startdate replaced from query param",
|
||||||
|
queryParams: map[string]string{"cf_startdate": "2024-01-01"},
|
||||||
|
checkResult: func(t *testing.T, result string) {
|
||||||
|
if strings.Contains(result, "[cf_startdate]") {
|
||||||
|
t.Errorf("Expected [cf_startdate] to be replaced, still present in:\n%s", result)
|
||||||
|
}
|
||||||
|
if !strings.Contains(result, "2024-01-01") {
|
||||||
|
t.Errorf("Expected date value in query, got:\n%s", result)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "missing param blanked to empty string inside JSON (double-quoted)",
|
||||||
|
queryParams: map[string]string{},
|
||||||
|
checkResult: func(t *testing.T, result string) {
|
||||||
|
// [cf_startdate] is surrounded by " in the JSON — should blank to ""
|
||||||
|
if strings.Contains(result, "[cf_startdate]") {
|
||||||
|
t.Errorf("Expected [cf_startdate] to be blanked, still present in:\n%s", result)
|
||||||
|
}
|
||||||
|
if strings.Contains(result, "NULL") && strings.Contains(result, "cf_startdate") {
|
||||||
|
t.Errorf("Expected empty string (not NULL) for double-quoted placeholder, got:\n%s", result)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
inputvars := make([]string, 0)
|
||||||
|
q := handler.extractInputVariables(sqlQuery, &inputvars)
|
||||||
|
|
||||||
|
req := createTestRequest("GET", "/test", tt.queryParams, nil, nil)
|
||||||
|
variables := make(map[string]interface{})
|
||||||
|
propQry := make(map[string]string)
|
||||||
|
|
||||||
|
q = handler.mergeQueryParams(req, q, variables, false, propQry)
|
||||||
|
|
||||||
|
// Simulate the variable replacement + blank-param loop (mirrors function_api.go)
|
||||||
|
for _, kw := range inputvars {
|
||||||
|
varName := kw[1 : len(kw)-1]
|
||||||
|
if val, ok := variables[varName]; ok {
|
||||||
|
if strVal := strings.TrimSpace(val.(string)); strVal != "" {
|
||||||
|
q = strings.ReplaceAll(q, kw, ValidSQL(strVal, "colvalue"))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
replacement := getReplacementForBlankParam(q, kw)
|
||||||
|
q = strings.ReplaceAll(q, kw, replacement)
|
||||||
|
}
|
||||||
|
|
||||||
|
tt.checkResult(t, q)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestGetReplacementForBlankParam tests the blank parameter replacement logic
|
// TestGetReplacementForBlankParam tests the blank parameter replacement logic
|
||||||
func TestGetReplacementForBlankParam(t *testing.T) {
|
func TestGetReplacementForBlankParam(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
|
|||||||
@@ -28,6 +28,10 @@ const (
|
|||||||
|
|
||||||
// Response hooks (before response is sent)
|
// Response hooks (before response is sent)
|
||||||
BeforeResponse HookType = "before_response"
|
BeforeResponse HookType = "before_response"
|
||||||
|
|
||||||
|
// BeforeOp fires immediately before every SQL operation (query, query list, SQL exec).
|
||||||
|
// It fires at each individual SQL-operation hook point, so it runs once per statement executed.
|
||||||
|
BeforeOp HookType = "before_op"
|
||||||
)
|
)
|
||||||
|
|
||||||
// HookContext contains all the data available to a hook
|
// HookContext contains all the data available to a hook
|
||||||
@@ -130,6 +134,16 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
|
||||||
|
// (BeforeQuery, BeforeQueryList, or BeforeSQLExec). BeforeOp always runs first so it can
|
||||||
|
// observe/veto every SQL operation regardless of type.
|
||||||
|
func (r *HookRegistry) ExecuteBeforeOp(hookType HookType, ctx *HookContext) error {
|
||||||
|
if err := r.Execute(BeforeOp, ctx); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return r.Execute(hookType, ctx)
|
||||||
|
}
|
||||||
|
|
||||||
// Clear removes all hooks for the specified type
|
// Clear removes all hooks for the specified type
|
||||||
func (r *HookRegistry) Clear(hookType HookType) {
|
func (r *HookRegistry) Clear(hookType HookType) {
|
||||||
delete(r.hooks, hookType)
|
delete(r.hooks, hookType)
|
||||||
|
|||||||
@@ -51,7 +51,7 @@ func (h *Handler) ParseParameters(r *http.Request) *RequestParameters {
|
|||||||
FieldFilters: make(map[string]string),
|
FieldFilters: make(map[string]string),
|
||||||
SearchFilters: make(map[string]string),
|
SearchFilters: make(map[string]string),
|
||||||
SearchOps: make(map[string]FilterOperator),
|
SearchOps: make(map[string]FilterOperator),
|
||||||
Limit: 20, // Default limit
|
Limit: 100000, // Default limit
|
||||||
Offset: 0, // Default offset
|
Offset: 0, // Default offset
|
||||||
ResponseFormat: "simple", // Default format
|
ResponseFormat: "simple", // Default format
|
||||||
ComplexAPI: false, // Default to simple API
|
ComplexAPI: false, // Default to simple API
|
||||||
@@ -259,7 +259,7 @@ func (h *Handler) ApplyFilters(sqlQuery string, params *RequestParameters) strin
|
|||||||
for colName, value := range params.SearchFilters {
|
for colName, value := range params.SearchFilters {
|
||||||
sval := strings.ReplaceAll(value, "'", "")
|
sval := strings.ReplaceAll(value, "'", "")
|
||||||
if sval != "" {
|
if sval != "" {
|
||||||
condition := fmt.Sprintf("%s ILIKE '%%%s%%'", ValidSQL(colName, "colname"), ValidSQL(sval, "colvalue"))
|
condition := fmt.Sprintf("CAST(%s AS TEXT) ILIKE '%%%s%%'", ValidSQL(colName, "colname"), ValidSQL(sval, "colvalue"))
|
||||||
sqlQuery = sqlQryWhere(sqlQuery, condition)
|
sqlQuery = sqlQryWhere(sqlQuery, condition)
|
||||||
logger.Debug("Applied search filter: %s", condition)
|
logger.Debug("Applied search filter: %s", condition)
|
||||||
}
|
}
|
||||||
@@ -307,11 +307,11 @@ func (h *Handler) buildFilterCondition(colName string, op FilterOperator) string
|
|||||||
|
|
||||||
switch operator {
|
switch operator {
|
||||||
case "contains", "contain", "like":
|
case "contains", "contain", "like":
|
||||||
return fmt.Sprintf("%s ILIKE '%%%s%%'", safCol, ValidSQL(value, "colvalue"))
|
return fmt.Sprintf("CAST(%s AS TEXT) ILIKE '%%%s%%'", safCol, ValidSQL(value, "colvalue"))
|
||||||
case "beginswith", "startswith":
|
case "beginswith", "startswith":
|
||||||
return fmt.Sprintf("%s ILIKE '%s%%'", safCol, ValidSQL(value, "colvalue"))
|
return fmt.Sprintf("CAST(%s AS TEXT) ILIKE '%s%%'", safCol, ValidSQL(value, "colvalue"))
|
||||||
case "endswith":
|
case "endswith":
|
||||||
return fmt.Sprintf("%s ILIKE '%%%s'", safCol, ValidSQL(value, "colvalue"))
|
return fmt.Sprintf("CAST(%s AS TEXT) ILIKE '%%%s'", safCol, ValidSQL(value, "colvalue"))
|
||||||
case "equals", "eq", "=":
|
case "equals", "eq", "=":
|
||||||
if IsNumeric(value) {
|
if IsNumeric(value) {
|
||||||
return fmt.Sprintf("%s = %s", safCol, ValidSQL(value, "colvalue"))
|
return fmt.Sprintf("%s = %s", safCol, ValidSQL(value, "colvalue"))
|
||||||
|
|||||||
@@ -274,7 +274,7 @@ func TestBuildFilterCondition(t *testing.T) {
|
|||||||
Value: "test",
|
Value: "test",
|
||||||
Logic: "AND",
|
Logic: "AND",
|
||||||
},
|
},
|
||||||
expected: "description ILIKE '%test%'",
|
expected: "CAST(description AS TEXT) ILIKE '%test%'",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Starts with operator",
|
name: "Starts with operator",
|
||||||
@@ -284,7 +284,7 @@ func TestBuildFilterCondition(t *testing.T) {
|
|||||||
Value: "john",
|
Value: "john",
|
||||||
Logic: "AND",
|
Logic: "AND",
|
||||||
},
|
},
|
||||||
expected: "name ILIKE 'john%'",
|
expected: "CAST(name AS TEXT) ILIKE 'john%'",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Ends with operator",
|
name: "Ends with operator",
|
||||||
@@ -294,7 +294,7 @@ func TestBuildFilterCondition(t *testing.T) {
|
|||||||
Value: "@example.com",
|
Value: "@example.com",
|
||||||
Logic: "AND",
|
Logic: "AND",
|
||||||
},
|
},
|
||||||
expected: "email ILIKE '%@example.com'",
|
expected: "CAST(email AS TEXT) ILIKE '%@example.com'",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "Between operator",
|
name: "Between operator",
|
||||||
|
|||||||
@@ -71,6 +71,16 @@ func (f *funcSpecSecurityContext) GetUserID() (int, bool) {
|
|||||||
return int(f.ctx.UserContext.UserID), true
|
return int(f.ctx.UserContext.UserID), true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetUserRef returns an opaque user identifier for row security lookups.
|
||||||
|
// It returns the full *security.UserContext so providers can read JWT claims
|
||||||
|
// (e.g. a UUID subject) instead of relying on the int user ID.
|
||||||
|
func (f *funcSpecSecurityContext) GetUserRef() (any, bool) {
|
||||||
|
if f.ctx.UserContext == nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
return f.ctx.UserContext, true
|
||||||
|
}
|
||||||
|
|
||||||
func (f *funcSpecSecurityContext) GetSchema() string {
|
func (f *funcSpecSecurityContext) GetSchema() string {
|
||||||
// funcspec doesn't have a schema concept, extract from SQL query or use default
|
// funcspec doesn't have a schema concept, extract from SQL query or use default
|
||||||
return "public"
|
return "public"
|
||||||
|
|||||||
+221
-57
@@ -5,16 +5,144 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
"os"
|
"os"
|
||||||
|
"regexp"
|
||||||
|
"runtime"
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
|
|
||||||
errortracking "github.com/bitechdev/ResolveSpec/pkg/errortracking"
|
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 Logger *zap.SugaredLogger
|
||||||
var errorTracker errortracking.Provider
|
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) {
|
func Init(dev bool) {
|
||||||
|
|
||||||
if dev {
|
if dev {
|
||||||
@@ -36,7 +164,16 @@ func UpdateLoggerPath(path string, dev bool) {
|
|||||||
UpdateLogger(&defaultConfig)
|
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) {
|
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 := zap.NewProductionConfig()
|
||||||
defaultConfig.OutputPaths = []string{"resolvespec.log"}
|
defaultConfig.OutputPaths = []string{"resolvespec.log"}
|
||||||
if config == nil {
|
if config == nil {
|
||||||
@@ -45,32 +182,45 @@ func UpdateLogger(config *zap.Config) {
|
|||||||
|
|
||||||
logger, err := config.Build()
|
logger, err := config.Build()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Print(err)
|
return err
|
||||||
return
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Logger = logger.Sugar()
|
old := swapLogger(logger.Sugar())
|
||||||
|
if old != nil {
|
||||||
|
_ = old.Sync()
|
||||||
|
}
|
||||||
Info("ResolveSpec Logger initialized")
|
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
|
// InitErrorTracking initializes the error tracking provider
|
||||||
func InitErrorTracking(provider errortracking.Provider) {
|
func InitErrorTracking(provider errortracking.Provider) {
|
||||||
|
stateMu.Lock()
|
||||||
errorTracker = provider
|
errorTracker = provider
|
||||||
if errorTracker != nil {
|
stateMu.Unlock()
|
||||||
|
if provider != nil {
|
||||||
Info("Error tracking initialized")
|
Info("Error tracking initialized")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetErrorTracker returns the current error tracking provider
|
// GetErrorTracker returns the current error tracking provider
|
||||||
func GetErrorTracker() errortracking.Provider {
|
func GetErrorTracker() errortracking.Provider {
|
||||||
return errorTracker
|
return getErrorTracker()
|
||||||
}
|
}
|
||||||
|
|
||||||
// CloseErrorTracking flushes and closes the error tracking provider
|
// CloseErrorTracking flushes and closes the error tracking provider
|
||||||
func CloseErrorTracking() error {
|
func CloseErrorTracking() error {
|
||||||
if errorTracker != nil {
|
if tracker := getErrorTracker(); tracker != nil {
|
||||||
errorTracker.Flush(5)
|
tracker.Flush(5)
|
||||||
return errorTracker.Close()
|
return tracker.Close()
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -98,53 +248,51 @@ func extractContext(args ...interface{}) (ctx context.Context, filteredArgs []in
|
|||||||
}
|
}
|
||||||
|
|
||||||
func Info(template string, args ...interface{}) {
|
func Info(template string, args ...interface{}) {
|
||||||
if Logger == nil {
|
_, args = extractContext(args...)
|
||||||
log.Printf(template, args...)
|
message := fmt.Sprintf(template, args...)
|
||||||
|
if lg := getLogger(); lg != nil {
|
||||||
|
lg.Infow(message, "process_id", pid)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
Logger.Infow(fmt.Sprintf(template, args...), "process_id", os.Getpid())
|
log.Printf("%s", sanitizeForStdlog(message))
|
||||||
}
|
|
||||||
|
|
||||||
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(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func Debug(template string, args ...interface{}) {
|
func Debug(template string, args ...interface{}) {
|
||||||
if Logger == nil {
|
_, args = extractContext(args...)
|
||||||
log.Printf(template, args...)
|
message := fmt.Sprintf(template, args...)
|
||||||
|
if lg := getLogger(); lg != nil {
|
||||||
|
lg.Debugw(message, "process_id", pid)
|
||||||
return
|
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
|
// CatchPanic - Handle panic
|
||||||
@@ -154,9 +302,11 @@ func CatchPanicCallback(location string, cb func(err any), args ...interface{})
|
|||||||
ctx, _ := extractContext(args...)
|
ctx, _ := extractContext(args...)
|
||||||
return func() {
|
return func() {
|
||||||
if err := recover(); err != nil {
|
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
|
Error("Panic in %s : %v", location, err, ctx) // Pass context implicitly
|
||||||
} else {
|
} else {
|
||||||
fmt.Printf("%s:PANIC->%+v", location, err)
|
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
|
// Send to error tracker
|
||||||
if errorTracker != nil {
|
if tracker != nil {
|
||||||
errorTracker.CapturePanic(ctx, err, callstack, map[string]interface{}{
|
tracker.CapturePanic(ctx, err, callstack, map[string]interface{}{
|
||||||
"location": location,
|
"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...)
|
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
|
// HandlePanic logs a panic and returns it as an error
|
||||||
// This should be called with the result of recover() from a deferred function
|
// This should be called with the result of recover() from a deferred function
|
||||||
// Example usage:
|
// Example usage:
|
||||||
@@ -195,15 +358,16 @@ func CatchPanic(location string, args ...interface{}) func() {
|
|||||||
// }
|
// }
|
||||||
// }()
|
// }()
|
||||||
func HandlePanic(methodName string, r any, args ...interface{}) error {
|
func HandlePanic(methodName string, r any, args ...interface{}) error {
|
||||||
|
tracker := getErrorTracker()
|
||||||
ctx, _ := extractContext(args...)
|
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
|
Error("Panic in %s: %v\nStack trace:\n%s", methodName, r, string(stack), ctx) // Pass context implicitly
|
||||||
|
|
||||||
// Send to error tracker
|
// Send to error tracker
|
||||||
if errorTracker != nil {
|
if tracker != nil {
|
||||||
errorTracker.CapturePanic(ctx, r, stack, map[string]interface{}{
|
tracker.CapturePanic(ctx, r, stack, map[string]interface{}{
|
||||||
"method": methodName,
|
"method": methodName,
|
||||||
"process_id": os.Getpid(),
|
"process_id": pid,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}()
|
||||||
|
}
|
||||||
@@ -2,6 +2,7 @@ package metrics
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
@@ -19,7 +20,7 @@ type Provider interface {
|
|||||||
DecRequestsInFlight()
|
DecRequestsInFlight()
|
||||||
|
|
||||||
// RecordDBQuery records metrics for a database query
|
// RecordDBQuery records metrics for a database query
|
||||||
RecordDBQuery(operation, table string, duration time.Duration, err error)
|
RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error)
|
||||||
|
|
||||||
// RecordCacheHit records a cache hit
|
// RecordCacheHit records a cache hit
|
||||||
RecordCacheHit(provider string)
|
RecordCacheHit(provider string)
|
||||||
@@ -46,21 +47,28 @@ type Provider interface {
|
|||||||
Handler() http.Handler
|
Handler() http.Handler
|
||||||
}
|
}
|
||||||
|
|
||||||
// globalProvider is the global metrics provider
|
// globalProvider is the global metrics provider, protected by globalProviderMu.
|
||||||
var globalProvider Provider
|
var (
|
||||||
|
globalProviderMu sync.RWMutex
|
||||||
|
globalProvider Provider
|
||||||
|
)
|
||||||
|
|
||||||
// SetProvider sets the global metrics provider
|
// SetProvider sets the global metrics provider.
|
||||||
func SetProvider(p Provider) {
|
func SetProvider(p Provider) {
|
||||||
|
globalProviderMu.Lock()
|
||||||
globalProvider = p
|
globalProvider = p
|
||||||
|
globalProviderMu.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetProvider returns the current metrics provider
|
// GetProvider returns the current metrics provider.
|
||||||
func GetProvider() Provider {
|
func GetProvider() Provider {
|
||||||
if globalProvider == nil {
|
globalProviderMu.RLock()
|
||||||
// Return no-op provider if none is set
|
p := globalProvider
|
||||||
|
globalProviderMu.RUnlock()
|
||||||
|
if p == nil {
|
||||||
return &NoOpProvider{}
|
return &NoOpProvider{}
|
||||||
}
|
}
|
||||||
return globalProvider
|
return p
|
||||||
}
|
}
|
||||||
|
|
||||||
// NoOpProvider is a no-op implementation of Provider
|
// NoOpProvider is a no-op implementation of Provider
|
||||||
@@ -69,7 +77,7 @@ type NoOpProvider struct{}
|
|||||||
func (n *NoOpProvider) RecordHTTPRequest(method, path, status string, duration time.Duration) {}
|
func (n *NoOpProvider) RecordHTTPRequest(method, path, status string, duration time.Duration) {}
|
||||||
func (n *NoOpProvider) IncRequestsInFlight() {}
|
func (n *NoOpProvider) IncRequestsInFlight() {}
|
||||||
func (n *NoOpProvider) DecRequestsInFlight() {}
|
func (n *NoOpProvider) DecRequestsInFlight() {}
|
||||||
func (n *NoOpProvider) RecordDBQuery(operation, table string, duration time.Duration, err error) {
|
func (n *NoOpProvider) RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error) {
|
||||||
}
|
}
|
||||||
func (n *NoOpProvider) RecordCacheHit(provider string) {}
|
func (n *NoOpProvider) RecordCacheHit(provider string) {}
|
||||||
func (n *NoOpProvider) RecordCacheMiss(provider string) {}
|
func (n *NoOpProvider) RecordCacheMiss(provider string) {}
|
||||||
|
|||||||
@@ -83,14 +83,14 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
|
|||||||
Help: "Database query duration in seconds",
|
Help: "Database query duration in seconds",
|
||||||
Buckets: cfg.DBQueryBuckets,
|
Buckets: cfg.DBQueryBuckets,
|
||||||
},
|
},
|
||||||
[]string{"operation", "table"},
|
[]string{"operation", "schema", "entity", "table"},
|
||||||
),
|
),
|
||||||
dbQueryTotal: promauto.NewCounterVec(
|
dbQueryTotal: promauto.NewCounterVec(
|
||||||
prometheus.CounterOpts{
|
prometheus.CounterOpts{
|
||||||
Name: metricName("db_queries_total"),
|
Name: metricName("db_queries_total"),
|
||||||
Help: "Total number of database queries",
|
Help: "Total number of database queries",
|
||||||
},
|
},
|
||||||
[]string{"operation", "table", "status"},
|
[]string{"operation", "schema", "entity", "table", "status"},
|
||||||
),
|
),
|
||||||
cacheHits: promauto.NewCounterVec(
|
cacheHits: promauto.NewCounterVec(
|
||||||
prometheus.CounterOpts{
|
prometheus.CounterOpts{
|
||||||
@@ -204,13 +204,13 @@ func (p *PrometheusProvider) DecRequestsInFlight() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// RecordDBQuery implements Provider interface
|
// RecordDBQuery implements Provider interface
|
||||||
func (p *PrometheusProvider) RecordDBQuery(operation, table string, duration time.Duration, err error) {
|
func (p *PrometheusProvider) RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error) {
|
||||||
status := "success"
|
status := "success"
|
||||||
if err != nil {
|
if err != nil {
|
||||||
status = "error"
|
status = "error"
|
||||||
}
|
}
|
||||||
p.dbQueryDuration.WithLabelValues(operation, table).Observe(duration.Seconds())
|
p.dbQueryDuration.WithLabelValues(operation, schema, entity, table).Observe(duration.Seconds())
|
||||||
p.dbQueryTotal.WithLabelValues(operation, table, status).Inc()
|
p.dbQueryTotal.WithLabelValues(operation, schema, entity, table, status).Inc()
|
||||||
}
|
}
|
||||||
|
|
||||||
// RecordCacheHit implements Provider interface
|
// RecordCacheHit implements Provider interface
|
||||||
|
|||||||
@@ -1,9 +1,12 @@
|
|||||||
package modelregistry
|
package modelregistry
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ModelRules defines the permissions and security settings for a model
|
// ModelRules defines the permissions and security settings for a model
|
||||||
@@ -51,6 +54,21 @@ var defaultRegistry = &DefaultModelRegistry{
|
|||||||
var registries = []*DefaultModelRegistry{defaultRegistry}
|
var registries = []*DefaultModelRegistry{defaultRegistry}
|
||||||
var registriesMutex sync.RWMutex
|
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
|
// NewModelRegistry creates a new model registry
|
||||||
func NewModelRegistry() *DefaultModelRegistry {
|
func NewModelRegistry() *DefaultModelRegistry {
|
||||||
return &DefaultModelRegistry{
|
return &DefaultModelRegistry{
|
||||||
@@ -59,11 +77,18 @@ func NewModelRegistry() *DefaultModelRegistry {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetDefaultRegistry returns the current default registry.
|
||||||
func GetDefaultRegistry() *DefaultModelRegistry {
|
func GetDefaultRegistry() *DefaultModelRegistry {
|
||||||
|
registriesMutex.RLock()
|
||||||
|
defer registriesMutex.RUnlock()
|
||||||
return defaultRegistry
|
return defaultRegistry
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetDefaultRegistry replaces the default registry. A nil registry is ignored.
|
||||||
func SetDefaultRegistry(registry *DefaultModelRegistry) {
|
func SetDefaultRegistry(registry *DefaultModelRegistry) {
|
||||||
|
if registry == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
registriesMutex.Lock()
|
registriesMutex.Lock()
|
||||||
defer registriesMutex.Unlock()
|
defer registriesMutex.Unlock()
|
||||||
|
|
||||||
@@ -90,59 +115,86 @@ func AddRegistry(registry *DefaultModelRegistry) {
|
|||||||
registries = append(registries, registry)
|
registries = append(registries, registry)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *DefaultModelRegistry) RegisterModel(name string, model interface{}) error {
|
// registriesSnapshot returns a copy of the registry list so callers can
|
||||||
r.mutex.Lock()
|
// iterate without holding registriesMutex.
|
||||||
defer r.mutex.Unlock()
|
func registriesSnapshot() []*DefaultModelRegistry {
|
||||||
|
registriesMutex.RLock()
|
||||||
|
defer registriesMutex.RUnlock()
|
||||||
|
return append([]*DefaultModelRegistry(nil), registries...)
|
||||||
|
}
|
||||||
|
|
||||||
if _, exists := r.models[name]; exists {
|
// validateModel checks the model is a struct (or pointer/slice/array of one)
|
||||||
return fmt.Errorf("model %s already registered", name)
|
// 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))
|
||||||
}
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
// Validate that model is a non-pointer struct
|
|
||||||
modelType := reflect.TypeOf(model)
|
modelType := reflect.TypeOf(model)
|
||||||
if modelType == nil {
|
if modelType == nil {
|
||||||
return fmt.Errorf("model cannot be nil")
|
return nil, fmt.Errorf("%w: model cannot be nil", ErrInvalidModel)
|
||||||
}
|
}
|
||||||
|
|
||||||
originalType := modelType
|
originalType := modelType
|
||||||
|
|
||||||
// Unwrap pointers, slices, and arrays to check the underlying type
|
// Unwrap pointers, slices, and arrays to check the underlying type
|
||||||
for modelType.Kind() == reflect.Ptr || 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()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate that the underlying type is a struct
|
|
||||||
if modelType.Kind() != reflect.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 a pointer/slice/array was passed, unwrap to the base struct
|
||||||
if originalType != modelType {
|
if originalType != modelType {
|
||||||
// Create a zero value of the struct type
|
|
||||||
model = reflect.New(modelType).Elem().Interface()
|
model = reflect.New(modelType).Elem().Interface()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Additional check: ensure model is not a pointer
|
if finalType := reflect.TypeOf(model); finalType.Kind() == reflect.Pointer {
|
||||||
finalType := reflect.TypeOf(model)
|
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())
|
||||||
if finalType.Kind() == reflect.Ptr {
|
}
|
||||||
return fmt.Errorf("model must be a non-pointer struct, got pointer to %s. Use MyModel{} instead of &MyModel{}", 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
|
r.mutex.Lock()
|
||||||
// Initialize with default rules if not already set
|
defer r.mutex.Unlock()
|
||||||
if _, exists := r.rules[name]; !exists {
|
|
||||||
r.rules[name] = DefaultModelRules()
|
if _, exists := r.models[name]; exists {
|
||||||
|
return fmt.Errorf("%w: %s", ErrModelExists, name)
|
||||||
}
|
}
|
||||||
|
r.models[name] = model
|
||||||
|
r.rules[name] = rules
|
||||||
return nil
|
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) {
|
func (r *DefaultModelRegistry) GetModel(name string) (interface{}, error) {
|
||||||
r.mutex.RLock()
|
r.mutex.RLock()
|
||||||
defer r.mutex.RUnlock()
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
model, exists := r.models[name]
|
model, exists := r.models[name]
|
||||||
if !exists {
|
if !exists {
|
||||||
return nil, fmt.Errorf("model %s not found", name)
|
return nil, fmt.Errorf("%w: %s", ErrModelNotFound, name)
|
||||||
}
|
}
|
||||||
|
|
||||||
return model, nil
|
return model, nil
|
||||||
@@ -152,7 +204,7 @@ func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
|
|||||||
r.mutex.RLock()
|
r.mutex.RLock()
|
||||||
defer r.mutex.RUnlock()
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
result := make(map[string]interface{})
|
result := make(map[string]interface{}, len(r.models))
|
||||||
for k, v := range r.models {
|
for k, v := range r.models {
|
||||||
result[k] = v
|
result[k] = v
|
||||||
}
|
}
|
||||||
@@ -162,9 +214,13 @@ func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
|
|||||||
func (r *DefaultModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) {
|
func (r *DefaultModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) {
|
||||||
// Try full name first
|
// Try full name first
|
||||||
fullName := fmt.Sprintf("%s.%s", schema, entity)
|
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
|
return model, nil
|
||||||
}
|
}
|
||||||
|
if !errors.Is(err, ErrModelNotFound) {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
// Fallback to entity name only
|
// Fallback to entity name only
|
||||||
return r.GetModel(entity)
|
return r.GetModel(entity)
|
||||||
@@ -175,9 +231,8 @@ func (r *DefaultModelRegistry) SetModelRules(name string, rules ModelRules) erro
|
|||||||
r.mutex.Lock()
|
r.mutex.Lock()
|
||||||
defer r.mutex.Unlock()
|
defer r.mutex.Unlock()
|
||||||
|
|
||||||
// Check if model exists
|
|
||||||
if _, exists := r.models[name]; !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
|
r.rules[name] = rules
|
||||||
@@ -190,12 +245,10 @@ func (r *DefaultModelRegistry) GetModelRules(name string) (ModelRules, error) {
|
|||||||
r.mutex.RLock()
|
r.mutex.RLock()
|
||||||
defer r.mutex.RUnlock()
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
// Check if model exists
|
|
||||||
if _, exists := r.models[name]; !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 {
|
if rules, exists := r.rules[name]; exists {
|
||||||
return rules, nil
|
return rules, nil
|
||||||
}
|
}
|
||||||
@@ -203,72 +256,62 @@ func (r *DefaultModelRegistry) GetModelRules(name string) (ModelRules, error) {
|
|||||||
return DefaultModelRules(), nil
|
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 {
|
func (r *DefaultModelRegistry) RegisterModelWithRules(name string, model interface{}, rules ModelRules) error {
|
||||||
// First register the model
|
return r.registerLocked(name, model, rules)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Global convenience functions using the default registry
|
// Global convenience functions using the default registry
|
||||||
|
|
||||||
// RegisterModel registers a model with the default global registry
|
// RegisterModel registers a model with the default global registry
|
||||||
func RegisterModel(model interface{}, name string) error {
|
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
|
// GetModelByName retrieves a model by searching through all registries in order
|
||||||
// Returns the first match found
|
// Returns the first match found
|
||||||
func GetModelByName(name string) (interface{}, error) {
|
func GetModelByName(name string) (interface{}, error) {
|
||||||
registriesMutex.RLock()
|
for _, registry := range registriesSnapshot() {
|
||||||
defer registriesMutex.RUnlock()
|
|
||||||
|
|
||||||
for _, registry := range registries {
|
|
||||||
if model, err := registry.GetModel(name); err == nil {
|
if model, err := registry.GetModel(name); err == nil {
|
||||||
return model, 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{})) {
|
func IterateModels(fn func(name string, model interface{})) {
|
||||||
defaultRegistry.mutex.RLock()
|
for name, model := range GetDefaultRegistry().GetAllModels() {
|
||||||
defer defaultRegistry.mutex.RUnlock()
|
callIsolated(name, model, fn)
|
||||||
|
|
||||||
for name, model := range defaultRegistry.models {
|
|
||||||
fn(name, model)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetModels returns a list of all models from all registries
|
func callIsolated(name string, model interface{}, fn func(name string, model interface{})) {
|
||||||
// Models are collected in registry order, with duplicates included
|
defer func() {
|
||||||
func GetModels() []interface{} {
|
if r := recover(); r != nil {
|
||||||
registriesMutex.RLock()
|
_ = logger.HandlePanic("modelregistry.IterateModels", r, "model", name)
|
||||||
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{}
|
var models []interface{}
|
||||||
seen := make(map[string]bool)
|
seen := make(map[string]bool)
|
||||||
|
|
||||||
for _, registry := range registries {
|
for _, registry := range registriesSnapshot() {
|
||||||
registry.mutex.RLock()
|
for name, model := range registry.GetAllModels() {
|
||||||
for name, model := range registry.models {
|
|
||||||
// Only add the first occurrence of each model name
|
|
||||||
if !seen[name] {
|
if !seen[name] {
|
||||||
models = append(models, model)
|
models = append(models, model)
|
||||||
seen[name] = true
|
seen[name] = true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
registry.mutex.RUnlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return models
|
return models
|
||||||
@@ -276,31 +319,31 @@ func GetModels() []interface{} {
|
|||||||
|
|
||||||
// SetModelRules sets the rules for a specific model in the default registry
|
// SetModelRules sets the rules for a specific model in the default registry
|
||||||
func SetModelRules(name string, rules ModelRules) error {
|
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
|
// GetModelRules retrieves the rules for a specific model from the default registry
|
||||||
func GetModelRules(name string) (ModelRules, error) {
|
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
|
// 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) {
|
func GetModelRulesByName(name string) (ModelRules, error) {
|
||||||
registriesMutex.RLock()
|
for _, registry := range registriesSnapshot() {
|
||||||
defer registriesMutex.RUnlock()
|
rules, err := registry.GetModelRules(name)
|
||||||
|
if err == nil {
|
||||||
for _, registry := range registries {
|
return rules, nil
|
||||||
if _, err := registry.GetModel(name); err == nil {
|
}
|
||||||
// Model found in this registry, get its rules
|
if !errors.Is(err, ErrModelNotFound) {
|
||||||
return registry.GetModelRules(name)
|
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
|
// RegisterModelWithRules registers a model with specific rules in the default registry
|
||||||
func RegisterModelWithRules(model interface{}, name string, rules ModelRules) error {
|
func RegisterModelWithRules(model interface{}, name string, rules ModelRules) error {
|
||||||
return defaultRegistry.RegisterModelWithRules(name, model, rules)
|
return GetDefaultRegistry().RegisterModelWithRules(name, model, rules)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -324,7 +324,7 @@ func (ebc *ExternalBrokerClient) Stop(ctx context.Context) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if ebc.client != nil && ebc.client.IsConnected() {
|
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
|
ebc.connected = false
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user