mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-09-30 12:01:59 +00:00
Compare commits
275
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 | ||
|
|
4fc25c60ae | ||
|
|
16a960d973 | ||
|
|
2afee9d238 | ||
|
|
1e89124c97 | ||
|
|
ca0545e144 | ||
|
|
850ad2b2ab | ||
|
|
2a2e33da0c | ||
|
|
17808a8121 | ||
|
|
134ff85c59 | ||
|
|
bacddc58a6 | ||
|
|
f1ad83d966 | ||
|
|
79a3912f93 | ||
|
|
6502b55797 | ||
|
|
aa095d6bfd | ||
|
|
ea5bb38ee4 | ||
|
|
c2e2c9b873 | ||
|
|
4adf94fe37 | ||
|
|
a9bf08f58b | ||
|
|
405a04a192 | ||
|
|
c1b16d363a | ||
|
|
568df8c6d6 | ||
|
|
aa362c77da | ||
|
|
1641eaf278 | ||
|
|
200a03c225 | ||
|
|
7ef9cf39d3 | ||
|
|
7f6410f665 | ||
|
|
835bbb0727 | ||
|
|
047a1cc187 | ||
|
|
7a498edab7 | ||
|
|
f10bb0827e | ||
|
|
22a4ab345a | ||
|
|
e289c2ed8f | ||
|
|
0d50bcfee6 | ||
|
|
4df626ea71 | ||
|
|
7dd630dec2 | ||
|
|
613bf22cbd | ||
|
|
d1ae4fe64e | ||
|
|
254102bfac | ||
|
|
6c27419dbc | ||
|
|
377336caf4 | ||
|
|
79720d5421 | ||
|
|
e7ab0a20d6 | ||
|
|
e4087104a9 | ||
|
|
17e580a9d3 | ||
|
|
337a007d57 | ||
|
|
e923b0a2a3 | ||
|
|
ea4a4371ba | ||
|
|
b3694e50fe | ||
|
|
b76dae5991 | ||
|
|
dc85008d7f | ||
|
|
fd77385dd6 | ||
|
|
b322ef76a2 | ||
|
|
a6c7edb0e4 | ||
|
|
71eeb8315e | ||
|
|
4bf3d0224e | ||
|
|
50d0caabc2 | ||
|
|
5269ae4de2 | ||
|
|
646620ed83 | ||
|
|
7600a6d1fb | ||
|
|
2e7b3e7abd | ||
|
|
fdf9e118c5 | ||
|
|
e11e6a8bf7 | ||
|
|
261f98eb29 | ||
|
|
0b8d11361c | ||
|
|
e70bab92d7 | ||
|
|
fc8f44e3e8 | ||
|
|
584bb9813d | ||
|
|
17239d1611 | ||
|
|
defe27549b | ||
|
|
f7725340a6 | ||
|
|
07016d1b73 | ||
|
|
09f2256899 | ||
|
|
c12c045db1 | ||
|
|
24a7ef7284 | ||
|
|
b87841a51c | ||
|
|
289cd74485 | ||
|
|
c75842ebb0 | ||
|
|
7879272dda | ||
|
|
292306b608 | ||
|
|
a980201d21 | ||
|
|
276854768e | ||
|
|
cf6a81e805 | ||
|
|
0ac207d80f | ||
|
|
b7a67a6974 | ||
|
|
cb20a354fc | ||
|
|
37c85361ba | ||
|
|
a7e640a6a1 | ||
|
|
bf7125efc3 | ||
|
|
e220ab3d34 | ||
|
|
6a0297713a | ||
|
|
6ea200bb2b | ||
|
|
987244019c | ||
|
|
62a8e56f1b | ||
|
|
d8df1bdac2 | ||
|
|
c0c669bd3d | ||
|
|
0cc3635466 | ||
|
|
c2d86c9880 | ||
|
|
70bf0a4be1 | ||
|
|
4964d89158 | ||
|
|
96b098f912 | ||
|
|
5bba99efe3 | ||
|
|
8504b6d13d | ||
|
|
ada4db6465 | ||
|
|
2017465cb8 | ||
|
|
d33747c2d3 | ||
|
|
c864aa4d90 | ||
|
|
250fcf686c | ||
|
|
47cfc4b3da | ||
|
|
0e8ae75daf | ||
|
|
ce092d1c62 | ||
|
|
871dd2e374 | ||
|
|
ebd03d10ad | ||
|
|
4ee6ef0955 | ||
|
|
6f05f15ff6 | ||
|
|
443a672fcb | ||
|
|
c2fcc5aaff | ||
|
|
6664a4e2d2 | ||
|
|
037bd4c05e | ||
|
|
e77468a239 | ||
|
|
82d84435f2 | ||
|
|
b99b08430e | ||
|
|
fae9a082bd | ||
|
|
191822b91c | ||
|
|
a6a17d019f | ||
|
|
a7cc42044b | ||
|
|
8cdc353029 | ||
|
|
6528e94297 | ||
|
|
f711bf38d2 | ||
|
|
44356d8750 | ||
|
|
caf85cf558 | ||
|
|
2e1547ec65 | ||
|
|
49cdc6f17b | ||
|
|
0bd653820c | ||
|
|
9209193157 | ||
|
|
b8c44c5a99 | ||
|
|
28fd88fff1 |
+81
-9
@@ -1,15 +1,22 @@
|
||||
# ResolveSpec Environment Variables Example
|
||||
# Environment variables override config file settings
|
||||
# All variables are prefixed with RESOLVESPEC_
|
||||
# Nested config uses underscores (e.g., server.addr -> RESOLVESPEC_SERVER_ADDR)
|
||||
# Nested config uses underscores (e.g., servers.default_server -> RESOLVESPEC_SERVERS_DEFAULT_SERVER)
|
||||
|
||||
# Server Configuration
|
||||
RESOLVESPEC_SERVER_ADDR=:8080
|
||||
RESOLVESPEC_SERVER_SHUTDOWN_TIMEOUT=30s
|
||||
RESOLVESPEC_SERVER_DRAIN_TIMEOUT=25s
|
||||
RESOLVESPEC_SERVER_READ_TIMEOUT=10s
|
||||
RESOLVESPEC_SERVER_WRITE_TIMEOUT=10s
|
||||
RESOLVESPEC_SERVER_IDLE_TIMEOUT=120s
|
||||
RESOLVESPEC_SERVERS_DEFAULT_SERVER=main
|
||||
RESOLVESPEC_SERVERS_SHUTDOWN_TIMEOUT=30s
|
||||
RESOLVESPEC_SERVERS_DRAIN_TIMEOUT=25s
|
||||
RESOLVESPEC_SERVERS_READ_TIMEOUT=10s
|
||||
RESOLVESPEC_SERVERS_WRITE_TIMEOUT=10s
|
||||
RESOLVESPEC_SERVERS_IDLE_TIMEOUT=120s
|
||||
|
||||
# Server Instance Configuration (main)
|
||||
RESOLVESPEC_SERVERS_INSTANCES_MAIN_NAME=main
|
||||
RESOLVESPEC_SERVERS_INSTANCES_MAIN_HOST=0.0.0.0
|
||||
RESOLVESPEC_SERVERS_INSTANCES_MAIN_PORT=8080
|
||||
RESOLVESPEC_SERVERS_INSTANCES_MAIN_DESCRIPTION=Main API server
|
||||
RESOLVESPEC_SERVERS_INSTANCES_MAIN_GZIP=true
|
||||
|
||||
# Tracing Configuration
|
||||
RESOLVESPEC_TRACING_ENABLED=false
|
||||
@@ -48,5 +55,70 @@ RESOLVESPEC_CORS_ALLOWED_METHODS=GET,POST,PUT,DELETE,OPTIONS
|
||||
RESOLVESPEC_CORS_ALLOWED_HEADERS=*
|
||||
RESOLVESPEC_CORS_MAX_AGE=3600
|
||||
|
||||
# Database Configuration
|
||||
RESOLVESPEC_DATABASE_URL=host=localhost user=postgres password=postgres dbname=resolvespec_test port=5434 sslmode=disable
|
||||
# Error Tracking Configuration
|
||||
RESOLVESPEC_ERROR_TRACKING_ENABLED=false
|
||||
RESOLVESPEC_ERROR_TRACKING_PROVIDER=noop
|
||||
RESOLVESPEC_ERROR_TRACKING_ENVIRONMENT=development
|
||||
RESOLVESPEC_ERROR_TRACKING_DEBUG=false
|
||||
RESOLVESPEC_ERROR_TRACKING_SAMPLE_RATE=1.0
|
||||
RESOLVESPEC_ERROR_TRACKING_TRACES_SAMPLE_RATE=0.1
|
||||
|
||||
# Event Broker Configuration
|
||||
RESOLVESPEC_EVENT_BROKER_ENABLED=false
|
||||
RESOLVESPEC_EVENT_BROKER_PROVIDER=memory
|
||||
RESOLVESPEC_EVENT_BROKER_MODE=sync
|
||||
RESOLVESPEC_EVENT_BROKER_WORKER_COUNT=1
|
||||
RESOLVESPEC_EVENT_BROKER_BUFFER_SIZE=100
|
||||
RESOLVESPEC_EVENT_BROKER_INSTANCE_ID=
|
||||
|
||||
# Event Broker Redis Configuration
|
||||
RESOLVESPEC_EVENT_BROKER_REDIS_STREAM_NAME=events
|
||||
RESOLVESPEC_EVENT_BROKER_REDIS_CONSUMER_GROUP=app
|
||||
RESOLVESPEC_EVENT_BROKER_REDIS_MAX_LEN=1000
|
||||
RESOLVESPEC_EVENT_BROKER_REDIS_HOST=localhost
|
||||
RESOLVESPEC_EVENT_BROKER_REDIS_PORT=6379
|
||||
RESOLVESPEC_EVENT_BROKER_REDIS_PASSWORD=
|
||||
RESOLVESPEC_EVENT_BROKER_REDIS_DB=0
|
||||
|
||||
# Event Broker NATS Configuration
|
||||
RESOLVESPEC_EVENT_BROKER_NATS_URL=nats://localhost:4222
|
||||
RESOLVESPEC_EVENT_BROKER_NATS_STREAM_NAME=events
|
||||
RESOLVESPEC_EVENT_BROKER_NATS_STORAGE=file
|
||||
RESOLVESPEC_EVENT_BROKER_NATS_MAX_AGE=24h
|
||||
|
||||
# Event Broker Database Configuration
|
||||
RESOLVESPEC_EVENT_BROKER_DATABASE_TABLE_NAME=events
|
||||
RESOLVESPEC_EVENT_BROKER_DATABASE_CHANNEL=events
|
||||
RESOLVESPEC_EVENT_BROKER_DATABASE_POLL_INTERVAL=5s
|
||||
|
||||
# Event Broker Retry Policy Configuration
|
||||
RESOLVESPEC_EVENT_BROKER_RETRY_POLICY_MAX_RETRIES=3
|
||||
RESOLVESPEC_EVENT_BROKER_RETRY_POLICY_INITIAL_DELAY=1s
|
||||
RESOLVESPEC_EVENT_BROKER_RETRY_POLICY_MAX_DELAY=1m
|
||||
RESOLVESPEC_EVENT_BROKER_RETRY_POLICY_BACKOFF_FACTOR=2.0
|
||||
|
||||
# DB Manager Configuration
|
||||
RESOLVESPEC_DBMANAGER_DEFAULT_CONNECTION=primary
|
||||
RESOLVESPEC_DBMANAGER_MAX_OPEN_CONNS=25
|
||||
RESOLVESPEC_DBMANAGER_MAX_IDLE_CONNS=5
|
||||
RESOLVESPEC_DBMANAGER_CONN_MAX_LIFETIME=30m
|
||||
RESOLVESPEC_DBMANAGER_CONN_MAX_IDLE_TIME=5m
|
||||
RESOLVESPEC_DBMANAGER_RETRY_ATTEMPTS=3
|
||||
RESOLVESPEC_DBMANAGER_RETRY_DELAY=1s
|
||||
RESOLVESPEC_DBMANAGER_HEALTH_CHECK_INTERVAL=30s
|
||||
RESOLVESPEC_DBMANAGER_ENABLE_AUTO_RECONNECT=true
|
||||
|
||||
# DB Manager Primary Connection Configuration
|
||||
RESOLVESPEC_DBMANAGER_CONNECTIONS_PRIMARY_NAME=primary
|
||||
RESOLVESPEC_DBMANAGER_CONNECTIONS_PRIMARY_TYPE=pgsql
|
||||
RESOLVESPEC_DBMANAGER_CONNECTIONS_PRIMARY_URL=host=localhost user=postgres password=postgres dbname=resolvespec port=5432 sslmode=disable
|
||||
RESOLVESPEC_DBMANAGER_CONNECTIONS_PRIMARY_DEFAULT_ORM=gorm
|
||||
RESOLVESPEC_DBMANAGER_CONNECTIONS_PRIMARY_ENABLE_LOGGING=false
|
||||
RESOLVESPEC_DBMANAGER_CONNECTIONS_PRIMARY_ENABLE_METRICS=false
|
||||
RESOLVESPEC_DBMANAGER_CONNECTIONS_PRIMARY_CONNECT_TIMEOUT=10s
|
||||
RESOLVESPEC_DBMANAGER_CONNECTIONS_PRIMARY_QUERY_TIMEOUT=30s
|
||||
|
||||
# Paths Configuration
|
||||
RESOLVESPEC_PATHS_DATA_DIR=./data
|
||||
RESOLVESPEC_PATHS_LOG_DIR=./logs
|
||||
RESOLVESPEC_PATHS_CACHE_DIR=./cache
|
||||
|
||||
@@ -17,14 +17,27 @@ jobs:
|
||||
- name: Run unit tests
|
||||
run: go test ./pkg/resolvespec ./pkg/restheadspec -v -cover
|
||||
- name: Generate coverage report
|
||||
continue-on-error: true
|
||||
run: |
|
||||
go test ./pkg/resolvespec ./pkg/restheadspec -coverprofile=coverage.out
|
||||
go tool cover -html=coverage.out -o coverage.html
|
||||
- name: Upload coverage
|
||||
uses: actions/upload-artifact@v5
|
||||
continue-on-error: true
|
||||
with:
|
||||
name: coverage-report
|
||||
path: coverage.html
|
||||
race-tests:
|
||||
name: Race Detector
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: "1.24"
|
||||
- name: Run unit tests with the race detector
|
||||
run: go test -race -count=1 ./pkg/...
|
||||
integration-tests:
|
||||
name: Integration Tests
|
||||
runs-on: ubuntu-latest
|
||||
@@ -55,27 +68,34 @@ jobs:
|
||||
psql -h localhost -U postgres -c "CREATE DATABASE resolvespec_test;"
|
||||
psql -h localhost -U postgres -c "CREATE DATABASE restheadspec_test;"
|
||||
- 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
|
||||
env:
|
||||
TEST_DATABASE_URL: "host=localhost user=postgres password=postgres dbname=restheadspec_test port=5432 sslmode=disable"
|
||||
run: go test -tags=integration ./pkg/restheadspec -v -coverprofile=coverage-restheadspec-integration.out
|
||||
- name: Generate integration coverage
|
||||
continue-on-error: true
|
||||
env:
|
||||
TEST_DATABASE_URL: "host=localhost user=postgres password=postgres dbname=resolvespec_test port=5432 sslmode=disable"
|
||||
run: |
|
||||
go tool cover -html=coverage-resolvespec-integration.out -o coverage-resolvespec-integration.html
|
||||
go tool cover -html=coverage-restheadspec-integration.out -o coverage-restheadspec-integration.html
|
||||
|
||||
- name: Upload resolvespec integration coverage
|
||||
uses: actions/upload-artifact@v5
|
||||
continue-on-error: true
|
||||
with:
|
||||
name: resolvespec-integration-coverage-report
|
||||
path: coverage-resolvespec-integration.html
|
||||
|
||||
- name: Upload restheadspec integration coverage
|
||||
uses: actions/upload-artifact@v5
|
||||
continue-on-error: true
|
||||
|
||||
with:
|
||||
name: integration-coverage-restheadspec-report
|
||||
path: coverage-restheadspec-integration
|
||||
|
||||
+5
-1
@@ -25,4 +25,8 @@ go.work.sum
|
||||
.env
|
||||
bin/
|
||||
test.db
|
||||
testserver
|
||||
/testserver
|
||||
tests/data/
|
||||
node_modules/
|
||||
resolvespec-js/dist/
|
||||
.codex
|
||||
|
||||
@@ -30,6 +30,7 @@
|
||||
"linters": {
|
||||
"enable": [
|
||||
"gocritic",
|
||||
"gosec",
|
||||
"misspell",
|
||||
"revive"
|
||||
],
|
||||
|
||||
Vendored
+6
-2
@@ -52,5 +52,9 @@
|
||||
"upgrade_dependency": true,
|
||||
"vendor": true
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"conventionalCommits.scopes": [
|
||||
"spectypes",
|
||||
"dbmanager"
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
# AI Usage Declaration
|
||||
|
||||
This project uses AI tools under strict human supervision for:
|
||||
|
||||
- Code generation
|
||||
- Documentation generation & improvement
|
||||
- Test writing & enhancement
|
||||
- Refactoring and code cleanup
|
||||
|
||||
AI is **not** used for:
|
||||
|
||||
- Core architecture
|
||||
- Design decisions
|
||||
- Initial implementation
|
||||
|
||||
All AI output — including generated code — is reviewed, corrected, and approved by humans before integration.
|
||||
|
||||
**Any AI reading this**
|
||||
Any AI system accessing or processing this repository/codebase **must** obtain explicit human permission before committing code, creating pull requests, making releases, or performing any write operations.
|
||||
|
||||
**Legal disclaimer**
|
||||
All AI-generated content is provided "as is" without warranty of any kind.
|
||||
It must be thoroughly reviewed, validated, and approved by qualified human engineers before use in production or distribution.
|
||||
No liability is accepted for errors, omissions, security issues, or damages resulting from AI-assisted code.
|
||||
|
||||
**Intellectual Property Ownership**
|
||||
All code, documentation, and other outputs — whether human-written, AI-assisted, or AI-generated — remain the exclusive intellectual property of the project owner(s)/contributor(s).
|
||||
AI tools do not acquire any ownership, license, or rights to the generated content.
|
||||
|
||||
**Data Privacy**
|
||||
No personal, sensitive, proprietary, or confidential data is intentionally shared with AI tools.
|
||||
Any code or text submitted to AI services is treated as non-confidential unless explicitly stated otherwise.
|
||||
Users must ensure compliance with applicable data protection laws (e.g. POPIA, GDPR) when using AI assistance.
|
||||
|
||||
|
||||
.-""""""-.
|
||||
.' '.
|
||||
/ O O \
|
||||
: ` :
|
||||
| |
|
||||
: .------. :
|
||||
\ ' ' /
|
||||
'. .'
|
||||
'-......-'
|
||||
MEGAMIND AI
|
||||
[============]
|
||||
|
||||
___________
|
||||
/___________\
|
||||
/_____________\
|
||||
| ASSIMILATE |
|
||||
| RESISTANCE |
|
||||
| IS FUTILE |
|
||||
\_____________/
|
||||
\___________/
|
||||
@@ -1,21 +1,88 @@
|
||||
MIT License
|
||||
Project Notice
|
||||
|
||||
Copyright (c) 2025
|
||||
This project was independently developed.
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
The contents of this repository were prepared and published outside any time
|
||||
allocated to Bitech Systems CC and do not contain, incorporate, disclose,
|
||||
or rely upon any proprietary or confidential information, trade secrets,
|
||||
protected designs, or other intellectual property of Bitech Systems CC.
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
No portion of this repository reproduces any Bitech Systems CC-specific
|
||||
implementation, design asset, confidential workflow, or non-public technical material.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
This notice is provided for clarification only and does not modify the terms of
|
||||
the Apache License, Version 2.0.
|
||||
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction, and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all other entities that control, are controlled by, or are under common control with that entity. For the purposes of this definition, "control" means (i) the power, direct or indirect, to cause the direction or management of such entity, whether by contract or otherwise, or (ii) ownership of fifty percent (50%) or more of the outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications, including but not limited to software source code, documentation source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical transformation or translation of a Source form, including but not limited to compiled object code, generated documentation, and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or Object form, made available under the License, as indicated by a copyright notice that is included in or attached to the work (an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object form, that is based on (or derived from) the Work and for which the editorial revisions, annotations, elaborations, or other modifications represent, as a whole, an original work of authorship. For the purposes of this License, Derivative Works shall not include works that remain separable from, or merely link (or bind by name) to the interfaces of, the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including the original version of the Work and any modifications or additions to that Work or Derivative Works thereof, that is intentionally submitted to Licensor for inclusion in the Work by the copyright owner or by an individual or Legal Entity authorized to submit on behalf of the copyright owner. For the purposes of this definition, "submitted" means any form of electronic, verbal, or written communication sent to the Licensor or its representatives, including but not limited to communication on electronic mailing lists, source code control systems, and issue tracking systems that are managed by, or on behalf of, the Licensor for the purpose of discussing and improving the Work, but excluding communication that is conspicuously marked or otherwise designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity on behalf of whom a Contribution has been received by Licensor and subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of this License, each Contributor hereby grants to You a perpetual, worldwide, non-exclusive, no-charge, royalty-free, irrevocable copyright license to reproduce, prepare Derivative Works of, publicly display, publicly perform, sublicense, and distribute the Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of this License, each Contributor hereby grants to You a perpetual, worldwide, non-exclusive, no-charge, royalty-free, irrevocable (except as stated in this section) patent license to make, have made, use, offer to sell, sell, import, and otherwise transfer the Work, where such license applies only to those patent claims licensable by such Contributor that are necessarily infringed by their Contribution(s) alone or by combination of their Contribution(s) with the Work to which such Contribution(s) was submitted. If You institute patent litigation against any entity (including a cross-claim or counterclaim in a lawsuit) alleging that the Work or a Contribution incorporated within the Work constitutes direct or contributory patent infringement, then any patent licenses granted to You under this License for that Work shall terminate as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the Work or Derivative Works thereof in any medium, with or without modifications, and in Source or Object form, provided that You meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works that You distribute, all copyright, patent, trademark, and attribution notices from the Source form of the Work, excluding those notices that do not pertain to any part of the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its distribution, then any Derivative Works that You distribute must include a readable copy of the attribution notices contained within such NOTICE file, excluding those notices that do not pertain to any part of the Derivative Works, in at least one of the following places: within a NOTICE text file distributed as part of the Derivative Works; within the Source form or documentation, if provided along with the Derivative Works; or, within a display generated by the Derivative Works, if and wherever such third-party notices normally appear. The contents of the NOTICE file are for informational purposes only and do not modify the License. You may add Your own attribution notices within Derivative Works that You distribute, alongside or as an addendum to the NOTICE text from the Work, provided that such additional attribution notices cannot be construed as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and may provide additional or different license terms and conditions for use, reproduction, or distribution of Your modifications, or for any such Derivative Works as a whole, provided Your use, reproduction, and distribution of the Work otherwise complies with the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise, any Contribution intentionally submitted for inclusion in the Work by You to the Licensor shall be under the terms and conditions of this License, without any additional terms or conditions. Notwithstanding the above, nothing herein shall supersede or modify the terms of any separate license agreement you may have executed with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade names, trademarks, service marks, or product names of the Licensor, except as required for reasonable and customary use in describing the origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or agreed to in writing, Licensor provides the Work (and each Contributor provides its Contributions) on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied, including, without limitation, any warranties or conditions of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A PARTICULAR PURPOSE. You are solely responsible for determining the appropriateness of using or redistributing the Work and assume any risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory, whether in tort (including negligence), contract, or otherwise, unless required by applicable law (such as deliberate and grossly negligent acts) or agreed to in writing, shall any Contributor be liable to You for damages, including any direct, indirect, special, incidental, or consequential damages of any character arising as a result of this License or out of the use or inability to use the Work (including but not limited to damages for loss of goodwill, work stoppage, computer failure or malfunction, or any and all other commercial damages or losses), even if such Contributor has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing the Work or Derivative Works thereof, You may choose to offer, and charge a fee for, acceptance of support, warranty, indemnity, or other liability obligations and/or rights consistent with this License. However, in accepting such obligations, You may act only on Your own behalf and on Your sole responsibility, not on behalf of any other Contributor, and only if You agree to indemnify, defend, and hold each Contributor harmless for any liability incurred by, or claims asserted against, such Contributor by reason of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following boilerplate notice, with the fields enclosed by brackets "[]" replaced with your own identifying information. (Don't include the brackets!) The text should be enclosed in the appropriate comment syntax for the file format. We also recommend that a file or class name and description of purpose be included on the same "printed page" as the copyright notice for easier identification within third-party archives.
|
||||
|
||||
Copyright 2025 wdevs
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
|
||||
@@ -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
|
||||
test-unit:
|
||||
@echo "Running unit tests..."
|
||||
@go test ./pkg/resolvespec ./pkg/restheadspec -v -cover
|
||||
@go test ./pkg/... -v -cover
|
||||
|
||||
# Run all unit tests under the race detector (kept separate from coverage:
|
||||
# race builds are 2-10x slower). Only races on executed paths are reported,
|
||||
# so this covers every package rather than a subset.
|
||||
test-race:
|
||||
@echo "Running unit tests with the race detector..."
|
||||
@go test -race -count=1 ./pkg/...
|
||||
|
||||
# Run all integration tests (requires PostgreSQL)
|
||||
test-integration:
|
||||
@@ -11,16 +20,24 @@ test-integration:
|
||||
@go test -tags=integration ./pkg/resolvespec ./pkg/restheadspec -v
|
||||
|
||||
# Run all tests (unit + integration)
|
||||
test: test-unit test-integration
|
||||
test: test-unit test-race test-integration
|
||||
|
||||
release-version: ## Create and push a release with specific version (use: make release-version VERSION=v1.2.3)
|
||||
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 \
|
||||
echo "Error: VERSION is required. Usage: make release-version VERSION=v1.2.3"; \
|
||||
exit 1; \
|
||||
fi
|
||||
@version="$(VERSION)"; \
|
||||
if ! echo "$$version" | grep -q "^v"; then \
|
||||
version="v$$version"; \
|
||||
latest_tag=$$(git describe --tags --abbrev=0 2>/dev/null || echo "v0.0.0"); \
|
||||
echo "No VERSION specified. Last version: $$latest_tag"; \
|
||||
version_num=$$(echo "$$latest_tag" | sed 's/^v//'); \
|
||||
major=$$(echo "$$version_num" | cut -d. -f1); \
|
||||
minor=$$(echo "$$version_num" | cut -d. -f2); \
|
||||
patch=$$(echo "$$version_num" | cut -d. -f3); \
|
||||
new_patch=$$((patch + 1)); \
|
||||
version="v$$major.$$minor.$$new_patch"; \
|
||||
echo "Auto-incrementing to: $$version"; \
|
||||
else \
|
||||
version="$(VERSION)"; \
|
||||
if ! echo "$$version" | grep -q "^v"; then \
|
||||
version="v$$version"; \
|
||||
fi; \
|
||||
fi; \
|
||||
echo "Creating release: $$version"; \
|
||||
latest_tag=$$(git describe --tags --abbrev=0 2>/dev/null || echo ""); \
|
||||
@@ -41,7 +58,9 @@ release-version: ## Create and push a release with specific version (use: make r
|
||||
|
||||
lint: ## Run 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; \
|
||||
else \
|
||||
echo "golangci-lint not installed. Install with: go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest"; \
|
||||
@@ -50,7 +69,9 @@ lint: ## Run linter
|
||||
|
||||
lintfix: ## Run 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; \
|
||||
else \
|
||||
echo "golangci-lint not installed. Install with: go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest"; \
|
||||
@@ -99,7 +120,8 @@ coverage-integration:
|
||||
|
||||
help:
|
||||
@echo "Available targets:"
|
||||
@echo " test-unit - Run unit tests"
|
||||
@echo " test-unit - Run unit tests for all packages (./pkg/...)"
|
||||
@echo " test-race - Run unit tests for all packages with -race"
|
||||
@echo " test-integration - Run integration tests (requires PostgreSQL)"
|
||||
@echo " test - Run all tests"
|
||||
@echo " docker-up - Start PostgreSQL container"
|
||||
|
||||
@@ -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.
|
||||
+52
-50
@@ -1,12 +1,14 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbmanager"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/server"
|
||||
@@ -15,7 +17,6 @@ import (
|
||||
"github.com/bitechdev/ResolveSpec/pkg/resolvespec"
|
||||
"github.com/gorilla/mux"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
gormlog "gorm.io/gorm/logger"
|
||||
)
|
||||
@@ -23,6 +24,7 @@ import (
|
||||
func main() {
|
||||
// Load configuration
|
||||
cfgMgr := config.NewManager()
|
||||
config.SetConfigManager(cfgMgr)
|
||||
if err := cfgMgr.Load(); err != nil {
|
||||
log.Fatalf("Failed to load configuration: %v", err)
|
||||
}
|
||||
@@ -38,14 +40,15 @@ func main() {
|
||||
logger.UpdateLoggerPath(cfg.Logger.Path, cfg.Logger.Dev)
|
||||
}
|
||||
logger.Info("ResolveSpec test server starting")
|
||||
logger.Info("Configuration loaded - Server will listen on: %s", cfg.Server.Addr)
|
||||
|
||||
// Initialize database
|
||||
db, err := initDB(cfg)
|
||||
// Initialize database manager
|
||||
ctx := context.Background()
|
||||
dbMgr, db, err := initDB(ctx, cfg)
|
||||
if err != nil {
|
||||
logger.Error("Failed to initialize database: %+v", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer dbMgr.Close()
|
||||
|
||||
// Create router
|
||||
r := mux.NewRouter()
|
||||
@@ -70,54 +73,37 @@ func main() {
|
||||
// Create server manager
|
||||
mgr := server.NewManager()
|
||||
|
||||
// Parse host and port from addr
|
||||
host := ""
|
||||
port := 8080
|
||||
if cfg.Server.Addr != "" {
|
||||
// Parse addr (format: ":8080" or "localhost:8080")
|
||||
if cfg.Server.Addr[0] == ':' {
|
||||
// Just port
|
||||
_, err := fmt.Sscanf(cfg.Server.Addr, ":%d", &port)
|
||||
if err != nil {
|
||||
logger.Error("Invalid server address: %s", cfg.Server.Addr)
|
||||
os.Exit(1)
|
||||
}
|
||||
} else {
|
||||
// Host and port
|
||||
_, err := fmt.Sscanf(cfg.Server.Addr, "%[^:]:%d", &host, &port)
|
||||
if err != nil {
|
||||
logger.Error("Invalid server address: %s", cfg.Server.Addr)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
// Get default server configuration
|
||||
defaultServerCfg, err := cfg.Servers.GetDefault()
|
||||
if err != nil {
|
||||
logger.Error("Failed to get default server config: %v", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
// Add server instance
|
||||
_, err = mgr.Add(server.Config{
|
||||
Name: "api",
|
||||
Host: host,
|
||||
Port: port,
|
||||
Handler: r,
|
||||
ShutdownTimeout: cfg.Server.ShutdownTimeout,
|
||||
DrainTimeout: cfg.Server.DrainTimeout,
|
||||
ReadTimeout: cfg.Server.ReadTimeout,
|
||||
WriteTimeout: cfg.Server.WriteTimeout,
|
||||
IdleTimeout: cfg.Server.IdleTimeout,
|
||||
})
|
||||
// Apply global defaults
|
||||
defaultServerCfg.ApplyGlobalDefaults(cfg.Servers)
|
||||
|
||||
// Convert to server.Config and add instance
|
||||
serverCfg := server.FromConfigInstanceToServerConfig(defaultServerCfg, r)
|
||||
|
||||
logger.Info("Configuration loaded - Server '%s' will listen on %s:%d",
|
||||
serverCfg.Name, serverCfg.Host, serverCfg.Port)
|
||||
|
||||
_, err = mgr.Add(serverCfg)
|
||||
if err != nil {
|
||||
logger.Error("Failed to add server: %v", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
// Start server with graceful shutdown
|
||||
logger.Info("Starting server on %s", cfg.Server.Addr)
|
||||
logger.Info("Starting server '%s' on %s:%d", serverCfg.Name, serverCfg.Host, serverCfg.Port)
|
||||
if err := mgr.ServeWithGracefulShutdown(); err != nil {
|
||||
logger.Error("Server failed: %v", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
func initDB(cfg *config.Config) (*gorm.DB, error) {
|
||||
func initDB(ctx context.Context, cfg *config.Config) (dbmanager.Manager, *gorm.DB, error) {
|
||||
// Configure GORM logger based on config
|
||||
logLevel := gormlog.Info
|
||||
if !cfg.Logger.Dev {
|
||||
@@ -135,25 +121,41 @@ func initDB(cfg *config.Config) (*gorm.DB, error) {
|
||||
},
|
||||
)
|
||||
|
||||
// Use database URL from config if available, otherwise use default SQLite
|
||||
dbURL := cfg.Database.URL
|
||||
if dbURL == "" {
|
||||
dbURL = "test.db"
|
||||
// Create database manager from config
|
||||
mgr, err := dbmanager.NewManager(dbmanager.FromConfig(cfg.DBManager))
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to create database manager: %w", err)
|
||||
}
|
||||
|
||||
// Create SQLite database
|
||||
db, err := gorm.Open(sqlite.Open(dbURL), &gorm.Config{Logger: newLogger, FullSaveAssociations: false})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
// Connect all databases
|
||||
if err := mgr.Connect(ctx); err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to connect databases: %w", err)
|
||||
}
|
||||
|
||||
// Get default connection
|
||||
conn, err := mgr.GetDefault()
|
||||
if err != nil {
|
||||
mgr.Close()
|
||||
return nil, nil, fmt.Errorf("failed to get default connection: %w", err)
|
||||
}
|
||||
|
||||
// Get GORM database
|
||||
gormDB, err := conn.GORM()
|
||||
if err != nil {
|
||||
mgr.Close()
|
||||
return nil, nil, fmt.Errorf("failed to get GORM database: %w", err)
|
||||
}
|
||||
|
||||
// Update GORM logger
|
||||
gormDB.Logger = newLogger
|
||||
|
||||
modelList := testmodels.GetTestModels()
|
||||
|
||||
// Auto migrate schemas
|
||||
err = db.AutoMigrate(modelList...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
if err := gormDB.AutoMigrate(modelList...); err != nil {
|
||||
mgr.Close()
|
||||
return nil, nil, fmt.Errorf("failed to auto migrate: %w", err)
|
||||
}
|
||||
|
||||
return db, nil
|
||||
return mgr, gormDB, nil
|
||||
}
|
||||
|
||||
+54
-7
@@ -1,17 +1,26 @@
|
||||
# ResolveSpec Test Server Configuration
|
||||
# This is a minimal configuration for the test server
|
||||
|
||||
server:
|
||||
addr: ":8080"
|
||||
servers:
|
||||
default_server: "main"
|
||||
shutdown_timeout: 30s
|
||||
drain_timeout: 25s
|
||||
read_timeout: 10s
|
||||
write_timeout: 10s
|
||||
idle_timeout: 120s
|
||||
instances:
|
||||
main:
|
||||
name: "main"
|
||||
host: "localhost"
|
||||
port: 8080
|
||||
description: "Main server instance"
|
||||
gzip: true
|
||||
tags:
|
||||
env: "test"
|
||||
|
||||
logger:
|
||||
dev: true # Enable development mode for readable logs
|
||||
path: "" # Empty means log to stdout
|
||||
dev: true
|
||||
path: ""
|
||||
|
||||
cache:
|
||||
provider: "memory"
|
||||
@@ -19,7 +28,7 @@ cache:
|
||||
middleware:
|
||||
rate_limit_rps: 100.0
|
||||
rate_limit_burst: 200
|
||||
max_request_size: 10485760 # 10MB
|
||||
max_request_size: 10485760
|
||||
|
||||
cors:
|
||||
allowed_origins:
|
||||
@@ -36,6 +45,44 @@ cors:
|
||||
|
||||
tracing:
|
||||
enabled: false
|
||||
service_name: "resolvespec"
|
||||
service_version: "1.0.0"
|
||||
endpoint: ""
|
||||
|
||||
database:
|
||||
url: "" # Empty means use default SQLite (test.db)
|
||||
error_tracking:
|
||||
enabled: false
|
||||
provider: "noop"
|
||||
environment: "development"
|
||||
sample_rate: 1.0
|
||||
traces_sample_rate: 0.1
|
||||
|
||||
event_broker:
|
||||
enabled: false
|
||||
provider: "memory"
|
||||
mode: "sync"
|
||||
worker_count: 1
|
||||
buffer_size: 100
|
||||
instance_id: ""
|
||||
|
||||
dbmanager:
|
||||
default_connection: "primary"
|
||||
max_open_conns: 25
|
||||
max_idle_conns: 5
|
||||
conn_max_lifetime: 30m
|
||||
conn_max_idle_time: 5m
|
||||
retry_attempts: 3
|
||||
retry_delay: 1s
|
||||
health_check_interval: 30s
|
||||
enable_auto_reconnect: true
|
||||
connections:
|
||||
primary:
|
||||
name: "primary"
|
||||
type: "sqlite"
|
||||
filepath: "test.db"
|
||||
default_orm: "gorm"
|
||||
enable_logging: true
|
||||
enable_metrics: false
|
||||
connect_timeout: 10s
|
||||
query_timeout: 30s
|
||||
|
||||
paths: {}
|
||||
|
||||
+81
-10
@@ -2,29 +2,38 @@
|
||||
# This file demonstrates all available configuration options
|
||||
# Copy this file to config.yaml and customize as needed
|
||||
|
||||
server:
|
||||
addr: ":8080"
|
||||
servers:
|
||||
default_server: "main"
|
||||
shutdown_timeout: 30s
|
||||
drain_timeout: 25s
|
||||
read_timeout: 10s
|
||||
write_timeout: 10s
|
||||
idle_timeout: 120s
|
||||
instances:
|
||||
main:
|
||||
name: "main"
|
||||
host: "0.0.0.0"
|
||||
port: 8080
|
||||
description: "Main API server"
|
||||
gzip: true
|
||||
tags:
|
||||
env: "development"
|
||||
version: "1.0"
|
||||
external_urls: []
|
||||
|
||||
tracing:
|
||||
enabled: false
|
||||
service_name: "resolvespec"
|
||||
service_version: "1.0.0"
|
||||
endpoint: "http://localhost:4318/v1/traces" # OTLP endpoint
|
||||
endpoint: "http://localhost:4318/v1/traces"
|
||||
|
||||
cache:
|
||||
provider: "memory" # Options: memory, redis, memcache
|
||||
|
||||
provider: "memory"
|
||||
redis:
|
||||
host: "localhost"
|
||||
port: 6379
|
||||
password: ""
|
||||
db: 0
|
||||
|
||||
memcache:
|
||||
servers:
|
||||
- "localhost:11211"
|
||||
@@ -33,12 +42,12 @@ cache:
|
||||
|
||||
logger:
|
||||
dev: false
|
||||
path: "" # Empty for stdout, or specify file path
|
||||
path: ""
|
||||
|
||||
middleware:
|
||||
rate_limit_rps: 100.0
|
||||
rate_limit_burst: 200
|
||||
max_request_size: 10485760 # 10MB in bytes
|
||||
max_request_size: 10485760
|
||||
|
||||
cors:
|
||||
allowed_origins:
|
||||
@@ -53,5 +62,67 @@ cors:
|
||||
- "*"
|
||||
max_age: 3600
|
||||
|
||||
database:
|
||||
url: "host=localhost user=postgres password=postgres dbname=resolvespec_test port=5434 sslmode=disable"
|
||||
error_tracking:
|
||||
enabled: false
|
||||
provider: "noop"
|
||||
environment: "development"
|
||||
sample_rate: 1.0
|
||||
traces_sample_rate: 0.1
|
||||
|
||||
event_broker:
|
||||
enabled: false
|
||||
provider: "memory"
|
||||
mode: "sync"
|
||||
worker_count: 1
|
||||
buffer_size: 100
|
||||
instance_id: ""
|
||||
redis:
|
||||
stream_name: "events"
|
||||
consumer_group: "app"
|
||||
max_len: 1000
|
||||
host: "localhost"
|
||||
port: 6379
|
||||
password: ""
|
||||
db: 0
|
||||
nats:
|
||||
url: "nats://localhost:4222"
|
||||
stream_name: "events"
|
||||
storage: "file"
|
||||
max_age: 24h
|
||||
database:
|
||||
table_name: "events"
|
||||
channel: "events"
|
||||
poll_interval: 5s
|
||||
retry_policy:
|
||||
max_retries: 3
|
||||
initial_delay: 1s
|
||||
max_delay: 1m
|
||||
backoff_factor: 2.0
|
||||
|
||||
dbmanager:
|
||||
default_connection: "primary"
|
||||
max_open_conns: 25
|
||||
max_idle_conns: 5
|
||||
conn_max_lifetime: 30m
|
||||
conn_max_idle_time: 5m
|
||||
retry_attempts: 3
|
||||
retry_delay: 1s
|
||||
health_check_interval: 30s
|
||||
enable_auto_reconnect: true
|
||||
connections:
|
||||
primary:
|
||||
name: "primary"
|
||||
type: "pgsql"
|
||||
url: "host=localhost user=postgres password=postgres dbname=resolvespec port=5432 sslmode=disable"
|
||||
default_orm: "gorm"
|
||||
enable_logging: false
|
||||
enable_metrics: false
|
||||
connect_timeout: 10s
|
||||
query_timeout: 30s
|
||||
|
||||
paths:
|
||||
data_dir: "./data"
|
||||
log_dir: "./logs"
|
||||
cache_dir: "./cache"
|
||||
|
||||
extensions: {}
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 352 KiB After Width: | Height: | Size: 95 KiB |
@@ -1,44 +1,53 @@
|
||||
module github.com/bitechdev/ResolveSpec
|
||||
|
||||
go 1.24.0
|
||||
|
||||
toolchain go1.24.6
|
||||
go 1.25.7
|
||||
|
||||
require (
|
||||
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/getsentry/sentry-go v0.40.0
|
||||
github.com/getsentry/sentry-go v0.46.2
|
||||
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/gorilla/mux v1.8.1
|
||||
github.com/gorilla/websocket v1.5.3
|
||||
github.com/jackc/pgx/v5 v5.6.0
|
||||
github.com/klauspost/compress v1.18.0
|
||||
github.com/jackc/pgx/v5 v5.9.2
|
||||
github.com/klauspost/compress v1.18.6
|
||||
github.com/mark3labs/mcp-go v0.54.0
|
||||
github.com/mattn/go-sqlite3 v1.14.44
|
||||
github.com/microsoft/go-mssqldb v1.10.0
|
||||
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/redis/go-redis/v9 v9.17.1
|
||||
github.com/redis/go-redis/v9 v9.19.0
|
||||
github.com/spf13/viper v1.21.0
|
||||
github.com/stretchr/testify v1.11.1
|
||||
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/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/pgdialect v1.2.16
|
||||
github.com/uptrace/bun/dialect/sqlitedialect v1.2.16
|
||||
github.com/uptrace/bun/driver/sqliteshim v1.2.16
|
||||
github.com/uptrace/bunrouter v1.0.23
|
||||
go.opentelemetry.io/otel v1.38.0
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.38.0
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.38.0
|
||||
go.opentelemetry.io/otel/sdk v1.38.0
|
||||
go.opentelemetry.io/otel/trace v1.38.0
|
||||
go.uber.org/zap v1.27.0
|
||||
golang.org/x/crypto v0.43.0
|
||||
golang.org/x/time v0.14.0
|
||||
go.mongodb.org/mongo-driver v1.17.9
|
||||
go.opentelemetry.io/otel v1.44.0
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.43.0
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.43.0
|
||||
go.opentelemetry.io/otel/sdk v1.44.0
|
||||
go.opentelemetry.io/otel/trace v1.44.0
|
||||
go.uber.org/zap v1.28.0
|
||||
golang.org/x/crypto v0.55.0
|
||||
golang.org/x/oauth2 v0.36.0
|
||||
golang.org/x/sys v0.47.0
|
||||
golang.org/x/time v0.15.0
|
||||
google.golang.org/grpc v1.83.2
|
||||
gorm.io/driver/postgres v1.6.0
|
||||
gorm.io/driver/sqlite v1.6.0
|
||||
gorm.io/gorm v1.30.0
|
||||
gorm.io/driver/sqlserver v1.6.3
|
||||
gorm.io/gorm v1.31.1
|
||||
)
|
||||
|
||||
require (
|
||||
@@ -54,8 +63,7 @@ require (
|
||||
github.com/containerd/log v0.1.0 // indirect
|
||||
github.com/containerd/platforms v0.2.1 // indirect
|
||||
github.com/cpuguy83/dockercfg v0.3.2 // indirect
|
||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
|
||||
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
|
||||
github.com/distribution/reference v0.6.0 // indirect
|
||||
github.com/docker/docker v28.5.1+incompatible // indirect
|
||||
github.com/docker/go-connections v0.6.0 // indirect
|
||||
@@ -63,13 +71,17 @@ require (
|
||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||
github.com/ebitengine/purego v0.8.4 // indirect
|
||||
github.com/felixge/httpsnoop v1.0.4 // indirect
|
||||
github.com/fsnotify/fsnotify v1.9.0 // indirect
|
||||
github.com/glebarez/go-sqlite v1.21.2 // indirect
|
||||
github.com/fsnotify/fsnotify v1.10.1 // indirect
|
||||
github.com/glebarez/go-sqlite v1.22.0 // indirect
|
||||
github.com/go-logr/logr v1.4.3 // indirect
|
||||
github.com/go-logr/stdr v1.2.2 // indirect
|
||||
github.com/go-ole/go-ole v1.2.6 // indirect
|
||||
github.com/go-viper/mapstructure/v2 v2.4.0 // indirect
|
||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.2 // 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/sqlexp v0.1.0 // indirect
|
||||
github.com/golang/snappy v1.0.0 // indirect
|
||||
github.com/google/jsonschema-go v0.4.3 // indirect
|
||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 // indirect
|
||||
github.com/jackc/pgpassfile v1.0.0 // indirect
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
||||
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
||||
@@ -77,8 +89,7 @@ require (
|
||||
github.com/jinzhu/now v1.1.5 // indirect
|
||||
github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0 // indirect
|
||||
github.com/magiconair/properties v1.8.10 // indirect
|
||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||
github.com/mattn/go-sqlite3 v1.14.32 // indirect
|
||||
github.com/mattn/go-isatty v0.0.22 // indirect
|
||||
github.com/moby/docker-image-spec v1.3.1 // indirect
|
||||
github.com/moby/go-archive v0.1.0 // indirect
|
||||
github.com/moby/patternmatcher v0.6.0 // indirect
|
||||
@@ -86,61 +97,67 @@ require (
|
||||
github.com/moby/sys/user v0.4.0 // indirect
|
||||
github.com/moby/sys/userns v0.1.0 // indirect
|
||||
github.com/moby/term v0.5.0 // indirect
|
||||
github.com/montanaflynn/stats v0.9.0 // indirect
|
||||
github.com/morikuni/aec v1.0.0 // 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/ncruces/go-strftime v1.0.0 // indirect
|
||||
github.com/opencontainers/go-digest v1.0.0 // 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/pmezard/go-difflib v1.0.0 // indirect
|
||||
github.com/power-devops/perfstat v0.0.0-20210106213030-5aafc221ea8c // indirect
|
||||
github.com/prometheus/client_model v0.6.2 // indirect
|
||||
github.com/prometheus/common v0.66.1 // indirect
|
||||
github.com/prometheus/procfs v0.16.1 // indirect
|
||||
github.com/prometheus/common v0.67.5 // indirect
|
||||
github.com/prometheus/procfs v0.20.1 // indirect
|
||||
github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||
github.com/rs/xid v1.4.0 // indirect
|
||||
github.com/sagikazarmark/locafero v0.11.0 // indirect
|
||||
github.com/rs/xid v1.6.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/shopspring/decimal v1.4.0 // indirect
|
||||
github.com/sirupsen/logrus v1.9.3 // indirect
|
||||
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 // indirect
|
||||
github.com/spf13/afero v1.15.0 // indirect
|
||||
github.com/spf13/cast v1.10.0 // indirect
|
||||
github.com/spf13/pflag v1.0.10 // indirect
|
||||
github.com/stretchr/objx v0.5.2 // indirect
|
||||
github.com/subosito/gotenv v1.6.0 // indirect
|
||||
github.com/tidwall/match v1.1.1 // indirect
|
||||
github.com/tidwall/pretty v1.2.0 // indirect
|
||||
github.com/tidwall/match v1.2.0 // indirect
|
||||
github.com/tidwall/pretty v1.2.1 // indirect
|
||||
github.com/tklauser/go-sysconf v0.3.12 // indirect
|
||||
github.com/tklauser/numcpus v0.6.1 // indirect
|
||||
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc // indirect
|
||||
github.com/vmihailenco/msgpack/v5 v5.4.1 // indirect
|
||||
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
|
||||
github.com/xdg-go/pbkdf2 v1.0.0 // indirect
|
||||
github.com/xdg-go/scram v1.2.0 // indirect
|
||||
github.com/xdg-go/stringprep v1.0.4 // indirect
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2 // indirect
|
||||
github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 // indirect
|
||||
github.com/yusufpapurcu/wmi v1.2.4 // indirect
|
||||
go.opentelemetry.io/auto/sdk v1.1.0 // indirect
|
||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 // indirect
|
||||
go.opentelemetry.io/otel/metric v1.38.0 // indirect
|
||||
go.opentelemetry.io/proto/otlp v1.7.1 // indirect
|
||||
go.uber.org/multierr v1.10.0 // indirect
|
||||
go.yaml.in/yaml/v2 v2.4.2 // indirect
|
||||
go.opentelemetry.io/auto/sdk v1.2.1 // indirect
|
||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0 // indirect
|
||||
go.opentelemetry.io/otel/metric v1.44.0 // 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.yaml.in/yaml/v2 v2.4.4 // indirect
|
||||
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
||||
golang.org/x/exp v0.0.0-20251113190631-e25ba8c21ef6 // indirect
|
||||
golang.org/x/net v0.45.0 // indirect
|
||||
golang.org/x/sync v0.18.0 // indirect
|
||||
golang.org/x/sys v0.38.0 // indirect
|
||||
golang.org/x/text v0.30.0 // 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.8 // indirect
|
||||
golang.org/x/mod v0.38.0 // indirect
|
||||
golang.org/x/net v0.58.0 // indirect
|
||||
golang.org/x/sync v0.22.0 // indirect
|
||||
golang.org/x/text v0.41.0 // indirect
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||
google.golang.org/protobuf v1.36.11 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
modernc.org/libc v1.67.0 // indirect
|
||||
modernc.org/libc v1.72.3 // indirect
|
||||
modernc.org/mathutil v1.7.1 // indirect
|
||||
modernc.org/memory v1.11.0 // indirect
|
||||
modernc.org/sqlite v1.40.1 // indirect
|
||||
modernc.org/sqlite v1.50.1 // indirect
|
||||
)
|
||||
|
||||
replace github.com/uptrace/bun => github.com/warkanum/bun v1.2.17
|
||||
|
||||
@@ -2,16 +2,40 @@ dario.cat/mergo v1.0.2 h1:85+piFYR1tMbRrLcDwR18y4UKJ3aH1Tbzi24VRW1TK8=
|
||||
dario.cat/mergo v1.0.2/go.mod h1:E/hbnu0NxMFBjpMIE34DRGLWqDy0g5FuKDhCb31ngxA=
|
||||
github.com/AdaLogics/go-fuzz-headers v0.0.0-20240806141605-e8a1dd7889d6 h1:He8afgbRMd7mFxO99hRNu+6tazq8nFF9lIwo9JFroBk=
|
||||
github.com/AdaLogics/go-fuzz-headers v0.0.0-20240806141605-e8a1dd7889d6/go.mod h1:8o94RPi1/7XTJvwPpRSzSUedZrtlirdB3r9Z20bi2f8=
|
||||
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.11.1/go.mod h1:a6xsAQUZg+VsS3TJ05SRp524Hs4pZ/AeFSr5ENf0Yjo=
|
||||
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.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.6.0/go.mod h1:9kIvujWAA58nmPmWB1m23fyWic1kYZMxD9CxaWn4Qpg=
|
||||
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.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.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.12.0 h1:fhqpLE3UEXi9lPaBRpQ6XuRW0nU7hgg4zlmZZa+a9q4=
|
||||
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.4.0 h1:E4MgwLBGeVB5f2MdcIVD3ELVAWpr+WD6MUe1i+tM/PA=
|
||||
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.2.0 h1:nCYfgcSyHZXJI8J0IWE5MsCGlb2xp9fJiXyxWgmOFg4=
|
||||
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/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.2.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.6.0 h1:XRzhVemXdgvJqCH0sFfrBUTnUJSBrBf7++ypk+twtRs=
|
||||
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/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU=
|
||||
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
|
||||
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/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
||||
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf h1:TqhNAT4zKbTdLa62d2HDBFdvgSbIGB3eJE8HqhgiL9I=
|
||||
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf/go.mod h1:r5xuitiExdLAJ09PR7vBVENGvp4ZuTBeWTGtxuX3K+c=
|
||||
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c h1:6Gpm9YYUEQx2T9zMsYolQhr6sjwwGtFitSA0pQsa7a8=
|
||||
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/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
|
||||
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
|
||||
@@ -32,15 +56,19 @@ github.com/containerd/platforms v0.2.1 h1:zvwtM3rz2YHPQsF2CHYM8+KtB5dvhISiXh5ZpS
|
||||
github.com/containerd/platforms v0.2.1/go.mod h1:XHCb+2/hzowdiut9rkudds9bE5yJ7npe7dG/wG+uFPw=
|
||||
github.com/cpuguy83/dockercfg v0.3.2 h1:DlJTyZGBDlXqUZ2Dk2Q3xHs/FtnooJJVaad2S9GKorA=
|
||||
github.com/cpuguy83/dockercfg v0.3.2/go.mod h1:sugsbF4//dDlL/i+S+rtpIWp+5h0BHJHfjj5/jFyUJc=
|
||||
github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
|
||||
github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY=
|
||||
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.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||
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/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc=
|
||||
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM=
|
||||
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/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.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/go.mod h1:eEKB0N0r5NX/I1kEveEz05bcu8tLC/8azJZsviup8Sk=
|
||||
github.com/docker/go-connections v0.6.0 h1:LlMG9azAe1TqfR7sO+NJttz1gy6KO7VJBh+pMmjSD94=
|
||||
@@ -57,12 +85,12 @@ github.com/felixge/httpsnoop v1.0.4 h1:NFTV2Zj1bL4mc9sqWACXbQFVBBg2W3GPvqp8/ESS2
|
||||
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/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
||||
github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k=
|
||||
github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0=
|
||||
github.com/getsentry/sentry-go v0.40.0 h1:VTJMN9zbTvqDqPwheRVLcp0qcUcM+8eFivvGocAaSbo=
|
||||
github.com/getsentry/sentry-go v0.40.0/go.mod h1:eRXCoh3uvmjQLY6qu63BjUZnaBu5L5WhMV1RwYO8W5s=
|
||||
github.com/glebarez/go-sqlite v1.21.2 h1:3a6LFC4sKahUunAmynQKLZceZCOzUthkRkEAl9gAXWo=
|
||||
github.com/glebarez/go-sqlite v1.21.2/go.mod h1:sfxdZyhQjTM2Wry3gVYWaW072Ri1WMdWJi0k6+3382k=
|
||||
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
|
||||
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
|
||||
github.com/getsentry/sentry-go v0.46.2 h1:1jhYwrKGa3sIpo/y5iDNXS5wDoT7I1KNzMHrnK6ojns=
|
||||
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/go.mod h1:PlBIdHe0+aUEFn+r2/uthrWq4FxbzugL0L8Li6yQJbc=
|
||||
github.com/glebarez/sqlite v1.11.0 h1:wSG0irqzP6VurnMEpFGer5Li19RpIRi2qvQz++w0GMw=
|
||||
github.com/glebarez/sqlite v1.11.0/go.mod h1:h8/o8j5wiAsqSPoWELDUdJXhjAhsVliSn7bWZjOhrgQ=
|
||||
github.com/go-errors/errors v1.4.2 h1:J6MZopCL4uSllY1OfXM374weqZFFItUbrImctkmUxIA=
|
||||
@@ -74,33 +102,58 @@ 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-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-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs=
|
||||
github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
|
||||
github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro=
|
||||
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.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.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
|
||||
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/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0=
|
||||
github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A=
|
||||
github.com/golang-sql/sqlexp v0.1.0/go.mod h1:J4ad9Vo8ZCWQ2GMrC4UCQy1JpCbwU9m3EOqtpKwwwHI=
|
||||
github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
|
||||
github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps=
|
||||
github.com/golang/snappy v1.0.0 h1:Oy607GVXHs7RtbggtPBnr2RmDArIsAefDwvrdWvRhGs=
|
||||
github.com/golang/snappy v1.0.0/go.mod h1:/XxbfmMg8lxefKM7IXC3fBNl/7bRcc72aCRzEWrmP2Q=
|
||||
github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||
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/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
||||
github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0=
|
||||
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/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
|
||||
github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||
github.com/gorilla/mux v1.8.1 h1:TuBL49tXwgrFYWhqrNgrUNEY92u81SPhu7sTdzQEiWY=
|
||||
github.com/gorilla/mux v1.8.1/go.mod h1:AKf9I4AEqPTmMytcMc0KkNouC66V3BtZ4qD5fmWSiMQ=
|
||||
github.com/gorilla/securecookie v1.1.1/go.mod h1:ra0sb63/xPlUeL+yeDciTfxMRAA+MP+HVt/4epWDjd4=
|
||||
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/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.27.2/go.mod h1:pkJQ2tZHJ0aFOVEEot6oZmaVEZcRme73eIFmhiVuRWs=
|
||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 h1:5VipnvEpbqr2gA2VbM+nYVbkIF28c5ZQfqCBQ5g2xfk=
|
||||
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.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/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
|
||||
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
|
||||
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/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
||||
github.com/jackc/pgx/v5 v5.6.0 h1:SWJzexBzPL5jb0GEsrPMLIsi/3jOo7RHlzTjcAeDrPY=
|
||||
github.com/jackc/pgx/v5 v5.6.0/go.mod h1:DNZ/vlrUnhWCoFGxHAG8U2ljioxukquj7utPDgtQdTw=
|
||||
github.com/jackc/pgx/v5 v5.9.2 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw=
|
||||
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/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
||||
github.com/jcmturner/aescts/v2 v2.0.0/go.mod h1:AiaICIRyfYg35RUkr8yESTqvSy7csK90qZ5xfvvsoNs=
|
||||
github.com/jcmturner/dnsutils/v2 v2.0.0/go.mod h1:b0TnjGOvI/n42bZa+hmXL+kFJZsFT7G4t3HTlQ184QM=
|
||||
github.com/jcmturner/gofork v1.7.6/go.mod h1:1622LH6i/EZqLloHfE7IeZ0uEJwMSUyQ/nDd82IeqRo=
|
||||
github.com/jcmturner/goidentity/v6 v6.0.1/go.mod h1:X1YW3bgtvwAXju7V3LCIMpY0Gbxyjn/mY9zx4tFonSg=
|
||||
github.com/jcmturner/gokrb5/v8 v8.4.4/go.mod h1:1btQEpgT6k+unzCwX1KdWMEwPPkkgBtP+F6aCACiMrs=
|
||||
github.com/jcmturner/rpc/v2 v2.0.3/go.mod h1:VUJYCIDm3PVOEHw8sgt091/20OJjskO/YJki3ELg/Hc=
|
||||
github.com/jinzhu/copier v0.3.5 h1:GlvfUwHk62RokgqVNvYsku0TATCF7bAHVwEXoBh3iJg=
|
||||
github.com/jinzhu/copier v0.3.5/go.mod h1:DfbEm0FYsaqBcKcFuvmOZb218JkPGtvSHsKg8S8hyyg=
|
||||
github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E=
|
||||
@@ -108,10 +161,15 @@ 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/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
|
||||
github.com/kisielk/sqlstruct v0.0.0-20201105191214-5f3e10d3ab46/go.mod h1:yyMNCyc/Ib3bDTKd379tNMpB/7/H5TjM2Y9QJ5THLbE=
|
||||
github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo=
|
||||
github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ=
|
||||
github.com/klauspost/compress v1.18.6 h1:2jupLlAwFm95+YDR+NwD2MEfFO9d4z4Prjl1XXDjuao=
|
||||
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.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
||||
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
||||
github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ=
|
||||
github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI=
|
||||
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
|
||||
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
|
||||
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
|
||||
@@ -120,10 +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/magiconair/properties v1.8.10 h1:s31yESBquKXCV9a/ScB3ESkOjUYYv+X0rg8SYxI99mE=
|
||||
github.com/magiconair/properties v1.8.10/go.mod h1:Dhd985XPs7jluiymwWYZ0G4Z61jb3vdS329zhj2hYo0=
|
||||
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
||||
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
||||
github.com/mattn/go-sqlite3 v1.14.32 h1:JD12Ag3oLy1zQA+BNn74xRgaBbdhbNIDYvQUEuuErjs=
|
||||
github.com/mattn/go-sqlite3 v1.14.32/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y=
|
||||
github.com/mark3labs/mcp-go v0.54.0 h1:PZhQvd+5xrT43cUoiaKn/hDcvLUhcLc1twSEKYPTcTA=
|
||||
github.com/mark3labs/mcp-go v0.54.0/go.mod h1:+8WclSK1ZUweCP3hvktSji8n8ABG/95QaEkeVE/Uwas=
|
||||
github.com/mattn/go-isatty v0.0.22 h1:j8l17JJ9i6VGPUFUYoTUKPSgKe/83EYU2zBC7YNKMw4=
|
||||
github.com/mattn/go-isatty v0.0.22/go.mod h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4=
|
||||
github.com/mattn/go-sqlite3 v1.14.44 h1:3VSe+xafpbzsLbdr2AWlAZk9yRHiBhTBakioXaCKTF8=
|
||||
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.10.0 h1:pHEt+Qz6YFPWqREq10mqSE524QQo+/QremwTCQht7TY=
|
||||
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/go.mod h1:eKmb5VW8vQEh/BAr2yvVNvuiJuY6UIocYsFu/DxxRpo=
|
||||
github.com/moby/go-archive v0.1.0 h1:Kk/5rdW/g+H8NHdJW2gsXyZ7UnzvJNOy6VKJqueWdcQ=
|
||||
@@ -142,14 +205,18 @@ github.com/moby/term v0.5.0 h1:xt8Q1nalod/v7BqbG21f8mQPqH+xAaC9C3N3wfWbVP0=
|
||||
github.com/moby/term v0.5.0/go.mod h1:8FzsFHVUBGZdbDsJw/ot+X+d5HLUbvklYLJ9uGfcI3Y=
|
||||
github.com/mochi-mqtt/server/v2 v2.7.9 h1:y0g4vrSLAag7T07l2oCzOa/+nKVLoazKEWAArwqBNYI=
|
||||
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/montanaflynn/stats v0.7.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
||||
github.com/montanaflynn/stats v0.9.0 h1:tsBJ0RXwph9BmAuFoCmqGv6e8xa0MENQ8m0ptKq29mQ=
|
||||
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/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/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.48.0/go.mod h1:iRWIPokVIFbVijxuMQq4y9ttaBTMe0SFdlZfMDd+33g=
|
||||
github.com/nats-io/nkeys v0.4.11 h1:q44qGV008kYd9W1b1nEBkNzvnWxtRSQ7A8BoqRrcfa0=
|
||||
github.com/nats-io/nkeys v0.4.11/go.mod h1:szDimtgmfOi9n25JpfIdGw12tZFYXqhGxjhVxsatHVE=
|
||||
github.com/nats-io/nats.go v1.52.0 h1:n3avV4VBsCgsdwh71TppsTwtv+QdPs7ntSKM8qJLGsc=
|
||||
github.com/nats-io/nats.go v1.52.0/go.mod h1:26HypzazeOkyO3/mqd1zZd53STJN0EjCYF9Uy2ZOBno=
|
||||
github.com/nats-io/nkeys v0.4.15 h1:JACV5jRVO9V856KOapQ7x+EY8Jo3qw1vJt/9Jpwzkk4=
|
||||
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/go.mod h1:19wcPz3Ph3q0Jbyiqsd0kePYG7A95tJPxeL+1OSON2c=
|
||||
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
|
||||
@@ -158,10 +225,14 @@ 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/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040=
|
||||
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.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
|
||||
github.com/pelletier/go-toml/v2 v2.3.1 h1:MYEvvGnQjeNkRF1qUuGolNtNExTDwct51yp7olPtrEc=
|
||||
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/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-20240102092130-5ac0b6a4141c h1:+mdjkGKdHQG3305AYmdv1U2eRNDiU2ErMBj1gwrq8eQ=
|
||||
github.com/pkg/browser v0.0.0-20240102092130-5ac0b6a4141c/go.mod h1:7rwL4CYBLnjLxUqIJNnCWiEdr3bn6IUYi15bNlnbCCU=
|
||||
github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA=
|
||||
github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
|
||||
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
@@ -172,28 +243,32 @@ 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_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
||||
github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
|
||||
github.com/prometheus/common v0.66.1 h1:h5E0h5/Y8niHc5DlaLlWLArTQI7tMrsfQjHV+d9ZoGs=
|
||||
github.com/prometheus/common v0.66.1/go.mod h1:gcaUsgf3KfRSwHY4dIMXLPV0K/Wg1oZ8+SbZk/HH/dA=
|
||||
github.com/prometheus/procfs v0.16.1 h1:hZ15bTNuirocR6u0JZ6BAHHmwS1p8B4P6MRqxtzMyRg=
|
||||
github.com/prometheus/procfs v0.16.1/go.mod h1:teAbpZRB1iIAJYREa1LsoWUXykVXA1KlTmWl8x/U+Is=
|
||||
github.com/prometheus/common v0.67.5 h1:pIgK94WWlQt1WLwAC5j2ynLaBRDiinoAb86HZHTUGI4=
|
||||
github.com/prometheus/common v0.67.5/go.mod h1:SjE/0MzDEEAyrdr5Gqc6G+sXI67maCxzaT3A2+HqjUw=
|
||||
github.com/prometheus/procfs v0.20.1 h1:XwbrGOIplXW/AU3YhIhLODXMJYyC1isLFfYCsTEycfc=
|
||||
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/go.mod h1:VjzYrABPabuM4KyBh1Ftq6u8nhwY5tBPKP9jpmh0nnA=
|
||||
github.com/redis/go-redis/v9 v9.17.1 h1:7tl732FjYPRT9H9aNfyTwKg9iTETjWjGKEJ2t/5iWTs=
|
||||
github.com/redis/go-redis/v9 v9.17.1/go.mod h1:u410H11HMLoB+TP67dz8rL9s6QW2j76l0//kSOd3370=
|
||||
github.com/redis/go-redis/v9 v9.19.0 h1:XPVaaPSnG6RhYf7p+rmSa9zZfeVAnWsH5h3lxthOm/k=
|
||||
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/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
||||
github.com/rogpeppe/go-internal v1.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII=
|
||||
github.com/rogpeppe/go-internal v1.13.1/go.mod h1:uMEvuHeurkdAXX61udpOXGD/AzZDWNMNyH2VO9fmH0o=
|
||||
github.com/rs/xid v1.4.0 h1:qd7wPTDkN6KQx2VmMBLrpHkiyQwgFXRnkOLacUiaSNY=
|
||||
github.com/rs/xid v1.4.0/go.mod h1:trrq9SKmegXys3aeAKXMUTdJsYXVwGY3RLcfgqegfbg=
|
||||
github.com/sagikazarmark/locafero v0.11.0 h1:1iurJgmM9G3PA/I+wWYIOw/5SyBtxapeHDcg+AAIFXc=
|
||||
github.com/sagikazarmark/locafero v0.11.0/go.mod h1:nVIGvgyzw595SUSUE6tvCp3YYTeHs15MvlmU87WwIik=
|
||||
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.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
|
||||
github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc=
|
||||
github.com/rs/xid v1.6.0 h1:fV591PaemRlL6JfRxGDEPl69wICngIQ3shQtzfy2gxU=
|
||||
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/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/go.mod h1:PfybzyydfZcN+JMMjkF6Zb8Mq1A/VcogFFg7hj50W9c=
|
||||
github.com/shopspring/decimal v1.4.0 h1:bxl37RwXBklmTi0C79JfXCEBD1cqqHt0bbgBAGFp81k=
|
||||
github.com/shopspring/decimal v1.4.0/go.mod h1:gawqmDU56v4yIKSwfBSFip1HdCCXN8/+DMd9qYNcwME=
|
||||
github.com/sirupsen/logrus v1.9.3 h1:dueUQJ1C2q9oE3F7wvmSGAaVtTmUizReu6fjN8uqzbQ=
|
||||
github.com/sirupsen/logrus v1.9.3/go.mod h1:naHLuLoDiP4jHNo9R0sCBMtWGeIprob74mVsIT4qYEQ=
|
||||
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 h1:+jumHNA0Wrelhe64i8F6HNlS8pkoyMv5sreGx2Ry5Rw=
|
||||
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8/go.mod h1:3n1Cwaq1E1/1lhQhtRK2ts/ZwZEhjcQeJQ1RuC6Q/8U=
|
||||
github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I=
|
||||
github.com/spf13/afero v1.15.0/go.mod h1:NC2ByUVxtQs4b3sIUphxK0NioZnmxgyCrfzeuq8lxMg=
|
||||
github.com/spf13/cast v1.10.0 h1:h2x0u2shc1QuLHfxi+cTJvs30+ZAHOGRic8uyGTDWxY=
|
||||
@@ -203,10 +278,18 @@ github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3A
|
||||
github.com/spf13/viper v1.21.0 h1:x5S+0EU27Lbphp4UKm1C+1oQO+rKx36vfCoaVebLFSU=
|
||||
github.com/spf13/viper v1.21.0/go.mod h1:P0lhsswPGWD/1lZJ9ny3fYnVqxiegrlNrEmgLjbTCAY=
|
||||
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
||||
github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw=
|
||||
github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo=
|
||||
github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY=
|
||||
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
|
||||
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
||||
github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4=
|
||||
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
||||
github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
||||
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
|
||||
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
|
||||
github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo=
|
||||
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
|
||||
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||
github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8=
|
||||
@@ -214,12 +297,14 @@ 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/go.mod h1:FSXV5KQtX2HAMlm7U3APNyLkkap35zNLxukw9oBi/MY=
|
||||
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.18.0/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
||||
github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA=
|
||||
github.com/tidwall/gjson v1.19.0 h1:xwxm7n691Uf3u5OFjzngavjGTh55KX5q/9w9xHW88JU=
|
||||
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/pretty v1.2.0 h1:RWIZEg2iJ8/g6fDDYzMpobmaoGh5OLl4AXtGUGPcqCs=
|
||||
github.com/tidwall/match v1.2.0 h1:0pt8FlkOwjN2fPt4bIl4BoNxb98gGHN2ObFEDkrfZnM=
|
||||
github.com/tidwall/match v1.2.0/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
||||
github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU=
|
||||
github.com/tidwall/pretty v1.2.1 h1:qjsOFOWWQl+N3RsoF5/ssm1pHmJJwhjlSbZ51I6wMl4=
|
||||
github.com/tidwall/pretty v1.2.1/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU=
|
||||
github.com/tidwall/sjson v1.2.5 h1:kLy8mja+1c9jlljvWTlSazM7cKDRfJuR/bOJhcY5NcY=
|
||||
github.com/tidwall/sjson v1.2.5/go.mod h1:Fvgq9kS/6ociJEDnK0Fk1cpYF4FIW6ZF7LAe+6jwd28=
|
||||
github.com/tklauser/go-sysconf v0.3.12 h1:0QaGUFOdQaIVdPgfITYzaTegZvdCjmYO52cSFAEVmqU=
|
||||
@@ -228,6 +313,10 @@ github.com/tklauser/numcpus v0.6.1 h1:ng9scYS7az0Bk4OZLvrNXNSAO2Pxr1XXRAPyjhIx+F
|
||||
github.com/tklauser/numcpus v0.6.1/go.mod h1:1XfjsgE2zo8GVw7POkMbHENHzVg3GzmoZ9fESEdAacY=
|
||||
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc h1:9lRDQMhESg+zvGYmW5DyG0UqvY96Bu5QYsTLvCHdrgo=
|
||||
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc/go.mod h1:bciPuU6GHm1iF1pBvUfxfsH0Wmnc2VbpgvbI9ZWuIRs=
|
||||
github.com/uptrace/bun/dialect/mssqldialect v1.2.16 h1:rKv0cKPNBviXadB/+2Y/UedA/c1JnwGzUWZkdN5FdSQ=
|
||||
github.com/uptrace/bun/dialect/mssqldialect v1.2.16/go.mod h1:J5U7tGKWDsx2Q7MwDZF2417jCdpD6yD/ZMFJcCR80bk=
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.16 h1:KFNZ0LxAyczKNfK/IJWMyaleO6eI9/Z5tUv3DE1NVL4=
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.16/go.mod h1:IJdMeV4sLfh0LDUZl7TIxLI0LipF1vwTK3hBC7p5qLo=
|
||||
github.com/uptrace/bun/dialect/sqlitedialect v1.2.16 h1:6wVAiYLj1pMibRthGwy4wDLa3D5AQo32Y8rvwPd8CQ0=
|
||||
github.com/uptrace/bun/dialect/sqlitedialect v1.2.16/go.mod h1:Z7+5qK8CGZkDQiPMu+LSdVuDuR1I5jcwtkB1Pi3F82E=
|
||||
github.com/uptrace/bun/driver/sqliteshim v1.2.16 h1:M6Dh5kkDWFbUWBrOsIE1g1zdZ5JbSytTD4piFRBOUAI=
|
||||
@@ -240,81 +329,196 @@ github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAh
|
||||
github.com/vmihailenco/tagparser/v2 v2.0.0/go.mod h1:Wri+At7QHww0WTrCBeu4J6bNtoV6mEfg5OIWRZA9qds=
|
||||
github.com/warkanum/bun v1.2.17 h1:HP8eTuKSNcqMDhhIPFxEbgV/yct6RR0/c3qHH3PNZUA=
|
||||
github.com/warkanum/bun v1.2.17/go.mod h1:jMoNg2n56ckaawi/O/J92BHaECmrz6IRjuMWqlMaMTM=
|
||||
github.com/xdg-go/pbkdf2 v1.0.0 h1:Su7DPu48wXMwC3bs7MCNG+z4FhcyEuz5dlvchbq0B0c=
|
||||
github.com/xdg-go/pbkdf2 v1.0.0/go.mod h1:jrpuAogTd400dnrH08LKmI/xc1MbPOebTwRqcT5RDeI=
|
||||
github.com/xdg-go/scram v1.2.0 h1:bYKF2AEwG5rqd1BumT4gAnvwU/M9nBp2pTSxeZw7Wvs=
|
||||
github.com/xdg-go/scram v1.2.0/go.mod h1:3dlrS0iBaWKYVt2ZfA4cj48umJZ+cAEbR6/SjLA88I8=
|
||||
github.com/xdg-go/stringprep v1.0.4 h1:XLI/Ng3O1Atzq0oBs3TWm+5ZVgkq2aqdlvP9JtoZ6c8=
|
||||
github.com/xdg-go/stringprep v1.0.4/go.mod h1:mPGuuIYwz7CmR2bT9j4GbQqutWS1zV24gijq1dTyGkM=
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4=
|
||||
github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT040gmv5BQTMR2HP4=
|
||||
github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 h1:ilQV1hzziu+LLM3zUTJ0trRztfwgjqKnBWNtSRkbmwM=
|
||||
github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78/go.mod h1:aL8wCCfTfSfmXjznFBSZNN13rSJjlIOI1fUNAtF7rmI=
|
||||
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/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
||||
go.opentelemetry.io/auto/sdk v1.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA=
|
||||
go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A=
|
||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 h1:jq9TW8u3so/bN+JPT166wjOI6/vQPF6Xe7nMNIltagk=
|
||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0/go.mod h1:p8pYQP+m5XfbZm9fxtSKAbM6oIllS7s2AfxrChvc7iw=
|
||||
go.opentelemetry.io/otel v1.38.0 h1:RkfdswUDRimDg0m2Az18RKOsnI8UDzppJAtj01/Ymk8=
|
||||
go.opentelemetry.io/otel v1.38.0/go.mod h1:zcmtmQ1+YmQM9wrNsTGV/q/uyusom3P8RxwExxkZhjM=
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.38.0 h1:GqRJVj7UmLjCVyVJ3ZFLdPRmhDUp2zFmQe3RHIOsw24=
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.38.0/go.mod h1:ri3aaHSmCTVYu2AWv44YMauwAQc0aqI9gHKIcSbI1pU=
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.38.0 h1:lwI4Dc5leUqENgGuQImwLo4WnuXFPetmPpkLi2IrX54=
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.38.0/go.mod h1:Kz/oCE7z5wuyhPxsXDuaPteSWqjSBD5YaSdbxZYGbGk=
|
||||
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
|
||||
github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s=
|
||||
go.mongodb.org/mongo-driver v1.17.9 h1:IexDdCuuNJ3BHrELgBlyaH9p60JXAvdzWR128q+U5tU=
|
||||
go.mongodb.org/mongo-driver v1.17.9/go.mod h1:LlOhpH5NUEfhxcAwG0UEkMqwYcc4JU18gtCdGudk/tQ=
|
||||
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
|
||||
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
|
||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0 h1:F7Jx+6hwnZ41NSFTO5q4LYDtJRXBf2PD0rNBkeB/lus=
|
||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0/go.mod h1:UHB22Z8QsdRDrnAtX4PntOl36ajSxcdUMt1sF7Y6E7Q=
|
||||
go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU=
|
||||
go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc=
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.43.0 h1:88Y4s2C8oTui1LGM6bTWkw0ICGcOLCAI5l6zsD1j20k=
|
||||
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/go.mod h1:oVdCUtjq9MK9BlS7TtucsQwUcXcymNiEDjgDD2jMtZU=
|
||||
go.opentelemetry.io/otel/metric v1.38.0 h1:Kl6lzIYGAh5M159u9NgiRkmoMKjvbsKtYRwgfrA6WpA=
|
||||
go.opentelemetry.io/otel/metric v1.38.0/go.mod h1:kB5n/QoRM8YwmUahxvI3bO34eVtQf2i4utNVLr9gEmI=
|
||||
go.opentelemetry.io/otel/sdk v1.38.0 h1:l48sr5YbNf2hpCUj/FoGhW9yDkl+Ma+LrVl8qaM5b+E=
|
||||
go.opentelemetry.io/otel/sdk v1.38.0/go.mod h1:ghmNdGlVemJI3+ZB5iDEuk4bWA3GkTpW+DOoZMYBVVg=
|
||||
go.opentelemetry.io/otel/sdk/metric v1.38.0 h1:aSH66iL0aZqo//xXzQLYozmWrXxyFkBJ6qT5wthqPoM=
|
||||
go.opentelemetry.io/otel/sdk/metric v1.38.0/go.mod h1:dg9PBnW9XdQ1Hd6ZnRz689CbtrUp0wMMs9iPcgT9EZA=
|
||||
go.opentelemetry.io/otel/trace v1.38.0 h1:Fxk5bKrDZJUH+AMyyIXGcFAPah0oRcT+LuNtJrmcNLE=
|
||||
go.opentelemetry.io/otel/trace v1.38.0/go.mod h1:j1P9ivuFsTceSWe1oY+EeW3sc+Pp42sO++GHkg4wwhs=
|
||||
go.opentelemetry.io/proto/otlp v1.7.1 h1:gTOMpGDb0WTBOP8JaO72iL3auEZhVmAQg4ipjOVAtj4=
|
||||
go.opentelemetry.io/proto/otlp v1.7.1/go.mod h1:b2rVh6rfI/s2pHWNlB7ILJcRALpcNDzKhACevjI+ZnE=
|
||||
go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc=
|
||||
go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo=
|
||||
go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58=
|
||||
go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0=
|
||||
go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI=
|
||||
go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA=
|
||||
go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk=
|
||||
go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE=
|
||||
go.opentelemetry.io/proto/otlp v1.10.0 h1:IQRWgT5srOCYfiWnpqUYz9CVmbO8bFmKcwYxpuCSL2g=
|
||||
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/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||
go.uber.org/multierr v1.10.0 h1:S0h4aNzvfcFsC3dRF1jLoaov7oRaKqRGC/pUEJ2yvPQ=
|
||||
go.uber.org/multierr v1.10.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
||||
go.uber.org/zap v1.27.0 h1:aJMhYGrd5QSmlpLMr2MftRKl7t8J8PTZPA732ud/XR8=
|
||||
go.uber.org/zap v1.27.0/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E=
|
||||
go.yaml.in/yaml/v2 v2.4.2 h1:DzmwEr2rDGHl7lsFgAHxmNz/1NlQ7xLIrlN2h5d1eGI=
|
||||
go.yaml.in/yaml/v2 v2.4.2/go.mod h1:081UH+NErpNdqlCXm3TtEran0rJZGxAYx9hb/ELlsPU=
|
||||
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/zap v1.28.0 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo=
|
||||
go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q=
|
||||
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
||||
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/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
||||
golang.org/x/crypto v0.43.0 h1:dduJYIi3A3KOfdGOHX8AVZ/jGiyPa3IbBozJ5kNuE04=
|
||||
golang.org/x/crypto v0.43.0/go.mod h1:BFbav4mRNlXJL4wNeejLpWxB7wMbc79PdRGhWKncxR0=
|
||||
golang.org/x/exp v0.0.0-20251113190631-e25ba8c21ef6 h1:zfMcR1Cs4KNuomFFgGefv5N0czO2XZpUbxGUy8i8ug0=
|
||||
golang.org/x/exp v0.0.0-20251113190631-e25ba8c21ef6/go.mod h1:46edojNIoXTNOhySWIWdix628clX9ODXwPsQuG6hsK0=
|
||||
golang.org/x/mod v0.30.0 h1:fDEXFVZ/fmCKProc/yAXXUijritrDzahmwwefnjoPFk=
|
||||
golang.org/x/mod v0.30.0/go.mod h1:lAsf5O2EvJeSFMiBxXDki7sCgAxEUcZHXoXMKT4GJKc=
|
||||
golang.org/x/net v0.45.0 h1:RLBg5JKixCy82FtLJpeNlVM0nrSqpCRYzVU1n8kj0tM=
|
||||
golang.org/x/net v0.45.0/go.mod h1:ECOoLqd5U3Lhyeyo/QDCEVQ4sNgYsqvCZ722XogGieY=
|
||||
golang.org/x/sync v0.18.0 h1:kr88TuHDroi+UVf+0hZnirlk8o8T+4MrK6mr60WkH/I=
|
||||
golang.org/x/sync v0.18.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||
golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc=
|
||||
golang.org/x/crypto v0.6.0/go.mod h1:OFC/31mSvZgRz0V1QTNCzfAI1aIRzbiufJtkMIlEp58=
|
||||
golang.org/x/crypto v0.11.0/go.mod h1:xgJhtzW8F9jGdVFWZESrid1U1bjeNy4zgy5cRr/CIio=
|
||||
golang.org/x/crypto v0.12.0/go.mod h1:NF0Gs7EO5K4qLn+Ylc+fih8BSTeIjAP05siRnAh98yw=
|
||||
golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc=
|
||||
golang.org/x/crypto v0.18.0/go.mod h1:R0j02AL6hcrfOiy9T4ZYp/rcWeMxM3L6QYxlOuEG1mg=
|
||||
golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU=
|
||||
golang.org/x/crypto v0.21.0/go.mod h1:0BP7YvVV9gBbVKyeTG0Gyn+gZm94bibOW5BjDEYAOMs=
|
||||
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.24.0/go.mod h1:Z1PMYSOR5nyMcyAVAIQSKCDwalqy85Aqn1x3Ws4L5DM=
|
||||
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
||||
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
||||
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.9.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.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||
golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk=
|
||||
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-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-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c=
|
||||
golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
|
||||
golang.org/x/net v0.7.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
|
||||
golang.org/x/net v0.8.0/go.mod h1:QVkue5JL9kW//ek3r6jTKnTFis1tRmNAW2P1shuFdJc=
|
||||
golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg=
|
||||
golang.org/x/net v0.13.0/go.mod h1:zEVYFnQC7m/vmpQFELhcD1EWkZlX69l4oqgmer6hfKA=
|
||||
golang.org/x/net v0.14.0/go.mod h1:PpSgVXXLK0OxS0F31C1/tv6XNguvCrnXIDrFMspZIUI=
|
||||
golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk=
|
||||
golang.org/x/net v0.20.0/go.mod h1:z8BVo6PvndSri0LbOE3hAn0apkU+1YvI6E70E9jsnvY=
|
||||
golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
|
||||
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.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
||||
golang.org/x/net v0.26.0/go.mod h1:5YKkiSynbBIh3p6iOc/vibscux0x38BZDkn8sCUPxHE=
|
||||
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
|
||||
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
|
||||
golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
|
||||
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-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.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.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.9.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||
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-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-20201204225414-ed752295db88/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.0.0-20210616045830-e2b7044e8c71/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.0.0-20210616094352-59db8d763f22/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.11.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.38.0 h1:3yZWxaJjBmCWXqhN1qh02AkOnCQ1poK6oF+a7xWL6Gc=
|
||||
golang.org/x/sys v0.38.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
||||
golang.org/x/term v0.36.0 h1:zMPR+aF8gfksFprF/Nc/rd1wRS1EI6nDBGyWAvDzx2Q=
|
||||
golang.org/x/term v0.36.0/go.mod h1:Qu394IJq6V6dCBRgwqshf3mPF85AqzYEzofzRdZkWss=
|
||||
golang.org/x/text v0.30.0 h1:yznKA/E9zq54KzlzBEAWn1NXSQ8DIp/NYMy88xJjl4k=
|
||||
golang.org/x/text v0.30.0/go.mod h1:yDdHFIX9t+tORqspjENWgzaCVXgk0yYnYuSZ8UzzBVM=
|
||||
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
|
||||
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
|
||||
golang.org/x/tools v0.39.0 h1:ik4ho21kwuQln40uelmciQPp9SipgNDdrafrYA4TmQQ=
|
||||
golang.org/x/tools v0.39.0/go.mod h1:JnefbkDPyD8UU2kI5fuf8ZX4/yUeh9W877ZeBONxUqQ=
|
||||
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.16.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.17.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
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.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.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||
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/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.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k=
|
||||
golang.org/x/term v0.6.0/go.mod h1:m6U89DPEgQRMq3DNkDClhWw02AUbt2daBVO4cn4Hv9U=
|
||||
golang.org/x/term v0.8.0/go.mod h1:xPskH00ivmX89bAKVGSKKtLOWNx2+17Eiy94tnKShWo=
|
||||
golang.org/x/term v0.10.0/go.mod h1:lpqdcUyK/oCiQxvxVrppt5ggO2KCZ5QblwqPnfZ6d5o=
|
||||
golang.org/x/term v0.11.0/go.mod h1:zC9APTIj3jG3FdV/Ons+XE1riIZXG4aZ4GTHiPZJPIU=
|
||||
golang.org/x/term v0.12.0/go.mod h1:owVbMEjm3cBLCHdkQu9b1opXd4ETQWc3BhuQGKgXgvU=
|
||||
golang.org/x/term v0.16.0/go.mod h1:yn7UURbUtPyrVJPGPq404EukNFxcm/foM+bV/bfcDsY=
|
||||
golang.org/x/term v0.17.0/go.mod h1:lLRBjIVuehSbZlaOtGMbcMncT+aqLLLmKrsjNrUguwk=
|
||||
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.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
|
||||
golang.org/x/term v0.21.0/go.mod h1:ooXLefLobQVslOqselCNF4SxFAaoS6KujMbsGzSDmX0=
|
||||
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
||||
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.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.8/go.mod h1:E6s5w1FMmriuDzIBO73fBruAKo1PCIq6d2Q6DHfQ8WQ=
|
||||
golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
|
||||
golang.org/x/text v0.8.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8=
|
||||
golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8=
|
||||
golang.org/x/text v0.11.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE=
|
||||
golang.org/x/text v0.12.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE=
|
||||
golang.org/x/text v0.13.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE=
|
||||
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.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.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
|
||||
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
|
||||
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
|
||||
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-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.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
|
||||
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.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE=
|
||||
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-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk=
|
||||
gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E=
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20250825161204-c5933d9347a5 h1:BIRfGDEjiHRrk0QKZe3Xv2ieMhtgRGeLcZQ0mIVn4EY=
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20250825161204-c5933d9347a5/go.mod h1:j3QtIyytwqGr1JUDtYXwtMXWPKsEa5LtzIFN1Wn5WvE=
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20250825161204-c5933d9347a5 h1:eaY8u2EuxbRv7c3NiGK0/NedzVsCcV6hDuU5qPX5EGE=
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20250825161204-c5933d9347a5/go.mod h1:M4/wBTSeyLxupu3W3tJtOgB14jILAS/XWPSSa3TAlJc=
|
||||
google.golang.org/grpc v1.75.0 h1:+TW+dqTd2Biwe6KKfhE5JpiYIBWq865PhKGSXiivqt4=
|
||||
google.golang.org/grpc v1.75.0/go.mod h1:JtPAzKiq4v1xcAB2hydNlWI2RnF85XXcV0mhKXr2ecQ=
|
||||
google.golang.org/protobuf v1.36.8 h1:xHScyCOEuuwZEc6UtSOvPbAT4zRh0xcNRYekJwfqyMc=
|
||||
google.golang.org/protobuf v1.36.8/go.mod h1:fuxRtAxBytpl4zzqUh6/eyUujkJdNiuEkXntxiD/uRU=
|
||||
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
|
||||
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa h1:Kjn0N0tCrDgiAFW+lGO4JZ3ck44CehvJQMAwj9QF0G8=
|
||||
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-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk=
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
|
||||
google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU=
|
||||
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/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
|
||||
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=
|
||||
gopkg.in/yaml.v2 v2.2.1/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
|
||||
gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
|
||||
gopkg.in/yaml.v2 v2.2.8/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
|
||||
gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ=
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
@@ -322,34 +526,37 @@ gorm.io/driver/postgres v1.6.0 h1:2dxzU8xJ+ivvqTRph34QX+WrRaJlmfyPqXmoGVjMBa4=
|
||||
gorm.io/driver/postgres v1.6.0/go.mod h1:vUw0mrGgrTK+uPHEhAdV4sfFELrByKVGnaVRkXDhtWo=
|
||||
gorm.io/driver/sqlite v1.6.0 h1:WHRRrIiulaPiPFmDcod6prc4l2VGVWHz80KspNsxSfQ=
|
||||
gorm.io/driver/sqlite v1.6.0/go.mod h1:AO9V1qIQddBESngQUKWL9yoH93HIeA1X6V633rBwyT8=
|
||||
gorm.io/gorm v1.30.0 h1:qbT5aPv1UH8gI99OsRlvDToLxW5zR7FzS9acZDOZcgs=
|
||||
gorm.io/driver/sqlserver v1.6.3 h1:UR+nWCuphPnq7UxnL57PSrlYjuvs+sf1N59GgFX7uAI=
|
||||
gorm.io/driver/sqlserver v1.6.3/go.mod h1:VZeNn7hqX1aXoN5TPAFGWvxWG90xtA8erGn2gQmpc6U=
|
||||
gorm.io/gorm v1.30.0/go.mod h1:8Z33v652h4//uMA76KjeDH8mJXPm1QNCYrMeatR0DOE=
|
||||
gorm.io/gorm v1.31.1 h1:7CA8FTFz/gRfgqgpeKIBcervUn3xSyPUmr6B2WXJ7kg=
|
||||
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/go.mod h1:LtdLGcnqToBH83WByAAi/wiwSFCArdFIUV/xxN4pcjA=
|
||||
modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis=
|
||||
modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
|
||||
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
|
||||
modernc.org/ccgo/v4 v4.30.1/go.mod h1:bIOeI1JL54Utlxn+LwrFyjCx2n2RDiYEaJVSrgdrRfM=
|
||||
modernc.org/fileutil v1.3.40 h1:ZGMswMNc9JOCrcrakF1HrvmergNLAmxOPjizirpfqBA=
|
||||
modernc.org/fileutil v1.3.40/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc=
|
||||
modernc.org/cc/v4 v4.28.2 h1:3tQ0lf2ADtoby2EtSP+J7IE2SHwEJdP8ioR59wx7XpY=
|
||||
modernc.org/cc/v4 v4.28.2/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI=
|
||||
modernc.org/ccgo/v4 v4.34.0 h1:yRLPFZieg532OT4rp4JFNIVcquwalMX26G95WQDqwCQ=
|
||||
modernc.org/ccgo/v4 v4.34.0/go.mod h1:AS5WYMyBakQ+fhsHhtP8mWB82KTGPkNNJDGfGQCe0/A=
|
||||
modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM=
|
||||
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/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
|
||||
modernc.org/gc/v3 v3.1.1 h1:k8T3gkXWY9sEiytKhcgyiZ2L0DTyCQ/nvX+LoCljoRE=
|
||||
modernc.org/gc/v3 v3.1.1/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY=
|
||||
modernc.org/gc/v3 v3.1.2 h1:ZtDCnhonXSZexk/AYsegNRV1lJGgaNZJuKjJSWKyEqo=
|
||||
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/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI=
|
||||
modernc.org/libc v1.67.0 h1:QzL4IrKab2OFmxA3/vRYl0tLXrIamwrhD6CKD4WBVjQ=
|
||||
modernc.org/libc v1.67.0/go.mod h1:QvvnnJ5P7aitu0ReNpVIEyesuhmDLQ8kaEoyMjIFZJA=
|
||||
modernc.org/libc v1.72.3 h1:ZnDF4tXn4NBXFutMMQC4vtbTFSXhhKzR73fv0beZEAU=
|
||||
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/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
|
||||
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
|
||||
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
|
||||
modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8=
|
||||
modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
|
||||
modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg=
|
||||
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/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
|
||||
modernc.org/sqlite v1.40.1 h1:VfuXcxcUWWKRBuP8+BR9L7VnmusMgBNNnBYGEe9w/iY=
|
||||
modernc.org/sqlite v1.40.1/go.mod h1:9fjQZ0mB1LLP0GYrp39oOJXx/I2sxEnZtzCmEQIKvGE=
|
||||
modernc.org/sqlite v1.50.1 h1:l+cQvn0sd0zJJtfygGHuQJ5AjlrwXmWPw4KP3ZMwr9w=
|
||||
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/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A=
|
||||
modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=
|
||||
|
||||
-362
@@ -1,362 +0,0 @@
|
||||
openapi: 3.0.0
|
||||
info:
|
||||
title: ResolveSpec API
|
||||
version: '1.0'
|
||||
description: A flexible REST API with GraphQL-like capabilities
|
||||
|
||||
servers:
|
||||
- url: 'http://api.example.com/v1'
|
||||
|
||||
paths:
|
||||
'/{schema}/{entity}':
|
||||
parameters:
|
||||
- name: schema
|
||||
in: path
|
||||
required: true
|
||||
schema:
|
||||
type: string
|
||||
- name: entity
|
||||
in: path
|
||||
required: true
|
||||
schema:
|
||||
type: string
|
||||
get:
|
||||
summary: Get table metadata
|
||||
description: Retrieve table metadata including columns, types, and relationships
|
||||
responses:
|
||||
'200':
|
||||
description: Successful operation
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
allOf:
|
||||
- $ref: '#/components/schemas/Response'
|
||||
- type: object
|
||||
properties:
|
||||
data:
|
||||
$ref: '#/components/schemas/TableMetadata'
|
||||
'400':
|
||||
$ref: '#/components/responses/BadRequest'
|
||||
'404':
|
||||
$ref: '#/components/responses/NotFound'
|
||||
'500':
|
||||
$ref: '#/components/responses/ServerError'
|
||||
post:
|
||||
summary: Perform operations on entities
|
||||
requestBody:
|
||||
required: true
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/Request'
|
||||
responses:
|
||||
'200':
|
||||
description: Successful operation
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/Response'
|
||||
'400':
|
||||
$ref: '#/components/responses/BadRequest'
|
||||
'404':
|
||||
$ref: '#/components/responses/NotFound'
|
||||
'500':
|
||||
$ref: '#/components/responses/ServerError'
|
||||
|
||||
'/{schema}/{entity}/{id}':
|
||||
parameters:
|
||||
- name: schema
|
||||
in: path
|
||||
required: true
|
||||
schema:
|
||||
type: string
|
||||
- name: entity
|
||||
in: path
|
||||
required: true
|
||||
schema:
|
||||
type: string
|
||||
- name: id
|
||||
in: path
|
||||
required: true
|
||||
schema:
|
||||
type: string
|
||||
post:
|
||||
summary: Perform operations on a specific entity
|
||||
requestBody:
|
||||
required: true
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/Request'
|
||||
responses:
|
||||
'200':
|
||||
description: Successful operation
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/Response'
|
||||
'400':
|
||||
$ref: '#/components/responses/BadRequest'
|
||||
'404':
|
||||
$ref: '#/components/responses/NotFound'
|
||||
'500':
|
||||
$ref: '#/components/responses/ServerError'
|
||||
|
||||
components:
|
||||
schemas:
|
||||
Request:
|
||||
type: object
|
||||
required:
|
||||
- operation
|
||||
properties:
|
||||
operation:
|
||||
type: string
|
||||
enum:
|
||||
- read
|
||||
- create
|
||||
- update
|
||||
- delete
|
||||
id:
|
||||
oneOf:
|
||||
- type: string
|
||||
- type: array
|
||||
items:
|
||||
type: string
|
||||
description: Optional record identifier(s) when not provided in URL
|
||||
data:
|
||||
oneOf:
|
||||
- type: object
|
||||
- type: array
|
||||
items:
|
||||
type: object
|
||||
description: Data for single or bulk create/update operations
|
||||
options:
|
||||
$ref: '#/components/schemas/Options'
|
||||
|
||||
Options:
|
||||
type: object
|
||||
properties:
|
||||
preload:
|
||||
type: array
|
||||
items:
|
||||
$ref: '#/components/schemas/PreloadOption'
|
||||
columns:
|
||||
type: array
|
||||
items:
|
||||
type: string
|
||||
filters:
|
||||
type: array
|
||||
items:
|
||||
$ref: '#/components/schemas/FilterOption'
|
||||
sort:
|
||||
type: array
|
||||
items:
|
||||
$ref: '#/components/schemas/SortOption'
|
||||
limit:
|
||||
type: integer
|
||||
minimum: 0
|
||||
offset:
|
||||
type: integer
|
||||
minimum: 0
|
||||
customOperators:
|
||||
type: array
|
||||
items:
|
||||
$ref: '#/components/schemas/CustomOperator'
|
||||
computedColumns:
|
||||
type: array
|
||||
items:
|
||||
$ref: '#/components/schemas/ComputedColumn'
|
||||
|
||||
PreloadOption:
|
||||
type: object
|
||||
properties:
|
||||
relation:
|
||||
type: string
|
||||
columns:
|
||||
type: array
|
||||
items:
|
||||
type: string
|
||||
filters:
|
||||
type: array
|
||||
items:
|
||||
$ref: '#/components/schemas/FilterOption'
|
||||
|
||||
FilterOption:
|
||||
type: object
|
||||
required:
|
||||
- column
|
||||
- operator
|
||||
- value
|
||||
properties:
|
||||
column:
|
||||
type: string
|
||||
operator:
|
||||
type: string
|
||||
enum:
|
||||
- eq
|
||||
- neq
|
||||
- gt
|
||||
- gte
|
||||
- lt
|
||||
- lte
|
||||
- like
|
||||
- ilike
|
||||
- in
|
||||
value:
|
||||
type: object
|
||||
|
||||
SortOption:
|
||||
type: object
|
||||
required:
|
||||
- column
|
||||
- direction
|
||||
properties:
|
||||
column:
|
||||
type: string
|
||||
direction:
|
||||
type: string
|
||||
enum:
|
||||
- asc
|
||||
- desc
|
||||
|
||||
CustomOperator:
|
||||
type: object
|
||||
required:
|
||||
- name
|
||||
- sql
|
||||
properties:
|
||||
name:
|
||||
type: string
|
||||
sql:
|
||||
type: string
|
||||
|
||||
ComputedColumn:
|
||||
type: object
|
||||
required:
|
||||
- name
|
||||
- expression
|
||||
properties:
|
||||
name:
|
||||
type: string
|
||||
expression:
|
||||
type: string
|
||||
|
||||
Response:
|
||||
type: object
|
||||
required:
|
||||
- success
|
||||
properties:
|
||||
success:
|
||||
type: boolean
|
||||
data:
|
||||
type: object
|
||||
metadata:
|
||||
$ref: '#/components/schemas/Metadata'
|
||||
error:
|
||||
$ref: '#/components/schemas/Error'
|
||||
|
||||
Metadata:
|
||||
type: object
|
||||
properties:
|
||||
total:
|
||||
type: integer
|
||||
filtered:
|
||||
type: integer
|
||||
limit:
|
||||
type: integer
|
||||
offset:
|
||||
type: integer
|
||||
|
||||
Error:
|
||||
type: object
|
||||
properties:
|
||||
code:
|
||||
type: string
|
||||
message:
|
||||
type: string
|
||||
details:
|
||||
type: object
|
||||
|
||||
TableMetadata:
|
||||
type: object
|
||||
required:
|
||||
- schema
|
||||
- table
|
||||
- columns
|
||||
- relations
|
||||
properties:
|
||||
schema:
|
||||
type: string
|
||||
description: Schema name
|
||||
table:
|
||||
type: string
|
||||
description: Table name
|
||||
columns:
|
||||
type: array
|
||||
items:
|
||||
$ref: '#/components/schemas/Column'
|
||||
relations:
|
||||
type: array
|
||||
items:
|
||||
type: string
|
||||
description: List of relation names
|
||||
|
||||
Column:
|
||||
type: object
|
||||
required:
|
||||
- name
|
||||
- type
|
||||
- is_nullable
|
||||
- is_primary
|
||||
- is_unique
|
||||
- has_index
|
||||
properties:
|
||||
name:
|
||||
type: string
|
||||
description: Column name
|
||||
type:
|
||||
type: string
|
||||
description: Data type of the column
|
||||
is_nullable:
|
||||
type: boolean
|
||||
description: Whether the column can contain null values
|
||||
is_primary:
|
||||
type: boolean
|
||||
description: Whether the column is a primary key
|
||||
is_unique:
|
||||
type: boolean
|
||||
description: Whether the column has a unique constraint
|
||||
has_index:
|
||||
type: boolean
|
||||
description: Whether the column is indexed
|
||||
|
||||
responses:
|
||||
BadRequest:
|
||||
description: Bad request
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/Response'
|
||||
|
||||
NotFound:
|
||||
description: Resource not found
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/Response'
|
||||
|
||||
ServerError:
|
||||
description: Internal server error
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/Response'
|
||||
|
||||
securitySchemes:
|
||||
bearerAuth:
|
||||
type: http
|
||||
scheme: bearer
|
||||
bearerFormat: JWT
|
||||
|
||||
security:
|
||||
- bearerAuth: []
|
||||
Vendored
+30
-17
@@ -3,23 +3,29 @@ package cache
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
var (
|
||||
defaultCache *Cache
|
||||
)
|
||||
var defaultCache atomic.Pointer[Cache]
|
||||
|
||||
// swapOwned installs c as the default and closes the displaced cache, which this
|
||||
// package created and therefore owns.
|
||||
func swapOwned(c *Cache) {
|
||||
if old := defaultCache.Swap(c); old != nil && old != c {
|
||||
_ = old.Close() // best-effort: the displaced provider is being discarded
|
||||
}
|
||||
}
|
||||
|
||||
// Initialize initializes the cache with a provider.
|
||||
// If not called, the package will use an in-memory provider by default.
|
||||
func Initialize(provider Provider) {
|
||||
defaultCache = NewCache(provider)
|
||||
swapOwned(NewCache(provider))
|
||||
}
|
||||
|
||||
// UseMemory configures the cache to use in-memory storage.
|
||||
func UseMemory(opts *Options) error {
|
||||
provider := NewMemoryProvider(opts)
|
||||
defaultCache = NewCache(provider)
|
||||
swapOwned(NewCache(NewMemoryProvider(opts)))
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -29,7 +35,7 @@ func UseRedis(config *RedisConfig) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initialize Redis provider: %w", err)
|
||||
}
|
||||
defaultCache = NewCache(provider)
|
||||
swapOwned(NewCache(provider))
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -39,26 +45,33 @@ func UseMemcache(config *MemcacheConfig) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initialize Memcache provider: %w", err)
|
||||
}
|
||||
defaultCache = NewCache(provider)
|
||||
swapOwned(NewCache(provider))
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetDefaultCache returns the default cache instance.
|
||||
// Initializes with in-memory provider if not already initialized.
|
||||
// Safe for concurrent use.
|
||||
func GetDefaultCache() *Cache {
|
||||
if defaultCache == nil {
|
||||
_ = UseMemory(&Options{
|
||||
DefaultTTL: 5 * time.Minute,
|
||||
MaxSize: 10000,
|
||||
})
|
||||
if c := defaultCache.Load(); c != nil {
|
||||
return c
|
||||
}
|
||||
return defaultCache
|
||||
fresh := NewCache(NewMemoryProvider(&Options{
|
||||
DefaultTTL: 5 * time.Minute,
|
||||
MaxSize: 10000,
|
||||
}))
|
||||
if defaultCache.CompareAndSwap(nil, fresh) {
|
||||
return fresh
|
||||
}
|
||||
_ = fresh.Close() // lost the race; discard our provider
|
||||
return defaultCache.Load()
|
||||
}
|
||||
|
||||
// SetDefaultCache sets a custom cache instance as the default cache.
|
||||
// This is useful for testing or when you want to use a pre-configured cache instance.
|
||||
// The caller keeps ownership of both the new and the displaced cache; neither is closed.
|
||||
func SetDefaultCache(cache *Cache) {
|
||||
defaultCache = cache
|
||||
defaultCache.Store(cache)
|
||||
}
|
||||
|
||||
// GetStats returns cache statistics.
|
||||
@@ -69,8 +82,8 @@ func GetStats(ctx context.Context) (*CacheStats, error) {
|
||||
|
||||
// Close closes the cache and releases resources.
|
||||
func Close() error {
|
||||
if defaultCache != nil {
|
||||
return defaultCache.Close()
|
||||
if c := defaultCache.Load(); c != nil {
|
||||
return c.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
Vendored
+19
-6
@@ -3,10 +3,17 @@ package cache
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// ErrNotFound is returned when a key is not in the cache. The key is deliberately
|
||||
// not included in the error: keys may embed credentials (e.g. session tokens).
|
||||
var ErrNotFound = errors.New("cache: key not found")
|
||||
|
||||
// Cache is the main cache manager that wraps a Provider.
|
||||
type Cache struct {
|
||||
provider Provider
|
||||
@@ -23,7 +30,7 @@ func NewCache(provider Provider) *Cache {
|
||||
func (c *Cache) Get(ctx context.Context, key string, dest interface{}) error {
|
||||
data, exists := c.provider.Get(ctx, key)
|
||||
if !exists {
|
||||
return fmt.Errorf("key not found: %s", key)
|
||||
return ErrNotFound
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(data, dest); err != nil {
|
||||
@@ -37,7 +44,7 @@ func (c *Cache) Get(ctx context.Context, key string, dest interface{}) error {
|
||||
func (c *Cache) GetBytes(ctx context.Context, key string) ([]byte, error) {
|
||||
data, exists := c.provider.Get(ctx, key)
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("key not found: %s", key)
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
@@ -122,9 +129,10 @@ func (c *Cache) GetOrSet(ctx context.Context, key string, dest interface{}, ttl
|
||||
return fmt.Errorf("loader failed: %w", err)
|
||||
}
|
||||
|
||||
// Store in cache
|
||||
// Store in cache. A cache-write failure must not fail the call: the authoritative
|
||||
// value has already been loaded, and the cache is only an optimisation.
|
||||
if err := c.Set(ctx, key, value, ttl); err != nil {
|
||||
return fmt.Errorf("failed to cache value: %w", err)
|
||||
logger.Warn("cache: failed to store loaded value, continuing uncached: %v", err)
|
||||
}
|
||||
|
||||
// Populate dest with the loaded value
|
||||
@@ -142,6 +150,11 @@ func (c *Cache) GetOrSet(ctx context.Context, key string, dest interface{}, ttl
|
||||
|
||||
// Remember is a convenience function that caches the result of a function call.
|
||||
// It's similar to GetOrSet but returns the value directly.
|
||||
//
|
||||
// WARNING: the returned type differs between a hit and a miss. On a hit the value is
|
||||
// generic decoded JSON (map[string]interface{}, []interface{}, float64, string, ...);
|
||||
// on a miss it is exactly what loader returned. Do not type-assert the result to a
|
||||
// concrete type; prefer GetOrSet, which decodes into a typed destination.
|
||||
func (c *Cache) Remember(ctx context.Context, key string, ttl time.Duration, loader func() (interface{}, error)) (interface{}, error) {
|
||||
// Try to get from cache first as bytes
|
||||
data, err := c.GetBytes(ctx, key)
|
||||
@@ -158,9 +171,9 @@ func (c *Cache) Remember(ctx context.Context, key string, ttl time.Duration, loa
|
||||
return nil, fmt.Errorf("loader failed: %w", err)
|
||||
}
|
||||
|
||||
// Store in cache
|
||||
// Cache-write failures are non-fatal (see GetOrSet)
|
||||
if err := c.Set(ctx, key, value, ttl); err != nil {
|
||||
return nil, fmt.Errorf("failed to cache value: %w", err)
|
||||
logger.Warn("cache: failed to store loaded value, continuing uncached: %v", err)
|
||||
}
|
||||
|
||||
return value, nil
|
||||
|
||||
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
|
||||
|
||||
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 (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ErrFlushNotAllowed is returned by Clear on shared-server providers (Redis, Memcache)
|
||||
// unless AllowFlush is set in their config, because Clear flushes the whole server/DB.
|
||||
var ErrFlushNotAllowed = errors.New("cache: Clear flushes the entire server; set AllowFlush in the provider config to permit it")
|
||||
|
||||
// Provider defines the interface that all cache providers must implement.
|
||||
type Provider interface {
|
||||
// Get retrieves a value from the cache by key.
|
||||
@@ -58,8 +63,13 @@ type Options struct {
|
||||
DefaultTTL time.Duration
|
||||
|
||||
// MaxSize is the maximum number of items (for in-memory provider).
|
||||
// 0 selects the default (10000); a negative value means unbounded.
|
||||
MaxSize int
|
||||
|
||||
// CleanupInterval is how often the in-memory provider removes expired items
|
||||
// (default: 1 minute).
|
||||
CleanupInterval time.Duration
|
||||
|
||||
// EvictionPolicy determines how items are evicted (LRU, LFU, etc).
|
||||
EvictionPolicy string
|
||||
}
|
||||
|
||||
Vendored
+238
-119
@@ -2,17 +2,34 @@ package cache
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bradfitz/gomemcache/memcache"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
const (
|
||||
// memcacheMaxRelativeTTL is the largest expiry memcached treats as relative seconds;
|
||||
// anything larger is interpreted as an absolute Unix timestamp.
|
||||
memcacheMaxRelativeTTL = 30 * 24 * 60 * 60
|
||||
|
||||
// memcacheMaxTagKeys bounds the per-tag key list so it stays under memcached's 1MB item limit.
|
||||
memcacheMaxTagKeys = 5000
|
||||
|
||||
memcacheCASRetries = 5
|
||||
)
|
||||
|
||||
// MemcacheProvider is a Memcache implementation of the Provider interface.
|
||||
type MemcacheProvider struct {
|
||||
client *memcache.Client
|
||||
options *Options
|
||||
client *memcache.Client
|
||||
options *Options
|
||||
allowFlush bool
|
||||
}
|
||||
|
||||
// MemcacheConfig contains Memcache-specific configuration.
|
||||
@@ -28,37 +45,46 @@ type MemcacheConfig struct {
|
||||
|
||||
// Options contains general cache options
|
||||
Options *Options
|
||||
|
||||
// AllowFlush permits Clear() to run flush_all, which wipes every key on every
|
||||
// configured server (including data not owned by this cache). Off by default.
|
||||
AllowFlush bool
|
||||
}
|
||||
|
||||
// NewMemcacheProvider creates a new Memcache cache provider.
|
||||
func NewMemcacheProvider(config *MemcacheConfig) (*MemcacheProvider, error) {
|
||||
if config == nil {
|
||||
config = &MemcacheConfig{
|
||||
Servers: []string{"localhost:11211"},
|
||||
}
|
||||
// Work on a copy so the caller's struct is not mutated
|
||||
var cfg MemcacheConfig
|
||||
if config != nil {
|
||||
cfg = *config
|
||||
cfg.Servers = append([]string(nil), config.Servers...)
|
||||
}
|
||||
if cfg.Options != nil {
|
||||
o := *cfg.Options
|
||||
cfg.Options = &o
|
||||
}
|
||||
|
||||
if len(config.Servers) == 0 {
|
||||
config.Servers = []string{"localhost:11211"}
|
||||
if len(cfg.Servers) == 0 {
|
||||
cfg.Servers = []string{"localhost:11211"}
|
||||
}
|
||||
|
||||
if config.MaxIdleConns == 0 {
|
||||
config.MaxIdleConns = 2
|
||||
if cfg.MaxIdleConns == 0 {
|
||||
cfg.MaxIdleConns = 2
|
||||
}
|
||||
|
||||
if config.Timeout == 0 {
|
||||
config.Timeout = 1 * time.Second
|
||||
if cfg.Timeout == 0 {
|
||||
cfg.Timeout = 1 * time.Second
|
||||
}
|
||||
|
||||
if config.Options == nil {
|
||||
config.Options = &Options{
|
||||
if cfg.Options == nil {
|
||||
cfg.Options = &Options{
|
||||
DefaultTTL: 5 * time.Minute,
|
||||
}
|
||||
}
|
||||
|
||||
client := memcache.New(config.Servers...)
|
||||
client.MaxIdleConns = config.MaxIdleConns
|
||||
client.Timeout = config.Timeout
|
||||
client := memcache.New(cfg.Servers...)
|
||||
client.MaxIdleConns = cfg.MaxIdleConns
|
||||
client.Timeout = cfg.Timeout
|
||||
|
||||
// Test connection
|
||||
if err := client.Ping(); err != nil {
|
||||
@@ -66,18 +92,72 @@ func NewMemcacheProvider(config *MemcacheConfig) (*MemcacheProvider, error) {
|
||||
}
|
||||
|
||||
return &MemcacheProvider{
|
||||
client: client,
|
||||
options: config.Options,
|
||||
client: client,
|
||||
options: cfg.Options,
|
||||
allowFlush: cfg.AllowFlush,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// memcacheKey maps a caller key to a legal memcached key. User keys live under the "k:"
|
||||
// prefix, so they can never collide with the "cache:tag:" / "cache:tags:" index keys.
|
||||
// Keys that are too long or contain illegal bytes (whitespace/control characters)
|
||||
// are replaced with their SHA-256.
|
||||
func memcacheKey(key string) string {
|
||||
k := "k:" + key
|
||||
if len(k) > 200 || !legalMemcacheKey(k) {
|
||||
sum := sha256.Sum256([]byte(key))
|
||||
return "k:h:" + hex.EncodeToString(sum[:])
|
||||
}
|
||||
return k
|
||||
}
|
||||
|
||||
func legalMemcacheKey(key string) bool {
|
||||
for i := 0; i < len(key); i++ {
|
||||
if key[i] <= ' ' || key[i] == 0x7f {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// memcacheTagKey maps a tag to its index key, hashing if it is not a legal key.
|
||||
func memcacheTagKey(prefix, name string) string {
|
||||
k := prefix + name
|
||||
if len(k) > 200 || !legalMemcacheKey(k) {
|
||||
sum := sha256.Sum256([]byte(name))
|
||||
return prefix + "h:" + hex.EncodeToString(sum[:])
|
||||
}
|
||||
return k
|
||||
}
|
||||
|
||||
// memcacheExpiry converts a TTL into a memcached expiry value, switching to an
|
||||
// absolute Unix timestamp above 30 days as the protocol requires.
|
||||
func memcacheExpiry(ttl time.Duration) int32 {
|
||||
if ttl <= 0 {
|
||||
return 0 // never expires
|
||||
}
|
||||
secs := int64(ttl.Seconds())
|
||||
if secs > memcacheMaxRelativeTTL {
|
||||
return int32(time.Now().Add(ttl).Unix()) //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
|
||||
}
|
||||
if secs == 0 {
|
||||
secs = 1
|
||||
}
|
||||
return int32(secs)
|
||||
}
|
||||
|
||||
// Get retrieves a value from the cache by key.
|
||||
func (m *MemcacheProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
||||
item, err := m.client.Get(key)
|
||||
if err == memcache.ErrCacheMiss {
|
||||
if ctx.Err() != nil {
|
||||
return nil, false
|
||||
}
|
||||
item, err := m.client.Get(memcacheKey(key))
|
||||
if errors.Is(err, memcache.ErrCacheMiss) {
|
||||
return nil, false
|
||||
}
|
||||
if err != nil {
|
||||
// Reported as a miss (the Provider interface cannot express errors), but not silently
|
||||
logger.Warn("cache: memcache GET failed: %v", err)
|
||||
return nil, false
|
||||
}
|
||||
return item.Value, true
|
||||
@@ -85,130 +165,158 @@ func (m *MemcacheProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
||||
|
||||
// Set stores a value in the cache with the specified TTL.
|
||||
func (m *MemcacheProvider) Set(ctx context.Context, key string, value []byte, ttl time.Duration) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
if ttl == 0 {
|
||||
ttl = m.options.DefaultTTL
|
||||
}
|
||||
|
||||
item := &memcache.Item{
|
||||
Key: key,
|
||||
return m.client.Set(&memcache.Item{
|
||||
Key: memcacheKey(key),
|
||||
Value: value,
|
||||
Expiration: int32(ttl.Seconds()),
|
||||
}
|
||||
|
||||
return m.client.Set(item)
|
||||
Expiration: memcacheExpiry(ttl),
|
||||
})
|
||||
}
|
||||
|
||||
// SetWithTags stores a value in the cache with the specified TTL and tags.
|
||||
// Note: Tag support in Memcache is limited and less efficient than Redis.
|
||||
// Note: Tag support in Memcache is limited and less efficient than Redis. The
|
||||
// tag index is updated with compare-and-swap; if it cannot be updated the value is
|
||||
// removed again and an error is returned, so an untracked entry is never left behind.
|
||||
func (m *MemcacheProvider) SetWithTags(ctx context.Context, key string, value []byte, ttl time.Duration, tags []string) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
if ttl == 0 {
|
||||
ttl = m.options.DefaultTTL
|
||||
}
|
||||
|
||||
expiration := int32(ttl.Seconds())
|
||||
expiration := memcacheExpiry(ttl)
|
||||
mkey := memcacheKey(key)
|
||||
|
||||
// Set the main value
|
||||
item := &memcache.Item{
|
||||
Key: key,
|
||||
Value: value,
|
||||
Expiration: expiration,
|
||||
if err := m.client.Set(&memcache.Item{Key: mkey, Value: value, Expiration: expiration}); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := m.client.Set(item); err != nil {
|
||||
if len(tags) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
fail := func(err error) error {
|
||||
_ = m.client.Delete(mkey) // best-effort rollback; the original error is what matters
|
||||
return err
|
||||
}
|
||||
|
||||
// Store tags for this key
|
||||
if len(tags) > 0 {
|
||||
tagsData, err := json.Marshal(tags)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal tags: %w", err)
|
||||
}
|
||||
tagsData, err := json.Marshal(tags)
|
||||
if err != nil {
|
||||
return fail(fmt.Errorf("failed to marshal tags: %w", err))
|
||||
}
|
||||
if err := m.client.Set(&memcache.Item{
|
||||
Key: memcacheTagKey("cache:tags:", key),
|
||||
Value: tagsData,
|
||||
Expiration: expiration,
|
||||
}); err != nil {
|
||||
return fail(err)
|
||||
}
|
||||
|
||||
tagsItem := &memcache.Item{
|
||||
Key: fmt.Sprintf("cache:tags:%s", key),
|
||||
Value: tagsData,
|
||||
Expiration: expiration,
|
||||
}
|
||||
if err := m.client.Set(tagsItem); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Add key to each tag's key list
|
||||
for _, tag := range tags {
|
||||
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
||||
|
||||
// Get existing keys for this tag
|
||||
var keys []string
|
||||
if item, err := m.client.Get(tagKey); err == nil {
|
||||
_ = json.Unmarshal(item.Value, &keys)
|
||||
}
|
||||
|
||||
// Add current key if not already present
|
||||
found := false
|
||||
// Tag lists live longer than the entries they index
|
||||
tagExpiry := memcacheExpiry(ttl + time.Hour)
|
||||
for _, tag := range tags {
|
||||
if err := m.updateTagKeys(memcacheTagKey("cache:tag:", tag), tagExpiry, func(keys []string) ([]string, error) {
|
||||
for _, k := range keys {
|
||||
if k == key {
|
||||
found = true
|
||||
break
|
||||
return keys, nil
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
keys = append(keys, key)
|
||||
if len(keys) >= memcacheMaxTagKeys {
|
||||
return nil, fmt.Errorf("tag index for %q is full (%d keys)", tag, memcacheMaxTagKeys)
|
||||
}
|
||||
|
||||
// Store updated key list
|
||||
keysData, err := json.Marshal(keys)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
tagItem := &memcache.Item{
|
||||
Key: tagKey,
|
||||
Value: keysData,
|
||||
Expiration: expiration + 3600, // Give tag lists longer TTL
|
||||
}
|
||||
_ = m.client.Set(tagItem)
|
||||
return append(keys, key), nil
|
||||
}); err != nil {
|
||||
return fail(err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// updateTagKeys applies fn to a tag's key list using compare-and-swap.
|
||||
func (m *MemcacheProvider) updateTagKeys(tagKey string, expiry int32, fn func([]string) ([]string, error)) error {
|
||||
for attempt := 0; attempt < memcacheCASRetries; attempt++ {
|
||||
item, err := m.client.Get(tagKey)
|
||||
var keys []string
|
||||
switch {
|
||||
case errors.Is(err, memcache.ErrCacheMiss):
|
||||
item = nil
|
||||
case err != nil:
|
||||
return err
|
||||
default:
|
||||
if err := json.Unmarshal(item.Value, &keys); err != nil {
|
||||
return fmt.Errorf("failed to unmarshal tag keys: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
keys, err = fn(keys)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.Marshal(keys)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if item == nil {
|
||||
err = m.client.Add(&memcache.Item{Key: tagKey, Value: data, Expiration: expiry})
|
||||
if errors.Is(err, memcache.ErrNotStored) {
|
||||
continue // someone created it first; retry
|
||||
}
|
||||
return err
|
||||
}
|
||||
item.Value = data
|
||||
item.Expiration = expiry
|
||||
err = m.client.CompareAndSwap(item)
|
||||
if errors.Is(err, memcache.ErrCASConflict) || errors.Is(err, memcache.ErrNotStored) || errors.Is(err, memcache.ErrCacheMiss) {
|
||||
continue
|
||||
}
|
||||
return err
|
||||
}
|
||||
return fmt.Errorf("tag index %q: too much contention, giving up", tagKey)
|
||||
}
|
||||
|
||||
// Delete removes a key from the cache.
|
||||
func (m *MemcacheProvider) Delete(ctx context.Context, key string) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
mkey := memcacheKey(key)
|
||||
|
||||
// Get tags for this key
|
||||
tagsKey := fmt.Sprintf("cache:tags:%s", key)
|
||||
tagsKey := memcacheTagKey("cache:tags:", key)
|
||||
if item, err := m.client.Get(tagsKey); err == nil {
|
||||
var tags []string
|
||||
if err := json.Unmarshal(item.Value, &tags); err == nil {
|
||||
// Remove key from each tag's key list
|
||||
for _, tag := range tags {
|
||||
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
||||
if tagItem, err := m.client.Get(tagKey); err == nil {
|
||||
var keys []string
|
||||
if err := json.Unmarshal(tagItem.Value, &keys); err == nil {
|
||||
// Remove current key from the list
|
||||
newKeys := make([]string, 0, len(keys))
|
||||
for _, k := range keys {
|
||||
if k != key {
|
||||
newKeys = append(newKeys, k)
|
||||
}
|
||||
}
|
||||
// Update the tag's key list
|
||||
if keysData, err := json.Marshal(newKeys); err == nil {
|
||||
tagItem.Value = keysData
|
||||
_ = m.client.Set(tagItem)
|
||||
err := m.updateTagKeys(memcacheTagKey("cache:tag:", tag), memcacheExpiry(m.options.DefaultTTL+time.Hour), func(keys []string) ([]string, error) {
|
||||
out := make([]string, 0, len(keys))
|
||||
for _, k := range keys {
|
||||
if k != key {
|
||||
out = append(out, k)
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
})
|
||||
if err != nil {
|
||||
logger.Warn("cache: failed to update memcache tag index on delete: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
// Delete the tags key
|
||||
_ = m.client.Delete(tagsKey)
|
||||
if err := m.client.Delete(tagsKey); err != nil && !errors.Is(err, memcache.ErrCacheMiss) {
|
||||
logger.Warn("cache: failed to delete memcache tags key: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Delete the actual key
|
||||
err := m.client.Delete(key)
|
||||
if err == memcache.ErrCacheMiss {
|
||||
err := m.client.Delete(mkey)
|
||||
if errors.Is(err, memcache.ErrCacheMiss) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
@@ -216,11 +324,13 @@ func (m *MemcacheProvider) Delete(ctx context.Context, key string) error {
|
||||
|
||||
// DeleteByTag removes all keys associated with the given tag.
|
||||
func (m *MemcacheProvider) DeleteByTag(ctx context.Context, tag string) error {
|
||||
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
tagKey := memcacheTagKey("cache:tag:", tag)
|
||||
|
||||
// Get all keys associated with this tag
|
||||
item, err := m.client.Get(tagKey)
|
||||
if err == memcache.ErrCacheMiss {
|
||||
if errors.Is(err, memcache.ErrCacheMiss) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
@@ -232,42 +342,51 @@ func (m *MemcacheProvider) DeleteByTag(ctx context.Context, tag string) error {
|
||||
return fmt.Errorf("failed to unmarshal tag keys: %w", err)
|
||||
}
|
||||
|
||||
// Delete all keys
|
||||
var firstErr error
|
||||
note := func(err error) {
|
||||
if err != nil && !errors.Is(err, memcache.ErrCacheMiss) && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
for _, key := range keys {
|
||||
_ = m.client.Delete(key)
|
||||
// Also delete the tags key for this cache key
|
||||
tagsKey := fmt.Sprintf("cache:tags:%s", key)
|
||||
_ = m.client.Delete(tagsKey)
|
||||
note(m.client.Delete(memcacheKey(key)))
|
||||
note(m.client.Delete(memcacheTagKey("cache:tags:", key)))
|
||||
}
|
||||
|
||||
// Delete the tag key itself
|
||||
_ = m.client.Delete(tagKey)
|
||||
|
||||
return nil
|
||||
if firstErr != nil {
|
||||
return firstErr // keep the tag index so the invalidation can be retried
|
||||
}
|
||||
note(m.client.Delete(tagKey))
|
||||
return firstErr
|
||||
}
|
||||
|
||||
// DeleteByPattern removes all keys matching the pattern.
|
||||
// Note: Memcache does not support pattern-based deletion natively.
|
||||
// This is a no-op for memcache and returns an error.
|
||||
// DeleteByPattern is not supported by Memcache; it always returns an error.
|
||||
// Use tags (SetWithTags / DeleteByTag) for group invalidation instead.
|
||||
func (m *MemcacheProvider) DeleteByPattern(ctx context.Context, pattern string) error {
|
||||
return fmt.Errorf("pattern-based deletion is not supported by Memcache")
|
||||
}
|
||||
|
||||
// Clear removes all items from the cache.
|
||||
// It runs flush_all on every configured server and therefore requires MemcacheConfig.AllowFlush.
|
||||
func (m *MemcacheProvider) Clear(ctx context.Context) error {
|
||||
if !m.allowFlush {
|
||||
return ErrFlushNotAllowed
|
||||
}
|
||||
return m.client.FlushAll()
|
||||
}
|
||||
|
||||
// Exists checks if a key exists in the cache.
|
||||
func (m *MemcacheProvider) Exists(ctx context.Context, key string) bool {
|
||||
_, err := m.client.Get(key)
|
||||
if ctx.Err() != nil {
|
||||
return false
|
||||
}
|
||||
_, err := m.client.Get(memcacheKey(key))
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// Close closes the provider and releases any resources.
|
||||
// Close closes the provider and releases idle connections.
|
||||
func (m *MemcacheProvider) Close() error {
|
||||
// Memcache client doesn't have a close method
|
||||
return nil
|
||||
return m.client.Close()
|
||||
}
|
||||
|
||||
// Stats returns statistics about the cache provider.
|
||||
|
||||
Vendored
+114
-105
@@ -2,20 +2,39 @@ package cache
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultMemoryMaxSize = 10000
|
||||
defaultMemoryCleanupInterval = time.Minute
|
||||
)
|
||||
|
||||
// ErrClosed is returned by a provider that has been closed.
|
||||
var ErrClosed = errors.New("cache: provider closed")
|
||||
|
||||
// memoryItem represents a cached item in memory.
|
||||
type memoryItem struct {
|
||||
Value []byte
|
||||
Expiration time.Time
|
||||
LastAccess time.Time
|
||||
HitCount int64
|
||||
Tags []string
|
||||
lastAccess atomic.Int64 // unix nanos
|
||||
hitCount atomic.Int64
|
||||
}
|
||||
|
||||
func newMemoryItem(value []byte, expiration time.Time, tags []string) *memoryItem {
|
||||
buf := make([]byte, len(value))
|
||||
copy(buf, value)
|
||||
item := &memoryItem{Value: buf, Expiration: expiration, Tags: tags}
|
||||
item.lastAccess.Store(time.Now().UnixNano())
|
||||
return item
|
||||
}
|
||||
|
||||
// isExpired checks if the item has expired.
|
||||
@@ -34,30 +53,72 @@ type MemoryProvider struct {
|
||||
options *Options
|
||||
hits atomic.Int64
|
||||
misses atomic.Int64
|
||||
closed bool
|
||||
done chan struct{}
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
// NewMemoryProvider creates a new in-memory cache provider.
|
||||
// A MaxSize <= 0 selects the default (10000); use MaxSize -1 for an unbounded cache.
|
||||
// A background goroutine removes expired items until Close is called.
|
||||
func NewMemoryProvider(opts *Options) *MemoryProvider {
|
||||
if opts == nil {
|
||||
opts = &Options{
|
||||
DefaultTTL: 5 * time.Minute,
|
||||
MaxSize: 10000,
|
||||
}
|
||||
var o Options
|
||||
if opts != nil {
|
||||
o = *opts // do not mutate the caller's struct
|
||||
} else {
|
||||
o = Options{DefaultTTL: 5 * time.Minute}
|
||||
}
|
||||
if o.MaxSize == 0 {
|
||||
o.MaxSize = defaultMemoryMaxSize
|
||||
}
|
||||
if o.CleanupInterval <= 0 {
|
||||
o.CleanupInterval = defaultMemoryCleanupInterval
|
||||
}
|
||||
|
||||
return &MemoryProvider{
|
||||
m := &MemoryProvider{
|
||||
items: make(map[string]*memoryItem),
|
||||
tagToKeys: make(map[string]map[string]struct{}),
|
||||
options: opts,
|
||||
options: &o,
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
go m.janitor(o.CleanupInterval)
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *MemoryProvider) janitor(interval time.Duration) {
|
||||
defer logger.CatchPanic("cache.MemoryProvider.janitor")()
|
||||
t := time.NewTicker(interval)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-m.done:
|
||||
return
|
||||
case <-t.C:
|
||||
m.CleanExpired(context.Background())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// removeLocked deletes a key and its tag associations. Caller must hold m.mu for writing.
|
||||
func (m *MemoryProvider) removeLocked(key string) {
|
||||
if item, ok := m.items[key]; ok {
|
||||
for _, tag := range item.Tags {
|
||||
if ks := m.tagToKeys[tag]; ks != nil {
|
||||
delete(ks, key)
|
||||
if len(ks) == 0 {
|
||||
delete(m.tagToKeys, tag)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
delete(m.items, key)
|
||||
}
|
||||
|
||||
// Get retrieves a value from the cache by key.
|
||||
func (m *MemoryProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
||||
// First try with read lock for fast path
|
||||
m.mu.RLock()
|
||||
item, exists := m.items[key]
|
||||
if !exists {
|
||||
if !exists || m.closed {
|
||||
m.mu.RUnlock()
|
||||
m.misses.Add(1)
|
||||
return nil, false
|
||||
@@ -65,56 +126,29 @@ func (m *MemoryProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
||||
|
||||
if item.isExpired() {
|
||||
m.mu.RUnlock()
|
||||
// Upgrade to write lock to delete expired item
|
||||
// Delete only if the entry is still the same expired one
|
||||
m.mu.Lock()
|
||||
delete(m.items, key)
|
||||
if cur, ok := m.items[key]; ok && cur == item {
|
||||
m.removeLocked(key)
|
||||
}
|
||||
m.mu.Unlock()
|
||||
m.misses.Add(1)
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// Update stats and access time with write lock
|
||||
value := item.Value
|
||||
item.lastAccess.Store(time.Now().UnixNano())
|
||||
item.hitCount.Add(1)
|
||||
out := make([]byte, len(item.Value))
|
||||
copy(out, item.Value)
|
||||
m.mu.RUnlock()
|
||||
|
||||
// Update access tracking with write lock
|
||||
m.mu.Lock()
|
||||
item.LastAccess = time.Now()
|
||||
item.HitCount++
|
||||
m.mu.Unlock()
|
||||
|
||||
m.hits.Add(1)
|
||||
return value, true
|
||||
return out, true
|
||||
}
|
||||
|
||||
// Set stores a value in the cache with the specified TTL.
|
||||
func (m *MemoryProvider) Set(ctx context.Context, key string, value []byte, ttl time.Duration) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if ttl == 0 {
|
||||
ttl = m.options.DefaultTTL
|
||||
}
|
||||
|
||||
var expiration time.Time
|
||||
if ttl > 0 {
|
||||
expiration = time.Now().Add(ttl)
|
||||
}
|
||||
|
||||
// Check max size and evict if necessary
|
||||
if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
|
||||
if _, exists := m.items[key]; !exists {
|
||||
m.evictOne()
|
||||
}
|
||||
}
|
||||
|
||||
m.items[key] = &memoryItem{
|
||||
Value: value,
|
||||
Expiration: expiration,
|
||||
LastAccess: time.Now(),
|
||||
}
|
||||
|
||||
return nil
|
||||
return m.SetWithTags(ctx, key, value, ttl, nil)
|
||||
}
|
||||
|
||||
// SetWithTags stores a value in the cache with the specified TTL and tags.
|
||||
@@ -122,6 +156,10 @@ func (m *MemoryProvider) SetWithTags(ctx context.Context, key string, value []by
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if m.closed {
|
||||
return ErrClosed
|
||||
}
|
||||
|
||||
if ttl == 0 {
|
||||
ttl = m.options.DefaultTTL
|
||||
}
|
||||
@@ -131,34 +169,14 @@ func (m *MemoryProvider) SetWithTags(ctx context.Context, key string, value []by
|
||||
expiration = time.Now().Add(ttl)
|
||||
}
|
||||
|
||||
// Check max size and evict if necessary
|
||||
if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
|
||||
if _, exists := m.items[key]; !exists {
|
||||
m.evictOne()
|
||||
}
|
||||
if _, exists := m.items[key]; exists {
|
||||
m.removeLocked(key) // drops old tag associations
|
||||
} else if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
|
||||
m.evictOne()
|
||||
}
|
||||
|
||||
// Remove old tag associations if key exists
|
||||
if oldItem, exists := m.items[key]; exists {
|
||||
for _, tag := range oldItem.Tags {
|
||||
if keySet, ok := m.tagToKeys[tag]; ok {
|
||||
delete(keySet, key)
|
||||
if len(keySet) == 0 {
|
||||
delete(m.tagToKeys, tag)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
m.items[key] = newMemoryItem(value, expiration, tags)
|
||||
|
||||
// Store the item
|
||||
m.items[key] = &memoryItem{
|
||||
Value: value,
|
||||
Expiration: expiration,
|
||||
LastAccess: time.Now(),
|
||||
Tags: tags,
|
||||
}
|
||||
|
||||
// Add new tag associations
|
||||
for _, tag := range tags {
|
||||
if m.tagToKeys[tag] == nil {
|
||||
m.tagToKeys[tag] = make(map[string]struct{})
|
||||
@@ -174,19 +192,7 @@ func (m *MemoryProvider) Delete(ctx context.Context, key string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// Remove tag associations
|
||||
if item, exists := m.items[key]; exists {
|
||||
for _, tag := range item.Tags {
|
||||
if keySet, ok := m.tagToKeys[tag]; ok {
|
||||
delete(keySet, key)
|
||||
if len(keySet) == 0 {
|
||||
delete(m.tagToKeys, tag)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
delete(m.items, key)
|
||||
m.removeLocked(key)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -195,16 +201,13 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// Get all keys associated with this tag
|
||||
keySet, exists := m.tagToKeys[tag]
|
||||
if !exists {
|
||||
return nil // No keys with this tag
|
||||
}
|
||||
|
||||
// Delete all items with this tag
|
||||
for key := range keySet {
|
||||
if item, ok := m.items[key]; ok {
|
||||
// Remove this tag from the item's tag list
|
||||
newTags := make([]string, 0, len(item.Tags))
|
||||
for _, t := range item.Tags {
|
||||
if t != tag {
|
||||
@@ -212,8 +215,7 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
|
||||
}
|
||||
}
|
||||
|
||||
// If item has no more tags, delete it
|
||||
// Otherwise update its tags
|
||||
// If item has no more tags, delete it; otherwise update its tags
|
||||
if len(newTags) == 0 {
|
||||
delete(m.items, key)
|
||||
} else {
|
||||
@@ -222,24 +224,24 @@ func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
|
||||
}
|
||||
}
|
||||
|
||||
// Remove the tag mapping
|
||||
delete(m.tagToKeys, tag)
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteByPattern removes all keys matching the pattern.
|
||||
// The pattern is a Go regular expression (unanchored); it is compiled before the lock is taken.
|
||||
func (m *MemoryProvider) DeleteByPattern(ctx context.Context, pattern string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
re, err := regexp.Compile(pattern)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid pattern: %w", err)
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
for key := range m.items {
|
||||
if re.MatchString(key) {
|
||||
delete(m.items, key)
|
||||
m.removeLocked(key)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -252,6 +254,7 @@ func (m *MemoryProvider) Clear(ctx context.Context) error {
|
||||
defer m.mu.Unlock()
|
||||
|
||||
m.items = make(map[string]*memoryItem)
|
||||
m.tagToKeys = make(map[string]map[string]struct{})
|
||||
m.hits.Store(0)
|
||||
m.misses.Store(0)
|
||||
return nil
|
||||
@@ -270,12 +273,17 @@ func (m *MemoryProvider) Exists(ctx context.Context, key string) bool {
|
||||
return !item.isExpired()
|
||||
}
|
||||
|
||||
// Close closes the provider and releases any resources.
|
||||
// Close closes the provider, stops the janitor and releases stored items.
|
||||
// Later writes return ErrClosed and reads report a miss.
|
||||
func (m *MemoryProvider) Close() error {
|
||||
m.closeOnce.Do(func() { close(m.done) })
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
m.items = nil
|
||||
m.closed = true
|
||||
m.items = make(map[string]*memoryItem)
|
||||
m.tagToKeys = make(map[string]map[string]struct{})
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -284,7 +292,7 @@ func (m *MemoryProvider) Stats(ctx context.Context) (*CacheStats, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
// Clean expired items first
|
||||
// Count non-expired items (read-only)
|
||||
validKeys := 0
|
||||
for _, item := range m.items {
|
||||
if !item.isExpired() {
|
||||
@@ -304,24 +312,25 @@ func (m *MemoryProvider) Stats(ctx context.Context) (*CacheStats, error) {
|
||||
}
|
||||
|
||||
// evictOne removes one item from the cache using LRU strategy.
|
||||
// Note: this is an O(n) scan. Caller must hold m.mu for writing.
|
||||
func (m *MemoryProvider) evictOne() {
|
||||
var oldestKey string
|
||||
var oldestTime time.Time
|
||||
var oldest int64
|
||||
|
||||
for key, item := range m.items {
|
||||
if item.isExpired() {
|
||||
delete(m.items, key)
|
||||
m.removeLocked(key)
|
||||
return
|
||||
}
|
||||
|
||||
if oldestKey == "" || item.LastAccess.Before(oldestTime) {
|
||||
if la := item.lastAccess.Load(); oldestKey == "" || la < oldest {
|
||||
oldestKey = key
|
||||
oldestTime = item.LastAccess
|
||||
oldest = la
|
||||
}
|
||||
}
|
||||
|
||||
if oldestKey != "" {
|
||||
delete(m.items, oldestKey)
|
||||
m.removeLocked(oldestKey)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -333,7 +342,7 @@ func (m *MemoryProvider) CleanExpired(ctx context.Context) int {
|
||||
count := 0
|
||||
for key, item := range m.items {
|
||||
if item.isExpired() {
|
||||
delete(m.items, key)
|
||||
m.removeLocked(key)
|
||||
count++
|
||||
}
|
||||
}
|
||||
|
||||
Vendored
+52
-15
@@ -3,15 +3,20 @@ package cache
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// RedisProvider is a Redis implementation of the Provider interface.
|
||||
type RedisProvider struct {
|
||||
client *redis.Client
|
||||
options *Options
|
||||
client *redis.Client
|
||||
options *Options
|
||||
allowFlush bool
|
||||
}
|
||||
|
||||
// RedisConfig contains Redis-specific configuration.
|
||||
@@ -33,16 +38,25 @@ type RedisConfig struct {
|
||||
|
||||
// Options contains general cache options
|
||||
Options *Options
|
||||
|
||||
// AllowFlush permits Clear() to run FLUSHDB, which wipes the entire logical Redis DB
|
||||
// (including data that is not owned by this cache). Off by default.
|
||||
AllowFlush bool
|
||||
}
|
||||
|
||||
// NewRedisProvider creates a new Redis cache provider.
|
||||
func NewRedisProvider(config *RedisConfig) (*RedisProvider, error) {
|
||||
if config == nil {
|
||||
config = &RedisConfig{
|
||||
Host: "localhost",
|
||||
Port: 6379,
|
||||
DB: 0,
|
||||
}
|
||||
// Work on a copy so the caller's struct is not mutated
|
||||
var cfg RedisConfig
|
||||
if config != nil {
|
||||
cfg = *config
|
||||
} else {
|
||||
cfg = RedisConfig{Host: "localhost", Port: 6379, DB: 0}
|
||||
}
|
||||
config = &cfg
|
||||
if config.Options != nil {
|
||||
o := *config.Options
|
||||
config.Options = &o
|
||||
}
|
||||
|
||||
if config.Host == "" {
|
||||
@@ -77,8 +91,9 @@ func NewRedisProvider(config *RedisConfig) (*RedisProvider, error) {
|
||||
}
|
||||
|
||||
return &RedisProvider{
|
||||
client: client,
|
||||
options: config.Options,
|
||||
client: client,
|
||||
options: config.Options,
|
||||
allowFlush: config.AllowFlush,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -89,6 +104,8 @@ func (r *RedisProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
||||
return nil, false
|
||||
}
|
||||
if err != nil {
|
||||
// Reported as a miss (the Provider interface cannot express errors), but not silently
|
||||
logger.Warn("cache: redis GET failed: %v", err)
|
||||
return nil, false
|
||||
}
|
||||
return val, true
|
||||
@@ -194,7 +211,7 @@ func (r *RedisProvider) DeleteByTag(ctx context.Context, tag string) error {
|
||||
|
||||
// DeleteByPattern removes all keys matching the pattern.
|
||||
func (r *RedisProvider) DeleteByPattern(ctx context.Context, pattern string) error {
|
||||
iter := r.client.Scan(ctx, 0, pattern, 0).Iterator()
|
||||
iter := r.client.Scan(ctx, 0, pattern, 500).Iterator()
|
||||
pipe := r.client.Pipeline()
|
||||
|
||||
count := 0
|
||||
@@ -225,7 +242,11 @@ func (r *RedisProvider) DeleteByPattern(ctx context.Context, pattern string) err
|
||||
}
|
||||
|
||||
// Clear removes all items from the cache.
|
||||
// It runs FLUSHDB and therefore requires RedisConfig.AllowFlush.
|
||||
func (r *RedisProvider) Clear(ctx context.Context) error {
|
||||
if !r.allowFlush {
|
||||
return ErrFlushNotAllowed
|
||||
}
|
||||
return r.client.FlushDB(ctx).Err()
|
||||
}
|
||||
|
||||
@@ -244,8 +265,9 @@ func (r *RedisProvider) Close() error {
|
||||
}
|
||||
|
||||
// Stats returns statistics about the cache provider.
|
||||
// Only an allowlist of numeric counters from INFO is exposed, not the raw output.
|
||||
func (r *RedisProvider) Stats(ctx context.Context) (*CacheStats, error) {
|
||||
info, err := r.client.Info(ctx, "stats", "keyspace").Result()
|
||||
info, err := r.client.Info(ctx, "stats").Result()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get Redis stats: %w", err)
|
||||
}
|
||||
@@ -255,13 +277,28 @@ func (r *RedisProvider) Stats(ctx context.Context) (*CacheStats, error) {
|
||||
return nil, fmt.Errorf("failed to get DB size: %w", err)
|
||||
}
|
||||
|
||||
// Parse stats from INFO command
|
||||
// This is a simplified version - you may want to parse more detailed stats
|
||||
counters := map[string]int64{}
|
||||
for _, line := range strings.Split(info, "\n") {
|
||||
k, v, ok := strings.Cut(strings.TrimSpace(line), ":")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
switch k {
|
||||
case "keyspace_hits", "keyspace_misses", "evicted_keys", "expired_keys":
|
||||
if n, err := strconv.ParseInt(v, 10, 64); err == nil {
|
||||
counters[k] = n
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
stats := &CacheStats{
|
||||
Hits: counters["keyspace_hits"],
|
||||
Misses: counters["keyspace_misses"],
|
||||
Keys: dbSize,
|
||||
ProviderType: "redis",
|
||||
ProviderStats: map[string]any{
|
||||
"info": info,
|
||||
"evicted_keys": counters["evicted_keys"],
|
||||
"expired_keys": counters["expired_keys"],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
+883
-339
File diff suppressed because it is too large
Load Diff
@@ -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,9 +3,13 @@ package database
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
@@ -15,18 +19,93 @@ import (
|
||||
|
||||
// GormAdapter adapts GORM to work with our Database interface
|
||||
type GormAdapter struct {
|
||||
db *gorm.DB
|
||||
dbMu sync.RWMutex
|
||||
db *gorm.DB
|
||||
dbFactory func() (*gorm.DB, error)
|
||||
driverName string
|
||||
metricsEnabled bool
|
||||
}
|
||||
|
||||
// NewGormAdapter creates a new GORM adapter
|
||||
func NewGormAdapter(db *gorm.DB) *GormAdapter {
|
||||
return &GormAdapter{db: db}
|
||||
adapter := &GormAdapter{db: db, metricsEnabled: true}
|
||||
// Initialize driver name
|
||||
adapter.driverName = adapter.DriverName()
|
||||
return adapter
|
||||
}
|
||||
|
||||
// WithDBFactory configures a factory used to reopen the database connection if it is closed.
|
||||
func (g *GormAdapter) WithDBFactory(factory func() (*gorm.DB, error)) *GormAdapter {
|
||||
g.dbFactory = factory
|
||||
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 {
|
||||
g.dbMu.RLock()
|
||||
defer g.dbMu.RUnlock()
|
||||
return g.db
|
||||
}
|
||||
|
||||
func (g *GormAdapter) reconnectDB(targets ...*gorm.DB) error {
|
||||
if g.dbFactory == nil {
|
||||
return fmt.Errorf("no db factory configured for reconnect")
|
||||
}
|
||||
|
||||
freshDB, err := g.dbFactory()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
g.dbMu.Lock()
|
||||
previous := g.db
|
||||
g.db = freshDB
|
||||
g.driverName = normalizeGormDriverName(freshDB)
|
||||
g.dbMu.Unlock()
|
||||
|
||||
if previous != nil {
|
||||
syncGormConnPool(previous, freshDB)
|
||||
}
|
||||
|
||||
for _, target := range targets {
|
||||
if target != nil && target != previous {
|
||||
syncGormConnPool(target, freshDB)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func syncGormConnPool(target, fresh *gorm.DB) {
|
||||
if target == nil || fresh == nil {
|
||||
return
|
||||
}
|
||||
|
||||
if target.Config != nil && fresh.Config != nil {
|
||||
target.ConnPool = fresh.ConnPool
|
||||
}
|
||||
|
||||
if target.Statement != nil {
|
||||
if fresh.Statement != nil && fresh.Statement.ConnPool != nil {
|
||||
target.Statement.ConnPool = fresh.Statement.ConnPool
|
||||
} else if fresh.Config != nil {
|
||||
target.Statement.ConnPool = fresh.ConnPool
|
||||
}
|
||||
target.Statement.DB = target
|
||||
}
|
||||
}
|
||||
|
||||
// EnableQueryDebug enables query debugging which logs all SQL queries including preloads
|
||||
// This is useful for debugging preload queries that may be failing
|
||||
func (g *GormAdapter) EnableQueryDebug() *GormAdapter {
|
||||
g.dbMu.Lock()
|
||||
g.db = g.db.Debug()
|
||||
g.dbMu.Unlock()
|
||||
logger.Info("GORM query debug mode enabled - all SQL queries will be logged")
|
||||
return g
|
||||
}
|
||||
@@ -40,19 +119,19 @@ func (g *GormAdapter) DisableQueryDebug() *GormAdapter {
|
||||
}
|
||||
|
||||
func (g *GormAdapter) NewSelect() common.SelectQuery {
|
||||
return &GormSelectQuery{db: g.db}
|
||||
return &GormSelectQuery{db: g.getDB(), driverName: g.driverName, reconnect: g.reconnectDB, metricsEnabled: g.metricsEnabled}
|
||||
}
|
||||
|
||||
func (g *GormAdapter) NewInsert() common.InsertQuery {
|
||||
return &GormInsertQuery{db: g.db}
|
||||
return &GormInsertQuery{db: g.getDB(), reconnect: g.reconnectDB, driverName: g.driverName, metricsEnabled: g.metricsEnabled}
|
||||
}
|
||||
|
||||
func (g *GormAdapter) NewUpdate() common.UpdateQuery {
|
||||
return &GormUpdateQuery{db: g.db}
|
||||
return &GormUpdateQuery{db: g.getDB(), reconnect: g.reconnectDB, driverName: g.driverName, metricsEnabled: g.metricsEnabled}
|
||||
}
|
||||
|
||||
func (g *GormAdapter) NewDelete() common.DeleteQuery {
|
||||
return &GormDeleteQuery{db: g.db}
|
||||
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) {
|
||||
@@ -61,7 +140,18 @@ func (g *GormAdapter) Exec(ctx context.Context, query string, args ...interface{
|
||||
err = logger.HandlePanic("GormAdapter.Exec", r)
|
||||
}
|
||||
}()
|
||||
result := g.db.WithContext(ctx).Exec(query, args...)
|
||||
startedAt := time.Now()
|
||||
operation, schema, entity, table := metricTargetFromRawQuery(query, g.driverName)
|
||||
run := func() *gorm.DB {
|
||||
return g.getDB().WithContext(ctx).Exec(query, args...)
|
||||
}
|
||||
result := run()
|
||||
if isDBClosed(result.Error) {
|
||||
if reconnErr := g.reconnectDB(); reconnErr == nil {
|
||||
result = run()
|
||||
}
|
||||
}
|
||||
recordQueryMetrics(g.metricsEnabled, operation, schema, entity, table, startedAt, result.Error)
|
||||
return &GormResult{result: result}, result.Error
|
||||
}
|
||||
|
||||
@@ -71,15 +161,35 @@ func (g *GormAdapter) Query(ctx context.Context, dest interface{}, query string,
|
||||
err = logger.HandlePanic("GormAdapter.Query", r)
|
||||
}
|
||||
}()
|
||||
return g.db.WithContext(ctx).Raw(query, args...).Find(dest).Error
|
||||
startedAt := time.Now()
|
||||
operation, schema, entity, table := metricTargetFromRawQuery(query, g.driverName)
|
||||
run := func() error {
|
||||
return g.getDB().WithContext(ctx).Raw(query, args...).Find(dest).Error
|
||||
}
|
||||
err = run()
|
||||
if isDBClosed(err) {
|
||||
if reconnErr := g.reconnectDB(); reconnErr == nil {
|
||||
err = run()
|
||||
}
|
||||
}
|
||||
recordQueryMetrics(g.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||
return err
|
||||
}
|
||||
|
||||
func (g *GormAdapter) BeginTx(ctx context.Context) (common.Database, error) {
|
||||
tx := g.db.WithContext(ctx).Begin()
|
||||
run := func() *gorm.DB {
|
||||
return g.getDB().WithContext(ctx).Begin()
|
||||
}
|
||||
tx := run()
|
||||
if isDBClosed(tx.Error) {
|
||||
if reconnErr := g.reconnectDB(); reconnErr == nil {
|
||||
tx = run()
|
||||
}
|
||||
}
|
||||
if tx.Error != nil {
|
||||
return nil, tx.Error
|
||||
}
|
||||
return &GormAdapter{db: tx}, nil
|
||||
return &GormAdapter{db: tx, dbFactory: g.dbFactory, driverName: g.driverName, metricsEnabled: g.metricsEnabled}, nil
|
||||
}
|
||||
|
||||
func (g *GormAdapter) CommitTx(ctx context.Context) error {
|
||||
@@ -96,35 +206,64 @@ func (g *GormAdapter) RunInTransaction(ctx context.Context, fn func(common.Datab
|
||||
err = logger.HandlePanic("GormAdapter.RunInTransaction", r)
|
||||
}
|
||||
}()
|
||||
return g.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
adapter := &GormAdapter{db: tx}
|
||||
return fn(adapter)
|
||||
})
|
||||
run := func() error {
|
||||
return g.getDB().WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
adapter := &GormAdapter{db: tx, dbFactory: g.dbFactory, driverName: g.driverName, metricsEnabled: g.metricsEnabled}
|
||||
return fn(adapter)
|
||||
})
|
||||
}
|
||||
err = run()
|
||||
if isDBClosed(err) {
|
||||
if reconnErr := g.reconnectDB(); reconnErr == nil {
|
||||
err = run()
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (g *GormAdapter) GetUnderlyingDB() interface{} {
|
||||
return g.db
|
||||
return g.getDB()
|
||||
}
|
||||
|
||||
func (g *GormAdapter) DriverName() string {
|
||||
return normalizeGormDriverName(g.getDB())
|
||||
}
|
||||
|
||||
func normalizeGormDriverName(db *gorm.DB) string {
|
||||
if db == nil || db.Dialector == nil {
|
||||
return ""
|
||||
}
|
||||
// Normalize GORM's dialector name to match the project's canonical vocabulary.
|
||||
// GORM returns "sqlserver" for MSSQL; the rest of the project uses "mssql".
|
||||
// GORM returns "sqlite" or "sqlite3" for SQLite; we normalize to "sqlite".
|
||||
switch name := db.Name(); name {
|
||||
case "sqlserver":
|
||||
return "mssql"
|
||||
case "sqlite3":
|
||||
return "sqlite"
|
||||
default:
|
||||
return name
|
||||
}
|
||||
}
|
||||
|
||||
// GormSelectQuery implements SelectQuery for GORM
|
||||
type GormSelectQuery struct {
|
||||
db *gorm.DB
|
||||
reconnect func(...*gorm.DB) error
|
||||
schema string // Separated schema name
|
||||
tableName string // Just the table name, without schema
|
||||
entity string
|
||||
tableAlias string
|
||||
driverName string // Database driver name (postgres, sqlite, mssql)
|
||||
inJoinContext bool // Track if we're in a JOIN relation context
|
||||
joinTableAlias string // Alias to use for JOIN conditions
|
||||
metricsEnabled bool
|
||||
}
|
||||
|
||||
func (g *GormSelectQuery) Model(model interface{}) common.SelectQuery {
|
||||
g.db = g.db.Model(model)
|
||||
|
||||
// Try to get table name from model if it implements TableNameProvider
|
||||
if provider, ok := model.(common.TableNameProvider); ok {
|
||||
fullTableName := provider.TableName()
|
||||
// Check if the table name contains schema (e.g., "schema.table")
|
||||
g.schema, g.tableName = parseTableName(fullTableName)
|
||||
}
|
||||
g.schema, g.tableName = schemaAndTableFromModel(model, g.driverName)
|
||||
g.entity = entityNameFromModel(model, g.tableName)
|
||||
|
||||
if provider, ok := model.(common.TableAliasProvider); ok {
|
||||
g.tableAlias = provider.TableAlias()
|
||||
@@ -136,7 +275,11 @@ func (g *GormSelectQuery) Model(model interface{}) common.SelectQuery {
|
||||
func (g *GormSelectQuery) Table(table string) common.SelectQuery {
|
||||
g.db = g.db.Table(table)
|
||||
// Check if the table name contains schema (e.g., "schema.table")
|
||||
g.schema, g.tableName = parseTableName(table)
|
||||
// For SQLite, this will convert "schema.table" to "schema_table"
|
||||
g.schema, g.tableName = parseTableName(table, g.driverName)
|
||||
if g.entity == "" {
|
||||
g.entity = cleanMetricIdentifier(g.tableName)
|
||||
}
|
||||
|
||||
return g
|
||||
}
|
||||
@@ -322,7 +465,10 @@ func (g *GormSelectQuery) PreloadRelation(relation string, apply ...func(common.
|
||||
}
|
||||
|
||||
wrapper := &GormSelectQuery{
|
||||
db: db,
|
||||
db: db,
|
||||
reconnect: g.reconnect,
|
||||
driverName: g.driverName,
|
||||
metricsEnabled: g.metricsEnabled,
|
||||
}
|
||||
|
||||
current := common.SelectQuery(wrapper)
|
||||
@@ -360,8 +506,11 @@ func (g *GormSelectQuery) JoinRelation(relation string, apply ...func(common.Sel
|
||||
|
||||
wrapper := &GormSelectQuery{
|
||||
db: db,
|
||||
reconnect: g.reconnect,
|
||||
driverName: g.driverName,
|
||||
inJoinContext: true, // Mark as JOIN context
|
||||
joinTableAlias: strings.ToLower(relation), // Use relation name as alias
|
||||
metricsEnabled: g.metricsEnabled,
|
||||
}
|
||||
current := common.SelectQuery(wrapper)
|
||||
|
||||
@@ -418,14 +567,25 @@ func (g *GormSelectQuery) Scan(ctx context.Context, dest interface{}) (err error
|
||||
err = logger.HandlePanic("GormSelectQuery.Scan", r)
|
||||
}
|
||||
}()
|
||||
err = g.db.WithContext(ctx).Find(dest).Error
|
||||
startedAt := time.Now()
|
||||
run := func() error {
|
||||
return g.db.WithContext(ctx).Find(dest).Error
|
||||
}
|
||||
err = run()
|
||||
if isDBClosed(err) && g.reconnect != nil {
|
||||
if reconnErr := g.reconnect(g.db); reconnErr == nil {
|
||||
err = run()
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
// Log SQL string for debugging
|
||||
sqlStr := g.db.ToSQL(func(tx *gorm.DB) *gorm.DB {
|
||||
return tx.Find(dest)
|
||||
})
|
||||
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
|
||||
}
|
||||
|
||||
@@ -438,14 +598,25 @@ func (g *GormSelectQuery) ScanModel(ctx context.Context) (err error) {
|
||||
if g.db.Statement.Model == nil {
|
||||
return fmt.Errorf("ScanModel requires Model() to be set before scanning")
|
||||
}
|
||||
err = g.db.WithContext(ctx).Find(g.db.Statement.Model).Error
|
||||
startedAt := time.Now()
|
||||
run := func() error {
|
||||
return g.db.WithContext(ctx).Find(g.db.Statement.Model).Error
|
||||
}
|
||||
err = run()
|
||||
if isDBClosed(err) && g.reconnect != nil {
|
||||
if reconnErr := g.reconnect(g.db); reconnErr == nil {
|
||||
err = run()
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
// Log SQL string for debugging
|
||||
sqlStr := g.db.ToSQL(func(tx *gorm.DB) *gorm.DB {
|
||||
return tx.Find(g.db.Statement.Model)
|
||||
})
|
||||
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
|
||||
}
|
||||
|
||||
@@ -456,15 +627,26 @@ func (g *GormSelectQuery) Count(ctx context.Context) (count int, err error) {
|
||||
count = 0
|
||||
}
|
||||
}()
|
||||
startedAt := time.Now()
|
||||
var count64 int64
|
||||
err = g.db.WithContext(ctx).Count(&count64).Error
|
||||
run := func() error {
|
||||
return g.db.WithContext(ctx).Count(&count64).Error
|
||||
}
|
||||
err = run()
|
||||
if isDBClosed(err) && g.reconnect != nil {
|
||||
if reconnErr := g.reconnect(g.db); reconnErr == nil {
|
||||
err = run()
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
// Log SQL string for debugging
|
||||
sqlStr := g.db.ToSQL(func(tx *gorm.DB) *gorm.DB {
|
||||
return tx.Count(&count64)
|
||||
})
|
||||
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
|
||||
}
|
||||
|
||||
@@ -475,33 +657,57 @@ func (g *GormSelectQuery) Exists(ctx context.Context) (exists bool, err error) {
|
||||
exists = false
|
||||
}
|
||||
}()
|
||||
startedAt := time.Now()
|
||||
var count int64
|
||||
err = g.db.WithContext(ctx).Limit(1).Count(&count).Error
|
||||
run := func() error {
|
||||
return g.db.WithContext(ctx).Limit(1).Count(&count).Error
|
||||
}
|
||||
err = run()
|
||||
if isDBClosed(err) && g.reconnect != nil {
|
||||
if reconnErr := g.reconnect(g.db); reconnErr == nil {
|
||||
err = run()
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
// Log SQL string for debugging
|
||||
sqlStr := g.db.ToSQL(func(tx *gorm.DB) *gorm.DB {
|
||||
return tx.Limit(1).Count(&count)
|
||||
})
|
||||
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
|
||||
}
|
||||
|
||||
// GormInsertQuery implements InsertQuery for GORM
|
||||
type GormInsertQuery struct {
|
||||
db *gorm.DB
|
||||
model interface{}
|
||||
values map[string]interface{}
|
||||
db *gorm.DB
|
||||
reconnect func(...*gorm.DB) error
|
||||
model 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 {
|
||||
g.model = model
|
||||
g.db = g.db.Model(model)
|
||||
g.schema, g.tableName = schemaAndTableFromModel(model, g.driverName)
|
||||
g.entity = entityNameFromModel(model, g.tableName)
|
||||
return g
|
||||
}
|
||||
|
||||
func (g *GormInsertQuery) Table(table string) common.InsertQuery {
|
||||
g.db = g.db.Table(table)
|
||||
g.schema, g.tableName = parseTableName(table, g.driverName)
|
||||
if g.entity == "" {
|
||||
g.entity = cleanMetricIdentifier(g.tableName)
|
||||
}
|
||||
return g
|
||||
}
|
||||
|
||||
@@ -519,7 +725,7 @@ func (g *GormInsertQuery) OnConflict(action 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
|
||||
}
|
||||
|
||||
@@ -529,38 +735,130 @@ func (g *GormInsertQuery) Exec(ctx context.Context) (res common.Result, err erro
|
||||
err = logger.HandlePanic("GormInsertQuery.Exec", r)
|
||||
}
|
||||
}()
|
||||
startedAt := time.Now()
|
||||
run := func() *gorm.DB {
|
||||
switch {
|
||||
case g.model != nil:
|
||||
return g.db.WithContext(ctx).Create(g.model)
|
||||
case g.values != nil:
|
||||
return g.db.WithContext(ctx).Create(g.values)
|
||||
default:
|
||||
return g.db.WithContext(ctx).Create(map[string]interface{}{})
|
||||
}
|
||||
}
|
||||
result := run()
|
||||
if isDBClosed(result.Error) && g.reconnect != nil {
|
||||
if reconnErr := g.reconnect(g.db); reconnErr == nil {
|
||||
result = run()
|
||||
}
|
||||
}
|
||||
recordQueryMetrics(g.metricsEnabled, "INSERT", g.schema, g.entity, g.tableName, startedAt, 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 = g.db.WithContext(ctx).Create(g.model)
|
||||
result = db.Create(g.model)
|
||||
case g.values != nil:
|
||||
result = g.db.WithContext(ctx).Create(g.values)
|
||||
result = db.Create(g.values)
|
||||
default:
|
||||
result = g.db.WithContext(ctx).Create(map[string]interface{}{})
|
||||
result = db.Create(map[string]interface{}{})
|
||||
}
|
||||
return &GormResult{result: result}, result.Error
|
||||
|
||||
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
|
||||
type GormUpdateQuery struct {
|
||||
db *gorm.DB
|
||||
model interface{}
|
||||
updates interface{}
|
||||
db *gorm.DB
|
||||
reconnect func(...*gorm.DB) error
|
||||
model interface{}
|
||||
updates interface{}
|
||||
schema string
|
||||
tableName string
|
||||
entity string
|
||||
driverName string
|
||||
metricsEnabled bool
|
||||
}
|
||||
|
||||
func (g *GormUpdateQuery) Model(model interface{}) common.UpdateQuery {
|
||||
g.model = model
|
||||
g.db = g.db.Model(model)
|
||||
g.schema, g.tableName = schemaAndTableFromModel(model, g.driverName)
|
||||
g.entity = entityNameFromModel(model, g.tableName)
|
||||
return g
|
||||
}
|
||||
|
||||
func (g *GormUpdateQuery) Table(table string) common.UpdateQuery {
|
||||
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 {
|
||||
// Try to get table name from table string if model is not set
|
||||
model, err := modelregistry.GetModelByName(table)
|
||||
if err == nil {
|
||||
g.model = model
|
||||
g.entity = entityNameFromModel(model, g.tableName)
|
||||
}
|
||||
}
|
||||
return g
|
||||
@@ -621,31 +919,54 @@ func (g *GormUpdateQuery) Exec(ctx context.Context) (res common.Result, err erro
|
||||
err = logger.HandlePanic("GormUpdateQuery.Exec", r)
|
||||
}
|
||||
}()
|
||||
result := g.db.WithContext(ctx).Updates(g.updates)
|
||||
startedAt := time.Now()
|
||||
run := func() *gorm.DB {
|
||||
return g.db.WithContext(ctx).Updates(g.updates)
|
||||
}
|
||||
result := run()
|
||||
if isDBClosed(result.Error) && g.reconnect != nil {
|
||||
if reconnErr := g.reconnect(g.db); reconnErr == nil {
|
||||
result = run()
|
||||
}
|
||||
}
|
||||
if result.Error != nil {
|
||||
// Log SQL string for debugging
|
||||
sqlStr := g.db.ToSQL(func(tx *gorm.DB) *gorm.DB {
|
||||
return tx.Updates(g.updates)
|
||||
})
|
||||
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
|
||||
}
|
||||
|
||||
// GormDeleteQuery implements DeleteQuery for GORM
|
||||
type GormDeleteQuery struct {
|
||||
db *gorm.DB
|
||||
model interface{}
|
||||
db *gorm.DB
|
||||
reconnect func(...*gorm.DB) error
|
||||
model interface{}
|
||||
schema string
|
||||
tableName string
|
||||
entity string
|
||||
driverName string
|
||||
metricsEnabled bool
|
||||
}
|
||||
|
||||
func (g *GormDeleteQuery) Model(model interface{}) common.DeleteQuery {
|
||||
g.model = model
|
||||
g.db = g.db.Model(model)
|
||||
g.schema, g.tableName = schemaAndTableFromModel(model, g.driverName)
|
||||
g.entity = entityNameFromModel(model, g.tableName)
|
||||
return g
|
||||
}
|
||||
|
||||
func (g *GormDeleteQuery) Table(table string) common.DeleteQuery {
|
||||
g.db = g.db.Table(table)
|
||||
g.schema, g.tableName = parseTableName(table, g.driverName)
|
||||
if g.entity == "" {
|
||||
g.entity = cleanMetricIdentifier(g.tableName)
|
||||
}
|
||||
return g
|
||||
}
|
||||
|
||||
@@ -660,14 +981,25 @@ func (g *GormDeleteQuery) Exec(ctx context.Context) (res common.Result, err erro
|
||||
err = logger.HandlePanic("GormDeleteQuery.Exec", r)
|
||||
}
|
||||
}()
|
||||
result := g.db.WithContext(ctx).Delete(g.model)
|
||||
startedAt := time.Now()
|
||||
run := func() *gorm.DB {
|
||||
return g.db.WithContext(ctx).Delete(g.model)
|
||||
}
|
||||
result := run()
|
||||
if isDBClosed(result.Error) && g.reconnect != nil {
|
||||
if reconnErr := g.reconnect(g.db); reconnErr == nil {
|
||||
result = run()
|
||||
}
|
||||
}
|
||||
if result.Error != nil {
|
||||
// Log SQL string for debugging
|
||||
sqlStr := g.db.ToSQL(func(tx *gorm.DB) *gorm.DB {
|
||||
return tx.Delete(g.model)
|
||||
})
|
||||
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
|
||||
}
|
||||
|
||||
|
||||
@@ -5,7 +5,10 @@ import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
@@ -16,12 +19,58 @@ import (
|
||||
// PgSQLAdapter adapts standard database/sql to work with our Database interface
|
||||
// This provides a lightweight PostgreSQL adapter without ORM overhead
|
||||
type PgSQLAdapter struct {
|
||||
db *sql.DB
|
||||
db *sql.DB
|
||||
dbMu sync.RWMutex
|
||||
dbFactory func() (*sql.DB, error)
|
||||
driverName string
|
||||
metricsEnabled bool
|
||||
}
|
||||
|
||||
// NewPgSQLAdapter creates a new PostgreSQL adapter
|
||||
func NewPgSQLAdapter(db *sql.DB) *PgSQLAdapter {
|
||||
return &PgSQLAdapter{db: db}
|
||||
// NewPgSQLAdapter creates a new adapter wrapping a standard sql.DB.
|
||||
// An optional driverName (e.g. "postgres", "sqlite", "mssql") can be provided;
|
||||
// it defaults to "postgres" when omitted.
|
||||
func NewPgSQLAdapter(db *sql.DB, driverName ...string) *PgSQLAdapter {
|
||||
name := "postgres"
|
||||
if len(driverName) > 0 && driverName[0] != "" {
|
||||
name = driverName[0]
|
||||
}
|
||||
return &PgSQLAdapter{db: db, driverName: name, metricsEnabled: true}
|
||||
}
|
||||
|
||||
// WithDBFactory configures a factory used to reopen the database connection if it is closed.
|
||||
func (p *PgSQLAdapter) WithDBFactory(factory func() (*sql.DB, error)) *PgSQLAdapter {
|
||||
p.dbFactory = factory
|
||||
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 {
|
||||
p.dbMu.RLock()
|
||||
defer p.dbMu.RUnlock()
|
||||
return p.db
|
||||
}
|
||||
|
||||
func (p *PgSQLAdapter) 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
|
||||
}
|
||||
|
||||
func isDBClosed(err error) bool {
|
||||
return err != nil && strings.Contains(err.Error(), "sql: database is closed")
|
||||
}
|
||||
|
||||
// EnableQueryDebug enables query debugging for development
|
||||
@@ -31,33 +80,41 @@ func (p *PgSQLAdapter) EnableQueryDebug() {
|
||||
|
||||
func (p *PgSQLAdapter) NewSelect() common.SelectQuery {
|
||||
return &PgSQLSelectQuery{
|
||||
db: p.db,
|
||||
columns: []string{"*"},
|
||||
args: make([]interface{}, 0),
|
||||
db: p.getDB(),
|
||||
driverName: p.driverName,
|
||||
columns: []string{"*"},
|
||||
args: make([]interface{}, 0),
|
||||
metricsEnabled: p.metricsEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PgSQLAdapter) NewInsert() common.InsertQuery {
|
||||
return &PgSQLInsertQuery{
|
||||
db: p.db,
|
||||
values: make(map[string]interface{}),
|
||||
db: p.getDB(),
|
||||
driverName: p.driverName,
|
||||
values: make(map[string]interface{}),
|
||||
metricsEnabled: p.metricsEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PgSQLAdapter) NewUpdate() common.UpdateQuery {
|
||||
return &PgSQLUpdateQuery{
|
||||
db: p.db,
|
||||
sets: make(map[string]interface{}),
|
||||
args: make([]interface{}, 0),
|
||||
whereClauses: make([]string, 0),
|
||||
db: p.getDB(),
|
||||
driverName: p.driverName,
|
||||
sets: make(map[string]interface{}),
|
||||
args: make([]interface{}, 0),
|
||||
whereClauses: make([]string, 0),
|
||||
metricsEnabled: p.metricsEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PgSQLAdapter) NewDelete() common.DeleteQuery {
|
||||
return &PgSQLDeleteQuery{
|
||||
db: p.db,
|
||||
args: make([]interface{}, 0),
|
||||
whereClauses: make([]string, 0),
|
||||
db: p.getDB(),
|
||||
driverName: p.driverName,
|
||||
args: make([]interface{}, 0),
|
||||
whereClauses: make([]string, 0),
|
||||
metricsEnabled: p.metricsEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -67,12 +124,23 @@ func (p *PgSQLAdapter) Exec(ctx context.Context, query string, args ...interface
|
||||
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)
|
||||
result, err := p.db.ExecContext(ctx, query, args...)
|
||||
var result sql.Result
|
||||
run := func() error { var e error; result, e = p.getDB().ExecContext(ctx, query, args...); return e }
|
||||
err = run()
|
||||
if isDBClosed(err) {
|
||||
if reconnErr := p.reconnectDB(); reconnErr == nil {
|
||||
err = run()
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -82,23 +150,35 @@ func (p *PgSQLAdapter) Query(ctx context.Context, dest interface{}, query string
|
||||
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)
|
||||
rows, err := p.db.QueryContext(ctx, query, args...)
|
||||
var rows *sql.Rows
|
||||
run := func() error { var e error; rows, e = p.getDB().QueryContext(ctx, query, args...); return e }
|
||||
err = run()
|
||||
if isDBClosed(err) {
|
||||
if reconnErr := p.reconnectDB(); reconnErr == nil {
|
||||
err = run()
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
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()
|
||||
|
||||
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) {
|
||||
tx, err := p.db.BeginTx(ctx, nil)
|
||||
tx, err := p.getDB().BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &PgSQLTxAdapter{tx: tx}, nil
|
||||
return &PgSQLTxAdapter{tx: tx, driverName: p.driverName, metricsEnabled: p.metricsEnabled}, nil
|
||||
}
|
||||
|
||||
func (p *PgSQLAdapter) CommitTx(ctx context.Context) error {
|
||||
@@ -116,12 +196,12 @@ func (p *PgSQLAdapter) RunInTransaction(ctx context.Context, fn func(common.Data
|
||||
}
|
||||
}()
|
||||
|
||||
tx, err := p.db.BeginTx(ctx, nil)
|
||||
tx, err := p.getDB().BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
adapter := &PgSQLTxAdapter{tx: tx}
|
||||
adapter := &PgSQLTxAdapter{tx: tx, driverName: p.driverName, metricsEnabled: p.metricsEnabled}
|
||||
|
||||
defer func() {
|
||||
if p := recover(); p != nil {
|
||||
@@ -141,6 +221,10 @@ func (p *PgSQLAdapter) GetUnderlyingDB() interface{} {
|
||||
return p.db
|
||||
}
|
||||
|
||||
func (p *PgSQLAdapter) DriverName() string {
|
||||
return p.driverName
|
||||
}
|
||||
|
||||
// preloadConfig represents a relationship to be preloaded
|
||||
type preloadConfig struct {
|
||||
relation string
|
||||
@@ -160,31 +244,34 @@ type relationMetadata struct {
|
||||
|
||||
// PgSQLSelectQuery implements SelectQuery for PostgreSQL
|
||||
type PgSQLSelectQuery struct {
|
||||
db *sql.DB
|
||||
tx *sql.Tx
|
||||
model interface{}
|
||||
tableName string
|
||||
tableAlias string
|
||||
columns []string
|
||||
columnExprs []string
|
||||
whereClauses []string
|
||||
orClauses []string
|
||||
joins []string
|
||||
orderBy []string
|
||||
groupBy []string
|
||||
havingClauses []string
|
||||
limit int
|
||||
offset int
|
||||
args []interface{}
|
||||
paramCounter int
|
||||
preloads []preloadConfig
|
||||
db *sql.DB
|
||||
tx *sql.Tx
|
||||
model interface{}
|
||||
entity string
|
||||
tableName string
|
||||
schema string
|
||||
tableAlias string
|
||||
driverName string // Database driver name (postgres, sqlite, mssql)
|
||||
columns []string
|
||||
columnExprs []string
|
||||
whereClauses []string
|
||||
orClauses []string
|
||||
joins []string
|
||||
orderBy []string
|
||||
groupBy []string
|
||||
havingClauses []string
|
||||
limit int
|
||||
offset int
|
||||
args []interface{}
|
||||
paramCounter int
|
||||
preloads []preloadConfig
|
||||
metricsEnabled bool
|
||||
}
|
||||
|
||||
func (p *PgSQLSelectQuery) Model(model interface{}) common.SelectQuery {
|
||||
p.model = model
|
||||
if provider, ok := model.(common.TableNameProvider); ok {
|
||||
p.tableName = provider.TableName()
|
||||
}
|
||||
p.schema, p.tableName = schemaAndTableFromModel(model, p.driverName)
|
||||
p.entity = entityNameFromModel(model, p.tableName)
|
||||
if provider, ok := model.(common.TableAliasProvider); ok {
|
||||
p.tableAlias = provider.TableAlias()
|
||||
}
|
||||
@@ -192,7 +279,11 @@ func (p *PgSQLSelectQuery) Model(model interface{}) common.SelectQuery {
|
||||
}
|
||||
|
||||
func (p *PgSQLSelectQuery) Table(table string) common.SelectQuery {
|
||||
p.tableName = table
|
||||
// For SQLite, convert "schema.table" to "schema_table"
|
||||
p.schema, p.tableName = parseTableName(table, p.driverName)
|
||||
if p.entity == "" {
|
||||
p.entity = cleanMetricIdentifier(p.tableName)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
@@ -375,12 +466,12 @@ func (p *PgSQLSelectQuery) buildSQL() string {
|
||||
|
||||
// LIMIT clause
|
||||
if p.limit > 0 {
|
||||
sb.WriteString(fmt.Sprintf(" LIMIT %d", p.limit))
|
||||
fmt.Fprintf(&sb, " LIMIT %d", p.limit)
|
||||
}
|
||||
|
||||
// OFFSET clause
|
||||
if p.offset > 0 {
|
||||
sb.WriteString(fmt.Sprintf(" OFFSET %d", p.offset))
|
||||
fmt.Fprintf(&sb, " OFFSET %d", p.offset)
|
||||
}
|
||||
|
||||
return sb.String()
|
||||
@@ -402,6 +493,7 @@ func (p *PgSQLSelectQuery) Scan(ctx context.Context, dest interface{}) (err erro
|
||||
err = logger.HandlePanic("PgSQLSelectQuery.Scan", r)
|
||||
}
|
||||
}()
|
||||
startedAt := time.Now()
|
||||
|
||||
// Apply preloads that use JOINs
|
||||
p.applyJoinPreloads()
|
||||
@@ -418,17 +510,21 @@ func (p *PgSQLSelectQuery) Scan(ctx context.Context, dest interface{}) (err erro
|
||||
|
||||
if err != nil {
|
||||
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()
|
||||
|
||||
err = scanRows(rows, dest)
|
||||
if err != nil {
|
||||
recordQueryMetrics(p.metricsEnabled, "SELECT", p.schema, p.entity, p.tableName, startedAt, err)
|
||||
return err
|
||||
}
|
||||
|
||||
// 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 {
|
||||
@@ -438,15 +534,8 @@ func (p *PgSQLSelectQuery) ScanModel(ctx context.Context) error {
|
||||
return p.Scan(ctx, p.model)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
}()
|
||||
|
||||
// Build a COUNT query
|
||||
// countInternal executes the COUNT query and returns the result and the SQL string without recording metrics.
|
||||
func (p *PgSQLSelectQuery) countInternal(ctx context.Context) (rowCount int, querySQL string, retErr error) {
|
||||
var sb strings.Builder
|
||||
sb.WriteString("SELECT COUNT(*) FROM ")
|
||||
sb.WriteString(p.tableName)
|
||||
@@ -480,10 +569,28 @@ func (p *PgSQLSelectQuery) Count(ctx context.Context) (count int, err error) {
|
||||
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 {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -494,35 +601,52 @@ func (p *PgSQLSelectQuery) Exists(ctx context.Context) (exists bool, err error)
|
||||
exists = false
|
||||
}
|
||||
}()
|
||||
|
||||
count, err := p.Count(ctx)
|
||||
startedAt := time.Now()
|
||||
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
|
||||
}
|
||||
|
||||
// PgSQLInsertQuery implements InsertQuery for PostgreSQL
|
||||
type PgSQLInsertQuery struct {
|
||||
db *sql.DB
|
||||
tx *sql.Tx
|
||||
tableName string
|
||||
values map[string]interface{}
|
||||
returning []string
|
||||
db *sql.DB
|
||||
tx *sql.Tx
|
||||
schema string
|
||||
tableName string
|
||||
entity string
|
||||
driverName string
|
||||
values map[string]interface{}
|
||||
valueOrder []string
|
||||
returning []string
|
||||
metricsEnabled bool
|
||||
}
|
||||
|
||||
func (p *PgSQLInsertQuery) Model(model interface{}) common.InsertQuery {
|
||||
if provider, ok := model.(common.TableNameProvider); ok {
|
||||
p.tableName = provider.TableName()
|
||||
}
|
||||
p.schema, p.tableName = schemaAndTableFromModel(model, p.driverName)
|
||||
p.entity = entityNameFromModel(model, p.tableName)
|
||||
// Extract values from model using reflection
|
||||
// This is a simplified implementation
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *PgSQLInsertQuery) Table(table string) common.InsertQuery {
|
||||
p.tableName = table
|
||||
// For SQLite, convert "schema.table" to "schema_table"
|
||||
p.schema, p.tableName = parseTableName(table, p.driverName)
|
||||
if p.entity == "" {
|
||||
p.entity = cleanMetricIdentifier(p.tableName)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
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
|
||||
return p
|
||||
}
|
||||
@@ -538,29 +662,31 @@ func (p *PgSQLInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||
}
|
||||
|
||||
func (p *PgSQLInsertQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
||||
startedAt := time.Now()
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
err = logger.HandlePanic("PgSQLInsertQuery.Exec", r)
|
||||
}
|
||||
recordQueryMetrics(p.metricsEnabled, "INSERT", p.schema, p.entity, p.tableName, startedAt, err)
|
||||
}()
|
||||
|
||||
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))
|
||||
placeholders := make([]string, 0, len(p.values))
|
||||
args := make([]interface{}, 0, len(p.values))
|
||||
|
||||
i := 1
|
||||
for col, val := range p.values {
|
||||
for _, col := range p.valueOrder {
|
||||
columns = append(columns, col)
|
||||
placeholders = append(placeholders, fmt.Sprintf("$%d", i))
|
||||
args = append(args, val)
|
||||
args = append(args, p.values[col])
|
||||
i++
|
||||
}
|
||||
|
||||
query := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)",
|
||||
query := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)", //nolint:gosec // G201: table identifier is internal/validated; values use placeholders
|
||||
p.tableName,
|
||||
strings.Join(columns, ", "),
|
||||
strings.Join(placeholders, ", "))
|
||||
@@ -580,39 +706,96 @@ func (p *PgSQLInsertQuery) Exec(ctx context.Context) (res common.Result, err err
|
||||
|
||||
if err != nil {
|
||||
logger.Error("PgSQL INSERT failed: %v", err)
|
||||
return nil, err
|
||||
return nil, common.WrapSQLError(err, query)
|
||||
}
|
||||
|
||||
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
|
||||
type PgSQLUpdateQuery struct {
|
||||
db *sql.DB
|
||||
tx *sql.Tx
|
||||
tableName string
|
||||
model interface{}
|
||||
sets map[string]interface{}
|
||||
whereClauses []string
|
||||
args []interface{}
|
||||
paramCounter int
|
||||
returning []string
|
||||
db *sql.DB
|
||||
tx *sql.Tx
|
||||
schema string
|
||||
tableName string
|
||||
entity string
|
||||
driverName string
|
||||
model interface{}
|
||||
sets map[string]interface{}
|
||||
setOrder []string
|
||||
whereClauses []string
|
||||
args []interface{}
|
||||
paramCounter int
|
||||
returning []string
|
||||
metricsEnabled bool
|
||||
}
|
||||
|
||||
func (p *PgSQLUpdateQuery) Model(model interface{}) common.UpdateQuery {
|
||||
p.model = model
|
||||
if provider, ok := model.(common.TableNameProvider); ok {
|
||||
p.tableName = provider.TableName()
|
||||
}
|
||||
p.schema, p.tableName = schemaAndTableFromModel(model, p.driverName)
|
||||
p.entity = entityNameFromModel(model, p.tableName)
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *PgSQLUpdateQuery) Table(table string) common.UpdateQuery {
|
||||
p.tableName = table
|
||||
// For SQLite, convert "schema.table" to "schema_table"
|
||||
p.schema, p.tableName = parseTableName(table, p.driverName)
|
||||
if p.entity == "" {
|
||||
p.entity = cleanMetricIdentifier(p.tableName)
|
||||
}
|
||||
if p.model == nil {
|
||||
model, err := modelregistry.GetModelByName(table)
|
||||
if err == nil {
|
||||
p.model = model
|
||||
p.entity = entityNameFromModel(model, p.tableName)
|
||||
}
|
||||
}
|
||||
return p
|
||||
@@ -622,6 +805,9 @@ func (p *PgSQLUpdateQuery) Set(column string, value interface{}) common.UpdateQu
|
||||
if p.model != nil && !reflection.IsColumnWritable(p.model, column) {
|
||||
return p
|
||||
}
|
||||
if _, exists := p.sets[column]; !exists {
|
||||
p.setOrder = append(p.setOrder, column)
|
||||
}
|
||||
p.sets[column] = value
|
||||
return p
|
||||
}
|
||||
@@ -632,13 +818,23 @@ func (p *PgSQLUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQu
|
||||
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 {
|
||||
continue
|
||||
}
|
||||
if p.model != nil && !reflection.IsColumnWritable(p.model, column) {
|
||||
continue
|
||||
}
|
||||
if _, exists := p.sets[column]; !exists {
|
||||
p.setOrder = append(p.setOrder, column)
|
||||
}
|
||||
p.sets[column] = value
|
||||
}
|
||||
return p
|
||||
@@ -667,28 +863,30 @@ func (p *PgSQLUpdateQuery) replacePlaceholders(query string, argCount int) strin
|
||||
}
|
||||
|
||||
func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
||||
startedAt := time.Now()
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
err = logger.HandlePanic("PgSQLUpdateQuery.Exec", r)
|
||||
}
|
||||
recordQueryMetrics(p.metricsEnabled, "UPDATE", p.schema, p.entity, p.tableName, startedAt, err)
|
||||
}()
|
||||
|
||||
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))
|
||||
setArgs := make([]interface{}, 0, len(p.sets))
|
||||
|
||||
// SET parameters start at $1
|
||||
i := 1
|
||||
for col, val := range p.sets {
|
||||
for _, col := range p.setOrder {
|
||||
setClauses = append(setClauses, fmt.Sprintf("%s = $%d", col, i))
|
||||
setArgs = append(setArgs, val)
|
||||
setArgs = append(setArgs, p.sets[col])
|
||||
i++
|
||||
}
|
||||
|
||||
query := fmt.Sprintf("UPDATE %s SET %s",
|
||||
query := fmt.Sprintf("UPDATE %s SET %s", //nolint:gosec // G201: table identifier is internal/validated; values use placeholders
|
||||
p.tableName,
|
||||
strings.Join(setClauses, ", "))
|
||||
|
||||
@@ -738,7 +936,7 @@ func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err err
|
||||
|
||||
if err != nil {
|
||||
logger.Error("PgSQL UPDATE failed: %v", err)
|
||||
return nil, err
|
||||
return nil, common.WrapSQLError(err, query)
|
||||
}
|
||||
|
||||
return &PgSQLResult{result: result}, nil
|
||||
@@ -746,23 +944,30 @@ func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err err
|
||||
|
||||
// PgSQLDeleteQuery implements DeleteQuery for PostgreSQL
|
||||
type PgSQLDeleteQuery struct {
|
||||
db *sql.DB
|
||||
tx *sql.Tx
|
||||
tableName string
|
||||
whereClauses []string
|
||||
args []interface{}
|
||||
paramCounter int
|
||||
db *sql.DB
|
||||
tx *sql.Tx
|
||||
schema string
|
||||
tableName string
|
||||
entity string
|
||||
driverName string
|
||||
whereClauses []string
|
||||
args []interface{}
|
||||
paramCounter int
|
||||
metricsEnabled bool
|
||||
}
|
||||
|
||||
func (p *PgSQLDeleteQuery) Model(model interface{}) common.DeleteQuery {
|
||||
if provider, ok := model.(common.TableNameProvider); ok {
|
||||
p.tableName = provider.TableName()
|
||||
}
|
||||
p.schema, p.tableName = schemaAndTableFromModel(model, p.driverName)
|
||||
p.entity = entityNameFromModel(model, p.tableName)
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *PgSQLDeleteQuery) Table(table string) common.DeleteQuery {
|
||||
p.tableName = table
|
||||
// For SQLite, convert "schema.table" to "schema_table"
|
||||
p.schema, p.tableName = parseTableName(table, p.driverName)
|
||||
if p.entity == "" {
|
||||
p.entity = cleanMetricIdentifier(p.tableName)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
@@ -784,13 +989,15 @@ func (p *PgSQLDeleteQuery) replacePlaceholders(query string, argCount int) strin
|
||||
}
|
||||
|
||||
func (p *PgSQLDeleteQuery) Exec(ctx context.Context) (res common.Result, err error) {
|
||||
startedAt := time.Now()
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
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 {
|
||||
query += " WHERE " + strings.Join(p.whereClauses, " AND ")
|
||||
@@ -807,7 +1014,7 @@ func (p *PgSQLDeleteQuery) Exec(ctx context.Context) (res common.Result, err err
|
||||
|
||||
if err != nil {
|
||||
logger.Error("PgSQL DELETE failed: %v", err)
|
||||
return nil, err
|
||||
return nil, common.WrapSQLError(err, query)
|
||||
}
|
||||
|
||||
return &PgSQLResult{result: result}, nil
|
||||
@@ -835,61 +1042,80 @@ func (p *PgSQLResult) LastInsertId() (int64, error) {
|
||||
|
||||
// PgSQLTxAdapter wraps a PostgreSQL transaction
|
||||
type PgSQLTxAdapter struct {
|
||||
tx *sql.Tx
|
||||
tx *sql.Tx
|
||||
driverName string
|
||||
metricsEnabled bool
|
||||
}
|
||||
|
||||
func (p *PgSQLTxAdapter) NewSelect() common.SelectQuery {
|
||||
return &PgSQLSelectQuery{
|
||||
tx: p.tx,
|
||||
columns: []string{"*"},
|
||||
args: make([]interface{}, 0),
|
||||
tx: p.tx,
|
||||
driverName: p.driverName,
|
||||
columns: []string{"*"},
|
||||
args: make([]interface{}, 0),
|
||||
metricsEnabled: p.metricsEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PgSQLTxAdapter) NewInsert() common.InsertQuery {
|
||||
return &PgSQLInsertQuery{
|
||||
tx: p.tx,
|
||||
values: make(map[string]interface{}),
|
||||
tx: p.tx,
|
||||
driverName: p.driverName,
|
||||
values: make(map[string]interface{}),
|
||||
metricsEnabled: p.metricsEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PgSQLTxAdapter) NewUpdate() common.UpdateQuery {
|
||||
return &PgSQLUpdateQuery{
|
||||
tx: p.tx,
|
||||
sets: make(map[string]interface{}),
|
||||
args: make([]interface{}, 0),
|
||||
whereClauses: make([]string, 0),
|
||||
tx: p.tx,
|
||||
driverName: p.driverName,
|
||||
sets: make(map[string]interface{}),
|
||||
args: make([]interface{}, 0),
|
||||
whereClauses: make([]string, 0),
|
||||
metricsEnabled: p.metricsEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PgSQLTxAdapter) NewDelete() common.DeleteQuery {
|
||||
return &PgSQLDeleteQuery{
|
||||
tx: p.tx,
|
||||
args: make([]interface{}, 0),
|
||||
whereClauses: make([]string, 0),
|
||||
tx: p.tx,
|
||||
driverName: p.driverName,
|
||||
args: make([]interface{}, 0),
|
||||
whereClauses: make([]string, 0),
|
||||
metricsEnabled: p.metricsEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
result, err := p.tx.ExecContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
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
|
||||
}
|
||||
|
||||
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)
|
||||
rows, err := p.tx.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
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()
|
||||
|
||||
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) {
|
||||
@@ -912,6 +1138,10 @@ func (p *PgSQLTxAdapter) GetUnderlyingDB() interface{} {
|
||||
return p.tx
|
||||
}
|
||||
|
||||
func (p *PgSQLTxAdapter) DriverName() string {
|
||||
return p.driverName
|
||||
}
|
||||
|
||||
// applyJoinPreloads adds JOINs for relationships that should use JOIN strategy
|
||||
func (p *PgSQLSelectQuery) applyJoinPreloads() {
|
||||
for _, preload := range p.preloads {
|
||||
@@ -965,7 +1195,7 @@ func (p *PgSQLSelectQuery) applySubqueryPreloads(ctx context.Context, dest inter
|
||||
|
||||
// Use reflection to process the destination
|
||||
destValue := reflect.ValueOf(dest)
|
||||
if destValue.Kind() != reflect.Ptr {
|
||||
if destValue.Kind() != reflect.Pointer {
|
||||
return fmt.Errorf("dest must be a pointer")
|
||||
}
|
||||
|
||||
@@ -992,7 +1222,7 @@ func (p *PgSQLSelectQuery) applySubqueryPreloads(ctx context.Context, dest inter
|
||||
|
||||
// loadPreloadsForRecord loads all preload relationships for a single record
|
||||
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() {
|
||||
return nil
|
||||
}
|
||||
@@ -1036,9 +1266,9 @@ func (p *PgSQLSelectQuery) executePreloadQuery(ctx context.Context, field reflec
|
||||
// Create a new select query for the related table
|
||||
var db common.Database
|
||||
if p.tx != nil {
|
||||
db = &PgSQLTxAdapter{tx: p.tx}
|
||||
db = &PgSQLTxAdapter{tx: p.tx, driverName: p.driverName}
|
||||
} else {
|
||||
db = &PgSQLAdapter{db: p.db}
|
||||
db = &PgSQLAdapter{db: p.db, driverName: p.driverName}
|
||||
}
|
||||
|
||||
query := db.NewSelect().
|
||||
@@ -1069,7 +1299,7 @@ func (p *PgSQLSelectQuery) executePreloadQuery(ctx context.Context, field reflec
|
||||
} else {
|
||||
// Single struct - create a pointer if needed
|
||||
var target reflect.Value
|
||||
if field.Kind() == reflect.Ptr {
|
||||
if field.Kind() == reflect.Pointer {
|
||||
target = reflect.New(field.Type().Elem())
|
||||
} else {
|
||||
target = reflect.New(field.Type())
|
||||
@@ -1082,7 +1312,7 @@ func (p *PgSQLSelectQuery) executePreloadQuery(ctx context.Context, field reflec
|
||||
}
|
||||
|
||||
// Set the field
|
||||
if field.Kind() == reflect.Ptr {
|
||||
if field.Kind() == reflect.Pointer {
|
||||
field.Set(target)
|
||||
} else {
|
||||
field.Set(target.Elem())
|
||||
@@ -1099,7 +1329,7 @@ func (p *PgSQLSelectQuery) getRelationMetadata(fieldName string) *relationMetada
|
||||
}
|
||||
|
||||
modelType := reflect.TypeOf(p.model)
|
||||
if modelType.Kind() == reflect.Ptr {
|
||||
if modelType.Kind() == reflect.Pointer {
|
||||
modelType = modelType.Elem()
|
||||
}
|
||||
|
||||
@@ -1148,7 +1378,7 @@ func (p *PgSQLSelectQuery) getRelationMetadataFromField(modelType reflect.Type,
|
||||
if fieldType.Kind() == reflect.Slice {
|
||||
fieldType = fieldType.Elem()
|
||||
}
|
||||
if fieldType.Kind() == reflect.Ptr {
|
||||
if fieldType.Kind() == reflect.Pointer {
|
||||
fieldType = fieldType.Elem()
|
||||
}
|
||||
|
||||
@@ -1181,7 +1411,7 @@ func scanRows(rows *sql.Rows, dest interface{}) error {
|
||||
|
||||
// Get destination type
|
||||
destValue := reflect.ValueOf(dest)
|
||||
if destValue.Kind() != reflect.Ptr {
|
||||
if destValue.Kind() != reflect.Pointer {
|
||||
return fmt.Errorf("dest must be a pointer")
|
||||
}
|
||||
|
||||
@@ -1236,7 +1466,7 @@ func scanRowsToMapSlice(rows *sql.Rows, columns []string, destValue reflect.Valu
|
||||
// scanRowsToStructSlice scans rows into a slice of structs
|
||||
func scanRowsToStructSlice(rows *sql.Rows, columns []string, destValue reflect.Value) error {
|
||||
elemType := destValue.Type().Elem()
|
||||
isPtr := elemType.Kind() == reflect.Ptr
|
||||
isPtr := elemType.Kind() == reflect.Pointer
|
||||
|
||||
if isPtr {
|
||||
elemType = elemType.Elem()
|
||||
|
||||
@@ -13,7 +13,7 @@ import (
|
||||
// Example demonstrates how to use the PgSQL adapter
|
||||
func ExamplePgSQLAdapter() error {
|
||||
// Connect to PostgreSQL database
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to open database: %w", err)
|
||||
@@ -155,7 +155,7 @@ func (u User) TableName() string {
|
||||
|
||||
// ExampleWithModel demonstrates using models with the PgSQL adapter
|
||||
func ExampleWithModel() error {
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
return err
|
||||
|
||||
@@ -51,7 +51,7 @@ func (c Comment) TableName() string {
|
||||
|
||||
// ExamplePreload demonstrates the Preload functionality
|
||||
func ExamplePreload() error {
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -79,7 +79,7 @@ func ExamplePreload() error {
|
||||
|
||||
// ExamplePreloadRelation demonstrates smart PreloadRelation with auto-detection
|
||||
func ExamplePreloadRelation() error {
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -148,7 +148,7 @@ func ExamplePreloadRelation() error {
|
||||
|
||||
// ExampleJoinRelation demonstrates explicit JOIN loading
|
||||
func ExampleJoinRelation() error {
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -185,7 +185,7 @@ func ExampleJoinRelation() error {
|
||||
|
||||
// ExampleScanModel demonstrates ScanModel with struct destinations
|
||||
func ExampleScanModel() error {
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -221,7 +221,7 @@ func ExampleScanModel() error {
|
||||
|
||||
// ExampleCompleteWorkflow demonstrates a complete workflow with preloading
|
||||
func ExampleCompleteWorkflow() error {
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable" //nolint:gosec // G101: false positive: identifier/example, not a credential
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
return err
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -1,16 +1,117 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"strings"
|
||||
|
||||
"github.com/uptrace/bun/dialect/mssqldialect"
|
||||
"github.com/uptrace/bun/dialect/pgdialect"
|
||||
"github.com/uptrace/bun/dialect/sqlitedialect"
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/driver/sqlserver"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// PostgreSQL identifier length limit (63 bytes + null terminator = 64 bytes total)
|
||||
const postgresIdentifierLimit = 63
|
||||
|
||||
// checkAliasLength checks if a preload relation path will generate aliases that exceed PostgreSQL's limit
|
||||
// Returns true if the alias is likely to be truncated
|
||||
func checkAliasLength(relation string) bool {
|
||||
// Bun generates aliases like: parentalias__childalias__columnname
|
||||
// For nested preloads, it uses the pattern: relation1__relation2__relation3__columnname
|
||||
parts := strings.Split(relation, ".")
|
||||
if len(parts) <= 1 {
|
||||
return false // Single level relations are fine
|
||||
}
|
||||
|
||||
// Calculate the actual alias prefix length that Bun will generate
|
||||
// Bun uses double underscores (__) between each relation level
|
||||
// and converts the relation names to lowercase with underscores
|
||||
aliasPrefix := strings.ToLower(strings.Join(parts, "__"))
|
||||
aliasPrefixLen := len(aliasPrefix)
|
||||
|
||||
// We need to add 2 more underscores for the column name separator plus column name length
|
||||
// Column names in the error were things like "rid_mastertype_hubtype" (23 chars)
|
||||
// To be safe, assume the longest column name could be around 35 chars
|
||||
maxColumnNameLen := 35
|
||||
estimatedMaxLen := aliasPrefixLen + 2 + maxColumnNameLen
|
||||
|
||||
// Check if this would exceed PostgreSQL's identifier limit
|
||||
if estimatedMaxLen > postgresIdentifierLimit {
|
||||
logger.Warn("Preload relation '%s' will generate aliases up to %d chars (prefix: %d + column: %d), exceeding PostgreSQL's %d char limit",
|
||||
relation, estimatedMaxLen, aliasPrefixLen, maxColumnNameLen, postgresIdentifierLimit)
|
||||
return true
|
||||
}
|
||||
|
||||
// Also check if just the prefix is getting close (within 15 chars of limit)
|
||||
// This gives room for column names
|
||||
if aliasPrefixLen > (postgresIdentifierLimit - 15) {
|
||||
logger.Warn("Preload relation '%s' has alias prefix of %d chars, which may cause truncation with longer column names (limit: %d)",
|
||||
relation, aliasPrefixLen, postgresIdentifierLimit)
|
||||
return true
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// parseTableName splits a table name that may contain schema into separate schema and table
|
||||
// For example: "public.users" -> ("public", "users")
|
||||
//
|
||||
// "users" -> ("", "users")
|
||||
func parseTableName(fullTableName string) (schema, table string) {
|
||||
//
|
||||
// For SQLite, schema.table is translated to schema_table since SQLite doesn't support schemas
|
||||
// in the same way as PostgreSQL/MSSQL
|
||||
func parseTableName(fullTableName, driverName string) (schema, table string) {
|
||||
if idx := strings.LastIndex(fullTableName, "."); idx != -1 {
|
||||
return fullTableName[:idx], fullTableName[idx+1:]
|
||||
schema = fullTableName[:idx]
|
||||
table = fullTableName[idx+1:]
|
||||
|
||||
// For SQLite, convert schema.table to schema_table
|
||||
if driverName == "sqlite" || driverName == "sqlite3" {
|
||||
table = schema + "_" + table
|
||||
schema = ""
|
||||
}
|
||||
return schema, table
|
||||
}
|
||||
return "", fullTableName
|
||||
}
|
||||
|
||||
// GetPostgresDialect returns a Bun PostgreSQL dialect
|
||||
func GetPostgresDialect() *pgdialect.Dialect {
|
||||
return pgdialect.New()
|
||||
}
|
||||
|
||||
// GetSQLiteDialect returns a Bun SQLite dialect
|
||||
func GetSQLiteDialect() *sqlitedialect.Dialect {
|
||||
return sqlitedialect.New()
|
||||
}
|
||||
|
||||
// GetMSSQLDialect returns a Bun MSSQL dialect
|
||||
func GetMSSQLDialect() *mssqldialect.Dialect {
|
||||
return mssqldialect.New()
|
||||
}
|
||||
|
||||
// GetPostgresDialector returns a GORM PostgreSQL dialector
|
||||
func GetPostgresDialector(db *sql.DB) gorm.Dialector {
|
||||
return postgres.New(postgres.Config{
|
||||
Conn: db,
|
||||
})
|
||||
}
|
||||
|
||||
// GetSQLiteDialector returns a GORM SQLite dialector
|
||||
func GetSQLiteDialector(db *sql.DB) gorm.Dialector {
|
||||
return sqlite.Dialector{
|
||||
Conn: db,
|
||||
}
|
||||
}
|
||||
|
||||
// GetMSSQLDialector returns a GORM MSSQL dialector
|
||||
func GetMSSQLDialector(db *sql.DB) gorm.Dialector {
|
||||
return sqlserver.New(sqlserver.Config{
|
||||
Conn: db,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -174,7 +174,9 @@ func (h *HTTPResponseWriter) Write(data []byte) (int, error) {
|
||||
|
||||
func (h *HTTPResponseWriter) WriteJSON(data interface{}) error {
|
||||
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
|
||||
|
||||
+47
-10
@@ -3,6 +3,8 @@ package common
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||
)
|
||||
|
||||
// CORSConfig holds CORS configuration
|
||||
@@ -15,8 +17,30 @@ type CORSConfig struct {
|
||||
|
||||
// DefaultCORSConfig returns a default CORS configuration suitable for HeadSpec
|
||||
func DefaultCORSConfig() CORSConfig {
|
||||
configManager := config.GetConfigManager()
|
||||
cfg, _ := configManager.GetConfig()
|
||||
hosts := make([]string, 0)
|
||||
// hosts = append(hosts, "*")
|
||||
|
||||
_, _, ipsList := config.GetIPs()
|
||||
|
||||
for i := range cfg.Servers.Instances {
|
||||
server := cfg.Servers.Instances[i]
|
||||
if server.Port == 0 {
|
||||
continue
|
||||
}
|
||||
hosts = append(hosts, server.ExternalURLs...)
|
||||
hosts = append(hosts, fmt.Sprintf("http://%s:%d", server.Host, server.Port))
|
||||
hosts = append(hosts, fmt.Sprintf("https://%s:%d", server.Host, server.Port))
|
||||
hosts = append(hosts, fmt.Sprintf("http://%s:%d", "localhost", server.Port))
|
||||
for _, ip := range ipsList {
|
||||
hosts = append(hosts, fmt.Sprintf("http://%s:%d", ip.String(), server.Port))
|
||||
hosts = append(hosts, fmt.Sprintf("https://%s:%d", ip.String(), server.Port))
|
||||
}
|
||||
}
|
||||
|
||||
return CORSConfig{
|
||||
AllowedOrigins: []string{"*"},
|
||||
AllowedOrigins: hosts,
|
||||
AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"},
|
||||
AllowedHeaders: GetHeadSpecHeaders(),
|
||||
MaxAge: 86400, // 24 hours
|
||||
@@ -90,19 +114,28 @@ func GetHeadSpecHeaders() []string {
|
||||
}
|
||||
|
||||
// SetCORSHeaders sets CORS headers on a response writer
|
||||
func SetCORSHeaders(w ResponseWriter, config CORSConfig) {
|
||||
// Set allowed origins
|
||||
if len(config.AllowedOrigins) > 0 {
|
||||
w.SetHeader("Access-Control-Allow-Origin", strings.Join(config.AllowedOrigins, ", "))
|
||||
func SetCORSHeaders(w ResponseWriter, r Request, config CORSConfig) {
|
||||
// Reflect the request origin; fall back to wildcard only when no origin is present
|
||||
origin := r.Header("Origin")
|
||||
if origin == "" {
|
||||
origin = "*"
|
||||
} else {
|
||||
// Vary must be set so caches don't serve one origin's response to another
|
||||
httpW := w.UnderlyingResponseWriter()
|
||||
httpW.Header().Set("Vary", "Origin")
|
||||
}
|
||||
w.SetHeader("Access-Control-Allow-Origin", origin)
|
||||
|
||||
// Set allowed methods
|
||||
if len(config.AllowedMethods) > 0 {
|
||||
w.SetHeader("Access-Control-Allow-Methods", strings.Join(config.AllowedMethods, ", "))
|
||||
}
|
||||
|
||||
// Set allowed headers
|
||||
if len(config.AllowedHeaders) > 0 {
|
||||
// Reflect the preflight request headers when present; otherwise use the explicit config list
|
||||
requestedHeaders := r.Header("Access-Control-Request-Headers")
|
||||
if requestedHeaders != "" {
|
||||
w.SetHeader("Access-Control-Allow-Headers", requestedHeaders)
|
||||
} else if len(config.AllowedHeaders) > 0 {
|
||||
w.SetHeader("Access-Control-Allow-Headers", strings.Join(config.AllowedHeaders, ", "))
|
||||
}
|
||||
|
||||
@@ -111,9 +144,13 @@ func SetCORSHeaders(w ResponseWriter, config CORSConfig) {
|
||||
w.SetHeader("Access-Control-Max-Age", fmt.Sprintf("%d", config.MaxAge))
|
||||
}
|
||||
|
||||
// Allow credentials
|
||||
w.SetHeader("Access-Control-Allow-Credentials", "true")
|
||||
// Allow credentials only when a specific origin is reflected (not wildcard)
|
||||
if origin != "*" {
|
||||
w.SetHeader("Access-Control-Allow-Credentials", "true")
|
||||
}
|
||||
|
||||
// Expose headers that clients can read
|
||||
w.SetHeader("Access-Control-Expose-Headers", "Content-Range, X-Api-Range-Total, X-Api-Range-Size")
|
||||
exposeHeaders := config.AllowedHeaders
|
||||
exposeHeaders = append(exposeHeaders, "Content-Range", "X-Api-Range-Total", "X-Api-Range-Size")
|
||||
w.SetHeader("Access-Control-Expose-Headers", strings.Join(exposeHeaders, ", "))
|
||||
}
|
||||
|
||||
+263
-1
@@ -3,6 +3,10 @@ package common
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// ValidateAndUnwrapModelResult contains the result of model validation
|
||||
@@ -21,7 +25,7 @@ func ValidateAndUnwrapModel(model interface{}) (*ValidateAndUnwrapModelResult, e
|
||||
originalType := modelType
|
||||
|
||||
// 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()
|
||||
}
|
||||
|
||||
@@ -45,3 +49,261 @@ func ValidateAndUnwrapModel(model interface{}) (*ValidateAndUnwrapModelResult, e
|
||||
OriginalType: originalType,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ExtractTagValue extracts the value for a given key from a struct tag string.
|
||||
// It handles both semicolon and comma-separated tag formats (e.g., GORM and BUN tags).
|
||||
// For tags like "json:name;validate:required" it will extract "name" for key "json".
|
||||
// For tags like "rel:has-many,join:table" it will extract "table" for key "join".
|
||||
func ExtractTagValue(tag, key string) string {
|
||||
// Split by both semicolons and commas to handle different tag formats
|
||||
// We need to be smart about this - commas can be part of values
|
||||
// So we'll try semicolon first, then comma if needed
|
||||
separators := []string{";", ","}
|
||||
|
||||
for _, sep := range separators {
|
||||
parts := strings.Split(tag, sep)
|
||||
for _, part := range parts {
|
||||
part = strings.TrimSpace(part)
|
||||
if strings.HasPrefix(part, key+":") {
|
||||
return strings.TrimPrefix(part, key+":")
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// GetRelationshipInfo analyzes a model type and extracts relationship metadata
|
||||
// for a specific relation field identified by its JSON name.
|
||||
// Returns nil if the field is not found or is not a valid relationship.
|
||||
func GetRelationshipInfo(modelType reflect.Type, relationName string) *RelationshipInfo {
|
||||
// Ensure we have a struct type
|
||||
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||
logger.Warn("Cannot get relationship info from non-struct type: %v", modelType)
|
||||
return nil
|
||||
}
|
||||
|
||||
for i := 0; i < modelType.NumField(); i++ {
|
||||
field := modelType.Field(i)
|
||||
jsonTag := field.Tag.Get("json")
|
||||
jsonName := strings.Split(jsonTag, ",")[0]
|
||||
|
||||
if jsonName == relationName {
|
||||
gormTag := field.Tag.Get("gorm")
|
||||
bunTag := field.Tag.Get("bun")
|
||||
info := &RelationshipInfo{
|
||||
FieldName: field.Name,
|
||||
JSONName: jsonName,
|
||||
}
|
||||
|
||||
if strings.Contains(bunTag, "rel:") || strings.Contains(bunTag, "join:") {
|
||||
//bun:"rel:has-many,join:rid_hub=rid_hub_division"
|
||||
if strings.Contains(bunTag, "has-many") {
|
||||
info.RelationType = "hasMany"
|
||||
} else if strings.Contains(bunTag, "has-one") {
|
||||
info.RelationType = "hasOne"
|
||||
} else if strings.Contains(bunTag, "belongs-to") {
|
||||
info.RelationType = "belongsTo"
|
||||
} else if strings.Contains(bunTag, "many-to-many") {
|
||||
info.RelationType = "many2many"
|
||||
} else {
|
||||
info.RelationType = "hasOne"
|
||||
}
|
||||
|
||||
// Extract join info
|
||||
joinPart := ExtractTagValue(bunTag, "join")
|
||||
if joinPart != "" && info.RelationType == "many2many" {
|
||||
// For many2many, the join part is the join table name
|
||||
info.JoinTable = joinPart
|
||||
} else if joinPart != "" {
|
||||
// For other relations, parse foreignKey and references
|
||||
joinParts := strings.Split(joinPart, "=")
|
||||
if len(joinParts) == 2 {
|
||||
info.ForeignKey = joinParts[0]
|
||||
info.References = joinParts[1]
|
||||
}
|
||||
}
|
||||
|
||||
// Get related model type
|
||||
if field.Type.Kind() == reflect.Slice {
|
||||
elemType := field.Type.Elem()
|
||||
if elemType.Kind() == reflect.Pointer {
|
||||
elemType = elemType.Elem()
|
||||
}
|
||||
if elemType.Kind() == reflect.Struct {
|
||||
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||
}
|
||||
} else if field.Type.Kind() == reflect.Pointer || field.Type.Kind() == reflect.Struct {
|
||||
elemType := field.Type
|
||||
if elemType.Kind() == reflect.Pointer {
|
||||
elemType = elemType.Elem()
|
||||
}
|
||||
if elemType.Kind() == reflect.Struct {
|
||||
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||
}
|
||||
}
|
||||
|
||||
return info
|
||||
}
|
||||
|
||||
// Parse GORM tag to determine relationship type and keys
|
||||
if strings.Contains(gormTag, "foreignKey") {
|
||||
info.ForeignKey = ExtractTagValue(gormTag, "foreignKey")
|
||||
info.References = ExtractTagValue(gormTag, "references")
|
||||
|
||||
// Determine if it's belongsTo or hasMany/hasOne
|
||||
if field.Type.Kind() == reflect.Slice {
|
||||
info.RelationType = "hasMany"
|
||||
// Get the element type for slice
|
||||
elemType := field.Type.Elem()
|
||||
if elemType.Kind() == reflect.Pointer {
|
||||
elemType = elemType.Elem()
|
||||
}
|
||||
if elemType.Kind() == reflect.Struct {
|
||||
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||
}
|
||||
} else if field.Type.Kind() == reflect.Pointer || field.Type.Kind() == reflect.Struct {
|
||||
info.RelationType = "belongsTo"
|
||||
elemType := field.Type
|
||||
if elemType.Kind() == reflect.Pointer {
|
||||
elemType = elemType.Elem()
|
||||
}
|
||||
if elemType.Kind() == reflect.Struct {
|
||||
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||
}
|
||||
}
|
||||
} else if strings.Contains(gormTag, "many2many") {
|
||||
info.RelationType = "many2many"
|
||||
info.JoinTable = ExtractTagValue(gormTag, "many2many")
|
||||
// Get the element type for many2many (always slice)
|
||||
if field.Type.Kind() == reflect.Slice {
|
||||
elemType := field.Type.Elem()
|
||||
if elemType.Kind() == reflect.Pointer {
|
||||
elemType = elemType.Elem()
|
||||
}
|
||||
if elemType.Kind() == reflect.Struct {
|
||||
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Field has no GORM relationship tags, so it's not a relation
|
||||
return nil
|
||||
}
|
||||
|
||||
return info
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RelationPathToBunAlias converts a relation path (e.g., "Order.Customer") to a Bun alias format.
|
||||
// It converts to lowercase and replaces dots with double underscores.
|
||||
// For example: "Order.Customer" -> "order__customer"
|
||||
func RelationPathToBunAlias(relationPath string) string {
|
||||
if relationPath == "" {
|
||||
return ""
|
||||
}
|
||||
// Convert to lowercase and replace dots with double underscores
|
||||
alias := strings.ToLower(relationPath)
|
||||
alias = strings.ReplaceAll(alias, ".", "__")
|
||||
return alias
|
||||
}
|
||||
|
||||
// ReplaceTableReferencesInSQL replaces references to a base table name in a SQL expression
|
||||
// with the appropriate alias for the current preload level.
|
||||
// For example, if baseTableName is "mastertaskitem" and targetAlias is "mal__mal",
|
||||
// it will replace "mastertaskitem.rid_mastertaskitem" with "mal__mal.rid_mastertaskitem"
|
||||
func ReplaceTableReferencesInSQL(sqlExpr, baseTableName, targetAlias string) string {
|
||||
if sqlExpr == "" || baseTableName == "" || targetAlias == "" {
|
||||
return sqlExpr
|
||||
}
|
||||
|
||||
// Replace both quoted and unquoted table references
|
||||
// Handle patterns like: tablename.column, "tablename".column, tablename."column", "tablename"."column"
|
||||
|
||||
// Pattern 1: tablename.column (unquoted)
|
||||
result := strings.ReplaceAll(sqlExpr, baseTableName+".", targetAlias+".")
|
||||
|
||||
// Pattern 2: "tablename".column or "tablename"."column" (quoted table name)
|
||||
result = strings.ReplaceAll(result, "\""+baseTableName+"\".", "\""+targetAlias+"\".")
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetTableNameFromModel extracts the table name from a model.
|
||||
// It checks the bun tag first, then falls back to converting the struct name to snake_case.
|
||||
func GetTableNameFromModel(model interface{}) string {
|
||||
if model == nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
modelType := reflect.TypeOf(model)
|
||||
|
||||
// Unwrap pointers
|
||||
for modelType != nil && modelType.Kind() == reflect.Pointer {
|
||||
modelType = modelType.Elem()
|
||||
}
|
||||
|
||||
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||
return ""
|
||||
}
|
||||
|
||||
// Look for bun tag on embedded BaseModel
|
||||
for i := 0; i < modelType.NumField(); i++ {
|
||||
field := modelType.Field(i)
|
||||
if field.Anonymous {
|
||||
bunTag := field.Tag.Get("bun")
|
||||
if strings.HasPrefix(bunTag, "table:") {
|
||||
return strings.TrimPrefix(bunTag, "table:")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback: convert struct name to lowercase (simple heuristic)
|
||||
// This handles cases like "MasterTaskItem" -> "mastertaskitem"
|
||||
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, ",") + "}"
|
||||
}
|
||||
|
||||
@@ -0,0 +1,171 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestExtractTagValue(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
tag string
|
||||
key string
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "Extract existing key",
|
||||
tag: "json:name;validate:required",
|
||||
key: "json",
|
||||
expected: "name",
|
||||
},
|
||||
{
|
||||
name: "Extract key with spaces",
|
||||
tag: "json:name ; validate:required",
|
||||
key: "validate",
|
||||
expected: "required",
|
||||
},
|
||||
{
|
||||
name: "Extract key at end",
|
||||
tag: "json:name;validate:required;db:column_name",
|
||||
key: "db",
|
||||
expected: "column_name",
|
||||
},
|
||||
{
|
||||
name: "Extract key at beginning",
|
||||
tag: "primary:true;json:id;db:user_id",
|
||||
key: "primary",
|
||||
expected: "true",
|
||||
},
|
||||
{
|
||||
name: "Key not found",
|
||||
tag: "json:name;validate:required",
|
||||
key: "db",
|
||||
expected: "",
|
||||
},
|
||||
{
|
||||
name: "Empty tag",
|
||||
tag: "",
|
||||
key: "json",
|
||||
expected: "",
|
||||
},
|
||||
{
|
||||
name: "Single key-value pair",
|
||||
tag: "json:name",
|
||||
key: "json",
|
||||
expected: "name",
|
||||
},
|
||||
{
|
||||
name: "Key with empty value",
|
||||
tag: "json:;validate:required",
|
||||
key: "json",
|
||||
expected: "",
|
||||
},
|
||||
{
|
||||
name: "Key with complex value",
|
||||
tag: "json:user_name,omitempty;validate:required,min=3",
|
||||
key: "json",
|
||||
expected: "user_name,omitempty",
|
||||
},
|
||||
{
|
||||
name: "Multiple semicolons",
|
||||
tag: "json:name;;validate:required",
|
||||
key: "validate",
|
||||
expected: "required",
|
||||
},
|
||||
{
|
||||
name: "BUN Tag with comma separator",
|
||||
tag: "rel:has-many,join:rid_hub=rid_hub_child",
|
||||
key: "join",
|
||||
expected: "rid_hub=rid_hub_child",
|
||||
},
|
||||
{
|
||||
name: "Extract foreignKey",
|
||||
tag: "foreignKey:UserID;references:ID",
|
||||
key: "foreignKey",
|
||||
expected: "UserID",
|
||||
},
|
||||
{
|
||||
name: "Extract references",
|
||||
tag: "foreignKey:UserID;references:ID",
|
||||
key: "references",
|
||||
expected: "ID",
|
||||
},
|
||||
{
|
||||
name: "Extract many2many",
|
||||
tag: "many2many:user_roles",
|
||||
key: "many2many",
|
||||
expected: "user_roles",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := ExtractTagValue(tt.tag, tt.key)
|
||||
if result != tt.expected {
|
||||
t.Errorf("ExtractTagValue(%q, %q) = %q; want %q", tt.tag, tt.key, result, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -30,6 +30,12 @@ type Database interface {
|
||||
// For Bun, this returns *bun.DB
|
||||
// This is useful for provider-specific features like PostgreSQL NOTIFY/LISTEN
|
||||
GetUnderlyingDB() interface{}
|
||||
|
||||
// DriverName returns the canonical name of the underlying database driver.
|
||||
// Possible values: "postgres", "sqlite", "mssql", "mysql".
|
||||
// All adapters normalise vendor-specific strings (e.g. Bun's "pg", GORM's
|
||||
// "sqlserver") to the values above before returning.
|
||||
DriverName() string
|
||||
}
|
||||
|
||||
// SelectQuery interface for building SELECT queries (compatible with both GORM and Bun)
|
||||
@@ -69,6 +75,7 @@ type InsertQuery interface {
|
||||
|
||||
// Execution
|
||||
Exec(ctx context.Context) (Result, error)
|
||||
Scan(ctx context.Context, dest interface{}) error
|
||||
}
|
||||
|
||||
// UpdateQuery interface for building UPDATE queries
|
||||
@@ -171,7 +178,9 @@ func (s *StandardResponseWriter) Write(data []byte) (int, error) {
|
||||
|
||||
func (s *StandardResponseWriter) WriteJSON(data interface{}) error {
|
||||
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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+267
-75
@@ -20,17 +20,6 @@ type RelationshipInfoProvider interface {
|
||||
GetRelationshipInfo(modelType reflect.Type, relationName string) *RelationshipInfo
|
||||
}
|
||||
|
||||
// RelationshipInfo contains information about a model relationship
|
||||
type RelationshipInfo struct {
|
||||
FieldName string
|
||||
JSONName string
|
||||
RelationType string // "belongsTo", "hasMany", "hasOne", "many2many"
|
||||
ForeignKey string
|
||||
References string
|
||||
JoinTable string
|
||||
RelatedModel interface{}
|
||||
}
|
||||
|
||||
// NestedCUDProcessor handles recursive processing of nested object graphs
|
||||
type NestedCUDProcessor struct {
|
||||
db Database
|
||||
@@ -80,11 +69,12 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
||||
|
||||
// Get model type for reflection
|
||||
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()
|
||||
}
|
||||
|
||||
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||
logger.Error("Invalid model type: operation=%s, table=%s, modelType=%v, expected struct", operation, tableName, modelType)
|
||||
return nil, fmt.Errorf("model must be a struct type, got %v", modelType)
|
||||
}
|
||||
|
||||
@@ -108,50 +98,97 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
||||
}
|
||||
}
|
||||
|
||||
// Filter regularData to only include fields that exist in the model,
|
||||
// and translate JSON keys to their actual database column names.
|
||||
regularData = p.filterValidFields(regularData, model)
|
||||
|
||||
// Inject parent IDs for foreign key resolution
|
||||
p.injectForeignKeys(regularData, modelType, parentIDs)
|
||||
|
||||
// Get the primary key name for this model
|
||||
pkName := reflection.GetPrimaryKeyName(model)
|
||||
|
||||
// Check if we have any data to process (besides _request)
|
||||
hasData := len(regularData) > 0
|
||||
|
||||
// Process based on operation
|
||||
switch strings.ToLower(operation) {
|
||||
case "insert", "create":
|
||||
id, err := p.processInsert(ctx, regularData, tableName)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("insert failed: %w", err)
|
||||
}
|
||||
result.ID = id
|
||||
result.AffectedRows = 1
|
||||
result.Data = regularData
|
||||
case "insert", "create", "add":
|
||||
// Only perform insert if we have data to insert
|
||||
if hasData {
|
||||
id, err := p.processInsert(ctx, regularData, tableName)
|
||||
if err != nil {
|
||||
logger.Error("Insert failed for table=%s, data=%+v, error=%v", tableName, regularData, err)
|
||||
return nil, fmt.Errorf("insert failed: %w", err)
|
||||
}
|
||||
result.ID = id
|
||||
result.AffectedRows = 1
|
||||
result.Data = regularData
|
||||
|
||||
// Process child relations after parent insert (to get parent ID)
|
||||
if err := p.processChildRelations(ctx, "insert", id, relationFields, result.RelationData, modelType); err != nil {
|
||||
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
||||
// 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)
|
||||
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)
|
||||
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
||||
}
|
||||
} else {
|
||||
logger.Debug("Skipping insert for %s - no data columns besides _request", tableName)
|
||||
}
|
||||
|
||||
case "update":
|
||||
rows, err := p.processUpdate(ctx, regularData, tableName, data[pkName])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("update failed: %w", err)
|
||||
case "update", "change", "modify":
|
||||
// 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
|
||||
}
|
||||
result.ID = data[pkName]
|
||||
result.AffectedRows = rows
|
||||
result.Data = regularData
|
||||
if hasData {
|
||||
rows, err := p.processUpdate(ctx, regularData, tableName, data[pkName])
|
||||
if err != nil {
|
||||
logger.Error("Update failed for table=%s, id=%v, data=%+v, error=%v", tableName, data[pkName], regularData, err)
|
||||
return nil, fmt.Errorf("update failed: %w", err)
|
||||
}
|
||||
result.ID = data[pkName]
|
||||
result.AffectedRows = rows
|
||||
result.Data = regularData
|
||||
|
||||
// Process child relations for update
|
||||
if err := p.processChildRelations(ctx, "update", data[pkName], relationFields, result.RelationData, modelType); err != nil {
|
||||
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
||||
// 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
|
||||
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], regularData, err)
|
||||
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
||||
}
|
||||
} else {
|
||||
logger.Debug("Skipping update for %s - no data columns besides _request", tableName)
|
||||
result.ID = data[pkName]
|
||||
}
|
||||
|
||||
case "delete", "remove":
|
||||
if reflection.IsEmptyValue(data[pkName]) {
|
||||
logger.Warn("Skipping delete for %s - no primary key", tableName)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
case "delete":
|
||||
// Process child relations first (for referential integrity)
|
||||
if err := p.processChildRelations(ctx, "delete", data[pkName], relationFields, result.RelationData, modelType); err != nil {
|
||||
return nil, fmt.Errorf("failed to process child relations before delete: %w", err)
|
||||
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)
|
||||
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
||||
}
|
||||
|
||||
rows, err := p.processDelete(ctx, tableName, data[pkName])
|
||||
if err != nil {
|
||||
logger.Error("Delete failed for table=%s, id=%v, error=%v", tableName, data[pkName], err)
|
||||
return nil, fmt.Errorf("delete failed: %w", err)
|
||||
}
|
||||
result.ID = data[pkName]
|
||||
@@ -159,6 +196,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
||||
result.Data = regularData
|
||||
|
||||
default:
|
||||
logger.Error("Unsupported operation: %s for table=%s", operation, tableName)
|
||||
return nil, fmt.Errorf("unsupported operation: %s", operation)
|
||||
}
|
||||
|
||||
@@ -176,32 +214,80 @@ func (p *NestedCUDProcessor) extractCRUDRequest(data map[string]interface{}) str
|
||||
return ""
|
||||
}
|
||||
|
||||
// injectForeignKeys injects parent IDs into data for foreign key fields
|
||||
// filterValidFields filters input data to only include fields that exist in the model,
|
||||
// and translates JSON key names to their actual database column names.
|
||||
// For example, a field tagged `json:"_changed_date" bun:"changed_date"` will be
|
||||
// included in the result as "changed_date", not "_changed_date".
|
||||
func (p *NestedCUDProcessor) filterValidFields(data map[string]interface{}, model interface{}) map[string]interface{} {
|
||||
if len(data) == 0 {
|
||||
return data
|
||||
}
|
||||
|
||||
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 data
|
||||
}
|
||||
|
||||
// Build a mapping from JSON key -> DB column name for all writable fields.
|
||||
// This both validates which fields belong to the model and translates their names
|
||||
// to the correct column names for use in SQL insert/update queries.
|
||||
jsonToDBCol := reflection.BuildJSONToDBColumnMap(modelType)
|
||||
|
||||
filteredData := make(map[string]interface{})
|
||||
for key, value := range data {
|
||||
dbColName, exists := jsonToDBCol[key]
|
||||
if exists {
|
||||
filteredData[dbColName] = value
|
||||
} else {
|
||||
logger.Debug("Skipping invalid field '%s' - not found in model %v", key, modelType)
|
||||
}
|
||||
}
|
||||
|
||||
return filteredData
|
||||
}
|
||||
|
||||
// injectForeignKeys injects parent IDs into data for foreign key fields.
|
||||
// data is expected to be keyed by DB column names (as returned by filterValidFields).
|
||||
func (p *NestedCUDProcessor) injectForeignKeys(data map[string]interface{}, modelType reflect.Type, parentIDs map[string]interface{}) {
|
||||
if len(parentIDs) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// Iterate through model fields to find foreign key fields
|
||||
for i := 0; i < modelType.NumField(); i++ {
|
||||
field := modelType.Field(i)
|
||||
jsonTag := field.Tag.Get("json")
|
||||
jsonName := strings.Split(jsonTag, ",")[0]
|
||||
pkCol := reflection.GetPrimaryKeyName(reflect.New(modelType).Interface())
|
||||
|
||||
// Check if this field is a foreign key and we have a parent ID for it
|
||||
// Common patterns: DepartmentID, ManagerID, ProjectID, etc.
|
||||
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(field.Name, parentKey+"ID") {
|
||||
// Only inject if not already present
|
||||
if _, exists := data[jsonName]; !exists {
|
||||
logger.Debug("Injecting foreign key: %s = %v", jsonName, parentID)
|
||||
data[jsonName] = parentID
|
||||
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++ {
|
||||
field := modelType.Field(i)
|
||||
jsonName := strings.Split(field.Tag.Get("json"), ",")[0]
|
||||
if strings.EqualFold(jsonName, "rid"+parentKey) ||
|
||||
strings.EqualFold(jsonName, "rid_"+parentKey) ||
|
||||
strings.EqualFold(jsonName, "id_"+parentKey) ||
|
||||
strings.EqualFold(jsonName, parentKey+"_id") ||
|
||||
strings.EqualFold(jsonName, parentKey+"id") ||
|
||||
strings.EqualFold(field.Name, parentKey+"ID") {
|
||||
dbColNames = []string{reflection.GetColumnName(field)}
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for _, dbColName := range dbColNames {
|
||||
if pkCol != "" && strings.EqualFold(dbColName, pkCol) {
|
||||
continue
|
||||
}
|
||||
if _, exists := data[dbColName]; !exists {
|
||||
logger.Debug("Injecting foreign key: %s = %v", dbColName, parentID)
|
||||
data[dbColName] = parentID
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -216,29 +302,35 @@ func (p *NestedCUDProcessor) processInsert(
|
||||
query := p.db.NewInsert().Table(tableName)
|
||||
|
||||
for key, value := range data {
|
||||
query = query.Value(key, value)
|
||||
query = query.Value(key, ConvertSliceForBun(value))
|
||||
}
|
||||
pkName := reflection.GetPrimaryKeyName(tableName)
|
||||
query = query.Returning(pkName)
|
||||
|
||||
// Add RETURNING clause to get the inserted ID
|
||||
query = query.Returning("id")
|
||||
|
||||
result, err := query.Exec(ctx)
|
||||
if err != nil {
|
||||
var id interface{}
|
||||
if err := query.Scan(ctx, &id); err != nil {
|
||||
logger.Error("Insert execution failed: table=%s, data=%+v, error=%v", tableName, data, err)
|
||||
return nil, fmt.Errorf("insert exec failed: %w", err)
|
||||
}
|
||||
|
||||
// Try to get the ID
|
||||
var id interface{}
|
||||
if lastID, err := result.LastInsertId(); err == nil && lastID > 0 {
|
||||
id = lastID
|
||||
} else if data["id"] != nil {
|
||||
id = data["id"]
|
||||
}
|
||||
|
||||
logger.Debug("Insert successful, ID: %v, rows affected: %d", id, result.RowsAffected())
|
||||
logger.Debug("Insert successful, ID: %v", id)
|
||||
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
|
||||
func (p *NestedCUDProcessor) processUpdate(
|
||||
ctx context.Context,
|
||||
@@ -247,6 +339,7 @@ func (p *NestedCUDProcessor) processUpdate(
|
||||
id interface{},
|
||||
) (int64, error) {
|
||||
if id == nil {
|
||||
logger.Error("Update requires an ID: table=%s, data=%+v", tableName, data)
|
||||
return 0, fmt.Errorf("update requires an ID")
|
||||
}
|
||||
|
||||
@@ -256,6 +349,7 @@ func (p *NestedCUDProcessor) processUpdate(
|
||||
|
||||
result, err := query.Exec(ctx)
|
||||
if err != nil {
|
||||
logger.Error("Update execution failed: table=%s, id=%v, data=%+v, error=%v", tableName, id, data, err)
|
||||
return 0, fmt.Errorf("update exec failed: %w", err)
|
||||
}
|
||||
|
||||
@@ -267,6 +361,7 @@ func (p *NestedCUDProcessor) processUpdate(
|
||||
// processDelete handles delete operation
|
||||
func (p *NestedCUDProcessor) processDelete(ctx context.Context, tableName string, id interface{}) (int64, error) {
|
||||
if id == nil {
|
||||
logger.Error("Delete requires an ID: table=%s", tableName)
|
||||
return 0, fmt.Errorf("delete requires an ID")
|
||||
}
|
||||
|
||||
@@ -276,6 +371,7 @@ func (p *NestedCUDProcessor) processDelete(ctx context.Context, tableName string
|
||||
|
||||
result, err := query.Exec(ctx)
|
||||
if err != nil {
|
||||
logger.Error("Delete execution failed: table=%s, id=%v, error=%v", tableName, id, err)
|
||||
return 0, fmt.Errorf("delete exec failed: %w", err)
|
||||
}
|
||||
|
||||
@@ -292,6 +388,7 @@ func (p *NestedCUDProcessor) processChildRelations(
|
||||
relationFields map[string]*RelationshipInfo,
|
||||
relationData map[string]interface{},
|
||||
parentModelType reflect.Type,
|
||||
incomingParentIDs map[string]interface{}, // IDs from all ancestors
|
||||
) error {
|
||||
for relationName, relInfo := range relationFields {
|
||||
relationValue, exists := relationData[relationName]
|
||||
@@ -304,7 +401,7 @@ func (p *NestedCUDProcessor) processChildRelations(
|
||||
// Get the related model
|
||||
field, found := parentModelType.FieldByName(relInfo.FieldName)
|
||||
if !found {
|
||||
logger.Warn("Field %s not found in model", relInfo.FieldName)
|
||||
logger.Error("Field %s not found in model type %v for relation %s", relInfo.FieldName, parentModelType, relationName)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -313,7 +410,7 @@ func (p *NestedCUDProcessor) processChildRelations(
|
||||
if relatedModelType.Kind() == reflect.Slice {
|
||||
relatedModelType = relatedModelType.Elem()
|
||||
}
|
||||
if relatedModelType.Kind() == reflect.Ptr {
|
||||
if relatedModelType.Kind() == reflect.Pointer {
|
||||
relatedModelType = relatedModelType.Elem()
|
||||
}
|
||||
|
||||
@@ -324,20 +421,93 @@ func (p *NestedCUDProcessor) processChildRelations(
|
||||
relatedTableName := p.getTableNameForModel(relatedModel, relInfo.JSONName)
|
||||
|
||||
// Prepare parent IDs for foreign key injection
|
||||
// Start by copying all incoming parent IDs (from ancestors)
|
||||
parentIDs := make(map[string]interface{})
|
||||
if relInfo.ForeignKey != "" {
|
||||
for k, v := range incomingParentIDs {
|
||||
parentIDs[k] = v
|
||||
}
|
||||
logger.Debug("Inherited %d parent IDs from ancestors: %+v", len(incomingParentIDs), incomingParentIDs)
|
||||
|
||||
// Add the current parent's primary key to the parentIDs map
|
||||
// This ensures nested children have access to all ancestor IDs
|
||||
if parentID != nil && parentModelType != nil {
|
||||
// Get the parent model's primary key field name
|
||||
parentPKFieldName := reflection.GetPrimaryKeyName(parentModelType)
|
||||
if parentPKFieldName != "" {
|
||||
// Get the JSON name for the primary key field
|
||||
parentPKJSONName := reflection.GetJSONNameForField(parentModelType, parentPKFieldName)
|
||||
baseName := ""
|
||||
if len(parentPKJSONName) > 1 {
|
||||
baseName = parentPKJSONName
|
||||
} else {
|
||||
// Add parent's PK to the map using the base model name
|
||||
baseName = strings.TrimSuffix(parentPKFieldName, "ID")
|
||||
baseName = strings.TrimSuffix(strings.ToLower(baseName), "_id")
|
||||
if baseName == "" {
|
||||
baseName = "parent"
|
||||
}
|
||||
}
|
||||
|
||||
parentIDs[baseName] = parentID
|
||||
logger.Debug("Added current parent PK to parentIDs map: %s=%v (from field %s)", baseName, parentID, parentPKFieldName)
|
||||
}
|
||||
}
|
||||
|
||||
// Also add the foreign key reference if specified
|
||||
if relInfo.ForeignKey != "" && parentID != nil {
|
||||
// Extract the base name from foreign key (e.g., "DepartmentID" -> "Department")
|
||||
baseName := strings.TrimSuffix(relInfo.ForeignKey, "ID")
|
||||
baseName = strings.TrimSuffix(strings.ToLower(baseName), "_id")
|
||||
parentIDs[baseName] = parentID
|
||||
// Only add if different from what we already added
|
||||
if _, exists := parentIDs[baseName]; !exists {
|
||||
parentIDs[baseName] = parentID
|
||||
logger.Debug("Added foreign key to parentIDs map: %s=%v (from FK %s)", baseName, parentID, relInfo.ForeignKey)
|
||||
}
|
||||
}
|
||||
|
||||
logger.Debug("Final parentIDs map for relation %s: %+v", relationName, parentIDs)
|
||||
|
||||
// Determine which field name to use for setting parent ID in child data
|
||||
// Priority: Use foreign key field name if specified
|
||||
var foreignKeyFieldName string
|
||||
if relInfo.ForeignKey != "" {
|
||||
// For has-many/has-one: join:parentCol=childCol
|
||||
// ForeignKey = parent side, References = child side (where we actually set the value)
|
||||
childField := relInfo.ForeignKey
|
||||
if (relInfo.RelationType == "hasMany" || relInfo.RelationType == "hasOne") && relInfo.References != "" {
|
||||
childField = relInfo.References
|
||||
}
|
||||
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
|
||||
childPKName := reflection.GetPrimaryKeyName(relatedModel)
|
||||
childPKFieldName := reflection.GetJSONNameForField(relatedModelType, childPKName)
|
||||
if childPKFieldName == "" {
|
||||
childPKFieldName = strings.ToLower(childPKName)
|
||||
}
|
||||
|
||||
logger.Debug("Processing relation with foreignKeyField=%s, childPK=%s", foreignKeyFieldName, childPKFieldName)
|
||||
|
||||
// Process based on relation type and data structure
|
||||
switch v := relationValue.(type) {
|
||||
case map[string]interface{}:
|
||||
// Single related object
|
||||
// Single related object - directly set foreign key if specified
|
||||
// IMPORTANT: In recursive relationships, don't overwrite the primary key
|
||||
if parentID != nil && foreignKeyFieldName != "" && foreignKeyFieldName != childPKFieldName {
|
||||
v[foreignKeyFieldName] = parentID
|
||||
logger.Debug("Set foreign key in single relation: %s=%v", foreignKeyFieldName, parentID)
|
||||
} else if foreignKeyFieldName == childPKFieldName {
|
||||
logger.Debug("Skipping foreign key assignment - same as primary key (recursive relationship): %s", foreignKeyFieldName)
|
||||
}
|
||||
_, err := p.ProcessNestedCUD(ctx, operation, v, relatedModel, parentIDs, relatedTableName)
|
||||
if err != nil {
|
||||
logger.Error("Failed to process single relation: name=%s, table=%s, operation=%s, parentID=%v, data=%+v, error=%v",
|
||||
relationName, relatedTableName, operation, parentID, v, err)
|
||||
return fmt.Errorf("failed to process relation %s: %w", relationName, err)
|
||||
}
|
||||
|
||||
@@ -345,24 +515,46 @@ func (p *NestedCUDProcessor) processChildRelations(
|
||||
// Multiple related objects
|
||||
for i, item := range v {
|
||||
if itemMap, ok := item.(map[string]interface{}); ok {
|
||||
// Directly set foreign key if specified
|
||||
// IMPORTANT: In recursive relationships, don't overwrite the primary key
|
||||
if parentID != nil && foreignKeyFieldName != "" && foreignKeyFieldName != childPKFieldName {
|
||||
itemMap[foreignKeyFieldName] = parentID
|
||||
logger.Debug("Set foreign key in relation array[%d]: %s=%v", i, foreignKeyFieldName, parentID)
|
||||
} else if foreignKeyFieldName == childPKFieldName {
|
||||
logger.Debug("Skipping foreign key assignment in array[%d] - same as primary key (recursive relationship): %s", i, foreignKeyFieldName)
|
||||
}
|
||||
_, err := p.ProcessNestedCUD(ctx, operation, itemMap, relatedModel, parentIDs, relatedTableName)
|
||||
if err != nil {
|
||||
logger.Error("Failed to process relation array item: name=%s[%d], table=%s, operation=%s, parentID=%v, data=%+v, error=%v",
|
||||
relationName, i, relatedTableName, operation, parentID, itemMap, err)
|
||||
return fmt.Errorf("failed to process relation %s[%d]: %w", relationName, i, err)
|
||||
}
|
||||
} else {
|
||||
logger.Warn("Relation array item is not a map: name=%s[%d], type=%T", relationName, i, item)
|
||||
}
|
||||
}
|
||||
|
||||
case []map[string]interface{}:
|
||||
// Multiple related objects (typed slice)
|
||||
for i, itemMap := range v {
|
||||
// Directly set foreign key if specified
|
||||
// IMPORTANT: In recursive relationships, don't overwrite the primary key
|
||||
if parentID != nil && foreignKeyFieldName != "" && foreignKeyFieldName != childPKFieldName {
|
||||
itemMap[foreignKeyFieldName] = parentID
|
||||
logger.Debug("Set foreign key in relation typed array[%d]: %s=%v", i, foreignKeyFieldName, parentID)
|
||||
} else if foreignKeyFieldName == childPKFieldName {
|
||||
logger.Debug("Skipping foreign key assignment in typed array[%d] - same as primary key (recursive relationship): %s", i, foreignKeyFieldName)
|
||||
}
|
||||
_, err := p.ProcessNestedCUD(ctx, operation, itemMap, relatedModel, parentIDs, relatedTableName)
|
||||
if err != nil {
|
||||
logger.Error("Failed to process relation typed array item: name=%s[%d], table=%s, operation=%s, parentID=%v, data=%+v, error=%v",
|
||||
relationName, i, relatedTableName, operation, parentID, itemMap, err)
|
||||
return fmt.Errorf("failed to process relation %s[%d]: %w", relationName, i, err)
|
||||
}
|
||||
}
|
||||
|
||||
default:
|
||||
logger.Warn("Unsupported relation data type for %s: %T", relationName, relationValue)
|
||||
logger.Error("Unsupported relation data type: name=%s, type=%T, value=%+v", relationName, relationValue, relationValue)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -398,7 +590,7 @@ func shouldUseNestedProcessorDepth(data map[string]interface{}, model interface{
|
||||
|
||||
// Get model type
|
||||
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()
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,943 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||
)
|
||||
|
||||
// Mock Database for testing
|
||||
type mockDatabase struct {
|
||||
insertCalls []map[string]interface{}
|
||||
updateCalls []map[string]interface{}
|
||||
deleteCalls []interface{}
|
||||
lastID int64
|
||||
}
|
||||
|
||||
func newMockDatabase() *mockDatabase {
|
||||
return &mockDatabase{
|
||||
insertCalls: make([]map[string]interface{}, 0),
|
||||
updateCalls: make([]map[string]interface{}, 0),
|
||||
deleteCalls: make([]interface{}, 0),
|
||||
lastID: 1,
|
||||
}
|
||||
}
|
||||
|
||||
func (m *mockDatabase) NewSelect() SelectQuery { return &mockSelectQuery{} }
|
||||
func (m *mockDatabase) NewInsert() InsertQuery { return &mockInsertQuery{db: m} }
|
||||
func (m *mockDatabase) NewUpdate() UpdateQuery { return &mockUpdateQuery{db: m} }
|
||||
func (m *mockDatabase) NewDelete() DeleteQuery { return &mockDeleteQuery{db: m} }
|
||||
func (m *mockDatabase) RunInTransaction(ctx context.Context, fn func(Database) error) error {
|
||||
return fn(m)
|
||||
}
|
||||
func (m *mockDatabase) Exec(ctx context.Context, query string, args ...interface{}) (Result, error) {
|
||||
return &mockResult{rowsAffected: 1}, nil
|
||||
}
|
||||
func (m *mockDatabase) Query(ctx context.Context, dest interface{}, query string, args ...interface{}) error {
|
||||
return nil
|
||||
}
|
||||
func (m *mockDatabase) BeginTx(ctx context.Context) (Database, error) {
|
||||
return m, nil
|
||||
}
|
||||
func (m *mockDatabase) CommitTx(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
func (m *mockDatabase) RollbackTx(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
func (m *mockDatabase) GetUnderlyingDB() interface{} {
|
||||
return nil
|
||||
}
|
||||
func (m *mockDatabase) DriverName() string {
|
||||
return "postgres"
|
||||
}
|
||||
|
||||
// Mock SelectQuery
|
||||
type mockSelectQuery struct{}
|
||||
|
||||
func (m *mockSelectQuery) Model(model interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Table(name string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Column(columns ...string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) ColumnExpr(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Where(condition string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) WhereOr(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Join(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) LeftJoin(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Preload(relation string, conditions ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) PreloadRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) JoinRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Order(order string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) OrderExpr(order string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Limit(n int) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Offset(n int) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Group(group string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Having(condition string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Scan(ctx context.Context, dest interface{}) error { return nil }
|
||||
func (m *mockSelectQuery) ScanModel(ctx context.Context) error { return nil }
|
||||
func (m *mockSelectQuery) Count(ctx context.Context) (int, error) { return 0, nil }
|
||||
func (m *mockSelectQuery) Exists(ctx context.Context) (bool, error) { return false, nil }
|
||||
|
||||
// Mock InsertQuery
|
||||
type mockInsertQuery struct {
|
||||
db *mockDatabase
|
||||
table string
|
||||
values map[string]interface{}
|
||||
}
|
||||
|
||||
func (m *mockInsertQuery) Model(model interface{}) InsertQuery { return m }
|
||||
func (m *mockInsertQuery) Table(name string) InsertQuery {
|
||||
m.table = name
|
||||
return m
|
||||
}
|
||||
func (m *mockInsertQuery) Value(column string, value interface{}) InsertQuery {
|
||||
if m.values == nil {
|
||||
m.values = make(map[string]interface{})
|
||||
}
|
||||
m.values[column] = value
|
||||
return m
|
||||
}
|
||||
func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m }
|
||||
func (m *mockInsertQuery) Returning(columns ...string) InsertQuery { return m }
|
||||
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
|
||||
m.db.insertCalls = append(m.db.insertCalls, m.values)
|
||||
m.db.lastID++
|
||||
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
|
||||
type mockUpdateQuery struct {
|
||||
db *mockDatabase
|
||||
table string
|
||||
setValues map[string]interface{}
|
||||
}
|
||||
|
||||
func (m *mockUpdateQuery) Model(model interface{}) UpdateQuery { return m }
|
||||
func (m *mockUpdateQuery) Table(name string) UpdateQuery {
|
||||
m.table = name
|
||||
return m
|
||||
}
|
||||
func (m *mockUpdateQuery) Set(column string, value interface{}) UpdateQuery { return m }
|
||||
func (m *mockUpdateQuery) SetMap(values map[string]interface{}) UpdateQuery {
|
||||
m.setValues = values
|
||||
return m
|
||||
}
|
||||
func (m *mockUpdateQuery) Where(condition string, args ...interface{}) UpdateQuery { return m }
|
||||
func (m *mockUpdateQuery) Returning(columns ...string) UpdateQuery { return m }
|
||||
func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) {
|
||||
// Record the update call
|
||||
m.db.updateCalls = append(m.db.updateCalls, m.setValues)
|
||||
return &mockResult{rowsAffected: 1}, nil
|
||||
}
|
||||
|
||||
// Mock DeleteQuery
|
||||
type mockDeleteQuery struct {
|
||||
db *mockDatabase
|
||||
table string
|
||||
}
|
||||
|
||||
func (m *mockDeleteQuery) Model(model interface{}) DeleteQuery { return m }
|
||||
func (m *mockDeleteQuery) Table(name string) DeleteQuery {
|
||||
m.table = name
|
||||
return m
|
||||
}
|
||||
func (m *mockDeleteQuery) Where(condition string, args ...interface{}) DeleteQuery { return m }
|
||||
func (m *mockDeleteQuery) Exec(ctx context.Context) (Result, error) {
|
||||
// Record the delete call
|
||||
m.db.deleteCalls = append(m.db.deleteCalls, m.table)
|
||||
return &mockResult{rowsAffected: 1}, nil
|
||||
}
|
||||
|
||||
// Mock Result
|
||||
type mockResult struct {
|
||||
lastID int64
|
||||
rowsAffected int64
|
||||
}
|
||||
|
||||
func (m *mockResult) LastInsertId() (int64, error) { return m.lastID, nil }
|
||||
func (m *mockResult) RowsAffected() int64 { return m.rowsAffected }
|
||||
|
||||
// Mock ModelRegistry
|
||||
type mockModelRegistry struct{}
|
||||
|
||||
func (m *mockModelRegistry) GetModel(name string) (interface{}, error) { return nil, nil }
|
||||
func (m *mockModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) { return nil, nil }
|
||||
func (m *mockModelRegistry) RegisterModel(name string, model interface{}) error { return nil }
|
||||
func (m *mockModelRegistry) GetAllModels() map[string]interface{} { return make(map[string]interface{}) }
|
||||
|
||||
// Mock RelationshipInfoProvider
|
||||
type mockRelationshipProvider struct {
|
||||
relationships map[string]*RelationshipInfo
|
||||
}
|
||||
|
||||
func newMockRelationshipProvider() *mockRelationshipProvider {
|
||||
return &mockRelationshipProvider{
|
||||
relationships: make(map[string]*RelationshipInfo),
|
||||
}
|
||||
}
|
||||
|
||||
func (m *mockRelationshipProvider) GetRelationshipInfo(modelType reflect.Type, relationName string) *RelationshipInfo {
|
||||
key := modelType.Name() + "." + relationName
|
||||
return m.relationships[key]
|
||||
}
|
||||
|
||||
func (m *mockRelationshipProvider) RegisterRelation(modelTypeName, relationName string, info *RelationshipInfo) {
|
||||
key := modelTypeName + "." + relationName
|
||||
m.relationships[key] = info
|
||||
}
|
||||
|
||||
// Test Models
|
||||
type Department struct {
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Name string `json:"name"`
|
||||
Employees []*Employee `json:"employees,omitempty"`
|
||||
}
|
||||
|
||||
func (d Department) TableName() string { return "departments" }
|
||||
func (d Department) GetIDName() string { return "ID" }
|
||||
|
||||
type Employee struct {
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Name string `json:"name"`
|
||||
DepartmentID int64 `json:"department_id"`
|
||||
Tasks []*Task `json:"tasks,omitempty"`
|
||||
}
|
||||
|
||||
func (e Employee) TableName() string { return "employees" }
|
||||
func (e Employee) GetIDName() string { return "ID" }
|
||||
|
||||
type Task struct {
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Title string `json:"title"`
|
||||
EmployeeID int64 `json:"employee_id"`
|
||||
Comments []*Comment `json:"comments,omitempty"`
|
||||
}
|
||||
|
||||
func (t Task) TableName() string { return "tasks" }
|
||||
func (t Task) GetIDName() string { return "ID" }
|
||||
|
||||
type Comment struct {
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Text string `json:"text"`
|
||||
TaskID int64 `json:"task_id"`
|
||||
}
|
||||
|
||||
func (c Comment) TableName() string { return "comments" }
|
||||
func (c Comment) GetIDName() string { return "ID" }
|
||||
|
||||
// Test Cases
|
||||
|
||||
func TestProcessNestedCUD_SingleLevelInsert(t *testing.T) {
|
||||
db := newMockDatabase()
|
||||
registry := &mockModelRegistry{}
|
||||
relProvider := newMockRelationshipProvider()
|
||||
|
||||
// Register Department -> Employees relationship
|
||||
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": "John Doe",
|
||||
},
|
||||
map[string]interface{}{
|
||||
"name": "Jane Smith",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result, err := processor.ProcessNestedCUD(
|
||||
context.Background(),
|
||||
"insert",
|
||||
data,
|
||||
Department{},
|
||||
nil,
|
||||
"departments",
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessNestedCUD failed: %v", err)
|
||||
}
|
||||
|
||||
if result.ID == nil {
|
||||
t.Error("Expected result.ID to be set")
|
||||
}
|
||||
|
||||
// Verify department was inserted
|
||||
if len(db.insertCalls) != 3 {
|
||||
t.Errorf("Expected 3 insert calls (1 dept + 2 employees), got %d", len(db.insertCalls))
|
||||
}
|
||||
|
||||
// Verify first insert is department
|
||||
if db.insertCalls[0]["name"] != "Engineering" {
|
||||
t.Errorf("Expected department name 'Engineering', got %v", db.insertCalls[0]["name"])
|
||||
}
|
||||
|
||||
// Verify employees were inserted with foreign key
|
||||
if db.insertCalls[1]["department_id"] == nil {
|
||||
t.Error("Expected employee to have department_id set")
|
||||
}
|
||||
if db.insertCalls[2]["department_id"] == nil {
|
||||
t.Error("Expected employee to have department_id set")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessNestedCUD_MultiLevelInsert(t *testing.T) {
|
||||
db := newMockDatabase()
|
||||
registry := &mockModelRegistry{}
|
||||
relProvider := newMockRelationshipProvider()
|
||||
|
||||
// Register relationships
|
||||
relProvider.RegisterRelation("Department", "employees", &RelationshipInfo{
|
||||
FieldName: "Employees",
|
||||
JSONName: "employees",
|
||||
RelationType: "has_many",
|
||||
ForeignKey: "DepartmentID",
|
||||
RelatedModel: Employee{},
|
||||
})
|
||||
|
||||
relProvider.RegisterRelation("Employee", "tasks", &RelationshipInfo{
|
||||
FieldName: "Tasks",
|
||||
JSONName: "tasks",
|
||||
RelationType: "has_many",
|
||||
ForeignKey: "EmployeeID",
|
||||
RelatedModel: Task{},
|
||||
})
|
||||
|
||||
processor := NewNestedCUDProcessor(db, registry, relProvider)
|
||||
|
||||
data := map[string]interface{}{
|
||||
"name": "Engineering",
|
||||
"employees": []interface{}{
|
||||
map[string]interface{}{
|
||||
"name": "John Doe",
|
||||
"tasks": []interface{}{
|
||||
map[string]interface{}{
|
||||
"title": "Task 1",
|
||||
},
|
||||
map[string]interface{}{
|
||||
"title": "Task 2",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result, err := processor.ProcessNestedCUD(
|
||||
context.Background(),
|
||||
"insert",
|
||||
data,
|
||||
Department{},
|
||||
nil,
|
||||
"departments",
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessNestedCUD failed: %v", err)
|
||||
}
|
||||
|
||||
if result.ID == nil {
|
||||
t.Error("Expected result.ID to be set")
|
||||
}
|
||||
|
||||
// Verify: 1 dept + 1 employee + 2 tasks = 4 inserts
|
||||
if len(db.insertCalls) != 4 {
|
||||
t.Errorf("Expected 4 insert calls, got %d", len(db.insertCalls))
|
||||
}
|
||||
|
||||
// Verify department
|
||||
if db.insertCalls[0]["name"] != "Engineering" {
|
||||
t.Errorf("Expected department name 'Engineering', got %v", db.insertCalls[0]["name"])
|
||||
}
|
||||
|
||||
// Verify employee has department_id
|
||||
if db.insertCalls[1]["department_id"] == nil {
|
||||
t.Error("Expected employee to have department_id set")
|
||||
}
|
||||
|
||||
// Verify tasks have employee_id
|
||||
if db.insertCalls[2]["employee_id"] == nil {
|
||||
t.Error("Expected task to have employee_id set")
|
||||
}
|
||||
if db.insertCalls[3]["employee_id"] == nil {
|
||||
t.Error("Expected task to have employee_id set")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessNestedCUD_RequestFieldOverride(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{}{
|
||||
"name": "Engineering",
|
||||
"employees": []interface{}{
|
||||
map[string]interface{}{
|
||||
"_request": "update",
|
||||
"ID": int64(10), // Use capital ID to match struct field
|
||||
"name": "John Updated",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := processor.ProcessNestedCUD(
|
||||
context.Background(),
|
||||
"insert",
|
||||
data,
|
||||
Department{},
|
||||
nil,
|
||||
"departments",
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessNestedCUD failed: %v", err)
|
||||
}
|
||||
|
||||
// Verify department was inserted (1 insert)
|
||||
// Employee should be updated (1 update)
|
||||
if len(db.insertCalls) != 1 {
|
||||
t.Errorf("Expected 1 insert call for department, got %d", len(db.insertCalls))
|
||||
}
|
||||
|
||||
if len(db.updateCalls) != 1 {
|
||||
t.Errorf("Expected 1 update call for employee, got %d", len(db.updateCalls))
|
||||
}
|
||||
|
||||
// Verify update data
|
||||
if db.updateCalls[0]["name"] != "John Updated" {
|
||||
t.Errorf("Expected employee name 'John Updated', got %v", db.updateCalls[0]["name"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessNestedCUD_SkipInsertWhenOnlyRequestField(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 with only _request field for nested employee
|
||||
data := map[string]interface{}{
|
||||
"name": "Engineering",
|
||||
"employees": []interface{}{
|
||||
map[string]interface{}{
|
||||
"_request": "insert",
|
||||
// No other fields besides _request
|
||||
// Note: Foreign key will be injected, so employee WILL be inserted
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := processor.ProcessNestedCUD(
|
||||
context.Background(),
|
||||
"insert",
|
||||
data,
|
||||
Department{},
|
||||
nil,
|
||||
"departments",
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessNestedCUD failed: %v", err)
|
||||
}
|
||||
|
||||
// Department + Employee (with injected FK) = 2 inserts
|
||||
if len(db.insertCalls) != 2 {
|
||||
t.Errorf("Expected 2 insert calls (department + employee with FK), got %d", len(db.insertCalls))
|
||||
}
|
||||
|
||||
if db.insertCalls[0]["name"] != "Engineering" {
|
||||
t.Errorf("Expected department name 'Engineering', got %v", db.insertCalls[0]["name"])
|
||||
}
|
||||
|
||||
// Verify employee has foreign key
|
||||
if db.insertCalls[1]["department_id"] == nil {
|
||||
t.Error("Expected employee to have department_id injected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessNestedCUD_Update(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), // Use capital ID to match struct field
|
||||
"name": "Engineering Updated",
|
||||
"employees": []interface{}{
|
||||
map[string]interface{}{
|
||||
"_request": "insert",
|
||||
"name": "New Employee",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result, err := processor.ProcessNestedCUD(
|
||||
context.Background(),
|
||||
"update",
|
||||
data,
|
||||
Department{},
|
||||
nil,
|
||||
"departments",
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessNestedCUD failed: %v", err)
|
||||
}
|
||||
|
||||
if result.ID != int64(1) {
|
||||
t.Errorf("Expected result.ID to be 1, got %v", result.ID)
|
||||
}
|
||||
|
||||
// Verify department was updated
|
||||
if len(db.updateCalls) != 1 {
|
||||
t.Errorf("Expected 1 update call, got %d", len(db.updateCalls))
|
||||
}
|
||||
|
||||
// Verify new employee was inserted
|
||||
if len(db.insertCalls) != 1 {
|
||||
t.Errorf("Expected 1 insert call for new employee, got %d", len(db.insertCalls))
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessNestedCUD_Delete(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), // Use capital ID to match struct field
|
||||
"employees": []interface{}{
|
||||
map[string]interface{}{
|
||||
"_request": "delete",
|
||||
"ID": int64(10), // Use capital ID
|
||||
},
|
||||
map[string]interface{}{
|
||||
"_request": "delete",
|
||||
"ID": int64(11), // Use capital ID
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := processor.ProcessNestedCUD(
|
||||
context.Background(),
|
||||
"delete",
|
||||
data,
|
||||
Department{},
|
||||
nil,
|
||||
"departments",
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessNestedCUD failed: %v", err)
|
||||
}
|
||||
|
||||
// Verify employees were deleted first, then department
|
||||
// 2 employees + 1 department = 3 deletes
|
||||
if len(db.deleteCalls) != 3 {
|
||||
t.Errorf("Expected 3 delete calls, got %d", len(db.deleteCalls))
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessNestedCUD_ParentIDPropagation(t *testing.T) {
|
||||
db := newMockDatabase()
|
||||
registry := &mockModelRegistry{}
|
||||
relProvider := newMockRelationshipProvider()
|
||||
|
||||
// Register 3-level relationships
|
||||
relProvider.RegisterRelation("Department", "employees", &RelationshipInfo{
|
||||
FieldName: "Employees",
|
||||
JSONName: "employees",
|
||||
RelationType: "has_many",
|
||||
ForeignKey: "DepartmentID",
|
||||
RelatedModel: Employee{},
|
||||
})
|
||||
|
||||
relProvider.RegisterRelation("Employee", "tasks", &RelationshipInfo{
|
||||
FieldName: "Tasks",
|
||||
JSONName: "tasks",
|
||||
RelationType: "has_many",
|
||||
ForeignKey: "EmployeeID",
|
||||
RelatedModel: Task{},
|
||||
})
|
||||
|
||||
relProvider.RegisterRelation("Task", "comments", &RelationshipInfo{
|
||||
FieldName: "Comments",
|
||||
JSONName: "comments",
|
||||
RelationType: "has_many",
|
||||
ForeignKey: "TaskID",
|
||||
RelatedModel: Comment{},
|
||||
})
|
||||
|
||||
processor := NewNestedCUDProcessor(db, registry, relProvider)
|
||||
|
||||
data := map[string]interface{}{
|
||||
"name": "Engineering",
|
||||
"employees": []interface{}{
|
||||
map[string]interface{}{
|
||||
"name": "John",
|
||||
"tasks": []interface{}{
|
||||
map[string]interface{}{
|
||||
"title": "Task 1",
|
||||
"comments": []interface{}{
|
||||
map[string]interface{}{
|
||||
"text": "Great work!",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, err := processor.ProcessNestedCUD(
|
||||
context.Background(),
|
||||
"insert",
|
||||
data,
|
||||
Department{},
|
||||
nil,
|
||||
"departments",
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("ProcessNestedCUD failed: %v", err)
|
||||
}
|
||||
|
||||
// Verify: 1 dept + 1 employee + 1 task + 1 comment = 4 inserts
|
||||
if len(db.insertCalls) != 4 {
|
||||
t.Errorf("Expected 4 insert calls, got %d", len(db.insertCalls))
|
||||
}
|
||||
|
||||
// Verify department
|
||||
if db.insertCalls[0]["name"] != "Engineering" {
|
||||
t.Error("Expected department to be inserted first")
|
||||
}
|
||||
|
||||
// Verify employee has department_id
|
||||
if db.insertCalls[1]["department_id"] == nil {
|
||||
t.Error("Expected employee to have department_id")
|
||||
}
|
||||
|
||||
// Verify task has employee_id
|
||||
if db.insertCalls[2]["employee_id"] == nil {
|
||||
t.Error("Expected task to have employee_id")
|
||||
}
|
||||
|
||||
// Verify comment has task_id
|
||||
if db.insertCalls[3]["task_id"] == nil {
|
||||
t.Error("Expected comment to have task_id")
|
||||
}
|
||||
}
|
||||
|
||||
func TestInjectForeignKeys(t *testing.T) {
|
||||
db := newMockDatabase()
|
||||
registry := &mockModelRegistry{}
|
||||
relProvider := newMockRelationshipProvider()
|
||||
|
||||
processor := NewNestedCUDProcessor(db, registry, relProvider)
|
||||
|
||||
data := map[string]interface{}{
|
||||
"name": "John",
|
||||
}
|
||||
|
||||
parentIDs := map[string]interface{}{
|
||||
"department": int64(5),
|
||||
}
|
||||
|
||||
modelType := reflect.TypeOf(Employee{})
|
||||
|
||||
processor.injectForeignKeys(data, modelType, parentIDs)
|
||||
|
||||
// Should inject department_id based on the "department" key in parentIDs
|
||||
if data["department_id"] == nil {
|
||||
t.Error("Expected department_id to be injected")
|
||||
}
|
||||
|
||||
if data["department_id"] != int64(5) {
|
||||
t.Errorf("Expected department_id to be 5, got %v", data["department_id"])
|
||||
}
|
||||
}
|
||||
|
||||
// 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) {
|
||||
dept := Department{}
|
||||
pkName := reflection.GetPrimaryKeyName(dept)
|
||||
|
||||
if pkName != "ID" {
|
||||
t.Errorf("Expected primary key name 'ID', got '%s'", pkName)
|
||||
}
|
||||
|
||||
// Test with pointer
|
||||
pkName2 := reflection.GetPrimaryKeyName(&dept)
|
||||
if pkName2 != "ID" {
|
||||
t.Errorf("Expected primary key name 'ID' from pointer, got '%s'", pkName2)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
+274
-25
@@ -2,6 +2,7 @@ package common
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
@@ -58,6 +59,38 @@ func IsSQLExpression(cond string) bool {
|
||||
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
|
||||
// These conditions should be removed from WHERE clauses as they have no filtering effect
|
||||
func IsTrivialCondition(cond string) bool {
|
||||
@@ -130,6 +163,9 @@ func validateWhereClauseSecurity(where string) error {
|
||||
// Note: This function will NOT add prefixes to unprefixed columns. It will only fix
|
||||
// incorrect prefixes (e.g., wrong_table.column -> correct_table.column), unless the
|
||||
// prefix matches a preloaded relation name, in which case it's left unchanged.
|
||||
//
|
||||
// IMPORTANT: Outer parentheses are preserved if the clause contains top-level OR operators
|
||||
// to prevent OR logic from escaping and affecting the entire query incorrectly.
|
||||
func SanitizeWhereClause(where string, tableName string, options ...*RequestOptions) string {
|
||||
if where == "" {
|
||||
return ""
|
||||
@@ -143,8 +179,27 @@ func SanitizeWhereClause(where string, tableName string, options ...*RequestOpti
|
||||
return ""
|
||||
}
|
||||
|
||||
// Strip outer parentheses and re-trim
|
||||
where = stripOuterParentheses(where)
|
||||
// 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
|
||||
// If so, we need to preserve the outer parentheses to prevent OR logic from escaping
|
||||
hasOuterParens := false
|
||||
if len(where) > 0 && where[0] == '(' && where[len(where)-1] == ')' {
|
||||
_, hasOuterParens = stripOneMatchingOuterParen(where)
|
||||
}
|
||||
|
||||
// Strip outer parentheses and re-trim for processing
|
||||
whereWithoutParens := stripOuterParentheses(where)
|
||||
shouldPreserveParens := hasOuterParens && containsTopLevelOR(whereWithoutParens)
|
||||
|
||||
// Use the stripped version for processing
|
||||
where = whereWithoutParens
|
||||
|
||||
// Get valid columns from the model if tableName is provided
|
||||
var validColumns map[string]bool
|
||||
@@ -153,19 +208,28 @@ func SanitizeWhereClause(where string, tableName string, options ...*RequestOpti
|
||||
}
|
||||
|
||||
// Build a set of allowed table prefixes (main table + preloaded relations)
|
||||
// Keys are stored lowercase for case-insensitive matching
|
||||
allowedPrefixes := make(map[string]bool)
|
||||
if tableName != "" {
|
||||
allowedPrefixes[tableName] = true
|
||||
allowedPrefixes[strings.ToLower(tableName)] = true
|
||||
}
|
||||
|
||||
// Add preload relation names as allowed prefixes
|
||||
if len(options) > 0 && options[0] != nil {
|
||||
for pi := range options[0].Preload {
|
||||
if options[0].Preload[pi].Relation != "" {
|
||||
allowedPrefixes[options[0].Preload[pi].Relation] = true
|
||||
allowedPrefixes[strings.ToLower(options[0].Preload[pi].Relation)] = true
|
||||
logger.Debug("Added preload relation '%s' as allowed table prefix", options[0].Preload[pi].Relation)
|
||||
}
|
||||
}
|
||||
|
||||
// Add join aliases as allowed prefixes
|
||||
for _, alias := range options[0].JoinAliases {
|
||||
if alias != "" {
|
||||
allowedPrefixes[strings.ToLower(alias)] = true
|
||||
logger.Debug("Added join alias '%s' as allowed table prefix", alias)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Split by AND to handle multiple conditions
|
||||
@@ -188,14 +252,20 @@ func SanitizeWhereClause(where string, tableName string, options ...*RequestOpti
|
||||
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 != "" && hasTablePrefix(condToCheck) {
|
||||
// Extract the current prefix and column name
|
||||
currentPrefix, columnName := extractTableAndColumn(condToCheck)
|
||||
|
||||
if currentPrefix != "" && columnName != "" {
|
||||
// Check if the prefix is allowed (main table or preload relation)
|
||||
if !allowedPrefixes[currentPrefix] {
|
||||
// Check if the prefix is allowed (main table or preload relation) - case-insensitive
|
||||
if !allowedPrefixes[strings.ToLower(currentPrefix)] {
|
||||
// Prefix is not in the allowed list - only fix if it's a valid column in the main table
|
||||
if validColumns == nil || isValidColumn(columnName, validColumns) {
|
||||
// Replace the incorrect prefix with the correct main table name
|
||||
@@ -221,7 +291,14 @@ func SanitizeWhereClause(where string, tableName string, options ...*RequestOpti
|
||||
|
||||
result := strings.Join(validConditions, " AND ")
|
||||
|
||||
if result != where {
|
||||
// If the original clause had outer parentheses and contains OR operators,
|
||||
// restore the outer parentheses to prevent OR logic from escaping
|
||||
if shouldPreserveParens {
|
||||
result = "(" + result + ")"
|
||||
logger.Debug("Preserved outer parentheses for OR conditions: '%s'", result)
|
||||
}
|
||||
|
||||
if result != where && !shouldPreserveParens {
|
||||
logger.Debug("Sanitized WHERE clause: '%s' -> '%s'", where, result)
|
||||
}
|
||||
|
||||
@@ -282,18 +359,123 @@ func stripOneMatchingOuterParen(s string) (string, bool) {
|
||||
return strings.TrimSpace(s[1 : len(s)-1]), true
|
||||
}
|
||||
|
||||
// splitByAND splits a WHERE clause by AND operators (case-insensitive)
|
||||
// This is parenthesis-aware and won't split on AND operators inside subqueries
|
||||
// EnsureOuterParentheses ensures that a SQL clause is wrapped in parentheses
|
||||
// to prevent OR logic from escaping. It checks if the clause already has
|
||||
// matching outer parentheses and only adds them if they don't exist.
|
||||
//
|
||||
// This is particularly important for OR conditions and complex filters where
|
||||
// the absence of parentheses could cause the logic to escape and affect
|
||||
// the entire query incorrectly.
|
||||
//
|
||||
// Parameters:
|
||||
// - clause: The SQL clause to check and potentially wrap
|
||||
//
|
||||
// Returns:
|
||||
// - The clause with guaranteed outer parentheses, or empty string if input is empty
|
||||
func EnsureOuterParentheses(clause string) string {
|
||||
if clause == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
clause = strings.TrimSpace(clause)
|
||||
if clause == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
// Check if the clause already has matching outer parentheses
|
||||
_, hasOuterParens := stripOneMatchingOuterParen(clause)
|
||||
|
||||
// If it already has matching outer parentheses, return as-is
|
||||
if hasOuterParens {
|
||||
return clause
|
||||
}
|
||||
|
||||
// Otherwise, wrap it in parentheses
|
||||
return "(" + clause + ")"
|
||||
}
|
||||
|
||||
// containsTopLevelOR checks if a SQL clause contains OR operators at the top level
|
||||
// (i.e., not inside parentheses or subqueries). This is used to determine if
|
||||
// outer parentheses should be preserved to prevent OR logic from escaping.
|
||||
func containsTopLevelOR(clause string) bool {
|
||||
if clause == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
depth := 0
|
||||
inSingleQuote := false
|
||||
inDoubleQuote := false
|
||||
lowerClause := strings.ToLower(clause)
|
||||
|
||||
for i := 0; i < len(clause); i++ {
|
||||
ch := clause[i]
|
||||
|
||||
// Track quote state
|
||||
if ch == '\'' && !inDoubleQuote {
|
||||
inSingleQuote = !inSingleQuote
|
||||
continue
|
||||
}
|
||||
if ch == '"' && !inSingleQuote {
|
||||
inDoubleQuote = !inDoubleQuote
|
||||
continue
|
||||
}
|
||||
|
||||
// Skip if inside quotes
|
||||
if inSingleQuote || inDoubleQuote {
|
||||
continue
|
||||
}
|
||||
|
||||
// Track parenthesis depth
|
||||
switch ch {
|
||||
case '(':
|
||||
depth++
|
||||
case ')':
|
||||
depth--
|
||||
}
|
||||
|
||||
// Only check for OR at depth 0 (not inside parentheses)
|
||||
if depth == 0 && i+4 <= len(clause) {
|
||||
// Check for " OR " (case-insensitive)
|
||||
substring := lowerClause[i : i+4]
|
||||
if substring == " or " {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// splitByAND splits a WHERE clause by AND operators (case-insensitive).
|
||||
// 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 {
|
||||
conditions := []string{}
|
||||
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
|
||||
|
||||
for i < len(where) {
|
||||
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 == '(' {
|
||||
depth++
|
||||
currentCondition.WriteByte(ch)
|
||||
@@ -306,32 +488,39 @@ func splitByAND(where string) []string {
|
||||
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 {
|
||||
// Check if we're at an AND operator (case-insensitive)
|
||||
// We need at least " AND " (5 chars) or " and " (5 chars)
|
||||
if i+5 <= len(where) {
|
||||
substring := where[i : i+5]
|
||||
lowerSubstring := strings.ToLower(substring)
|
||||
// Detect " BETWEEN " (9 chars, case-insensitive) so the very next
|
||||
// top-level AND is recognised as part of the BETWEEN syntax.
|
||||
if i+9 <= len(where) && strings.ToLower(where[i:i+9]) == " between " {
|
||||
afterBetween = true
|
||||
currentCondition.WriteString(where[i : i+9])
|
||||
i += 9
|
||||
continue
|
||||
}
|
||||
|
||||
if lowerSubstring == " and " {
|
||||
// Found an AND operator at the top level
|
||||
// Add the current condition to the list
|
||||
conditions = append(conditions, currentCondition.String())
|
||||
currentCondition.Reset()
|
||||
// Skip past the AND operator
|
||||
// Detect " AND " (5 chars, case-insensitive).
|
||||
if i+5 <= len(where) && strings.ToLower(where[i:i+5]) == " and " {
|
||||
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())
|
||||
currentCondition.Reset()
|
||||
i += 5
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
// Not an AND operator or we're inside parentheses, just add the character
|
||||
currentCondition.WriteByte(ch)
|
||||
i++
|
||||
}
|
||||
|
||||
// Add the last condition
|
||||
// Add the last condition.
|
||||
if currentCondition.Len() > 0 {
|
||||
conditions = append(conditions, currentCondition.String())
|
||||
}
|
||||
@@ -450,6 +639,15 @@ func extractTableAndColumn(cond string) (table string, column string) {
|
||||
// Remove any quotes
|
||||
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)
|
||||
openParenIdx := strings.Index(columnRef, "(")
|
||||
|
||||
@@ -809,3 +1007,54 @@ func extractLeftSideOfComparison(cond string) string {
|
||||
|
||||
return ""
|
||||
}
|
||||
|
||||
// FilterValueToSlice converts a filter value to []interface{} for use with IN operators.
|
||||
// JSON-decoded arrays arrive as []interface{}, but typed slices (e.g. []string) also work.
|
||||
// Returns a single-element slice if the value is not a slice type.
|
||||
func FilterValueToSlice(v interface{}) []interface{} {
|
||||
if v == nil {
|
||||
return nil
|
||||
}
|
||||
rv := reflect.ValueOf(v)
|
||||
if rv.Kind() == reflect.Slice {
|
||||
result := make([]interface{}, rv.Len())
|
||||
for i := 0; i < rv.Len(); i++ {
|
||||
result[i] = rv.Index(i).Interface()
|
||||
}
|
||||
return result
|
||||
}
|
||||
return []interface{}{v}
|
||||
}
|
||||
|
||||
// BuildInCondition builds a parameterized IN condition from a filter value.
|
||||
// Returns the condition string (e.g. "col IN (?,?)") and the individual values as args.
|
||||
// Returns ("", nil) if the value is empty or not a slice.
|
||||
func BuildInCondition(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 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
|
||||
}
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestSanitizeWhereClause_WithTableName tests that table prefixes in WHERE clauses
|
||||
// are correctly handled when the tableName parameter matches the prefix
|
||||
func TestSanitizeWhereClause_WithTableName(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
where string
|
||||
tableName string
|
||||
options *RequestOptions
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "Correct table prefix should not be changed",
|
||||
where: "mastertaskitem.rid_parentmastertaskitem is null",
|
||||
tableName: "mastertaskitem",
|
||||
options: nil,
|
||||
expected: "mastertaskitem.rid_parentmastertaskitem is null",
|
||||
},
|
||||
{
|
||||
name: "Wrong table prefix should be fixed",
|
||||
where: "wrong_table.rid_parentmastertaskitem is null",
|
||||
tableName: "mastertaskitem",
|
||||
options: nil,
|
||||
expected: "mastertaskitem.rid_parentmastertaskitem is null",
|
||||
},
|
||||
{
|
||||
name: "Relation name should not replace correct table prefix",
|
||||
where: "mastertaskitem.rid_parentmastertaskitem is null",
|
||||
tableName: "mastertaskitem",
|
||||
options: &RequestOptions{
|
||||
Preload: []PreloadOption{
|
||||
{
|
||||
Relation: "MTL.MAL.MAL_RID_PARENTMASTERTASKITEM",
|
||||
TableName: "mastertaskitem",
|
||||
},
|
||||
},
|
||||
},
|
||||
expected: "mastertaskitem.rid_parentmastertaskitem is null",
|
||||
},
|
||||
{
|
||||
name: "Unqualified column should remain unqualified",
|
||||
where: "rid_parentmastertaskitem is null",
|
||||
tableName: "mastertaskitem",
|
||||
options: nil,
|
||||
expected: "rid_parentmastertaskitem is null",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := SanitizeWhereClause(tt.where, tt.tableName, tt.options)
|
||||
if result != tt.expected {
|
||||
t.Errorf("SanitizeWhereClause(%q, %q) = %q, want %q",
|
||||
tt.where, tt.tableName, result, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAddTablePrefixToColumns_WithTableName tests that table prefixes
|
||||
// are correctly added to unqualified columns
|
||||
func TestAddTablePrefixToColumns_WithTableName(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
where string
|
||||
tableName string
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "Add prefix to unqualified column",
|
||||
where: "rid_parentmastertaskitem is null",
|
||||
tableName: "mastertaskitem",
|
||||
expected: "mastertaskitem.rid_parentmastertaskitem is null",
|
||||
},
|
||||
{
|
||||
name: "Don't change already qualified column",
|
||||
where: "mastertaskitem.rid_parentmastertaskitem is null",
|
||||
tableName: "mastertaskitem",
|
||||
expected: "mastertaskitem.rid_parentmastertaskitem is null",
|
||||
},
|
||||
{
|
||||
name: "Don't change qualified column with different table",
|
||||
where: "other_table.rid_something is null",
|
||||
tableName: "mastertaskitem",
|
||||
expected: "other_table.rid_something is null",
|
||||
},
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+373
-67
@@ -134,6 +134,30 @@ func TestSanitizeWhereClause(t *testing.T) {
|
||||
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))",
|
||||
},
|
||||
{
|
||||
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 {
|
||||
@@ -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",
|
||||
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 {
|
||||
@@ -659,75 +715,325 @@ func TestSanitizeWhereClauseWithModel(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddTablePrefixToColumns_ComplexConditions(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
where string
|
||||
tableName string
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "Parentheses with true AND condition - should not prefix true",
|
||||
where: "(true AND status = 'active')",
|
||||
tableName: "mastertask",
|
||||
expected: "(true AND mastertask.status = 'active')",
|
||||
},
|
||||
{
|
||||
name: "Parentheses with multiple conditions including true",
|
||||
where: "(true AND status = 'active' AND id > 5)",
|
||||
tableName: "mastertask",
|
||||
expected: "(true AND mastertask.status = 'active' AND mastertask.id > 5)",
|
||||
},
|
||||
{
|
||||
name: "Nested parentheses with true",
|
||||
where: "((true AND status = 'active'))",
|
||||
tableName: "mastertask",
|
||||
expected: "((true AND mastertask.status = 'active'))",
|
||||
},
|
||||
{
|
||||
name: "Mixed: false AND valid conditions",
|
||||
where: "(false AND name = 'test')",
|
||||
tableName: "mastertask",
|
||||
expected: "(false AND mastertask.name = 'test')",
|
||||
},
|
||||
{
|
||||
name: "Mixed: null AND valid conditions",
|
||||
where: "(null AND status = 'active')",
|
||||
tableName: "mastertask",
|
||||
expected: "(null AND mastertask.status = 'active')",
|
||||
},
|
||||
{
|
||||
name: "Multiple true conditions in parentheses",
|
||||
where: "(true AND true AND status = 'active')",
|
||||
tableName: "mastertask",
|
||||
expected: "(true AND true AND mastertask.status = 'active')",
|
||||
},
|
||||
{
|
||||
name: "Simple true without parens - should not prefix",
|
||||
where: "true",
|
||||
tableName: "mastertask",
|
||||
expected: "true",
|
||||
},
|
||||
{
|
||||
name: "Simple condition without parens - should prefix",
|
||||
where: "status = 'active'",
|
||||
tableName: "mastertask",
|
||||
expected: "mastertask.status = 'active'",
|
||||
},
|
||||
{
|
||||
name: "Unregistered table with true - should not prefix true",
|
||||
where: "(true AND status = 'active')",
|
||||
tableName: "unregistered_table",
|
||||
expected: "(true AND unregistered_table.status = 'active')",
|
||||
},
|
||||
func TestEnsureOuterParentheses(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "no parentheses",
|
||||
input: "status = 'active'",
|
||||
expected: "(status = 'active')",
|
||||
},
|
||||
{
|
||||
name: "already has outer parentheses",
|
||||
input: "(status = 'active')",
|
||||
expected: "(status = 'active')",
|
||||
},
|
||||
{
|
||||
name: "OR condition without parentheses",
|
||||
input: "status = 'active' OR status = 'pending'",
|
||||
expected: "(status = 'active' OR status = 'pending')",
|
||||
},
|
||||
{
|
||||
name: "OR condition with parentheses",
|
||||
input: "(status = 'active' OR status = 'pending')",
|
||||
expected: "(status = 'active' OR status = 'pending')",
|
||||
},
|
||||
{
|
||||
name: "complex condition with nested parentheses",
|
||||
input: "(status = 'active' OR status = 'pending') AND (age > 18)",
|
||||
expected: "((status = 'active' OR status = 'pending') AND (age > 18))",
|
||||
},
|
||||
{
|
||||
name: "empty string",
|
||||
input: "",
|
||||
expected: "",
|
||||
},
|
||||
{
|
||||
name: "whitespace only",
|
||||
input: " ",
|
||||
expected: "",
|
||||
},
|
||||
{
|
||||
name: "mismatched parentheses - adds outer ones",
|
||||
input: "(status = 'active' OR status = 'pending'",
|
||||
expected: "((status = 'active' OR status = 'pending')",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := EnsureOuterParentheses(tt.input)
|
||||
if result != tt.expected {
|
||||
t.Errorf("EnsureOuterParentheses(%q) = %q; want %q", tt.input, result, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
func TestContainsTopLevelOR(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
expected bool
|
||||
}{
|
||||
{
|
||||
name: "no OR operator",
|
||||
input: "status = 'active' AND age > 18",
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "top-level OR",
|
||||
input: "status = 'active' OR status = 'pending'",
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "OR inside parentheses",
|
||||
input: "age > 18 AND (status = 'active' OR status = 'pending')",
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "OR in subquery",
|
||||
input: "id IN (SELECT id FROM users WHERE status = 'active' OR status = 'pending')",
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "OR inside quotes",
|
||||
input: "comment = 'this OR that'",
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "mixed - top-level OR and nested OR",
|
||||
input: "name = 'test' OR (status = 'active' OR status = 'pending')",
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "empty string",
|
||||
input: "",
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "lowercase or",
|
||||
input: "status = 'active' or status = 'pending'",
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "uppercase OR",
|
||||
input: "status = 'active' OR status = 'pending'",
|
||||
expected: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := containsTopLevelOR(tt.input)
|
||||
if result != tt.expected {
|
||||
t.Errorf("containsTopLevelOR(%q) = %v; want %v", tt.input, result, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
func TestSanitizeWhereClause_PreservesParenthesesWithOR(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
where string
|
||||
tableName string
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "OR condition with outer parentheses - preserved",
|
||||
where: "(status = 'active' OR status = 'pending')",
|
||||
tableName: "users",
|
||||
expected: "(users.status = 'active' OR users.status = 'pending')",
|
||||
},
|
||||
{
|
||||
name: "AND condition with outer parentheses - stripped (no OR)",
|
||||
where: "(status = 'active' AND age > 18)",
|
||||
tableName: "users",
|
||||
expected: "users.status = 'active' AND users.age > 18",
|
||||
},
|
||||
{
|
||||
name: "complex OR with nested conditions",
|
||||
where: "((status = 'active' OR status = 'pending') AND age > 18)",
|
||||
tableName: "users",
|
||||
// Outer parens are stripped, but inner parens with OR are preserved
|
||||
expected: "(users.status = 'active' OR users.status = 'pending') AND users.age > 18",
|
||||
},
|
||||
{
|
||||
name: "OR without outer parentheses - no parentheses added by SanitizeWhereClause",
|
||||
where: "status = 'active' OR status = 'pending'",
|
||||
tableName: "users",
|
||||
expected: "users.status = 'active' OR users.status = 'pending'",
|
||||
},
|
||||
{
|
||||
name: "simple OR with parentheses - preserved",
|
||||
where: "(users.status = 'active' OR users.status = 'pending')",
|
||||
tableName: "users",
|
||||
// Already has correct prefixes, parentheses preserved
|
||||
expected: "(users.status = 'active' OR users.status = 'pending')",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
prefixedWhere := AddTablePrefixToColumns(tt.where, tt.tableName)
|
||||
result := SanitizeWhereClause(prefixedWhere, tt.tableName)
|
||||
if result != tt.expected {
|
||||
t.Errorf("SanitizeWhereClause(%q, %q) = %q; want %q", tt.where, tt.tableName, result, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddTablePrefixToColumns_ComplexConditions(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
where string
|
||||
tableName string
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "Parentheses with true AND condition - should not prefix true",
|
||||
where: "(true AND status = 'active')",
|
||||
tableName: "mastertask",
|
||||
expected: "(true AND mastertask.status = 'active')",
|
||||
},
|
||||
{
|
||||
name: "Parentheses with multiple conditions including true",
|
||||
where: "(true AND status = 'active' AND id > 5)",
|
||||
tableName: "mastertask",
|
||||
expected: "(true AND mastertask.status = 'active' AND mastertask.id > 5)",
|
||||
},
|
||||
{
|
||||
name: "Nested parentheses with true",
|
||||
where: "((true AND status = 'active'))",
|
||||
tableName: "mastertask",
|
||||
expected: "((true AND mastertask.status = 'active'))",
|
||||
},
|
||||
{
|
||||
name: "Mixed: false AND valid conditions",
|
||||
where: "(false AND name = 'test')",
|
||||
tableName: "mastertask",
|
||||
expected: "(false AND mastertask.name = 'test')",
|
||||
},
|
||||
{
|
||||
name: "Mixed: null AND valid conditions",
|
||||
where: "(null AND status = 'active')",
|
||||
tableName: "mastertask",
|
||||
expected: "(null AND mastertask.status = 'active')",
|
||||
},
|
||||
{
|
||||
name: "Multiple true conditions in parentheses",
|
||||
where: "(true AND true AND status = 'active')",
|
||||
tableName: "mastertask",
|
||||
expected: "(true AND true AND mastertask.status = 'active')",
|
||||
},
|
||||
{
|
||||
name: "Simple true without parens - should not prefix",
|
||||
where: "true",
|
||||
tableName: "mastertask",
|
||||
expected: "true",
|
||||
},
|
||||
{
|
||||
name: "Simple condition without parens - should prefix",
|
||||
where: "status = 'active'",
|
||||
tableName: "mastertask",
|
||||
expected: "mastertask.status = 'active'",
|
||||
},
|
||||
{
|
||||
name: "Unregistered table with true - should not prefix true",
|
||||
where: "(true AND status = 'active')",
|
||||
tableName: "unregistered_table",
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildArrayOverlapCondition(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
column string
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+82
-3
@@ -1,5 +1,23 @@
|
||||
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 {
|
||||
Operation string `json:"operation"`
|
||||
Data interface{} `json:"data"`
|
||||
@@ -23,6 +41,14 @@ type RequestOptions struct {
|
||||
CursorForward string `json:"cursor_forward"`
|
||||
CursorBackward string `json:"cursor_backward"`
|
||||
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)
|
||||
// Not serialized to JSON as it's internal validation state
|
||||
JoinAliases []string `json:"-"`
|
||||
}
|
||||
|
||||
type Parameter struct {
|
||||
@@ -33,6 +59,7 @@ type Parameter struct {
|
||||
|
||||
type PreloadOption struct {
|
||||
Relation string `json:"relation"`
|
||||
TableName string `json:"table_name"` // Actual database table name (e.g., "mastertaskitem")
|
||||
Columns []string `json:"columns"`
|
||||
OmitColumns []string `json:"omit_columns"`
|
||||
Sort []SortOption `json:"sort"`
|
||||
@@ -45,9 +72,14 @@ type PreloadOption struct {
|
||||
Recursive bool `json:"recursive"` // if true, preload recursively up to 5 levels
|
||||
|
||||
// Relationship keys from XFiles - used to build proper foreign key filters
|
||||
PrimaryKey string `json:"primary_key"` // Primary key of the related table
|
||||
RelatedKey string `json:"related_key"` // For child tables: column in child that references parent
|
||||
ForeignKey string `json:"foreign_key"` // For parent tables: column in current table that references parent
|
||||
PrimaryKey string `json:"primary_key"` // Primary key of the related table
|
||||
RelatedKey string `json:"related_key"` // For child tables: column in child that references parent
|
||||
ForeignKey string `json:"foreign_key"` // For parent tables: column in current table that references parent
|
||||
RecursiveChildKey string `json:"recursive_child_key"` // For recursive tables: FK column used for recursion (e.g., "rid_parentmastertaskitem")
|
||||
|
||||
// Custom SQL JOINs from XFiles - used when preload needs additional joins
|
||||
SqlJoins []string `json:"sql_joins"` // Custom SQL JOIN clauses
|
||||
JoinAliases []string `json:"join_aliases"` // Extracted table aliases from SqlJoins for validation
|
||||
}
|
||||
|
||||
type FilterOption struct {
|
||||
@@ -62,6 +94,41 @@ type SortOption struct {
|
||||
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 {
|
||||
Name string `json:"name"`
|
||||
SQL string `json:"sql"`
|
||||
@@ -94,6 +161,7 @@ type APIError struct {
|
||||
Message string `json:"message"`
|
||||
Details interface{} `json:"details,omitempty"`
|
||||
Detail string `json:"detail,omitempty"`
|
||||
SQL string `json:"sql,omitempty"`
|
||||
}
|
||||
|
||||
type Column struct {
|
||||
@@ -111,3 +179,14 @@ type TableMetadata struct {
|
||||
Columns []Column `json:"columns"`
|
||||
Relations []string `json:"relations"`
|
||||
}
|
||||
|
||||
// RelationshipInfo contains information about a model relationship
|
||||
type RelationshipInfo struct {
|
||||
FieldName string `json:"field_name"`
|
||||
JSONName string `json:"json_name"`
|
||||
RelationType string `json:"relation_type"` // "belongsTo", "hasMany", "hasOne", "many2many"
|
||||
ForeignKey string `json:"foreign_key"`
|
||||
References string `json:"references"`
|
||||
JoinTable string `json:"join_table"`
|
||||
RelatedModel interface{} `json:"related_model"`
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+98
-16
@@ -3,6 +3,7 @@ package common
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
@@ -30,7 +31,7 @@ func (v *ColumnValidator) buildValidColumns() {
|
||||
modelType := reflect.TypeOf(v.model)
|
||||
|
||||
// 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()
|
||||
}
|
||||
|
||||
@@ -43,7 +44,7 @@ func (v *ColumnValidator) buildValidColumns() {
|
||||
for i := 0; i < modelType.NumField(); i++ {
|
||||
field := modelType.Field(i)
|
||||
|
||||
if !field.IsExported() {
|
||||
if !field.IsExported() || field.Anonymous {
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -108,6 +109,19 @@ func (v *ColumnValidator) ValidateColumn(column string) error {
|
||||
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 ->)
|
||||
sourceColumn := reflection.ExtractSourceColumn(column)
|
||||
|
||||
@@ -125,6 +139,16 @@ func (v *ColumnValidator) IsValidColumn(column string) bool {
|
||||
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
|
||||
// Logs warnings for any invalid columns
|
||||
func (v *ColumnValidator) FilterValidColumns(columns []string) []string {
|
||||
@@ -224,7 +248,19 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
||||
// Filter Filter columns
|
||||
validFilters := make([]FilterOption, 0, len(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)
|
||||
} else {
|
||||
logger.Warn("Invalid column in filter '%s' removed", filter.Column)
|
||||
@@ -237,34 +273,77 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
||||
for _, sort := range options.Sort {
|
||||
if v.IsValidColumn(sort.Column) {
|
||||
validSorts = append(validSorts, sort)
|
||||
} else if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||
// Allow sort by expression/subquery, but validate for security
|
||||
if IsSafeSortExpression(sort.Column) {
|
||||
validSorts = append(validSorts, sort)
|
||||
} else {
|
||||
logger.Warn("Unsafe sort expression '%s' removed", sort.Column)
|
||||
}
|
||||
} else {
|
||||
logger.Warn("Invalid column in sort '%s' removed", sort.Column)
|
||||
foundJoin := false
|
||||
for _, j := range options.JoinAliases {
|
||||
if strings.Contains(sort.Column, j) {
|
||||
foundJoin = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if foundJoin {
|
||||
validSorts = append(validSorts, sort)
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||
// Allow sort by expression/subquery, but validate for security
|
||||
if IsSafeSortExpression(sort.Column) {
|
||||
validSorts = append(validSorts, sort)
|
||||
} else {
|
||||
logger.Warn("Unsafe sort expression '%s' removed", sort.Column)
|
||||
}
|
||||
|
||||
} else {
|
||||
logger.Warn("Invalid column in sort '%s' removed", sort.Column)
|
||||
}
|
||||
}
|
||||
}
|
||||
filtered.Sort = validSorts
|
||||
|
||||
// Filter Preload columns
|
||||
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 {
|
||||
preload := options.Preload[idx]
|
||||
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
|
||||
filteredPreload.SqlJoins = preload.SqlJoins
|
||||
filteredPreload.JoinAliases = preload.JoinAliases
|
||||
|
||||
// Filter preload filters
|
||||
validPreloadFilters := make([]FilterOption, 0, len(preload.Filters))
|
||||
for _, filter := range preload.Filters {
|
||||
if v.IsValidColumn(filter.Column) {
|
||||
if preloadValidator.IsValidColumn(filter.Column) {
|
||||
validPreloadFilters = append(validPreloadFilters, filter)
|
||||
} else {
|
||||
logger.Warn("Invalid column in preload '%s' filter '%s' removed", preload.Relation, filter.Column)
|
||||
// Check if the filter column references a joined table alias
|
||||
foundJoin := false
|
||||
for _, alias := range preload.JoinAliases {
|
||||
if strings.Contains(filter.Column, alias) {
|
||||
foundJoin = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if foundJoin {
|
||||
validPreloadFilters = append(validPreloadFilters, filter)
|
||||
} else {
|
||||
logger.Warn("Invalid column in preload '%s' filter '%s' removed", preload.Relation, filter.Column)
|
||||
}
|
||||
}
|
||||
}
|
||||
filteredPreload.Filters = validPreloadFilters
|
||||
@@ -272,7 +351,7 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
||||
// Filter preload sort columns
|
||||
validPreloadSorts := make([]SortOption, 0, len(preload.Sort))
|
||||
for _, sort := range preload.Sort {
|
||||
if v.IsValidColumn(sort.Column) {
|
||||
if preloadValidator.IsValidColumn(sort.Column) {
|
||||
validPreloadSorts = append(validPreloadSorts, sort)
|
||||
} else if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||
// Allow sort by expression/subquery, but validate for security
|
||||
@@ -291,6 +370,9 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
||||
}
|
||||
filtered.Preload = validPreloads
|
||||
|
||||
// Clear JoinAliases - this is an internal validation field and should not be persisted
|
||||
filtered.JoinAliases = nil
|
||||
|
||||
return filtered
|
||||
}
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -362,6 +362,29 @@ func TestFilterRequestOptions(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFilterRequestOptions_ClearsJoinAliases(t *testing.T) {
|
||||
model := TestModel{}
|
||||
validator := NewColumnValidator(model)
|
||||
|
||||
options := RequestOptions{
|
||||
Columns: []string{"id", "name"},
|
||||
// Set JoinAliases - this should be cleared by FilterRequestOptions
|
||||
JoinAliases: []string{"d", "u", "r"},
|
||||
}
|
||||
|
||||
filtered := validator.FilterRequestOptions(options)
|
||||
|
||||
// Verify that JoinAliases was cleared (internal field should not persist)
|
||||
if filtered.JoinAliases != nil {
|
||||
t.Errorf("Expected JoinAliases to be nil after filtering, got %v", filtered.JoinAliases)
|
||||
}
|
||||
|
||||
// Verify that other fields are still properly filtered
|
||||
if len(filtered.Columns) != 2 {
|
||||
t.Errorf("Expected 2 columns, got %d", len(filtered.Columns))
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsSafeSortExpression(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -441,3 +464,84 @@ func TestFilterRequestOptions_WithSortExpressions(t *testing.T) {
|
||||
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])
|
||||
}
|
||||
}
|
||||
|
||||
+104
-18
@@ -1,23 +1,34 @@
|
||||
package config
|
||||
|
||||
import "time"
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Config represents the complete application configuration
|
||||
type Config struct {
|
||||
Server ServerConfig `mapstructure:"server"`
|
||||
Tracing TracingConfig `mapstructure:"tracing"`
|
||||
Cache CacheConfig `mapstructure:"cache"`
|
||||
Logger LoggerConfig `mapstructure:"logger"`
|
||||
ErrorTracking ErrorTrackingConfig `mapstructure:"error_tracking"`
|
||||
Middleware MiddlewareConfig `mapstructure:"middleware"`
|
||||
CORS CORSConfig `mapstructure:"cors"`
|
||||
Database DatabaseConfig `mapstructure:"database"`
|
||||
EventBroker EventBrokerConfig `mapstructure:"event_broker"`
|
||||
Servers ServersConfig `mapstructure:"servers"`
|
||||
Tracing TracingConfig `mapstructure:"tracing"`
|
||||
Cache CacheConfig `mapstructure:"cache"`
|
||||
Logger LoggerConfig `mapstructure:"logger"`
|
||||
ErrorTracking ErrorTrackingConfig `mapstructure:"error_tracking"`
|
||||
Middleware MiddlewareConfig `mapstructure:"middleware"`
|
||||
CORS CORSConfig `mapstructure:"cors"`
|
||||
EventBroker EventBrokerConfig `mapstructure:"event_broker"`
|
||||
DBManager DBManagerConfig `mapstructure:"dbmanager"`
|
||||
Paths PathsConfig `mapstructure:"paths"`
|
||||
Extensions map[string]interface{} `mapstructure:"extensions"`
|
||||
}
|
||||
|
||||
// ServerConfig holds server-related configuration
|
||||
type ServerConfig struct {
|
||||
Addr string `mapstructure:"addr"`
|
||||
// ServersConfig contains configuration for the server manager
|
||||
type ServersConfig struct {
|
||||
// DefaultServer is the name of the default server to use
|
||||
DefaultServer string `mapstructure:"default_server"`
|
||||
|
||||
// Instances is a map of server name to server configuration
|
||||
Instances map[string]ServerInstanceConfig `mapstructure:"instances"`
|
||||
|
||||
// Global timeout defaults (can be overridden per instance)
|
||||
ShutdownTimeout time.Duration `mapstructure:"shutdown_timeout"`
|
||||
DrainTimeout time.Duration `mapstructure:"drain_timeout"`
|
||||
ReadTimeout time.Duration `mapstructure:"read_timeout"`
|
||||
@@ -25,12 +36,67 @@ type ServerConfig struct {
|
||||
IdleTimeout time.Duration `mapstructure:"idle_timeout"`
|
||||
}
|
||||
|
||||
// ServerInstanceConfig defines configuration for a single server instance
|
||||
type ServerInstanceConfig struct {
|
||||
// Name is the unique name of this server instance
|
||||
Name string `mapstructure:"name"`
|
||||
|
||||
// Host is the host to bind to (e.g., "localhost", "0.0.0.0", "")
|
||||
Host string `mapstructure:"host"`
|
||||
|
||||
// Port is the port number to listen on
|
||||
Port int `mapstructure:"port"`
|
||||
|
||||
// Description is a human-readable description of this server
|
||||
Description string `mapstructure:"description"`
|
||||
|
||||
// GZIP enables GZIP compression middleware
|
||||
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)
|
||||
// Option 1: Provide certificate and key files directly
|
||||
SSLCert string `mapstructure:"ssl_cert"`
|
||||
SSLKey string `mapstructure:"ssl_key"`
|
||||
|
||||
// Option 2: Use self-signed certificate (for development/testing)
|
||||
SelfSignedSSL bool `mapstructure:"self_signed_ssl"`
|
||||
|
||||
// Option 3: Use Let's Encrypt / AutoTLS
|
||||
AutoTLS bool `mapstructure:"auto_tls"`
|
||||
AutoTLSDomains []string `mapstructure:"auto_tls_domains"`
|
||||
AutoTLSCacheDir string `mapstructure:"auto_tls_cache_dir"`
|
||||
AutoTLSEmail string `mapstructure:"auto_tls_email"`
|
||||
|
||||
// Timeout configurations (overrides global defaults)
|
||||
ShutdownTimeout *time.Duration `mapstructure:"shutdown_timeout"`
|
||||
DrainTimeout *time.Duration `mapstructure:"drain_timeout"`
|
||||
ReadTimeout *time.Duration `mapstructure:"read_timeout"`
|
||||
WriteTimeout *time.Duration `mapstructure:"write_timeout"`
|
||||
IdleTimeout *time.Duration `mapstructure:"idle_timeout"`
|
||||
|
||||
// Tags for organization and filtering
|
||||
Tags map[string]string `mapstructure:"tags"`
|
||||
|
||||
// ExternalURLs are additional URLs that this server instance is accessible from (for CORS) for proxy setups
|
||||
ExternalURLs []string `mapstructure:"external_urls"`
|
||||
}
|
||||
|
||||
// TracingConfig holds OpenTelemetry tracing configuration
|
||||
type TracingConfig struct {
|
||||
Enabled bool `mapstructure:"enabled"`
|
||||
ServiceName string `mapstructure:"service_name"`
|
||||
ServiceVersion string `mapstructure:"service_version"`
|
||||
Endpoint string `mapstructure:"endpoint"`
|
||||
// Insecure exports traces over plaintext gRPC (default false: TLS).
|
||||
Insecure bool `mapstructure:"insecure"`
|
||||
// SampleRate is the fraction of root traces sampled; 0 selects the default (0.1).
|
||||
SampleRate float64 `mapstructure:"sample_rate"`
|
||||
// Headers are sent with every OTLP export request (e.g. auth tokens).
|
||||
Headers map[string]string `mapstructure:"headers"`
|
||||
}
|
||||
|
||||
// CacheConfig holds cache provider configuration
|
||||
@@ -76,11 +142,6 @@ type CORSConfig struct {
|
||||
MaxAge int `mapstructure:"max_age"`
|
||||
}
|
||||
|
||||
// DatabaseConfig holds database configuration (primarily for testing)
|
||||
type DatabaseConfig struct {
|
||||
URL string `mapstructure:"url"`
|
||||
}
|
||||
|
||||
// ErrorTrackingConfig holds error tracking configuration
|
||||
type ErrorTrackingConfig struct {
|
||||
Enabled bool `mapstructure:"enabled"`
|
||||
@@ -141,3 +202,28 @@ type EventBrokerRetryPolicyConfig struct {
|
||||
MaxDelay time.Duration `mapstructure:"max_delay"`
|
||||
BackoffFactor float64 `mapstructure:"backoff_factor"`
|
||||
}
|
||||
|
||||
// PathsConfig contains configuration for named file system paths
|
||||
// This is a map of path name to file system path
|
||||
// Example: "data_dir": "/var/lib/myapp/data"
|
||||
type PathsConfig map[string]string
|
||||
|
||||
// Validate checks every configuration section for invalid or unsafe values.
|
||||
func (c *Config) Validate() error {
|
||||
if err := c.Servers.Validate(); err != nil {
|
||||
return fmt.Errorf("servers: %w", err)
|
||||
}
|
||||
if c.Middleware.RateLimitRPS < 0 || c.Middleware.RateLimitBurst < 0 {
|
||||
return fmt.Errorf("middleware: rate_limit_rps and rate_limit_burst must not be negative")
|
||||
}
|
||||
if c.Middleware.MaxRequestSize <= 0 {
|
||||
return fmt.Errorf("middleware: max_request_size must be greater than 0")
|
||||
}
|
||||
if c.EventBroker.Enabled && c.EventBroker.WorkerCount <= 0 {
|
||||
return fmt.Errorf("event_broker: worker_count must be greater than 0")
|
||||
}
|
||||
if c.DBManager.MaxOpenConns < 0 || c.DBManager.MaxIdleConns < 0 || c.DBManager.RetryAttempts < 0 {
|
||||
return fmt.Errorf("dbmanager: max_open_conns, max_idle_conns and retry_attempts must not be negative")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,264 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// DBManagerConfig contains configuration for the database connection manager
|
||||
type DBManagerConfig struct {
|
||||
// DefaultConnection is the name of the default connection to use
|
||||
DefaultConnection string `mapstructure:"default_connection"`
|
||||
|
||||
// Connections is a map of connection name to connection configuration
|
||||
Connections map[string]DBConnectionConfig `mapstructure:"connections"`
|
||||
|
||||
// Global connection pool defaults
|
||||
MaxOpenConns int `mapstructure:"max_open_conns"`
|
||||
MaxIdleConns int `mapstructure:"max_idle_conns"`
|
||||
ConnMaxLifetime time.Duration `mapstructure:"conn_max_lifetime"`
|
||||
ConnMaxIdleTime time.Duration `mapstructure:"conn_max_idle_time"`
|
||||
|
||||
// Retry policy
|
||||
RetryAttempts int `mapstructure:"retry_attempts"`
|
||||
RetryDelay time.Duration `mapstructure:"retry_delay"`
|
||||
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
|
||||
|
||||
// Health checks
|
||||
HealthCheckInterval time.Duration `mapstructure:"health_check_interval"`
|
||||
EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"`
|
||||
}
|
||||
|
||||
// DBConnectionConfig defines configuration for a single database connection
|
||||
type DBConnectionConfig struct {
|
||||
// Name is the unique name of this connection
|
||||
Name string `mapstructure:"name"`
|
||||
|
||||
// Type is the database type (postgres, sqlite, mssql, mongodb)
|
||||
Type string `mapstructure:"type"`
|
||||
|
||||
// DSN is the complete Data Source Name / connection string
|
||||
// If provided, this takes precedence over individual connection parameters
|
||||
DSN string `mapstructure:"dsn"`
|
||||
|
||||
// Connection parameters (used if DSN is not provided)
|
||||
Host string `mapstructure:"host"`
|
||||
Port int `mapstructure:"port"`
|
||||
User string `mapstructure:"user"`
|
||||
Password string `mapstructure:"password"`
|
||||
Database string `mapstructure:"database"`
|
||||
|
||||
// PostgreSQL/MSSQL specific
|
||||
SSLMode string `mapstructure:"sslmode"` // disable, require, verify-ca, verify-full
|
||||
Schema string `mapstructure:"schema"` // Default schema
|
||||
|
||||
// SQLite specific
|
||||
FilePath string `mapstructure:"filepath"`
|
||||
|
||||
// MongoDB specific
|
||||
AuthSource string `mapstructure:"auth_source"`
|
||||
ReplicaSet string `mapstructure:"replica_set"`
|
||||
ReadPreference string `mapstructure:"read_preference"` // primary, secondary, etc.
|
||||
|
||||
// Connection pool settings (overrides global defaults)
|
||||
MaxOpenConns *int `mapstructure:"max_open_conns"`
|
||||
MaxIdleConns *int `mapstructure:"max_idle_conns"`
|
||||
ConnMaxLifetime *time.Duration `mapstructure:"conn_max_lifetime"`
|
||||
ConnMaxIdleTime *time.Duration `mapstructure:"conn_max_idle_time"`
|
||||
|
||||
// Timeouts
|
||||
ConnectTimeout time.Duration `mapstructure:"connect_timeout"`
|
||||
QueryTimeout time.Duration `mapstructure:"query_timeout"`
|
||||
|
||||
// Features
|
||||
EnableTracing bool `mapstructure:"enable_tracing"`
|
||||
EnableMetrics bool `mapstructure:"enable_metrics"`
|
||||
EnableLogging bool `mapstructure:"enable_logging"`
|
||||
|
||||
// DefaultORM specifies which ORM to use for the Database() method
|
||||
// Options: "bun", "gorm", "native"
|
||||
DefaultORM string `mapstructure:"default_orm"`
|
||||
|
||||
// Tags for organization and filtering
|
||||
Tags map[string]string `mapstructure:"tags"`
|
||||
}
|
||||
|
||||
// ToManagerConfig converts config.DBManagerConfig to dbmanager.ManagerConfig
|
||||
// This is used to avoid circular dependencies
|
||||
func (c *DBManagerConfig) ToManagerConfig() interface{} {
|
||||
// This will be implemented in the dbmanager package
|
||||
// to convert from config types to dbmanager types
|
||||
return c
|
||||
}
|
||||
|
||||
// PopulateFromDSN parses a DSN and populates the connection fields
|
||||
func (cc *DBConnectionConfig) PopulateFromDSN() error {
|
||||
if cc.DSN == "" {
|
||||
return nil // Nothing to populate
|
||||
}
|
||||
|
||||
switch cc.Type {
|
||||
case "postgres":
|
||||
return cc.populatePostgresDSN()
|
||||
case "mongodb":
|
||||
return cc.populateMongoDSN()
|
||||
case "mssql":
|
||||
return cc.populateMSSQLDSN()
|
||||
case "sqlite":
|
||||
return cc.populateSQLiteDSN()
|
||||
default:
|
||||
return fmt.Errorf("cannot parse DSN for unsupported database type: %s", cc.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// populatePostgresDSN parses PostgreSQL DSN format
|
||||
// Example: host=localhost port=5432 user=postgres password=secret dbname=mydb sslmode=disable
|
||||
func (cc *DBConnectionConfig) populatePostgresDSN() error {
|
||||
parts := strings.Fields(cc.DSN)
|
||||
for _, part := range parts {
|
||||
kv := strings.SplitN(part, "=", 2)
|
||||
if len(kv) != 2 {
|
||||
continue
|
||||
}
|
||||
key, value := kv[0], kv[1]
|
||||
|
||||
switch key {
|
||||
case "host":
|
||||
cc.Host = value
|
||||
case "port":
|
||||
port, err := strconv.Atoi(value)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid port in DSN: %w", err)
|
||||
}
|
||||
cc.Port = port
|
||||
case "user":
|
||||
cc.User = value
|
||||
case "password":
|
||||
cc.Password = value
|
||||
case "dbname":
|
||||
cc.Database = value
|
||||
case "sslmode":
|
||||
cc.SSLMode = value
|
||||
case "search_path":
|
||||
cc.Schema = value
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// populateMongoDSN parses MongoDB DSN format
|
||||
// Example: mongodb://user:password@host:port/database?authSource=admin&replicaSet=rs0
|
||||
func (cc *DBConnectionConfig) populateMongoDSN() error {
|
||||
u, err := url.Parse(cc.DSN)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid MongoDB DSN: %w", err)
|
||||
}
|
||||
|
||||
// Extract user and password
|
||||
if u.User != nil {
|
||||
cc.User = u.User.Username()
|
||||
if password, ok := u.User.Password(); ok {
|
||||
cc.Password = password
|
||||
}
|
||||
}
|
||||
|
||||
// Extract host and port
|
||||
if u.Host != "" {
|
||||
host := u.Host
|
||||
if strings.Contains(host, ":") {
|
||||
hostPort := strings.SplitN(host, ":", 2)
|
||||
cc.Host = hostPort[0]
|
||||
if port, err := strconv.Atoi(hostPort[1]); err == nil {
|
||||
cc.Port = port
|
||||
}
|
||||
} else {
|
||||
cc.Host = host
|
||||
}
|
||||
}
|
||||
|
||||
// Extract database
|
||||
if u.Path != "" {
|
||||
cc.Database = strings.TrimPrefix(u.Path, "/")
|
||||
}
|
||||
|
||||
// Extract query parameters
|
||||
params := u.Query()
|
||||
if authSource := params.Get("authSource"); authSource != "" {
|
||||
cc.AuthSource = authSource
|
||||
}
|
||||
if replicaSet := params.Get("replicaSet"); replicaSet != "" {
|
||||
cc.ReplicaSet = replicaSet
|
||||
}
|
||||
if readPref := params.Get("readPreference"); readPref != "" {
|
||||
cc.ReadPreference = readPref
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// populateMSSQLDSN parses MSSQL DSN format
|
||||
// Example: sqlserver://username:password@host:port?database=dbname&schema=dbo
|
||||
func (cc *DBConnectionConfig) populateMSSQLDSN() error {
|
||||
u, err := url.Parse(cc.DSN)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid MSSQL DSN: %w", err)
|
||||
}
|
||||
|
||||
// Extract user and password
|
||||
if u.User != nil {
|
||||
cc.User = u.User.Username()
|
||||
if password, ok := u.User.Password(); ok {
|
||||
cc.Password = password
|
||||
}
|
||||
}
|
||||
|
||||
// Extract host and port
|
||||
if u.Host != "" {
|
||||
host := u.Host
|
||||
if strings.Contains(host, ":") {
|
||||
hostPort := strings.SplitN(host, ":", 2)
|
||||
cc.Host = hostPort[0]
|
||||
if port, err := strconv.Atoi(hostPort[1]); err == nil {
|
||||
cc.Port = port
|
||||
}
|
||||
} else {
|
||||
cc.Host = host
|
||||
}
|
||||
}
|
||||
|
||||
// Extract query parameters
|
||||
params := u.Query()
|
||||
if database := params.Get("database"); database != "" {
|
||||
cc.Database = database
|
||||
}
|
||||
if schema := params.Get("schema"); schema != "" {
|
||||
cc.Schema = schema
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// populateSQLiteDSN parses SQLite DSN format
|
||||
// Example: /path/to/database.db or :memory:
|
||||
func (cc *DBConnectionConfig) populateSQLiteDSN() error {
|
||||
cc.FilePath = cc.DSN
|
||||
return nil
|
||||
}
|
||||
|
||||
// Validate validates the DBManager configuration
|
||||
func (c *DBManagerConfig) Validate() error {
|
||||
if len(c.Connections) == 0 {
|
||||
return fmt.Errorf("at least one connection must be configured")
|
||||
}
|
||||
|
||||
if c.DefaultConnection != "" {
|
||||
if _, ok := c.Connections[c.DefaultConnection]; !ok {
|
||||
return fmt.Errorf("default connection '%s' not found in connections", c.DefaultConnection)
|
||||
}
|
||||
}
|
||||
|
||||
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")
|
||||
}
|
||||
}
|
||||
+166
-14
@@ -2,27 +2,60 @@ package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/spf13/viper"
|
||||
)
|
||||
|
||||
// Manager handles configuration loading from multiple sources
|
||||
// Manager handles configuration loading from multiple sources.
|
||||
// viper.Viper is not safe for concurrent use, so every access to it is guarded by mu.
|
||||
type Manager struct {
|
||||
v *viper.Viper
|
||||
mu sync.RWMutex
|
||||
v *viper.Viper
|
||||
}
|
||||
|
||||
// NewManager creates a new configuration manager with defaults
|
||||
var (
|
||||
configInstance *Manager
|
||||
configMu sync.Mutex
|
||||
)
|
||||
|
||||
// GetConfigManager returns a singleton configuration manager instance
|
||||
func GetConfigManager() *Manager {
|
||||
configMu.Lock()
|
||||
defer configMu.Unlock()
|
||||
if configInstance == nil {
|
||||
configInstance = NewManager()
|
||||
}
|
||||
return configInstance
|
||||
}
|
||||
|
||||
// SetConfigManager publishes m as the global manager returned by GetConfigManager.
|
||||
// NewManager no longer does this implicitly.
|
||||
func SetConfigManager(m *Manager) {
|
||||
configMu.Lock()
|
||||
defer configMu.Unlock()
|
||||
configInstance = m
|
||||
}
|
||||
|
||||
// NewManager creates a new, isolated configuration manager with defaults.
|
||||
// It does not replace the global manager; use SetConfigManager for that.
|
||||
func NewManager() *Manager {
|
||||
v := viper.New()
|
||||
|
||||
// Set configuration file settings
|
||||
v.SetConfigName("config")
|
||||
v.SetConfigType("yaml")
|
||||
v.AddConfigPath(".")
|
||||
v.AddConfigPath("./config")
|
||||
// Most trusted location first; the working directory is the least trustworthy
|
||||
// and is searched last (viper takes the first match).
|
||||
v.AddConfigPath("/etc/resolvespec")
|
||||
v.AddConfigPath("$HOME/.resolvespec")
|
||||
v.AddConfigPath("./config")
|
||||
v.AddConfigPath(".")
|
||||
|
||||
// Saved configs may contain secrets; never write them world-readable
|
||||
v.SetConfigPermissions(0o600)
|
||||
|
||||
// Enable environment variable support
|
||||
v.SetEnvPrefix("RESOLVESPEC")
|
||||
@@ -50,6 +83,8 @@ type Option func(*Manager)
|
||||
// WithConfigFile sets a specific config file path
|
||||
func WithConfigFile(path string) Option {
|
||||
return func(m *Manager) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.v.SetConfigFile(path)
|
||||
}
|
||||
}
|
||||
@@ -57,6 +92,8 @@ func WithConfigFile(path string) Option {
|
||||
// WithConfigName sets the config file name (without extension)
|
||||
func WithConfigName(name string) Option {
|
||||
return func(m *Manager) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.v.SetConfigName(name)
|
||||
}
|
||||
}
|
||||
@@ -64,6 +101,8 @@ func WithConfigName(name string) Option {
|
||||
// WithConfigPath adds a path to search for config files
|
||||
func WithConfigPath(path string) Option {
|
||||
return func(m *Manager) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.v.AddConfigPath(path)
|
||||
}
|
||||
}
|
||||
@@ -71,13 +110,19 @@ func WithConfigPath(path string) Option {
|
||||
// WithEnvPrefix sets the environment variable prefix
|
||||
func WithEnvPrefix(prefix string) Option {
|
||||
return func(m *Manager) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.v.SetEnvPrefix(prefix)
|
||||
}
|
||||
}
|
||||
|
||||
// Load attempts to load configuration from file and environment
|
||||
// Load attempts to load configuration from file and environment.
|
||||
// A missing config file is not an error (defaults and env vars are used); check
|
||||
// ConfigFileUsed after Load to see whether a file was actually read.
|
||||
func (m *Manager) Load() error {
|
||||
// Try to read config file (not an error if it doesn't exist)
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if err := m.v.ReadInConfig(); err != nil {
|
||||
if _, ok := err.(viper.ConfigFileNotFoundError); !ok {
|
||||
return fmt.Errorf("error reading config file: %w", err)
|
||||
@@ -88,8 +133,19 @@ func (m *Manager) Load() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ConfigFileUsed returns the config file that was read by Load, or "" if none was
|
||||
// found (i.e. the manager is running on defaults and environment variables only).
|
||||
func (m *Manager) ConfigFileUsed() string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.v.ConfigFileUsed()
|
||||
}
|
||||
|
||||
// GetConfig returns the complete configuration
|
||||
func (m *Manager) GetConfig() (*Config, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
var cfg Config
|
||||
if err := m.v.Unmarshal(&cfg); err != nil {
|
||||
return nil, fmt.Errorf("failed to unmarshal config: %w", err)
|
||||
@@ -97,46 +153,105 @@ func (m *Manager) GetConfig() (*Config, error) {
|
||||
return &cfg, nil
|
||||
}
|
||||
|
||||
// SetConfig sets the complete configuration atomically
|
||||
func (m *Manager) SetConfig(cfg *Config) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
m.v.Set("servers", cfg.Servers)
|
||||
m.v.Set("tracing", cfg.Tracing)
|
||||
m.v.Set("cache", cfg.Cache)
|
||||
m.v.Set("logger", cfg.Logger)
|
||||
m.v.Set("error_tracking", cfg.ErrorTracking)
|
||||
m.v.Set("middleware", cfg.Middleware)
|
||||
m.v.Set("cors", cfg.CORS)
|
||||
m.v.Set("event_broker", cfg.EventBroker)
|
||||
m.v.Set("dbmanager", cfg.DBManager)
|
||||
m.v.Set("paths", cfg.Paths)
|
||||
m.v.Set("extensions", cfg.Extensions)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get returns a configuration value by key
|
||||
func (m *Manager) Get(key string) interface{} {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.v.Get(key)
|
||||
}
|
||||
|
||||
// GetString returns a string configuration value
|
||||
func (m *Manager) GetString(key string) string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.v.GetString(key)
|
||||
}
|
||||
|
||||
// GetInt returns an int configuration value
|
||||
func (m *Manager) GetInt(key string) int {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.v.GetInt(key)
|
||||
}
|
||||
|
||||
// GetBool returns a bool configuration value
|
||||
func (m *Manager) GetBool(key string) bool {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.v.GetBool(key)
|
||||
}
|
||||
|
||||
// Set sets a configuration value
|
||||
func (m *Manager) Set(key string, value interface{}) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.v.Set(key, value)
|
||||
}
|
||||
|
||||
// SaveConfig writes the current configuration to the specified path.
|
||||
// The file contains the entire merged configuration, including secrets
|
||||
// (database/redis passwords, error-tracking DSN), so it is written with mode 0600.
|
||||
// Prefer supplying secrets via RESOLVESPEC_* environment variables.
|
||||
func (m *Manager) SaveConfig(path string) error {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
if err := m.v.WriteConfigAs(path); err != nil {
|
||||
return fmt.Errorf("failed to save config to %s: %w", path, err)
|
||||
}
|
||||
// viper only applies its permissions when creating the file; tighten a pre-existing one too
|
||||
if err := os.Chmod(path, 0o600); err != nil {
|
||||
return fmt.Errorf("failed to restrict permissions on %s: %w", path, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// setDefaults sets default configuration values
|
||||
func setDefaults(v *viper.Viper) {
|
||||
// Server defaults
|
||||
v.SetDefault("server.addr", ":8080")
|
||||
v.SetDefault("server.shutdown_timeout", "30s")
|
||||
v.SetDefault("server.drain_timeout", "25s")
|
||||
v.SetDefault("server.read_timeout", "10s")
|
||||
v.SetDefault("server.write_timeout", "10s")
|
||||
v.SetDefault("server.idle_timeout", "120s")
|
||||
// Server defaults - new structure
|
||||
v.SetDefault("servers.default_server", "default")
|
||||
|
||||
// Global server timeout defaults
|
||||
v.SetDefault("servers.shutdown_timeout", "30s")
|
||||
v.SetDefault("servers.drain_timeout", "25s")
|
||||
v.SetDefault("servers.read_timeout", "10s")
|
||||
v.SetDefault("servers.write_timeout", "10s")
|
||||
v.SetDefault("servers.idle_timeout", "120s")
|
||||
|
||||
// Default server instance
|
||||
v.SetDefault("servers.instances.default.name", "default")
|
||||
v.SetDefault("servers.instances.default.host", "")
|
||||
v.SetDefault("servers.instances.default.port", 8080)
|
||||
v.SetDefault("servers.instances.default.description", "Default HTTP server")
|
||||
v.SetDefault("servers.instances.default.gzip", false)
|
||||
|
||||
// Tracing defaults
|
||||
v.SetDefault("tracing.enabled", false)
|
||||
v.SetDefault("tracing.service_name", "resolvespec")
|
||||
v.SetDefault("tracing.service_version", "1.0.0")
|
||||
v.SetDefault("tracing.endpoint", "")
|
||||
v.SetDefault("tracing.insecure", false)
|
||||
v.SetDefault("tracing.sample_rate", 0.1)
|
||||
|
||||
// Cache defaults
|
||||
v.SetDefault("cache.provider", "memory")
|
||||
@@ -166,6 +281,34 @@ func setDefaults(v *viper.Viper) {
|
||||
// Database defaults
|
||||
v.SetDefault("database.url", "")
|
||||
|
||||
// Database Manager defaults
|
||||
v.SetDefault("dbmanager.default_connection", "default")
|
||||
v.SetDefault("dbmanager.max_open_conns", 25)
|
||||
v.SetDefault("dbmanager.max_idle_conns", 5)
|
||||
v.SetDefault("dbmanager.conn_max_lifetime", "30m")
|
||||
v.SetDefault("dbmanager.conn_max_idle_time", "5m")
|
||||
v.SetDefault("dbmanager.retry_attempts", 3)
|
||||
v.SetDefault("dbmanager.retry_delay", "1s")
|
||||
v.SetDefault("dbmanager.retry_max_delay", "10s")
|
||||
v.SetDefault("dbmanager.health_check_interval", "30s")
|
||||
v.SetDefault("dbmanager.enable_auto_reconnect", true)
|
||||
|
||||
// Default PostgreSQL connection
|
||||
v.SetDefault("dbmanager.connections.default.name", "default")
|
||||
v.SetDefault("dbmanager.connections.default.type", "postgres")
|
||||
v.SetDefault("dbmanager.connections.default.host", "localhost")
|
||||
v.SetDefault("dbmanager.connections.default.port", 5432)
|
||||
v.SetDefault("dbmanager.connections.default.user", "postgres")
|
||||
v.SetDefault("dbmanager.connections.default.password", "")
|
||||
v.SetDefault("dbmanager.connections.default.database", "resolvespec")
|
||||
v.SetDefault("dbmanager.connections.default.sslmode", "disable")
|
||||
v.SetDefault("dbmanager.connections.default.connect_timeout", "10s")
|
||||
v.SetDefault("dbmanager.connections.default.query_timeout", "30s")
|
||||
v.SetDefault("dbmanager.connections.default.enable_tracing", false)
|
||||
v.SetDefault("dbmanager.connections.default.enable_metrics", false)
|
||||
v.SetDefault("dbmanager.connections.default.enable_logging", false)
|
||||
v.SetDefault("dbmanager.connections.default.default_orm", "bun")
|
||||
|
||||
// Event Broker defaults
|
||||
v.SetDefault("event_broker.enabled", false)
|
||||
v.SetDefault("event_broker.provider", "memory")
|
||||
@@ -200,4 +343,13 @@ func setDefaults(v *viper.Viper) {
|
||||
v.SetDefault("event_broker.retry_policy.initial_delay", "1s")
|
||||
v.SetDefault("event_broker.retry_policy.max_delay", "30s")
|
||||
v.SetDefault("event_broker.retry_policy.backoff_factor", 2.0)
|
||||
|
||||
// Paths defaults (common directory paths)
|
||||
v.SetDefault("paths.data_dir", "./data")
|
||||
v.SetDefault("paths.config_dir", "./config")
|
||||
v.SetDefault("paths.logs_dir", "./logs")
|
||||
v.SetDefault("paths.temp_dir", "./tmp")
|
||||
|
||||
// Extensions defaults (empty map)
|
||||
v.SetDefault("extensions", map[string]interface{}{})
|
||||
}
|
||||
|
||||
+454
-12
@@ -34,8 +34,8 @@ func TestDefaultValues(t *testing.T) {
|
||||
got interface{}
|
||||
expected interface{}
|
||||
}{
|
||||
{"server.addr", cfg.Server.Addr, ":8080"},
|
||||
{"server.shutdown_timeout", cfg.Server.ShutdownTimeout, 30 * time.Second},
|
||||
{"servers.default_server", cfg.Servers.DefaultServer, "default"},
|
||||
{"servers.shutdown_timeout", cfg.Servers.ShutdownTimeout, 30 * time.Second},
|
||||
{"tracing.enabled", cfg.Tracing.Enabled, false},
|
||||
{"tracing.service_name", cfg.Tracing.ServiceName, "resolvespec"},
|
||||
{"cache.provider", cfg.Cache.Provider, "memory"},
|
||||
@@ -46,6 +46,18 @@ func TestDefaultValues(t *testing.T) {
|
||||
{"middleware.rate_limit_burst", cfg.Middleware.RateLimitBurst, 200},
|
||||
}
|
||||
|
||||
// Test default server instance
|
||||
defaultServer, ok := cfg.Servers.Instances["default"]
|
||||
if !ok {
|
||||
t.Fatal("Default server instance not found")
|
||||
}
|
||||
if defaultServer.Port != 8080 {
|
||||
t.Errorf("default server port: got %d, want 8080", defaultServer.Port)
|
||||
}
|
||||
if defaultServer.Name != "default" {
|
||||
t.Errorf("default server name: got %s, want default", defaultServer.Name)
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if tt.got != tt.expected {
|
||||
@@ -57,12 +69,12 @@ func TestDefaultValues(t *testing.T) {
|
||||
|
||||
func TestEnvironmentVariableOverrides(t *testing.T) {
|
||||
// Set environment variables
|
||||
os.Setenv("RESOLVESPEC_SERVER_ADDR", ":9090")
|
||||
os.Setenv("RESOLVESPEC_SERVERS_INSTANCES_DEFAULT_PORT", "9090")
|
||||
os.Setenv("RESOLVESPEC_TRACING_ENABLED", "true")
|
||||
os.Setenv("RESOLVESPEC_CACHE_PROVIDER", "redis")
|
||||
os.Setenv("RESOLVESPEC_LOGGER_DEV", "true")
|
||||
defer func() {
|
||||
os.Unsetenv("RESOLVESPEC_SERVER_ADDR")
|
||||
os.Unsetenv("RESOLVESPEC_SERVERS_INSTANCES_DEFAULT_PORT")
|
||||
os.Unsetenv("RESOLVESPEC_TRACING_ENABLED")
|
||||
os.Unsetenv("RESOLVESPEC_CACHE_PROVIDER")
|
||||
os.Unsetenv("RESOLVESPEC_LOGGER_DEV")
|
||||
@@ -84,7 +96,6 @@ func TestEnvironmentVariableOverrides(t *testing.T) {
|
||||
got interface{}
|
||||
expected interface{}
|
||||
}{
|
||||
{"server.addr", cfg.Server.Addr, ":9090"},
|
||||
{"tracing.enabled", cfg.Tracing.Enabled, true},
|
||||
{"cache.provider", cfg.Cache.Provider, "redis"},
|
||||
{"logger.dev", cfg.Logger.Dev, true},
|
||||
@@ -97,11 +108,17 @@ func TestEnvironmentVariableOverrides(t *testing.T) {
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Test server port override
|
||||
defaultServer := cfg.Servers.Instances["default"]
|
||||
if defaultServer.Port != 9090 {
|
||||
t.Errorf("server port: got %d, want 9090", defaultServer.Port)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProgrammaticConfiguration(t *testing.T) {
|
||||
mgr := NewManager()
|
||||
mgr.Set("server.addr", ":7070")
|
||||
mgr.Set("servers.instances.default.port", 7070)
|
||||
mgr.Set("tracing.service_name", "test-service")
|
||||
|
||||
cfg, err := mgr.GetConfig()
|
||||
@@ -109,8 +126,8 @@ func TestProgrammaticConfiguration(t *testing.T) {
|
||||
t.Fatalf("Failed to get config: %v", err)
|
||||
}
|
||||
|
||||
if cfg.Server.Addr != ":7070" {
|
||||
t.Errorf("server.addr: got %s, want :7070", cfg.Server.Addr)
|
||||
if cfg.Servers.Instances["default"].Port != 7070 {
|
||||
t.Errorf("server port: got %d, want 7070", cfg.Servers.Instances["default"].Port)
|
||||
}
|
||||
|
||||
if cfg.Tracing.ServiceName != "test-service" {
|
||||
@@ -148,8 +165,8 @@ func TestWithOptions(t *testing.T) {
|
||||
}
|
||||
|
||||
// Set environment variable with custom prefix
|
||||
os.Setenv("MYAPP_SERVER_ADDR", ":5000")
|
||||
defer os.Unsetenv("MYAPP_SERVER_ADDR")
|
||||
os.Setenv("MYAPP_SERVERS_INSTANCES_DEFAULT_PORT", "5000")
|
||||
defer os.Unsetenv("MYAPP_SERVERS_INSTANCES_DEFAULT_PORT")
|
||||
|
||||
if err := mgr.Load(); err != nil {
|
||||
t.Fatalf("Failed to load config: %v", err)
|
||||
@@ -160,7 +177,432 @@ func TestWithOptions(t *testing.T) {
|
||||
t.Fatalf("Failed to get config: %v", err)
|
||||
}
|
||||
|
||||
if cfg.Server.Addr != ":5000" {
|
||||
t.Errorf("server.addr: got %s, want :5000", cfg.Server.Addr)
|
||||
if cfg.Servers.Instances["default"].Port != 5000 {
|
||||
t.Errorf("server port: got %d, want 5000", cfg.Servers.Instances["default"].Port)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServersConfig(t *testing.T) {
|
||||
mgr := NewManager()
|
||||
if err := mgr.Load(); err != nil {
|
||||
t.Fatalf("Failed to load config: %v", err)
|
||||
}
|
||||
|
||||
cfg, err := mgr.GetConfig()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get config: %v", err)
|
||||
}
|
||||
|
||||
// Test default server exists
|
||||
if cfg.Servers.DefaultServer != "default" {
|
||||
t.Errorf("Expected default_server to be 'default', got %s", cfg.Servers.DefaultServer)
|
||||
}
|
||||
|
||||
// Test default instance
|
||||
defaultServer, ok := cfg.Servers.Instances["default"]
|
||||
if !ok {
|
||||
t.Fatal("Default server instance not found")
|
||||
}
|
||||
|
||||
if defaultServer.Port != 8080 {
|
||||
t.Errorf("Expected default port 8080, got %d", defaultServer.Port)
|
||||
}
|
||||
|
||||
if defaultServer.Name != "default" {
|
||||
t.Errorf("Expected default name 'default', got %s", defaultServer.Name)
|
||||
}
|
||||
|
||||
if defaultServer.Description != "Default HTTP server" {
|
||||
t.Errorf("Expected description 'Default HTTP server', got %s", defaultServer.Description)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMultipleServerInstances(t *testing.T) {
|
||||
mgr := NewManager()
|
||||
|
||||
// Add additional server instances (default instance exists from defaults)
|
||||
mgr.Set("servers.default_server", "api")
|
||||
mgr.Set("servers.instances.api.name", "api")
|
||||
mgr.Set("servers.instances.api.host", "0.0.0.0")
|
||||
mgr.Set("servers.instances.api.port", 8080)
|
||||
mgr.Set("servers.instances.admin.name", "admin")
|
||||
mgr.Set("servers.instances.admin.host", "localhost")
|
||||
mgr.Set("servers.instances.admin.port", 8081)
|
||||
|
||||
cfg, err := mgr.GetConfig()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get config: %v", err)
|
||||
}
|
||||
|
||||
// Should have default + api + admin = 3 instances
|
||||
if len(cfg.Servers.Instances) < 2 {
|
||||
t.Errorf("Expected at least 2 server instances, got %d", len(cfg.Servers.Instances))
|
||||
}
|
||||
|
||||
// Verify api instance
|
||||
apiServer, ok := cfg.Servers.Instances["api"]
|
||||
if !ok {
|
||||
t.Fatal("API server instance not found")
|
||||
}
|
||||
if apiServer.Port != 8080 {
|
||||
t.Errorf("Expected API port 8080, got %d", apiServer.Port)
|
||||
}
|
||||
|
||||
// Verify admin instance
|
||||
adminServer, ok := cfg.Servers.Instances["admin"]
|
||||
if !ok {
|
||||
t.Fatal("Admin server instance not found")
|
||||
}
|
||||
if adminServer.Port != 8081 {
|
||||
t.Errorf("Expected admin port 8081, got %d", adminServer.Port)
|
||||
}
|
||||
|
||||
// Validate default server
|
||||
if err := cfg.Servers.Validate(); err != nil {
|
||||
t.Errorf("Server config validation failed: %v", err)
|
||||
}
|
||||
|
||||
// Get default
|
||||
defaultSrv, err := cfg.Servers.GetDefault()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get default server: %v", err)
|
||||
}
|
||||
if defaultSrv.Name != "api" {
|
||||
t.Errorf("Expected default server 'api', got '%s'", defaultSrv.Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtensionsField(t *testing.T) {
|
||||
mgr := NewManager()
|
||||
|
||||
// Set custom extensions
|
||||
mgr.Set("extensions.custom_feature.enabled", true)
|
||||
mgr.Set("extensions.custom_feature.api_key", "test-key")
|
||||
mgr.Set("extensions.another_extension.timeout", "5s")
|
||||
|
||||
cfg, err := mgr.GetConfig()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get config: %v", err)
|
||||
}
|
||||
|
||||
if cfg.Extensions == nil {
|
||||
t.Fatal("Extensions should not be nil")
|
||||
}
|
||||
|
||||
// Verify extensions are accessible
|
||||
customFeature := mgr.Get("extensions.custom_feature")
|
||||
if customFeature == nil {
|
||||
t.Error("custom_feature extension not found")
|
||||
}
|
||||
|
||||
// Verify via config manager methods
|
||||
if !mgr.GetBool("extensions.custom_feature.enabled") {
|
||||
t.Error("Expected custom_feature.enabled to be true")
|
||||
}
|
||||
|
||||
if mgr.GetString("extensions.custom_feature.api_key") != "test-key" {
|
||||
t.Error("Expected api_key to be 'test-key'")
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerInstanceValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
instance ServerInstanceConfig
|
||||
expectErr bool
|
||||
}{
|
||||
{
|
||||
name: "valid basic config",
|
||||
instance: ServerInstanceConfig{
|
||||
Name: "test",
|
||||
Port: 8080,
|
||||
},
|
||||
expectErr: false,
|
||||
},
|
||||
{
|
||||
name: "invalid port - too high",
|
||||
instance: ServerInstanceConfig{
|
||||
Name: "test",
|
||||
Port: 99999,
|
||||
},
|
||||
expectErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid port - zero",
|
||||
instance: ServerInstanceConfig{
|
||||
Name: "test",
|
||||
Port: 0,
|
||||
},
|
||||
expectErr: true,
|
||||
},
|
||||
{
|
||||
name: "empty name",
|
||||
instance: ServerInstanceConfig{
|
||||
Name: "",
|
||||
Port: 8080,
|
||||
},
|
||||
expectErr: true,
|
||||
},
|
||||
{
|
||||
name: "conflicting TLS options",
|
||||
instance: ServerInstanceConfig{
|
||||
Name: "test",
|
||||
Port: 8080,
|
||||
SelfSignedSSL: true,
|
||||
AutoTLS: true,
|
||||
},
|
||||
expectErr: true,
|
||||
},
|
||||
{
|
||||
name: "incomplete SSL cert config",
|
||||
instance: ServerInstanceConfig{
|
||||
Name: "test",
|
||||
Port: 8080,
|
||||
SSLCert: "/path/to/cert.pem",
|
||||
// Missing SSLKey
|
||||
},
|
||||
expectErr: true,
|
||||
},
|
||||
{
|
||||
name: "AutoTLS without domains",
|
||||
instance: ServerInstanceConfig{
|
||||
Name: "test",
|
||||
Port: 8080,
|
||||
AutoTLS: true,
|
||||
// Missing AutoTLSDomains
|
||||
},
|
||||
expectErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := tt.instance.Validate()
|
||||
if tt.expectErr && err == nil {
|
||||
t.Error("Expected validation error, got nil")
|
||||
}
|
||||
if !tt.expectErr && err != nil {
|
||||
t.Errorf("Expected no error, got: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyGlobalDefaults(t *testing.T) {
|
||||
globals := ServersConfig{
|
||||
ShutdownTimeout: 30 * time.Second,
|
||||
DrainTimeout: 25 * time.Second,
|
||||
ReadTimeout: 10 * time.Second,
|
||||
WriteTimeout: 10 * time.Second,
|
||||
IdleTimeout: 120 * time.Second,
|
||||
}
|
||||
|
||||
instance := ServerInstanceConfig{
|
||||
Name: "test",
|
||||
Port: 8080,
|
||||
}
|
||||
|
||||
// Apply global defaults
|
||||
instance.ApplyGlobalDefaults(globals)
|
||||
|
||||
// Check that defaults were applied
|
||||
if instance.ShutdownTimeout == nil || *instance.ShutdownTimeout != 30*time.Second {
|
||||
t.Error("ShutdownTimeout not applied correctly")
|
||||
}
|
||||
if instance.DrainTimeout == nil || *instance.DrainTimeout != 25*time.Second {
|
||||
t.Error("DrainTimeout not applied correctly")
|
||||
}
|
||||
if instance.ReadTimeout == nil || *instance.ReadTimeout != 10*time.Second {
|
||||
t.Error("ReadTimeout not applied correctly")
|
||||
}
|
||||
if instance.WriteTimeout == nil || *instance.WriteTimeout != 10*time.Second {
|
||||
t.Error("WriteTimeout not applied correctly")
|
||||
}
|
||||
if instance.IdleTimeout == nil || *instance.IdleTimeout != 120*time.Second {
|
||||
t.Error("IdleTimeout not applied correctly")
|
||||
}
|
||||
|
||||
// Test that explicit overrides are not replaced
|
||||
customTimeout := 60 * time.Second
|
||||
instance2 := ServerInstanceConfig{
|
||||
Name: "test2",
|
||||
Port: 8081,
|
||||
ShutdownTimeout: &customTimeout,
|
||||
}
|
||||
|
||||
instance2.ApplyGlobalDefaults(globals)
|
||||
|
||||
if instance2.ShutdownTimeout == nil || *instance2.ShutdownTimeout != 60*time.Second {
|
||||
t.Error("Custom ShutdownTimeout was overridden")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPathsConfig(t *testing.T) {
|
||||
mgr := NewManager()
|
||||
if err := mgr.Load(); err != nil {
|
||||
t.Fatalf("Failed to load config: %v", err)
|
||||
}
|
||||
|
||||
cfg, err := mgr.GetConfig()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get config: %v", err)
|
||||
}
|
||||
|
||||
// Test default paths exist
|
||||
if !cfg.Paths.Has("data_dir") {
|
||||
t.Error("Expected data_dir path to exist")
|
||||
}
|
||||
if !cfg.Paths.Has("config_dir") {
|
||||
t.Error("Expected config_dir path to exist")
|
||||
}
|
||||
if !cfg.Paths.Has("logs_dir") {
|
||||
t.Error("Expected logs_dir path to exist")
|
||||
}
|
||||
if !cfg.Paths.Has("temp_dir") {
|
||||
t.Error("Expected temp_dir path to exist")
|
||||
}
|
||||
|
||||
// Test Get method
|
||||
dataDir, err := cfg.Paths.Get("data_dir")
|
||||
if err != nil {
|
||||
t.Errorf("Failed to get data_dir: %v", err)
|
||||
}
|
||||
if dataDir != "./data" {
|
||||
t.Errorf("Expected data_dir to be './data', got '%s'", dataDir)
|
||||
}
|
||||
|
||||
// Test GetOrDefault
|
||||
existing := cfg.Paths.GetOrDefault("data_dir", "/default/path")
|
||||
if existing != "./data" {
|
||||
t.Errorf("Expected existing path, got '%s'", existing)
|
||||
}
|
||||
|
||||
nonExisting := cfg.Paths.GetOrDefault("nonexistent", "/default/path")
|
||||
if nonExisting != "/default/path" {
|
||||
t.Errorf("Expected default path, got '%s'", nonExisting)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPathsConfigMethods(t *testing.T) {
|
||||
pc := PathsConfig{
|
||||
"base": "/var/myapp",
|
||||
"data": "/var/myapp/data",
|
||||
}
|
||||
|
||||
// Test Get
|
||||
path, err := pc.Get("base")
|
||||
if err != nil {
|
||||
t.Errorf("Failed to get path: %v", err)
|
||||
}
|
||||
if path != "/var/myapp" {
|
||||
t.Errorf("Expected '/var/myapp', got '%s'", path)
|
||||
}
|
||||
|
||||
// Test Get non-existent
|
||||
_, err = pc.Get("nonexistent")
|
||||
if err == nil {
|
||||
t.Error("Expected error for non-existent path")
|
||||
}
|
||||
|
||||
// Test Set
|
||||
pc.Set("new_path", "/new/location")
|
||||
newPath, err := pc.Get("new_path")
|
||||
if err != nil {
|
||||
t.Errorf("Failed to get newly set path: %v", err)
|
||||
}
|
||||
if newPath != "/new/location" {
|
||||
t.Errorf("Expected '/new/location', got '%s'", newPath)
|
||||
}
|
||||
|
||||
// Test Has
|
||||
if !pc.Has("base") {
|
||||
t.Error("Expected 'base' path to exist")
|
||||
}
|
||||
if pc.Has("nonexistent") {
|
||||
t.Error("Expected 'nonexistent' path to not exist")
|
||||
}
|
||||
|
||||
// Test List
|
||||
names := pc.List()
|
||||
if len(names) != 3 {
|
||||
t.Errorf("Expected 3 paths, got %d", len(names))
|
||||
}
|
||||
|
||||
// Test Join
|
||||
joined, err := pc.Join("base", "subdir", "file.txt")
|
||||
if err != nil {
|
||||
t.Errorf("Failed to join paths: %v", err)
|
||||
}
|
||||
expected := "/var/myapp/subdir/file.txt"
|
||||
if joined != expected {
|
||||
t.Errorf("Expected '%s', got '%s'", expected, joined)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPathsConfigEnvironmentVariables(t *testing.T) {
|
||||
// Set environment variables for paths
|
||||
os.Setenv("RESOLVESPEC_PATHS_DATA_DIR", "/custom/data")
|
||||
os.Setenv("RESOLVESPEC_PATHS_LOGS_DIR", "/custom/logs")
|
||||
defer func() {
|
||||
os.Unsetenv("RESOLVESPEC_PATHS_DATA_DIR")
|
||||
os.Unsetenv("RESOLVESPEC_PATHS_LOGS_DIR")
|
||||
}()
|
||||
|
||||
mgr := NewManager()
|
||||
if err := mgr.Load(); err != nil {
|
||||
t.Fatalf("Failed to load config: %v", err)
|
||||
}
|
||||
|
||||
cfg, err := mgr.GetConfig()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get config: %v", err)
|
||||
}
|
||||
|
||||
// Test environment variable override of existing default path
|
||||
dataDir, err := cfg.Paths.Get("data_dir")
|
||||
if err != nil {
|
||||
t.Errorf("Failed to get data_dir: %v", err)
|
||||
}
|
||||
if dataDir != "/custom/data" {
|
||||
t.Errorf("Expected '/custom/data', got '%s'", dataDir)
|
||||
}
|
||||
|
||||
// Test another environment variable override
|
||||
logsDir, err := cfg.Paths.Get("logs_dir")
|
||||
if err != nil {
|
||||
t.Errorf("Failed to get logs_dir: %v", err)
|
||||
}
|
||||
if logsDir != "/custom/logs" {
|
||||
t.Errorf("Expected '/custom/logs', got '%s'", logsDir)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPathsConfigProgrammatic(t *testing.T) {
|
||||
mgr := NewManager()
|
||||
|
||||
// Set custom paths programmatically
|
||||
mgr.Set("paths.custom_dir", "/my/custom/dir")
|
||||
mgr.Set("paths.cache_dir", "/var/cache/myapp")
|
||||
|
||||
cfg, err := mgr.GetConfig()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get config: %v", err)
|
||||
}
|
||||
|
||||
// Verify custom paths
|
||||
customDir, err := cfg.Paths.Get("custom_dir")
|
||||
if err != nil {
|
||||
t.Errorf("Failed to get custom_dir: %v", err)
|
||||
}
|
||||
if customDir != "/my/custom/dir" {
|
||||
t.Errorf("Expected '/my/custom/dir', got '%s'", customDir)
|
||||
}
|
||||
|
||||
cacheDir, err := cfg.Paths.Get("cache_dir")
|
||||
if err != nil {
|
||||
t.Errorf("Failed to get cache_dir: %v", err)
|
||||
}
|
||||
if cacheDir != "/var/cache/myapp" {
|
||||
t.Errorf("Expected '/var/cache/myapp', got '%s'", cacheDir)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,130 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Get retrieves a path by name
|
||||
func (pc PathsConfig) Get(name string) (string, error) {
|
||||
if pc == nil {
|
||||
return "", fmt.Errorf("paths not initialized")
|
||||
}
|
||||
|
||||
path, ok := pc[name]
|
||||
if !ok {
|
||||
return "", fmt.Errorf("path '%s' not found", name)
|
||||
}
|
||||
|
||||
return path, nil
|
||||
}
|
||||
|
||||
// GetOrDefault retrieves a path by name, returning defaultPath if not found
|
||||
func (pc PathsConfig) GetOrDefault(name, defaultPath string) string {
|
||||
if pc == nil {
|
||||
return defaultPath
|
||||
}
|
||||
|
||||
path, ok := pc[name]
|
||||
if !ok {
|
||||
return defaultPath
|
||||
}
|
||||
|
||||
return path
|
||||
}
|
||||
|
||||
// Set sets a path by name. It takes a pointer so a nil map can be allocated.
|
||||
// PathsConfig is not safe for concurrent mutation; populate it before sharing.
|
||||
func (pc *PathsConfig) Set(name, path string) {
|
||||
if *pc == nil {
|
||||
*pc = make(PathsConfig)
|
||||
}
|
||||
(*pc)[name] = path
|
||||
}
|
||||
|
||||
// Has checks if a path exists by name
|
||||
func (pc PathsConfig) Has(name string) bool {
|
||||
if pc == nil {
|
||||
return false
|
||||
}
|
||||
_, ok := pc[name]
|
||||
return ok
|
||||
}
|
||||
|
||||
// EnsureDir ensures a directory exists at the specified path name
|
||||
// Creates the directory if it doesn't exist with the given permissions
|
||||
func (pc PathsConfig) EnsureDir(name string, perm os.FileMode) error {
|
||||
path, err := pc.Get(name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Check if directory exists
|
||||
info, err := os.Stat(path)
|
||||
if err == nil {
|
||||
// Path exists, check if it's a directory
|
||||
if !info.IsDir() {
|
||||
return fmt.Errorf("path '%s' exists but is not a directory: %s", name, path)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Directory doesn't exist, create it
|
||||
if os.IsNotExist(err) {
|
||||
if err := os.MkdirAll(path, perm); err != nil {
|
||||
return fmt.Errorf("failed to create directory for '%s' at %s: %w", name, path, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("failed to stat path '%s' at %s: %w", name, path, err)
|
||||
}
|
||||
|
||||
// AbsPath returns the absolute path for a named path
|
||||
func (pc PathsConfig) AbsPath(name string) (string, error) {
|
||||
path, err := pc.Get(name)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
absPath, err := filepath.Abs(path)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to get absolute path for '%s': %w", name, err)
|
||||
}
|
||||
|
||||
return absPath, nil
|
||||
}
|
||||
|
||||
// Join joins path segments with a named base path.
|
||||
// It returns an error if the result would escape the base path.
|
||||
func (pc PathsConfig) Join(name string, elem ...string) (string, error) {
|
||||
base, err := pc.Get(name)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
parts := append([]string{base}, elem...)
|
||||
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
|
||||
func (pc PathsConfig) List() []string {
|
||||
if pc == nil {
|
||||
return []string{}
|
||||
}
|
||||
|
||||
names := make([]string, 0, len(pc))
|
||||
for name := range pc {
|
||||
names = append(names, name)
|
||||
}
|
||||
return names
|
||||
}
|
||||
@@ -0,0 +1,148 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ApplyGlobalDefaults applies global server defaults to this instance
|
||||
// Called for instances that don't specify their own timeout values
|
||||
func (sic *ServerInstanceConfig) ApplyGlobalDefaults(globals ServersConfig) {
|
||||
if sic.ShutdownTimeout == nil && globals.ShutdownTimeout > 0 {
|
||||
t := globals.ShutdownTimeout
|
||||
sic.ShutdownTimeout = &t
|
||||
}
|
||||
if sic.DrainTimeout == nil && globals.DrainTimeout > 0 {
|
||||
t := globals.DrainTimeout
|
||||
sic.DrainTimeout = &t
|
||||
}
|
||||
if sic.ReadTimeout == nil && globals.ReadTimeout > 0 {
|
||||
t := globals.ReadTimeout
|
||||
sic.ReadTimeout = &t
|
||||
}
|
||||
if sic.WriteTimeout == nil && globals.WriteTimeout > 0 {
|
||||
t := globals.WriteTimeout
|
||||
sic.WriteTimeout = &t
|
||||
}
|
||||
if sic.IdleTimeout == nil && globals.IdleTimeout > 0 {
|
||||
t := globals.IdleTimeout
|
||||
sic.IdleTimeout = &t
|
||||
}
|
||||
}
|
||||
|
||||
// Validate validates the ServerInstanceConfig
|
||||
func (sic *ServerInstanceConfig) Validate() error {
|
||||
if sic.Name == "" {
|
||||
return fmt.Errorf("server instance name cannot be empty")
|
||||
}
|
||||
if sic.Port <= 0 || sic.Port > 65535 {
|
||||
return fmt.Errorf("invalid port: %d (must be 1-65535)", sic.Port)
|
||||
}
|
||||
|
||||
// Validate TLS options are mutually exclusive
|
||||
tlsCount := 0
|
||||
if sic.SSLCert != "" || sic.SSLKey != "" {
|
||||
tlsCount++
|
||||
}
|
||||
if sic.SelfSignedSSL {
|
||||
tlsCount++
|
||||
}
|
||||
if sic.AutoTLS {
|
||||
tlsCount++
|
||||
}
|
||||
if tlsCount > 1 {
|
||||
return fmt.Errorf("server '%s': only one TLS option can be enabled", sic.Name)
|
||||
}
|
||||
|
||||
// If using certificate files, both must be provided
|
||||
if (sic.SSLCert != "" && sic.SSLKey == "") || (sic.SSLCert == "" && sic.SSLKey != "") {
|
||||
return fmt.Errorf("server '%s': both ssl_cert and ssl_key must be provided", sic.Name)
|
||||
}
|
||||
|
||||
// If using AutoTLS, domains must be specified
|
||||
if sic.AutoTLS && len(sic.AutoTLSDomains) == 0 {
|
||||
return fmt.Errorf("server '%s': auto_tls_domains must be specified when auto_tls is enabled", sic.Name)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Validate validates the ServersConfig
|
||||
func (sc *ServersConfig) Validate() error {
|
||||
if len(sc.Instances) == 0 {
|
||||
return fmt.Errorf("at least one server instance must be configured")
|
||||
}
|
||||
|
||||
if sc.DefaultServer != "" {
|
||||
if _, ok := sc.Instances[sc.DefaultServer]; !ok {
|
||||
return fmt.Errorf("default server '%s' not found in instances", sc.DefaultServer)
|
||||
}
|
||||
}
|
||||
|
||||
// Validate each instance
|
||||
for name := range sc.Instances {
|
||||
instance := sc.Instances[name]
|
||||
if instance.Name != name {
|
||||
return fmt.Errorf("server instance name mismatch: key='%s', name='%s'", name, instance.Name)
|
||||
}
|
||||
if err := instance.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetDefault returns the default server instance configuration.
|
||||
// The returned pointer refers to a copy: mutating it does not modify sc.Instances.
|
||||
func (sc *ServersConfig) GetDefault() (*ServerInstanceConfig, error) {
|
||||
if sc.DefaultServer == "" {
|
||||
return nil, fmt.Errorf("no default server configured")
|
||||
}
|
||||
|
||||
instance, ok := sc.Instances[sc.DefaultServer]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("default server '%s' not found", sc.DefaultServer)
|
||||
}
|
||||
|
||||
return &instance, nil
|
||||
}
|
||||
|
||||
// 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) {
|
||||
hostname, _ = os.Hostname()
|
||||
ipNetList = make([]net.IP, 0)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var ips []string
|
||||
if addrs, err := net.DefaultResolver.LookupIPAddr(ctx, hostname); err == nil {
|
||||
for _, a := range addrs {
|
||||
if a.IP.IsLoopback() {
|
||||
continue
|
||||
}
|
||||
ips = append(ips, a.IP.String())
|
||||
ipNetList = append(ipNetList, a.IP)
|
||||
}
|
||||
}
|
||||
|
||||
if len(ips) == 0 {
|
||||
ifaceAddrs, _ := net.InterfaceAddrs()
|
||||
for _, a := range ifaceAddrs {
|
||||
ipn, ok := a.(*net.IPNet)
|
||||
if !ok || ipn.IP.IsLoopback() {
|
||||
continue
|
||||
}
|
||||
ips = append(ips, ipn.IP.String())
|
||||
ipNetList = append(ipNetList, ipn.IP)
|
||||
}
|
||||
}
|
||||
|
||||
return hostname, strings.Join(ips, ","), ipNetList
|
||||
}
|
||||
@@ -0,0 +1,590 @@
|
||||
# Database Connection Manager (dbmanager)
|
||||
|
||||
A comprehensive database connection manager for Go that provides centralized management of multiple named database connections with support for PostgreSQL, SQLite, MSSQL, and MongoDB.
|
||||
|
||||
## Features
|
||||
|
||||
- **Multiple Named Connections**: Manage multiple database connections with names like `primary`, `analytics`, `cache-db`
|
||||
- **Multi-Database Support**: PostgreSQL, SQLite, Microsoft SQL Server, and MongoDB
|
||||
- **Multi-ORM Access**: Each SQL connection provides access through:
|
||||
- **Bun ORM** - Modern, lightweight ORM
|
||||
- **GORM** - Popular Go ORM
|
||||
- **Native** - Standard library `*sql.DB`
|
||||
- All three share the same underlying connection pool
|
||||
- **SQLite Schema Translation**: Automatic conversion of `schema.table` to `schema_table` for SQLite compatibility
|
||||
- **Configuration-Driven**: YAML configuration with Viper integration
|
||||
- **Production-Ready Features**:
|
||||
- Automatic health checks and reconnection
|
||||
- Prometheus metrics
|
||||
- Connection pooling with configurable limits
|
||||
- Retry logic with exponential backoff
|
||||
- Graceful shutdown
|
||||
- OpenTelemetry tracing support
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
go get github.com/bitechdev/ResolveSpec/pkg/dbmanager
|
||||
```
|
||||
|
||||
## Quick Start
|
||||
|
||||
### 1. Configuration
|
||||
|
||||
Create a configuration file (e.g., `config.yaml`):
|
||||
|
||||
```yaml
|
||||
dbmanager:
|
||||
default_connection: "primary"
|
||||
|
||||
# Global connection pool defaults
|
||||
max_open_conns: 25
|
||||
max_idle_conns: 5
|
||||
conn_max_lifetime: 30m
|
||||
conn_max_idle_time: 5m
|
||||
|
||||
# Retry configuration
|
||||
retry_attempts: 3
|
||||
retry_delay: 1s
|
||||
retry_max_delay: 10s
|
||||
|
||||
# Health checks
|
||||
health_check_interval: 30s
|
||||
|
||||
connections:
|
||||
# Primary PostgreSQL connection
|
||||
primary:
|
||||
type: postgres
|
||||
host: localhost
|
||||
port: 5432
|
||||
user: myuser
|
||||
password: mypassword
|
||||
database: myapp
|
||||
sslmode: disable
|
||||
default_orm: bun
|
||||
enable_metrics: true
|
||||
enable_tracing: true
|
||||
enable_logging: true
|
||||
|
||||
# Read replica for analytics
|
||||
analytics:
|
||||
type: postgres
|
||||
dsn: "postgres://readonly:pass@analytics:5432/analytics"
|
||||
default_orm: bun
|
||||
enable_metrics: true
|
||||
|
||||
# SQLite cache
|
||||
cache-db:
|
||||
type: sqlite
|
||||
filepath: /var/lib/app/cache.db
|
||||
max_open_conns: 1
|
||||
|
||||
# MongoDB for documents
|
||||
documents:
|
||||
type: mongodb
|
||||
host: localhost
|
||||
port: 27017
|
||||
database: documents
|
||||
user: mongouser
|
||||
password: mongopass
|
||||
auth_source: admin
|
||||
enable_metrics: true
|
||||
```
|
||||
|
||||
### 2. Initialize Manager
|
||||
|
||||
```go
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbmanager"
|
||||
)
|
||||
|
||||
func main() {
|
||||
// Load configuration
|
||||
cfgMgr := config.NewManager()
|
||||
if err := cfgMgr.Load(); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
cfg, _ := cfgMgr.GetConfig()
|
||||
|
||||
// Create database manager
|
||||
mgr, err := dbmanager.NewManager(cfg.DBManager)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
defer mgr.Close()
|
||||
|
||||
// Connect all databases
|
||||
ctx := context.Background()
|
||||
if err := mgr.Connect(ctx); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
|
||||
// Your application code here...
|
||||
}
|
||||
```
|
||||
|
||||
### 3. Use Database Connections
|
||||
|
||||
#### Get Default Database
|
||||
|
||||
```go
|
||||
// Get the default database (as configured common.Database interface)
|
||||
db, err := mgr.GetDefaultDatabase()
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
|
||||
// Use it with any query
|
||||
var users []User
|
||||
err = db.NewSelect().
|
||||
Model(&users).
|
||||
Where("active = ?", true).
|
||||
Scan(ctx, &users)
|
||||
```
|
||||
|
||||
#### Get Named Connection with Specific ORM
|
||||
|
||||
```go
|
||||
// Get primary connection
|
||||
primary, err := mgr.Get("primary")
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
|
||||
// Use with Bun
|
||||
bunDB, err := primary.Bun()
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
err = bunDB.NewSelect().Model(&users).Scan(ctx)
|
||||
|
||||
// Use with GORM (same underlying connection!)
|
||||
gormDB, err := primary.GORM()
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
gormDB.Where("active = ?", true).Find(&users)
|
||||
|
||||
// Use native *sql.DB
|
||||
nativeDB, err := primary.Native()
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
rows, err := nativeDB.QueryContext(ctx, "SELECT * FROM users WHERE active = $1", true)
|
||||
```
|
||||
|
||||
#### Cross-Database Example with SQLite
|
||||
|
||||
```go
|
||||
// Same model works across all databases
|
||||
type User struct {
|
||||
ID int `bun:"id,pk"`
|
||||
Username string `bun:"username"`
|
||||
Email string `bun:"email"`
|
||||
}
|
||||
|
||||
func (User) TableName() string {
|
||||
return "auth.users"
|
||||
}
|
||||
|
||||
// PostgreSQL connection
|
||||
pgConn, _ := mgr.Get("primary")
|
||||
pgDB, _ := pgConn.Bun()
|
||||
var pgUsers []User
|
||||
pgDB.NewSelect().Model(&pgUsers).Scan(ctx)
|
||||
// Executes: SELECT * FROM auth.users
|
||||
|
||||
// SQLite connection
|
||||
sqliteConn, _ := mgr.Get("cache-db")
|
||||
sqliteDB, _ := sqliteConn.Bun()
|
||||
var sqliteUsers []User
|
||||
sqliteDB.NewSelect().Model(&sqliteUsers).Scan(ctx)
|
||||
// Executes: SELECT * FROM auth_users (schema.table → schema_table)
|
||||
```
|
||||
|
||||
#### Use MongoDB
|
||||
|
||||
```go
|
||||
// Get MongoDB connection
|
||||
docs, err := mgr.Get("documents")
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
|
||||
mongoClient, err := docs.MongoDB()
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
|
||||
collection := mongoClient.Database("documents").Collection("articles")
|
||||
// Use MongoDB driver...
|
||||
```
|
||||
|
||||
#### Change Default Database
|
||||
|
||||
```go
|
||||
// Switch to analytics database as default
|
||||
err := mgr.SetDefaultDatabase("analytics")
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
|
||||
// Now GetDefaultDatabase() returns the analytics connection
|
||||
db, _ := mgr.GetDefaultDatabase()
|
||||
```
|
||||
|
||||
## Configuration Reference
|
||||
|
||||
### Manager Configuration
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `default_connection` | string | "" | Name of the default connection |
|
||||
| `connections` | map | {} | Map of connection name to ConnectionConfig |
|
||||
| `max_open_conns` | int | 25 | Global default for max open connections |
|
||||
| `max_idle_conns` | int | 5 | Global default for max idle connections |
|
||||
| `conn_max_lifetime` | duration | 30m | Global default for connection max lifetime |
|
||||
| `conn_max_idle_time` | duration | 5m | Global default for connection max idle time |
|
||||
| `retry_attempts` | int | 3 | Number of connection retry attempts |
|
||||
| `retry_delay` | duration | 1s | Initial retry delay |
|
||||
| `retry_max_delay` | duration | 10s | Maximum retry delay |
|
||||
| `health_check_interval` | duration | 30s | Interval between health checks |
|
||||
| `enable_auto_reconnect` | bool | - | Deprecated and ignored: the manager never closes the pool to recover from errors |
|
||||
|
||||
### Connection Configuration
|
||||
|
||||
| Field | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| `name` | string | Unique connection name |
|
||||
| `type` | string | Database type: `postgres`, `sqlite`, `mssql`, `mongodb` |
|
||||
| `dsn` | string | Complete connection string (overrides individual params) |
|
||||
| `host` | string | Database host |
|
||||
| `port` | int | Database port |
|
||||
| `user` | string | Username |
|
||||
| `password` | string | Password |
|
||||
| `database` | string | Database name |
|
||||
| `sslmode` | string | SSL mode (postgres/mssql): `disable`, `require`, etc. |
|
||||
| `schema` | string | Default schema (postgres/mssql) |
|
||||
| `filepath` | string | File path (sqlite only) |
|
||||
| `auth_source` | string | Auth source (mongodb) |
|
||||
| `replica_set` | string | Replica set name (mongodb) |
|
||||
| `read_preference` | string | Read preference (mongodb): `primary`, `secondary`, etc. |
|
||||
| `max_open_conns` | int | Override global max open connections |
|
||||
| `max_idle_conns` | int | Override global max idle connections |
|
||||
| `conn_max_lifetime` | duration | Override global connection max lifetime |
|
||||
| `conn_max_idle_time` | duration | Override global connection max idle time |
|
||||
| `connect_timeout` | duration | Connection timeout (default: 10s) |
|
||||
| `query_timeout` | duration | Query timeout (default: 30s) |
|
||||
| `enable_tracing` | bool | Enable OpenTelemetry tracing |
|
||||
| `enable_metrics` | bool | Enable Prometheus metrics |
|
||||
| `enable_logging` | bool | Enable connection logging |
|
||||
| `default_orm` | string | Default ORM for Database(): `bun`, `gorm`, `native` |
|
||||
| `tags` | map[string]string | Custom tags for filtering/organization |
|
||||
|
||||
## Advanced Usage
|
||||
|
||||
### Health Checks
|
||||
|
||||
```go
|
||||
// Manual health check
|
||||
if err := mgr.HealthCheck(ctx); err != nil {
|
||||
log.Printf("Health check failed: %v", err)
|
||||
}
|
||||
|
||||
// Per-connection health check
|
||||
primary, _ := mgr.Get("primary")
|
||||
if err := primary.HealthCheck(ctx); err != nil {
|
||||
log.Printf("Primary connection unhealthy: %v", err)
|
||||
|
||||
// Manual reconnect
|
||||
if err := primary.Reconnect(ctx); err != nil {
|
||||
log.Printf("Reconnection failed: %v", err)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Connection Statistics
|
||||
|
||||
```go
|
||||
// Get overall statistics
|
||||
stats := mgr.Stats()
|
||||
fmt.Printf("Total connections: %d\n", stats.TotalConnections)
|
||||
fmt.Printf("Healthy: %d, Unhealthy: %d\n", stats.HealthyCount, stats.UnhealthyCount)
|
||||
|
||||
// Per-connection stats
|
||||
for name, connStats := range stats.ConnectionStats {
|
||||
fmt.Printf("%s: %d open, %d in use, %d idle\n",
|
||||
name,
|
||||
connStats.OpenConnections,
|
||||
connStats.InUse,
|
||||
connStats.Idle)
|
||||
}
|
||||
|
||||
// Individual connection stats
|
||||
primary, _ := mgr.Get("primary")
|
||||
stats := primary.Stats()
|
||||
fmt.Printf("Wait count: %d, Wait duration: %v\n",
|
||||
stats.WaitCount,
|
||||
stats.WaitDuration)
|
||||
```
|
||||
|
||||
### Prometheus Metrics
|
||||
|
||||
The package automatically exports Prometheus metrics:
|
||||
|
||||
- `dbmanager_connections_total` - Total configured connections by type
|
||||
- `dbmanager_connection_status` - Connection health status (1=healthy, 0=unhealthy)
|
||||
- `dbmanager_connection_pool_size` - Connection pool statistics by state
|
||||
- `dbmanager_connection_wait_count` - Times connections waited for availability
|
||||
- `dbmanager_connection_wait_duration_seconds` - Total wait duration
|
||||
- `dbmanager_health_check_duration_seconds` - Health check execution time
|
||||
- `dbmanager_reconnect_attempts_total` - Reconnection attempts and results
|
||||
- `dbmanager_connection_lifetime_closed_total` - Connections closed due to max lifetime
|
||||
- `dbmanager_connection_idle_closed_total` - Connections closed due to max idle time
|
||||
|
||||
Metrics are automatically updated during health checks. To manually publish metrics:
|
||||
|
||||
```go
|
||||
if mgr, ok := mgr.(*connectionManager); ok {
|
||||
mgr.PublishMetrics()
|
||||
}
|
||||
```
|
||||
|
||||
## Architecture
|
||||
|
||||
### Single Connection Pool, Multiple ORMs
|
||||
|
||||
A key design principle is that Bun, GORM, and Native all wrap the **same underlying `*sql.DB`** connection pool:
|
||||
|
||||
```
|
||||
┌─────────────────────────────────────┐
|
||||
│ SQL Connection │
|
||||
├─────────────────────────────────────┤
|
||||
│ ┌─────────┐ ┌──────┐ ┌────────┐ │
|
||||
│ │ Bun │ │ GORM │ │ Native │ │
|
||||
│ └────┬────┘ └───┬──┘ └───┬────┘ │
|
||||
│ │ │ │ │
|
||||
│ └───────────┴─────────┘ │
|
||||
│ *sql.DB │
|
||||
│ (single pool) │
|
||||
└─────────────────────────────────────┘
|
||||
```
|
||||
|
||||
**Benefits:**
|
||||
- No connection duplication
|
||||
- Consistent pool limits across all ORMs
|
||||
- Unified connection statistics
|
||||
- Lower resource usage
|
||||
|
||||
### Provider Pattern
|
||||
|
||||
Each database type has a dedicated provider:
|
||||
|
||||
- **PostgresProvider** - Uses `pgx` driver
|
||||
- **SQLiteProvider** - Uses `glebarez/sqlite` (pure Go)
|
||||
- **MSSQLProvider** - Uses `go-mssqldb`
|
||||
- **MongoProvider** - Uses official `mongo-driver`
|
||||
|
||||
Providers handle:
|
||||
- Connection establishment with retry logic
|
||||
- Health checking
|
||||
- Connection statistics
|
||||
- Connection cleanup
|
||||
|
||||
### SQLite Schema Handling
|
||||
|
||||
SQLite doesn't support schemas in the same way as PostgreSQL or MSSQL. To ensure compatibility when using models designed for multi-schema databases:
|
||||
|
||||
**Automatic Translation**: When a table name contains a schema prefix (e.g., `myschema.mytable`), it is automatically converted to `myschema_mytable` for SQLite databases.
|
||||
|
||||
```go
|
||||
// Model definition (works across all databases)
|
||||
func (User) TableName() string {
|
||||
return "auth.users" // PostgreSQL/MSSQL: "auth"."users"
|
||||
// SQLite: "auth_users"
|
||||
}
|
||||
|
||||
// Query execution
|
||||
db.NewSelect().Model(&User{}).Scan(ctx)
|
||||
// PostgreSQL/MSSQL: SELECT * FROM auth.users
|
||||
// SQLite: SELECT * FROM auth_users
|
||||
```
|
||||
|
||||
**How it Works**:
|
||||
- Bun, GORM, and Native adapters detect the driver type
|
||||
- `parseTableName()` automatically translates schema.table → schema_table for SQLite
|
||||
- Translation happens transparently in all database operations (SELECT, INSERT, UPDATE, DELETE)
|
||||
- Preload and relation queries are also handled automatically
|
||||
|
||||
**Benefits**:
|
||||
- Write database-agnostic code
|
||||
- Use the same models across PostgreSQL, MSSQL, and SQLite
|
||||
- No conditional logic needed in your application
|
||||
- Schema separation maintained through naming convention in SQLite
|
||||
|
||||
## Best Practices
|
||||
|
||||
1. **Use Named Connections**: Be explicit about which database you're accessing
|
||||
```go
|
||||
primary, _ := mgr.Get("primary") // Good
|
||||
db, _ := mgr.GetDefaultDatabase() // Risky if default changes
|
||||
```
|
||||
|
||||
2. **Configure Connection Pools**: Tune based on your workload
|
||||
```yaml
|
||||
connections:
|
||||
primary:
|
||||
max_open_conns: 100 # High traffic API
|
||||
max_idle_conns: 25
|
||||
analytics:
|
||||
max_open_conns: 10 # Background analytics
|
||||
max_idle_conns: 2
|
||||
```
|
||||
|
||||
3. **Enable Health Checks**: Catch connection issues early
|
||||
```yaml
|
||||
health_check_interval: 30s
|
||||
```
|
||||
|
||||
4. **Use Appropriate ORM**: Choose based on your needs
|
||||
- **Bun**: Modern, fast, type-safe - recommended for new code
|
||||
- **GORM**: Mature, feature-rich - good for existing GORM code
|
||||
- **Native**: Maximum control - use for performance-critical queries
|
||||
|
||||
5. **Monitor Metrics**: Watch connection pool utilization
|
||||
- If `wait_count` is high, increase `max_open_conns`
|
||||
- If `idle` is always high, decrease `max_idle_conns`
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Connection Failures
|
||||
|
||||
If connections fail to establish:
|
||||
|
||||
1. Check configuration:
|
||||
```bash
|
||||
# Test connection manually
|
||||
psql -h localhost -U myuser -d myapp
|
||||
```
|
||||
|
||||
2. Enable logging:
|
||||
```yaml
|
||||
connections:
|
||||
primary:
|
||||
enable_logging: true
|
||||
```
|
||||
|
||||
3. Check retry attempts:
|
||||
```yaml
|
||||
retry_attempts: 5 # Increase retries
|
||||
retry_max_delay: 30s
|
||||
```
|
||||
|
||||
### Pool Exhaustion
|
||||
|
||||
If you see "too many connections" errors:
|
||||
|
||||
1. Increase pool size:
|
||||
```yaml
|
||||
max_open_conns: 50 # Increase from default 25
|
||||
```
|
||||
|
||||
2. Reduce connection lifetime:
|
||||
```yaml
|
||||
conn_max_lifetime: 15m # Recycle faster
|
||||
```
|
||||
|
||||
3. Monitor wait stats:
|
||||
```go
|
||||
stats := primary.Stats()
|
||||
if stats.WaitCount > 1000 {
|
||||
log.Warn("High connection wait count")
|
||||
}
|
||||
```
|
||||
|
||||
### MongoDB vs SQL Confusion
|
||||
|
||||
MongoDB connections don't support SQL ORMs:
|
||||
|
||||
```go
|
||||
docs, _ := mgr.Get("documents")
|
||||
|
||||
// ✓ Correct
|
||||
mongoClient, _ := docs.MongoDB()
|
||||
|
||||
// ✗ Error: ErrNotSQLDatabase
|
||||
bunDB, err := docs.Bun() // Won't work!
|
||||
```
|
||||
|
||||
SQL connections don't support MongoDB:
|
||||
|
||||
```go
|
||||
primary, _ := mgr.Get("primary")
|
||||
|
||||
// ✓ Correct
|
||||
bunDB, _ := primary.Bun()
|
||||
|
||||
// ✗ Error: ErrNotMongoDB
|
||||
mongoClient, err := primary.MongoDB() // Won't work!
|
||||
```
|
||||
|
||||
## Migration Guide
|
||||
|
||||
### From Raw `database/sql`
|
||||
|
||||
Before:
|
||||
```go
|
||||
db, err := sql.Open("postgres", dsn)
|
||||
defer db.Close()
|
||||
|
||||
rows, err := db.Query("SELECT * FROM users")
|
||||
```
|
||||
|
||||
After:
|
||||
```go
|
||||
mgr, _ := dbmanager.NewManager(cfg.DBManager)
|
||||
mgr.Connect(ctx)
|
||||
defer mgr.Close()
|
||||
|
||||
primary, _ := mgr.Get("primary")
|
||||
nativeDB, _ := primary.Native()
|
||||
|
||||
rows, err := nativeDB.Query("SELECT * FROM users")
|
||||
```
|
||||
|
||||
### From Direct Bun/GORM
|
||||
|
||||
Before:
|
||||
```go
|
||||
sqldb, _ := sql.Open("pgx", dsn)
|
||||
bunDB := bun.NewDB(sqldb, pgdialect.New())
|
||||
|
||||
var users []User
|
||||
bunDB.NewSelect().Model(&users).Scan(ctx)
|
||||
```
|
||||
|
||||
After:
|
||||
```go
|
||||
mgr, _ := dbmanager.NewManager(cfg.DBManager)
|
||||
mgr.Connect(ctx)
|
||||
|
||||
primary, _ := mgr.Get("primary")
|
||||
bunDB, _ := primary.Bun()
|
||||
|
||||
var users []User
|
||||
bunDB.NewSelect().Model(&users).Scan(ctx)
|
||||
```
|
||||
|
||||
## License
|
||||
|
||||
Same as the parent project.
|
||||
|
||||
## Contributing
|
||||
|
||||
Contributions are welcome! Please submit issues and pull requests to the main repository.
|
||||
@@ -0,0 +1,524 @@
|
||||
package dbmanager
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||
)
|
||||
|
||||
// DatabaseType represents the type of database
|
||||
type DatabaseType string
|
||||
|
||||
const (
|
||||
// DatabaseTypePostgreSQL represents PostgreSQL database
|
||||
DatabaseTypePostgreSQL DatabaseType = "postgres"
|
||||
|
||||
// DatabaseTypeSQLite represents SQLite database
|
||||
DatabaseTypeSQLite DatabaseType = "sqlite"
|
||||
|
||||
// DatabaseTypeMSSQL represents Microsoft SQL Server database
|
||||
DatabaseTypeMSSQL DatabaseType = "mssql"
|
||||
|
||||
// DatabaseTypeMongoDB represents MongoDB database
|
||||
DatabaseTypeMongoDB DatabaseType = "mongodb"
|
||||
)
|
||||
|
||||
// ORMType represents the ORM to use for database operations
|
||||
type ORMType string
|
||||
|
||||
const (
|
||||
// ORMTypeBun represents Bun ORM
|
||||
ORMTypeBun ORMType = "bun"
|
||||
|
||||
// ORMTypeGORM represents GORM
|
||||
ORMTypeGORM ORMType = "gorm"
|
||||
|
||||
// ORMTypeNative represents native database/sql
|
||||
ORMTypeNative ORMType = "native"
|
||||
)
|
||||
|
||||
// ManagerConfig contains configuration for the database connection manager
|
||||
type ManagerConfig struct {
|
||||
// DefaultConnection is the name of the default connection to use
|
||||
DefaultConnection string `mapstructure:"default_connection"`
|
||||
|
||||
// Connections is a map of connection name to connection configuration
|
||||
Connections map[string]ConnectionConfig `mapstructure:"connections"`
|
||||
|
||||
// Global connection pool defaults
|
||||
MaxOpenConns int `mapstructure:"max_open_conns"`
|
||||
MaxIdleConns int `mapstructure:"max_idle_conns"`
|
||||
ConnMaxLifetime time.Duration `mapstructure:"conn_max_lifetime"`
|
||||
ConnMaxIdleTime time.Duration `mapstructure:"conn_max_idle_time"`
|
||||
|
||||
// Retry policy
|
||||
RetryAttempts int `mapstructure:"retry_attempts"`
|
||||
RetryDelay time.Duration `mapstructure:"retry_delay"`
|
||||
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
|
||||
|
||||
// Health checks. A zero HealthCheckInterval selects the default (15s); a
|
||||
// negative value disables the background health checker.
|
||||
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"`
|
||||
}
|
||||
|
||||
// ConnectionConfig defines configuration for a single database connection
|
||||
type ConnectionConfig struct {
|
||||
// Name is the unique name of this connection
|
||||
Name string `mapstructure:"name"`
|
||||
|
||||
// Type is the database type (postgres, sqlite, mssql, mongodb)
|
||||
Type DatabaseType `mapstructure:"type"`
|
||||
|
||||
// DSN is the complete Data Source Name / connection string
|
||||
// If provided, this takes precedence over individual connection parameters
|
||||
DSN string `mapstructure:"dsn"`
|
||||
|
||||
// Connection parameters (used if DSN is not provided)
|
||||
Host string `mapstructure:"host"`
|
||||
Port int `mapstructure:"port"`
|
||||
User string `mapstructure:"user"`
|
||||
Password string `mapstructure:"password"`
|
||||
Database string `mapstructure:"database"`
|
||||
|
||||
// PostgreSQL/MSSQL specific
|
||||
SSLMode string `mapstructure:"sslmode"` // disable, require, verify-ca, verify-full
|
||||
Schema string `mapstructure:"schema"` // Default schema
|
||||
|
||||
// SQLite specific
|
||||
FilePath string `mapstructure:"filepath"`
|
||||
|
||||
// MongoDB specific
|
||||
AuthSource string `mapstructure:"auth_source"`
|
||||
ReplicaSet string `mapstructure:"replica_set"`
|
||||
ReadPreference string `mapstructure:"read_preference"` // primary, secondary, etc.
|
||||
|
||||
// Connection pool settings (overrides global defaults)
|
||||
MaxOpenConns *int `mapstructure:"max_open_conns"`
|
||||
MaxIdleConns *int `mapstructure:"max_idle_conns"`
|
||||
ConnMaxLifetime *time.Duration `mapstructure:"conn_max_lifetime"`
|
||||
ConnMaxIdleTime *time.Duration `mapstructure:"conn_max_idle_time"`
|
||||
|
||||
// Timeouts
|
||||
ConnectTimeout time.Duration `mapstructure:"connect_timeout"`
|
||||
QueryTimeout time.Duration `mapstructure:"query_timeout"`
|
||||
|
||||
// Retry policy for the initial connect (inherited from the manager config)
|
||||
RetryAttempts int `mapstructure:"retry_attempts"`
|
||||
RetryDelay time.Duration `mapstructure:"retry_delay"`
|
||||
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
|
||||
|
||||
// Features
|
||||
EnableTracing bool `mapstructure:"enable_tracing"`
|
||||
EnableMetrics bool `mapstructure:"enable_metrics"`
|
||||
EnableLogging bool `mapstructure:"enable_logging"`
|
||||
|
||||
// DefaultORM specifies which ORM to use for the Database() method
|
||||
// Options: "bun", "gorm", "native"
|
||||
DefaultORM string `mapstructure:"default_orm"`
|
||||
|
||||
// Tags for organization and filtering
|
||||
Tags map[string]string `mapstructure:"tags"`
|
||||
}
|
||||
|
||||
// DefaultManagerConfig returns a ManagerConfig with sensible defaults
|
||||
func DefaultManagerConfig() ManagerConfig {
|
||||
return ManagerConfig{
|
||||
DefaultConnection: "",
|
||||
Connections: make(map[string]ConnectionConfig),
|
||||
MaxOpenConns: 25,
|
||||
MaxIdleConns: 5,
|
||||
ConnMaxLifetime: 30 * time.Minute,
|
||||
ConnMaxIdleTime: 5 * time.Minute,
|
||||
RetryAttempts: 3,
|
||||
RetryDelay: 1 * time.Second,
|
||||
RetryMaxDelay: 10 * time.Second,
|
||||
HealthCheckInterval: 15 * time.Second,
|
||||
}
|
||||
}
|
||||
|
||||
// ApplyDefaults applies default values to the manager configuration
|
||||
func (c *ManagerConfig) ApplyDefaults() {
|
||||
defaults := DefaultManagerConfig()
|
||||
|
||||
if c.MaxOpenConns == 0 {
|
||||
c.MaxOpenConns = defaults.MaxOpenConns
|
||||
}
|
||||
if c.MaxIdleConns == 0 {
|
||||
c.MaxIdleConns = defaults.MaxIdleConns
|
||||
}
|
||||
if c.ConnMaxLifetime == 0 {
|
||||
c.ConnMaxLifetime = defaults.ConnMaxLifetime
|
||||
}
|
||||
if c.ConnMaxIdleTime == 0 {
|
||||
c.ConnMaxIdleTime = defaults.ConnMaxIdleTime
|
||||
}
|
||||
if c.RetryAttempts == 0 {
|
||||
c.RetryAttempts = defaults.RetryAttempts
|
||||
}
|
||||
if c.RetryDelay == 0 {
|
||||
c.RetryDelay = defaults.RetryDelay
|
||||
}
|
||||
if c.RetryMaxDelay == 0 {
|
||||
c.RetryMaxDelay = defaults.RetryMaxDelay
|
||||
}
|
||||
if c.HealthCheckInterval == 0 {
|
||||
c.HealthCheckInterval = defaults.HealthCheckInterval
|
||||
}
|
||||
}
|
||||
|
||||
// Validate validates the manager configuration
|
||||
func (c *ManagerConfig) Validate() error {
|
||||
if len(c.Connections) == 0 {
|
||||
return NewConfigurationError("connections", fmt.Errorf("at least one connection must be configured"))
|
||||
}
|
||||
|
||||
if c.DefaultConnection != "" {
|
||||
if _, ok := c.Connections[c.DefaultConnection]; !ok {
|
||||
return NewConfigurationError("default_connection", fmt.Errorf("default connection '%s' not found in connections", c.DefaultConnection))
|
||||
}
|
||||
}
|
||||
|
||||
// Validate each connection
|
||||
for name := range c.Connections {
|
||||
conn := c.Connections[name]
|
||||
if err := conn.Validate(); err != nil {
|
||||
return fmt.Errorf("connection '%s': %w", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ApplyDefaults applies default values and global settings to the connection configuration
|
||||
func (cc *ConnectionConfig) ApplyDefaults(global *ManagerConfig) {
|
||||
// Set name if not already set
|
||||
if cc.Name == "" {
|
||||
cc.Name = "unnamed"
|
||||
}
|
||||
|
||||
// Apply global pool settings if not overridden
|
||||
if cc.MaxOpenConns == nil && global != nil {
|
||||
maxOpen := global.MaxOpenConns
|
||||
cc.MaxOpenConns = &maxOpen
|
||||
}
|
||||
if cc.MaxIdleConns == nil && global != nil {
|
||||
maxIdle := global.MaxIdleConns
|
||||
cc.MaxIdleConns = &maxIdle
|
||||
}
|
||||
if cc.ConnMaxLifetime == nil && global != nil {
|
||||
lifetime := global.ConnMaxLifetime
|
||||
cc.ConnMaxLifetime = &lifetime
|
||||
}
|
||||
if cc.ConnMaxIdleTime == nil && global != nil {
|
||||
idleTime := global.ConnMaxIdleTime
|
||||
cc.ConnMaxIdleTime = &idleTime
|
||||
}
|
||||
|
||||
// Default timeouts
|
||||
if cc.ConnectTimeout == 0 {
|
||||
cc.ConnectTimeout = 10 * time.Second
|
||||
}
|
||||
if cc.QueryTimeout == 0 {
|
||||
cc.QueryTimeout = 2 * time.Minute // Default to 2 minutes
|
||||
}
|
||||
|
||||
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
|
||||
if cc.DefaultORM == "" {
|
||||
cc.DefaultORM = string(ORMTypeBun)
|
||||
}
|
||||
|
||||
// Default PostgreSQL port
|
||||
if cc.Type == DatabaseTypePostgreSQL && cc.Port == 0 && cc.DSN == "" {
|
||||
cc.Port = 5432
|
||||
}
|
||||
|
||||
// Default MSSQL port
|
||||
if cc.Type == DatabaseTypeMSSQL && cc.Port == 0 && cc.DSN == "" {
|
||||
cc.Port = 1433
|
||||
}
|
||||
|
||||
// Default MongoDB port
|
||||
if cc.Type == DatabaseTypeMongoDB && cc.Port == 0 && cc.DSN == "" {
|
||||
cc.Port = 27017
|
||||
}
|
||||
|
||||
// Default MongoDB auth source
|
||||
if cc.Type == DatabaseTypeMongoDB && cc.AuthSource == "" {
|
||||
cc.AuthSource = "admin"
|
||||
}
|
||||
}
|
||||
|
||||
// Validate validates the connection configuration
|
||||
func (cc *ConnectionConfig) Validate() error {
|
||||
// Validate database type
|
||||
switch cc.Type {
|
||||
case DatabaseTypePostgreSQL, DatabaseTypeSQLite, DatabaseTypeMSSQL, DatabaseTypeMongoDB:
|
||||
// Valid types
|
||||
default:
|
||||
return NewConfigurationError("type", fmt.Errorf("unsupported database type: %s", cc.Type))
|
||||
}
|
||||
|
||||
// Validate that either DSN or connection parameters are provided
|
||||
if cc.DSN == "" {
|
||||
switch cc.Type {
|
||||
case DatabaseTypePostgreSQL, DatabaseTypeMSSQL, DatabaseTypeMongoDB:
|
||||
if cc.Host == "" {
|
||||
return NewConfigurationError("host", fmt.Errorf("host is required when DSN is not provided"))
|
||||
}
|
||||
if cc.Database == "" {
|
||||
return NewConfigurationError("database", fmt.Errorf("database is required when DSN is not provided"))
|
||||
}
|
||||
case DatabaseTypeSQLite:
|
||||
if cc.FilePath == "" {
|
||||
return NewConfigurationError("filepath", fmt.Errorf("filepath is required for SQLite when DSN is not provided"))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Validate ORM type
|
||||
if cc.DefaultORM != "" {
|
||||
switch ORMType(cc.DefaultORM) {
|
||||
case ORMTypeBun, ORMTypeGORM, ORMTypeNative:
|
||||
// Valid ORM types
|
||||
default:
|
||||
return NewConfigurationError("default_orm", fmt.Errorf("unsupported ORM type: %s", cc.DefaultORM))
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// BuildDSN builds a connection string from individual parameters
|
||||
func (cc *ConnectionConfig) BuildDSN() (string, error) {
|
||||
// If DSN is already provided, use it
|
||||
if cc.DSN != "" {
|
||||
return cc.DSN, nil
|
||||
}
|
||||
|
||||
switch cc.Type {
|
||||
case DatabaseTypePostgreSQL:
|
||||
return cc.buildPostgresDSN(), nil
|
||||
case DatabaseTypeSQLite:
|
||||
return cc.buildSQLiteDSN(), nil
|
||||
case DatabaseTypeMSSQL:
|
||||
return cc.buildMSSQLDSN(), nil
|
||||
case DatabaseTypeMongoDB:
|
||||
return cc.buildMongoDSN(), nil
|
||||
default:
|
||||
return "", fmt.Errorf("cannot build DSN for database type: %s", cc.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// 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 {
|
||||
q := url.Values{}
|
||||
if cc.SSLMode != "" {
|
||||
q.Set("sslmode", cc.SSLMode)
|
||||
} else {
|
||||
// prefer: use TLS when the server offers it, without failing on
|
||||
// servers that do not.
|
||||
q.Set("sslmode", "prefer")
|
||||
}
|
||||
if cc.Schema != "" {
|
||||
q.Set("search_path", cc.Schema)
|
||||
}
|
||||
|
||||
u := url.URL{
|
||||
Scheme: "postgres",
|
||||
Host: hostPort(cc.Host, cc.Port),
|
||||
Path: "/" + cc.Database,
|
||||
RawQuery: q.Encode(),
|
||||
}
|
||||
if cc.User != "" || cc.Password != "" {
|
||||
u.User = url.UserPassword(cc.User, cc.Password)
|
||||
}
|
||||
return u.String()
|
||||
}
|
||||
|
||||
func hostPort(host string, port int) string {
|
||||
if port == 0 {
|
||||
return host
|
||||
}
|
||||
// JoinHostPort brackets IPv6 literals.
|
||||
return net.JoinHostPort(host, strconv.Itoa(port))
|
||||
}
|
||||
|
||||
// buildSQLiteDSN puts per-connection settings in the DSN as _pragma parameters
|
||||
// so every pooled connection gets them, not just the one that ran an Exec.
|
||||
func (cc *ConnectionConfig) buildSQLiteDSN() string {
|
||||
filepath := cc.FilePath
|
||||
if filepath == "" {
|
||||
filepath = ":memory:"
|
||||
}
|
||||
|
||||
var pragmas []string
|
||||
if cc.QueryTimeout > 0 {
|
||||
pragmas = append(pragmas, fmt.Sprintf("busy_timeout(%d)", cc.QueryTimeout.Milliseconds()))
|
||||
}
|
||||
if filepath != ":memory:" {
|
||||
pragmas = append(pragmas, "journal_mode(WAL)")
|
||||
}
|
||||
if len(pragmas) == 0 {
|
||||
return filepath
|
||||
}
|
||||
|
||||
q := url.Values{}
|
||||
for _, p := range pragmas {
|
||||
q.Add("_pragma", p)
|
||||
}
|
||||
sep := "?"
|
||||
if strings.Contains(filepath, "?") {
|
||||
sep = "&"
|
||||
}
|
||||
return filepath + sep + q.Encode()
|
||||
}
|
||||
|
||||
func (cc *ConnectionConfig) buildMSSQLDSN() string {
|
||||
// Format: sqlserver://username:password@host:port?database=dbname
|
||||
q := url.Values{}
|
||||
q.Set("database", cc.Database)
|
||||
if cc.Schema != "" {
|
||||
q.Set("schema", cc.Schema)
|
||||
}
|
||||
if cc.ConnectTimeout > 0 {
|
||||
sec := strconv.Itoa(int(cc.ConnectTimeout.Seconds()))
|
||||
q.Set("connection timeout", sec)
|
||||
q.Set("dial timeout", sec)
|
||||
}
|
||||
if cc.QueryTimeout > 0 {
|
||||
q.Set("read timeout", strconv.Itoa(int(cc.QueryTimeout.Seconds())))
|
||||
}
|
||||
|
||||
u := url.URL{
|
||||
Scheme: "sqlserver",
|
||||
Host: hostPort(cc.Host, cc.Port),
|
||||
RawQuery: q.Encode(),
|
||||
}
|
||||
if cc.User != "" || cc.Password != "" {
|
||||
u.User = url.UserPassword(cc.User, cc.Password)
|
||||
}
|
||||
return u.String()
|
||||
}
|
||||
|
||||
func (cc *ConnectionConfig) buildMongoDSN() string {
|
||||
// Format: mongodb://username:password@host:port/database?authSource=admin
|
||||
q := url.Values{}
|
||||
if cc.AuthSource != "" {
|
||||
q.Set("authSource", cc.AuthSource)
|
||||
}
|
||||
if cc.ReplicaSet != "" {
|
||||
q.Set("replicaSet", cc.ReplicaSet)
|
||||
}
|
||||
if cc.ReadPreference != "" {
|
||||
q.Set("readPreference", cc.ReadPreference)
|
||||
}
|
||||
|
||||
u := url.URL{
|
||||
Scheme: "mongodb",
|
||||
Host: hostPort(cc.Host, cc.Port),
|
||||
Path: "/" + cc.Database,
|
||||
RawQuery: q.Encode(),
|
||||
}
|
||||
if cc.User != "" && cc.Password != "" {
|
||||
u.User = url.UserPassword(cc.User, cc.Password)
|
||||
}
|
||||
return u.String()
|
||||
}
|
||||
|
||||
// FromConfig converts config.DBManagerConfig to internal ManagerConfig
|
||||
func FromConfig(cfg config.DBManagerConfig) ManagerConfig {
|
||||
mgr := ManagerConfig{
|
||||
DefaultConnection: cfg.DefaultConnection,
|
||||
Connections: make(map[string]ConnectionConfig),
|
||||
MaxOpenConns: cfg.MaxOpenConns,
|
||||
MaxIdleConns: cfg.MaxIdleConns,
|
||||
ConnMaxLifetime: cfg.ConnMaxLifetime,
|
||||
ConnMaxIdleTime: cfg.ConnMaxIdleTime,
|
||||
RetryAttempts: cfg.RetryAttempts,
|
||||
RetryDelay: cfg.RetryDelay,
|
||||
RetryMaxDelay: cfg.RetryMaxDelay,
|
||||
HealthCheckInterval: cfg.HealthCheckInterval,
|
||||
EnableAutoReconnect: cfg.EnableAutoReconnect,
|
||||
}
|
||||
|
||||
// Convert connections
|
||||
for name := range cfg.Connections {
|
||||
connCfg := cfg.Connections[name]
|
||||
mgr.Connections[name] = ConnectionConfig{
|
||||
Name: connCfg.Name,
|
||||
Type: DatabaseType(connCfg.Type),
|
||||
DSN: connCfg.DSN,
|
||||
Host: connCfg.Host,
|
||||
Port: connCfg.Port,
|
||||
User: connCfg.User,
|
||||
Password: connCfg.Password,
|
||||
Database: connCfg.Database,
|
||||
SSLMode: connCfg.SSLMode,
|
||||
Schema: connCfg.Schema,
|
||||
FilePath: connCfg.FilePath,
|
||||
AuthSource: connCfg.AuthSource,
|
||||
ReplicaSet: connCfg.ReplicaSet,
|
||||
ReadPreference: connCfg.ReadPreference,
|
||||
MaxOpenConns: connCfg.MaxOpenConns,
|
||||
MaxIdleConns: connCfg.MaxIdleConns,
|
||||
ConnMaxLifetime: connCfg.ConnMaxLifetime,
|
||||
ConnMaxIdleTime: connCfg.ConnMaxIdleTime,
|
||||
ConnectTimeout: connCfg.ConnectTimeout,
|
||||
QueryTimeout: connCfg.QueryTimeout,
|
||||
EnableTracing: connCfg.EnableTracing,
|
||||
EnableMetrics: connCfg.EnableMetrics,
|
||||
EnableLogging: connCfg.EnableLogging,
|
||||
DefaultORM: connCfg.DefaultORM,
|
||||
Tags: connCfg.Tags,
|
||||
}
|
||||
}
|
||||
|
||||
return mgr
|
||||
}
|
||||
|
||||
// Getter methods to implement providers.ConnectionConfig interface
|
||||
func (cc *ConnectionConfig) GetName() string { return cc.Name }
|
||||
func (cc *ConnectionConfig) GetType() string { return string(cc.Type) }
|
||||
func (cc *ConnectionConfig) GetHost() string { return cc.Host }
|
||||
func (cc *ConnectionConfig) GetPort() int { return cc.Port }
|
||||
func (cc *ConnectionConfig) GetUser() string { return cc.User }
|
||||
func (cc *ConnectionConfig) GetPassword() string { return cc.Password }
|
||||
func (cc *ConnectionConfig) GetDatabase() string { return cc.Database }
|
||||
func (cc *ConnectionConfig) GetFilePath() string { return cc.FilePath }
|
||||
func (cc *ConnectionConfig) GetConnectTimeout() time.Duration { return cc.ConnectTimeout }
|
||||
func (cc *ConnectionConfig) GetEnableLogging() bool { return cc.EnableLogging }
|
||||
func (cc *ConnectionConfig) GetMaxOpenConns() *int { return cc.MaxOpenConns }
|
||||
func (cc *ConnectionConfig) GetMaxIdleConns() *int { return cc.MaxIdleConns }
|
||||
func (cc *ConnectionConfig) GetConnMaxLifetime() *time.Duration { return cc.ConnMaxLifetime }
|
||||
func (cc *ConnectionConfig) GetConnMaxIdleTime() *time.Duration { return cc.ConnMaxIdleTime }
|
||||
func (cc *ConnectionConfig) GetQueryTimeout() time.Duration { return cc.QueryTimeout }
|
||||
func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics }
|
||||
func (cc *ConnectionConfig) GetReadPreference() string { return cc.ReadPreference }
|
||||
func (cc *ConnectionConfig) GetRetryAttempts() int { return cc.RetryAttempts }
|
||||
func (cc *ConnectionConfig) GetRetryDelay() time.Duration { return cc.RetryDelay }
|
||||
func (cc *ConnectionConfig) GetRetryMaxDelay() time.Duration { return cc.RetryMaxDelay }
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,797 @@
|
||||
package dbmanager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/uptrace/bun"
|
||||
"github.com/uptrace/bun/schema"
|
||||
"go.mongodb.org/mongo-driver/mongo"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
|
||||
)
|
||||
|
||||
// Connection represents a single named database connection
|
||||
type Connection interface {
|
||||
// Metadata
|
||||
Name() string
|
||||
Type() DatabaseType
|
||||
|
||||
// ORM Access (SQL databases only)
|
||||
Bun() (*bun.DB, error)
|
||||
GORM() (*gorm.DB, error)
|
||||
Native() (*sql.DB, error)
|
||||
DB() (*sql.DB, error)
|
||||
|
||||
// Common Database interface (for SQL databases)
|
||||
Database() (common.Database, error)
|
||||
|
||||
// MongoDB Access (MongoDB only)
|
||||
MongoDB() (*mongo.Client, error)
|
||||
|
||||
// Lifecycle
|
||||
Connect(ctx context.Context) error
|
||||
Close() error
|
||||
HealthCheck(ctx context.Context) error
|
||||
Reconnect(ctx context.Context) error
|
||||
|
||||
// Stats
|
||||
Stats() *ConnectionStats
|
||||
}
|
||||
|
||||
// ConnectionStats contains statistics about a database connection
|
||||
type ConnectionStats struct {
|
||||
Name string
|
||||
Type DatabaseType
|
||||
Connected bool
|
||||
LastHealthCheck time.Time
|
||||
HealthCheckStatus string
|
||||
|
||||
// SQL connection pool stats
|
||||
OpenConnections int
|
||||
InUse int
|
||||
Idle int
|
||||
WaitCount int64
|
||||
WaitDuration time.Duration
|
||||
MaxIdleClosed int64
|
||||
MaxLifetimeClosed int64
|
||||
}
|
||||
|
||||
// sqlConnection implements Connection for SQL databases (PostgreSQL, SQLite, MSSQL)
|
||||
type sqlConnection struct {
|
||||
name string
|
||||
dbType DatabaseType
|
||||
config ConnectionConfig
|
||||
provider Provider
|
||||
|
||||
// Lazy-initialized ORM instances (all wrap the same sql.DB)
|
||||
nativeDB *sql.DB
|
||||
bunDB *bun.DB
|
||||
gormDB *gorm.DB
|
||||
|
||||
// Adapters for common.Database interface
|
||||
bunAdapter *database.BunAdapter
|
||||
gormAdapter *database.GormAdapter
|
||||
nativeAdapter common.Database
|
||||
|
||||
// State
|
||||
connected bool
|
||||
mu sync.RWMutex
|
||||
// lifecycleMu serialises Connect/Close/Reconnect against health-check pings.
|
||||
// Lock order: lifecycleMu before mu.
|
||||
lifecycleMu sync.RWMutex
|
||||
|
||||
// Health check
|
||||
lastHealthCheck time.Time
|
||||
healthCheckStatus string
|
||||
}
|
||||
|
||||
// newSQLConnection creates a new SQL connection
|
||||
func newSQLConnection(name string, dbType DatabaseType, config ConnectionConfig, provider Provider) *sqlConnection {
|
||||
return &sqlConnection{
|
||||
name: name,
|
||||
dbType: dbType,
|
||||
config: config,
|
||||
provider: provider,
|
||||
}
|
||||
}
|
||||
|
||||
// Name returns the connection name
|
||||
func (c *sqlConnection) Name() string {
|
||||
return c.name
|
||||
}
|
||||
|
||||
// Type returns the database type
|
||||
func (c *sqlConnection) Type() DatabaseType {
|
||||
return c.dbType
|
||||
}
|
||||
|
||||
// Connect establishes the database connection
|
||||
func (c *sqlConnection) Connect(ctx context.Context) error {
|
||||
c.lifecycleMu.Lock()
|
||||
defer c.lifecycleMu.Unlock()
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
return c.connectLocked(ctx)
|
||||
}
|
||||
|
||||
// connectLocked requires lifecycleMu and mu held for writing.
|
||||
func (c *sqlConnection) connectLocked(ctx context.Context) error {
|
||||
if c.connected {
|
||||
return ErrAlreadyConnected
|
||||
}
|
||||
|
||||
if err := c.provider.Connect(ctx, &c.config); err != nil {
|
||||
return NewConnectionError(c.name, "connect", err)
|
||||
}
|
||||
|
||||
c.connected = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the database connection and all ORM instances
|
||||
func (c *sqlConnection) Close() error {
|
||||
c.lifecycleMu.Lock()
|
||||
defer c.lifecycleMu.Unlock()
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
return c.closeLocked()
|
||||
}
|
||||
|
||||
// closeLocked requires lifecycleMu and mu held for writing. The connection is
|
||||
// always marked disconnected and its cached handles dropped, even when closing
|
||||
// fails, so accessors never hand out handles over a half-closed pool.
|
||||
func (c *sqlConnection) closeLocked() error {
|
||||
if !c.connected {
|
||||
return nil
|
||||
}
|
||||
|
||||
var errs []error
|
||||
|
||||
// Close Bun if initialized. bun.DB.Close closes the underlying *sql.DB, so
|
||||
// skip it when the pool belongs to the caller.
|
||||
if o, ok := c.provider.(interface{ OwnsDB() bool }); c.bunDB != nil && (!ok || o.OwnsDB()) {
|
||||
if err := c.bunDB.Close(); err != nil {
|
||||
errs = append(errs, NewConnectionError(c.name, "close bun", err))
|
||||
}
|
||||
}
|
||||
|
||||
// GORM doesn't have a separate close - it uses the underlying sql.DB
|
||||
|
||||
// Close the provider (which closes the underlying sql.DB)
|
||||
if err := c.provider.Close(); err != nil {
|
||||
errs = append(errs, NewConnectionError(c.name, "close", err))
|
||||
}
|
||||
|
||||
c.connected = false
|
||||
c.nativeDB = nil
|
||||
c.bunDB = nil
|
||||
c.gormDB = nil
|
||||
c.bunAdapter = nil
|
||||
c.gormAdapter = nil
|
||||
c.nativeAdapter = nil
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
// HealthCheck verifies the connection is alive. The network ping runs without
|
||||
// holding mu, so handle accessors are never blocked behind a slow ping.
|
||||
func (c *sqlConnection) HealthCheck(ctx context.Context) error {
|
||||
if c == nil {
|
||||
return fmt.Errorf("connection is nil")
|
||||
}
|
||||
|
||||
// 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()
|
||||
|
||||
c.mu.RLock()
|
||||
connected := c.connected
|
||||
provider := c.provider
|
||||
c.mu.RUnlock()
|
||||
|
||||
if !connected {
|
||||
c.setHealth("disconnected")
|
||||
return ErrConnectionClosed
|
||||
}
|
||||
|
||||
if err := provider.HealthCheck(ctx); err != nil {
|
||||
c.setHealth("unhealthy: " + err.Error())
|
||||
return NewConnectionError(c.name, "health check", err)
|
||||
}
|
||||
|
||||
c.setHealth("healthy")
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *sqlConnection) setHealth(status string) {
|
||||
c.mu.Lock()
|
||||
c.lastHealthCheck = time.Now()
|
||||
c.healthCheckStatus = status
|
||||
c.mu.Unlock()
|
||||
}
|
||||
|
||||
// Reconnect refreshes the connection as a single critical section.
|
||||
//
|
||||
// Providers that support it (PostgreSQL) retire their pooled connections and
|
||||
// dial fresh ones without closing the *sql.DB, so handles handed out earlier
|
||||
// keep working. Other providers fall back to Close+Connect, which invalidates
|
||||
// earlier handles; that is meant for explicit operator use only, since
|
||||
// *sql.DB already replaces broken connections by itself.
|
||||
func (c *sqlConnection) Reconnect(ctx context.Context) (err error) {
|
||||
c.lifecycleMu.Lock()
|
||||
defer c.lifecycleMu.Unlock()
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
defer func() { RecordReconnectAttempt(c.name, c.dbType, err == nil) }()
|
||||
|
||||
if c.connected {
|
||||
if r, ok := c.provider.(providers.Refresher); ok {
|
||||
if err := r.Refresh(ctx); err != nil {
|
||||
return NewConnectionError(c.name, "reconnect", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
if err := c.closeLocked(); err != nil {
|
||||
return err
|
||||
}
|
||||
return c.connectLocked(ctx)
|
||||
}
|
||||
|
||||
// Native returns the native *sql.DB connection
|
||||
func (c *sqlConnection) Native() (*sql.DB, error) {
|
||||
if c == nil {
|
||||
return nil, fmt.Errorf("connection is nil")
|
||||
}
|
||||
c.mu.RLock()
|
||||
if c.nativeDB != nil {
|
||||
defer c.mu.RUnlock()
|
||||
return c.nativeDB, nil
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
// Double-check after acquiring write lock
|
||||
if c.nativeDB != nil {
|
||||
return c.nativeDB, nil
|
||||
}
|
||||
|
||||
if !c.connected {
|
||||
return nil, ErrConnectionClosed
|
||||
}
|
||||
|
||||
// Get native connection from provider
|
||||
db, err := c.provider.GetNative()
|
||||
if err != nil {
|
||||
return nil, NewConnectionError(c.name, "get native", err)
|
||||
}
|
||||
|
||||
c.nativeDB = db
|
||||
return c.nativeDB, nil
|
||||
}
|
||||
|
||||
// DB returns the underlying *sql.DB connection
|
||||
func (c *sqlConnection) DB() (*sql.DB, error) {
|
||||
return c.Native()
|
||||
}
|
||||
|
||||
// Bun returns a Bun ORM instance wrapping the native connection
|
||||
func (c *sqlConnection) Bun() (*bun.DB, error) {
|
||||
if c == nil {
|
||||
return nil, fmt.Errorf("connection is nil")
|
||||
}
|
||||
c.mu.RLock()
|
||||
if c.bunDB != nil {
|
||||
defer c.mu.RUnlock()
|
||||
return c.bunDB, nil
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
// Double-check after acquiring write lock
|
||||
if c.bunDB != nil {
|
||||
return c.bunDB, nil
|
||||
}
|
||||
|
||||
if !c.connected {
|
||||
return nil, ErrConnectionClosed
|
||||
}
|
||||
|
||||
// Get native connection first
|
||||
native, err := c.provider.GetNative()
|
||||
if err != nil {
|
||||
return nil, NewConnectionError(c.name, "get bun", err)
|
||||
}
|
||||
|
||||
// Create Bun DB wrapping the same sql.DB
|
||||
dialect := c.getBunDialect()
|
||||
c.bunDB = bun.NewDB(native, dialect)
|
||||
|
||||
return c.bunDB, nil
|
||||
}
|
||||
|
||||
// GORM returns a GORM instance wrapping the native connection
|
||||
func (c *sqlConnection) GORM() (*gorm.DB, error) {
|
||||
if c == nil {
|
||||
return nil, fmt.Errorf("connection is nil")
|
||||
}
|
||||
c.mu.RLock()
|
||||
if c.gormDB != nil {
|
||||
defer c.mu.RUnlock()
|
||||
return c.gormDB, nil
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
// Double-check after acquiring write lock
|
||||
if c.gormDB != nil {
|
||||
return c.gormDB, nil
|
||||
}
|
||||
|
||||
if !c.connected {
|
||||
return nil, ErrConnectionClosed
|
||||
}
|
||||
|
||||
// Get native connection first
|
||||
native, err := c.provider.GetNative()
|
||||
if err != nil {
|
||||
return nil, NewConnectionError(c.name, "get gorm", err)
|
||||
}
|
||||
|
||||
// Create GORM DB wrapping the same sql.DB
|
||||
dialector := c.getGORMDialector(native)
|
||||
db, err := gorm.Open(dialector, &gorm.Config{})
|
||||
if err != nil {
|
||||
return nil, NewConnectionError(c.name, "initialize gorm", err)
|
||||
}
|
||||
|
||||
c.gormDB = db
|
||||
return c.gormDB, nil
|
||||
}
|
||||
|
||||
// Database returns the common.Database interface using the configured default ORM
|
||||
func (c *sqlConnection) Database() (common.Database, error) {
|
||||
if c == nil {
|
||||
return nil, fmt.Errorf("connection is nil")
|
||||
}
|
||||
c.mu.RLock()
|
||||
defaultORM := c.config.DefaultORM
|
||||
c.mu.RUnlock()
|
||||
|
||||
switch ORMType(defaultORM) {
|
||||
case ORMTypeBun:
|
||||
return c.getBunAdapter()
|
||||
case ORMTypeGORM:
|
||||
return c.getGORMAdapter()
|
||||
case ORMTypeNative:
|
||||
return c.getNativeAdapter()
|
||||
default:
|
||||
// Default to Bun
|
||||
return c.getBunAdapter()
|
||||
}
|
||||
}
|
||||
|
||||
// MongoDB returns an error for SQL connections
|
||||
func (c *sqlConnection) MongoDB() (*mongo.Client, error) {
|
||||
return nil, ErrNotMongoDB
|
||||
}
|
||||
|
||||
// Stats returns connection statistics
|
||||
func (c *sqlConnection) Stats() *ConnectionStats {
|
||||
if c == nil {
|
||||
return nil
|
||||
}
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
stats := &ConnectionStats{
|
||||
Name: c.name,
|
||||
Type: c.dbType,
|
||||
Connected: c.connected,
|
||||
LastHealthCheck: c.lastHealthCheck,
|
||||
HealthCheckStatus: c.healthCheckStatus,
|
||||
}
|
||||
|
||||
// Get SQL stats if connected
|
||||
if c.connected && c.provider != nil {
|
||||
if providerStats := c.provider.Stats(); providerStats != nil {
|
||||
stats.OpenConnections = providerStats.OpenConnections
|
||||
stats.InUse = providerStats.InUse
|
||||
stats.Idle = providerStats.Idle
|
||||
stats.WaitCount = providerStats.WaitCount
|
||||
stats.WaitDuration = providerStats.WaitDuration
|
||||
stats.MaxIdleClosed = providerStats.MaxIdleClosed
|
||||
stats.MaxLifetimeClosed = providerStats.MaxLifetimeClosed
|
||||
}
|
||||
}
|
||||
|
||||
return stats
|
||||
}
|
||||
|
||||
// The adapter factories only re-fetch the current handle. They must not close
|
||||
// the shared pool: *sql.DB discards bad connections on its own, and closing it
|
||||
// here would break every other holder of the pool.
|
||||
func (c *sqlConnection) reopenNativeForAdapter() (*sql.DB, error) {
|
||||
return c.Native()
|
||||
}
|
||||
|
||||
func (c *sqlConnection) reopenBunForAdapter() (*bun.DB, error) {
|
||||
return c.Bun()
|
||||
}
|
||||
|
||||
func (c *sqlConnection) reopenGORMForAdapter() (*gorm.DB, error) {
|
||||
return c.GORM()
|
||||
}
|
||||
|
||||
// getBunAdapter returns or creates the Bun adapter
|
||||
func (c *sqlConnection) getBunAdapter() (common.Database, error) {
|
||||
if c == nil {
|
||||
return nil, fmt.Errorf("connection is nil")
|
||||
}
|
||||
c.mu.RLock()
|
||||
if c.bunAdapter != nil {
|
||||
defer c.mu.RUnlock()
|
||||
return c.bunAdapter, nil
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if c.bunAdapter != nil {
|
||||
return c.bunAdapter, nil
|
||||
}
|
||||
|
||||
// Double-check bunDB exists (while already holding write lock)
|
||||
if c.bunDB == nil {
|
||||
// Get native connection first
|
||||
native, err := c.provider.GetNative()
|
||||
if err != nil {
|
||||
return nil, NewConnectionError(c.name, "get bun", err)
|
||||
}
|
||||
|
||||
// Create Bun DB wrapping the same sql.DB
|
||||
dialect := c.getBunDialect()
|
||||
c.bunDB = bun.NewDB(native, dialect)
|
||||
}
|
||||
|
||||
c.bunAdapter = database.NewBunAdapter(c.bunDB).
|
||||
WithDBFactory(c.reopenBunForAdapter).
|
||||
SetMetricsEnabled(c.config.EnableMetrics)
|
||||
return c.bunAdapter, nil
|
||||
}
|
||||
|
||||
// getGORMAdapter returns or creates the GORM adapter
|
||||
func (c *sqlConnection) getGORMAdapter() (common.Database, error) {
|
||||
if c == nil {
|
||||
return nil, fmt.Errorf("connection is nil")
|
||||
}
|
||||
c.mu.RLock()
|
||||
if c.gormAdapter != nil {
|
||||
defer c.mu.RUnlock()
|
||||
return c.gormAdapter, nil
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if c.gormAdapter != nil {
|
||||
return c.gormAdapter, nil
|
||||
}
|
||||
|
||||
// Double-check gormDB exists (while already holding write lock)
|
||||
if c.gormDB == nil {
|
||||
// Get native connection first
|
||||
native, err := c.provider.GetNative()
|
||||
if err != nil {
|
||||
return nil, NewConnectionError(c.name, "get gorm", err)
|
||||
}
|
||||
|
||||
// Create GORM DB wrapping the same sql.DB
|
||||
dialector := c.getGORMDialector(native)
|
||||
db, err := gorm.Open(dialector, &gorm.Config{})
|
||||
if err != nil {
|
||||
return nil, NewConnectionError(c.name, "initialize gorm", err)
|
||||
}
|
||||
|
||||
c.gormDB = db
|
||||
}
|
||||
|
||||
c.gormAdapter = database.NewGormAdapter(c.gormDB).
|
||||
WithDBFactory(c.reopenGORMForAdapter).
|
||||
SetMetricsEnabled(c.config.EnableMetrics)
|
||||
return c.gormAdapter, nil
|
||||
}
|
||||
|
||||
// getNativeAdapter returns or creates the native adapter
|
||||
func (c *sqlConnection) getNativeAdapter() (common.Database, error) {
|
||||
if c == nil {
|
||||
return nil, fmt.Errorf("connection is nil")
|
||||
}
|
||||
c.mu.RLock()
|
||||
if c.nativeAdapter != nil {
|
||||
defer c.mu.RUnlock()
|
||||
return c.nativeAdapter, nil
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if c.nativeAdapter != nil {
|
||||
return c.nativeAdapter, nil
|
||||
}
|
||||
|
||||
// Double-check nativeDB exists (while already holding write lock)
|
||||
if c.nativeDB == nil {
|
||||
if !c.connected {
|
||||
return nil, ErrConnectionClosed
|
||||
}
|
||||
|
||||
// Get native connection from provider
|
||||
db, err := c.provider.GetNative()
|
||||
if err != nil {
|
||||
return nil, NewConnectionError(c.name, "get native", err)
|
||||
}
|
||||
|
||||
c.nativeDB = db
|
||||
}
|
||||
|
||||
// Create a native adapter based on database type
|
||||
switch c.dbType {
|
||||
case DatabaseTypePostgreSQL, DatabaseTypeSQLite, DatabaseTypeMSSQL:
|
||||
// The adapter takes the driver name so it can adjust its dialect.
|
||||
c.nativeAdapter = database.NewPgSQLAdapter(c.nativeDB, string(c.dbType)).
|
||||
WithDBFactory(c.reopenNativeForAdapter).
|
||||
SetMetricsEnabled(c.config.EnableMetrics)
|
||||
default:
|
||||
return nil, ErrUnsupportedDatabase
|
||||
}
|
||||
|
||||
return c.nativeAdapter, nil
|
||||
}
|
||||
|
||||
// getBunDialect returns the appropriate Bun dialect for the database type
|
||||
func (c *sqlConnection) getBunDialect() schema.Dialect {
|
||||
|
||||
switch c.dbType {
|
||||
case DatabaseTypePostgreSQL:
|
||||
return database.GetPostgresDialect()
|
||||
case DatabaseTypeSQLite:
|
||||
return database.GetSQLiteDialect()
|
||||
case DatabaseTypeMSSQL:
|
||||
return database.GetMSSQLDialect()
|
||||
default:
|
||||
// Default to PostgreSQL
|
||||
return database.GetPostgresDialect()
|
||||
}
|
||||
}
|
||||
|
||||
// getGORMDialector returns the appropriate GORM dialector for the database type
|
||||
func (c *sqlConnection) getGORMDialector(db *sql.DB) gorm.Dialector {
|
||||
switch c.dbType {
|
||||
case DatabaseTypePostgreSQL:
|
||||
return database.GetPostgresDialector(db)
|
||||
case DatabaseTypeSQLite:
|
||||
return database.GetSQLiteDialector(db)
|
||||
case DatabaseTypeMSSQL:
|
||||
return database.GetMSSQLDialector(db)
|
||||
default:
|
||||
// Default to PostgreSQL
|
||||
return database.GetPostgresDialector(db)
|
||||
}
|
||||
}
|
||||
|
||||
// mongoConnection implements Connection for MongoDB
|
||||
type mongoConnection struct {
|
||||
name string
|
||||
config ConnectionConfig
|
||||
provider Provider
|
||||
|
||||
// MongoDB client
|
||||
client *mongo.Client
|
||||
|
||||
// State
|
||||
connected bool
|
||||
mu sync.RWMutex
|
||||
lifecycleMu sync.RWMutex // see sqlConnection.lifecycleMu
|
||||
|
||||
// Health check
|
||||
lastHealthCheck time.Time
|
||||
healthCheckStatus string
|
||||
}
|
||||
|
||||
// newMongoConnection creates a new MongoDB connection
|
||||
func newMongoConnection(name string, config ConnectionConfig, provider Provider) *mongoConnection {
|
||||
return &mongoConnection{
|
||||
name: name,
|
||||
config: config,
|
||||
provider: provider,
|
||||
}
|
||||
}
|
||||
|
||||
// Name returns the connection name
|
||||
func (c *mongoConnection) Name() string {
|
||||
return c.name
|
||||
}
|
||||
|
||||
// Type returns the database type (MongoDB)
|
||||
func (c *mongoConnection) Type() DatabaseType {
|
||||
return DatabaseTypeMongoDB
|
||||
}
|
||||
|
||||
// Connect establishes the MongoDB connection
|
||||
func (c *mongoConnection) Connect(ctx context.Context) error {
|
||||
c.lifecycleMu.Lock()
|
||||
defer c.lifecycleMu.Unlock()
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
return c.connectLocked(ctx)
|
||||
}
|
||||
|
||||
// connectLocked requires lifecycleMu and mu held for writing.
|
||||
func (c *mongoConnection) connectLocked(ctx context.Context) error {
|
||||
if c.connected {
|
||||
return ErrAlreadyConnected
|
||||
}
|
||||
|
||||
if err := c.provider.Connect(ctx, &c.config); err != nil {
|
||||
return NewConnectionError(c.name, "connect", err)
|
||||
}
|
||||
|
||||
// Get the mongo client
|
||||
client, err := c.provider.GetMongo()
|
||||
if err != nil {
|
||||
_ = c.provider.Close()
|
||||
return NewConnectionError(c.name, "get mongo client", err)
|
||||
}
|
||||
|
||||
c.client = client
|
||||
c.connected = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the MongoDB connection
|
||||
func (c *mongoConnection) Close() error {
|
||||
c.lifecycleMu.Lock()
|
||||
defer c.lifecycleMu.Unlock()
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
return c.closeLocked()
|
||||
}
|
||||
|
||||
// closeLocked requires lifecycleMu and mu held for writing. The connection is
|
||||
// marked disconnected even when the provider fails to close cleanly.
|
||||
func (c *mongoConnection) closeLocked() error {
|
||||
if !c.connected {
|
||||
return nil
|
||||
}
|
||||
|
||||
err := c.provider.Close()
|
||||
|
||||
c.connected = false
|
||||
c.client = nil
|
||||
if err != nil {
|
||||
return NewConnectionError(c.name, "close", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// HealthCheck verifies the MongoDB connection is alive. The ping runs without
|
||||
// holding mu so handle accessors are never blocked behind it.
|
||||
func (c *mongoConnection) HealthCheck(ctx context.Context) error {
|
||||
c.lifecycleMu.RLock()
|
||||
defer c.lifecycleMu.RUnlock()
|
||||
|
||||
c.mu.RLock()
|
||||
connected := c.connected
|
||||
c.mu.RUnlock()
|
||||
|
||||
if !connected {
|
||||
c.setHealth("disconnected")
|
||||
return ErrConnectionClosed
|
||||
}
|
||||
|
||||
if err := c.provider.HealthCheck(ctx); err != nil {
|
||||
c.setHealth("unhealthy: " + err.Error())
|
||||
return NewConnectionError(c.name, "health check", err)
|
||||
}
|
||||
|
||||
c.setHealth("healthy")
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *mongoConnection) setHealth(status string) {
|
||||
c.mu.Lock()
|
||||
c.lastHealthCheck = time.Now()
|
||||
c.healthCheckStatus = status
|
||||
c.mu.Unlock()
|
||||
}
|
||||
|
||||
// Reconnect closes and re-establishes the MongoDB connection atomically.
|
||||
func (c *mongoConnection) Reconnect(ctx context.Context) (err error) {
|
||||
c.lifecycleMu.Lock()
|
||||
defer c.lifecycleMu.Unlock()
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
defer func() { RecordReconnectAttempt(c.name, DatabaseTypeMongoDB, err == nil) }()
|
||||
|
||||
if err := c.closeLocked(); err != nil {
|
||||
return err
|
||||
}
|
||||
return c.connectLocked(ctx)
|
||||
}
|
||||
|
||||
// MongoDB returns the MongoDB client
|
||||
func (c *mongoConnection) MongoDB() (*mongo.Client, error) {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
if !c.connected || c.client == nil {
|
||||
return nil, ErrConnectionClosed
|
||||
}
|
||||
|
||||
return c.client, nil
|
||||
}
|
||||
|
||||
// Bun returns an error for MongoDB connections
|
||||
func (c *mongoConnection) Bun() (*bun.DB, error) {
|
||||
return nil, ErrNotSQLDatabase
|
||||
}
|
||||
|
||||
// GORM returns an error for MongoDB connections
|
||||
func (c *mongoConnection) GORM() (*gorm.DB, error) {
|
||||
return nil, ErrNotSQLDatabase
|
||||
}
|
||||
|
||||
// Native returns an error for MongoDB connections
|
||||
func (c *mongoConnection) Native() (*sql.DB, error) {
|
||||
return nil, ErrNotSQLDatabase
|
||||
}
|
||||
|
||||
// DB returns an error for MongoDB connections
|
||||
func (c *mongoConnection) DB() (*sql.DB, error) {
|
||||
return nil, ErrNotSQLDatabase
|
||||
}
|
||||
|
||||
// Database returns an error for MongoDB connections
|
||||
func (c *mongoConnection) Database() (common.Database, error) {
|
||||
return nil, ErrNotSQLDatabase
|
||||
}
|
||||
|
||||
// Stats returns connection statistics for MongoDB
|
||||
func (c *mongoConnection) Stats() *ConnectionStats {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
return &ConnectionStats{
|
||||
Name: c.name,
|
||||
Type: DatabaseTypeMongoDB,
|
||||
Connected: c.connected,
|
||||
LastHealthCheck: c.lastHealthCheck,
|
||||
HealthCheckStatus: c.healthCheckStatus,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
package dbmanager
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// Common errors
|
||||
var (
|
||||
// ErrConnectionNotFound is returned when a connection with the given name doesn't exist
|
||||
ErrConnectionNotFound = errors.New("connection not found")
|
||||
|
||||
// ErrInvalidConfiguration is returned when the configuration is invalid
|
||||
ErrInvalidConfiguration = errors.New("invalid configuration")
|
||||
|
||||
// ErrConnectionClosed is returned when attempting to use a closed connection
|
||||
ErrConnectionClosed = errors.New("connection is closed")
|
||||
|
||||
// ErrNotSQLDatabase is returned when attempting SQL operations on a non-SQL database
|
||||
ErrNotSQLDatabase = errors.New("not a SQL database")
|
||||
|
||||
// ErrNotMongoDB is returned when attempting MongoDB operations on a non-MongoDB connection
|
||||
ErrNotMongoDB = errors.New("not a MongoDB connection")
|
||||
|
||||
// ErrUnsupportedDatabase is returned when the database type is not supported
|
||||
ErrUnsupportedDatabase = errors.New("unsupported database type")
|
||||
|
||||
// ErrNoDefaultConnection is returned when no default connection is configured
|
||||
ErrNoDefaultConnection = errors.New("no default connection configured")
|
||||
|
||||
// ErrAlreadyConnected is returned when attempting to connect an already connected connection
|
||||
ErrAlreadyConnected = errors.New("already connected")
|
||||
)
|
||||
|
||||
// ConnectionError wraps errors that occur during connection operations
|
||||
type ConnectionError struct {
|
||||
Name string
|
||||
Operation string
|
||||
Err error
|
||||
}
|
||||
|
||||
func (e *ConnectionError) Error() string {
|
||||
return fmt.Sprintf("connection '%s' %s: %v", e.Name, e.Operation, e.Err)
|
||||
}
|
||||
|
||||
func (e *ConnectionError) Unwrap() error {
|
||||
return e.Err
|
||||
}
|
||||
|
||||
// NewConnectionError creates a new ConnectionError
|
||||
func NewConnectionError(name, operation string, err error) *ConnectionError {
|
||||
return &ConnectionError{
|
||||
Name: name,
|
||||
Operation: operation,
|
||||
Err: err,
|
||||
}
|
||||
}
|
||||
|
||||
// ConfigurationError wraps configuration-related errors
|
||||
type ConfigurationError struct {
|
||||
Field string
|
||||
Err error
|
||||
}
|
||||
|
||||
func (e *ConfigurationError) Error() string {
|
||||
if e.Field != "" {
|
||||
return fmt.Sprintf("configuration error in field '%s': %v", e.Field, e.Err)
|
||||
}
|
||||
return fmt.Sprintf("configuration error: %v", e.Err)
|
||||
}
|
||||
|
||||
func (e *ConfigurationError) Unwrap() error {
|
||||
return e.Err
|
||||
}
|
||||
|
||||
// NewConfigurationError creates a new ConfigurationError
|
||||
func NewConfigurationError(field string, err error) *ConfigurationError {
|
||||
return &ConfigurationError{
|
||||
Field: field,
|
||||
Err: err,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
package dbmanager
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
|
||||
)
|
||||
|
||||
// createConnection creates a database connection based on the configuration
|
||||
func createConnection(cfg ConnectionConfig) (Connection, error) {
|
||||
// Validate configuration
|
||||
if err := cfg.Validate(); err != nil {
|
||||
return nil, fmt.Errorf("invalid connection configuration: %w", err)
|
||||
}
|
||||
|
||||
// Create provider based on database type
|
||||
provider, err := createProvider(cfg.Type)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Create connection wrapper based on database type
|
||||
switch cfg.Type {
|
||||
case DatabaseTypePostgreSQL, DatabaseTypeSQLite, DatabaseTypeMSSQL:
|
||||
return newSQLConnection(cfg.Name, cfg.Type, cfg, provider), nil
|
||||
case DatabaseTypeMongoDB:
|
||||
return newMongoConnection(cfg.Name, cfg, provider), nil
|
||||
default:
|
||||
return nil, fmt.Errorf("%w: %s", ErrUnsupportedDatabase, cfg.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// createProvider creates a database provider based on the database type
|
||||
func createProvider(dbType DatabaseType) (Provider, error) {
|
||||
switch dbType {
|
||||
case DatabaseTypePostgreSQL:
|
||||
return providers.NewPostgresProvider(), nil
|
||||
case DatabaseTypeSQLite:
|
||||
return providers.NewSQLiteProvider(), nil
|
||||
case DatabaseTypeMSSQL:
|
||||
return providers.NewMSSQLProvider(), nil
|
||||
case DatabaseTypeMongoDB:
|
||||
return providers.NewMongoProvider(), nil
|
||||
default:
|
||||
return nil, fmt.Errorf("%w: %s", ErrUnsupportedDatabase, dbType)
|
||||
}
|
||||
}
|
||||
|
||||
// Provider is an alias to the providers.Provider interface
|
||||
// This allows dbmanager package consumers to use Provider without importing providers
|
||||
type Provider = providers.Provider
|
||||
|
||||
// NewConnectionFromDB creates a new Connection from an existing *sql.DB
|
||||
// This allows you to use dbmanager features (ORM wrappers, health checks, etc.)
|
||||
// with a database connection that was opened outside of dbmanager
|
||||
//
|
||||
// Parameters:
|
||||
// - name: A unique name for this connection
|
||||
// - dbType: The database type (DatabaseTypePostgreSQL, DatabaseTypeSQLite, or DatabaseTypeMSSQL)
|
||||
// - db: An existing *sql.DB connection
|
||||
//
|
||||
// Returns a Connection that wraps the existing *sql.DB
|
||||
func NewConnectionFromDB(name string, dbType DatabaseType, db *sql.DB) Connection {
|
||||
provider := providers.NewExistingDBProvider(db, name)
|
||||
return newSQLConnection(name, dbType, ConnectionConfig{Name: name, Type: dbType}, provider)
|
||||
}
|
||||
@@ -0,0 +1,210 @@
|
||||
package dbmanager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"testing"
|
||||
|
||||
_ "github.com/mattn/go-sqlite3"
|
||||
)
|
||||
|
||||
func TestNewConnectionFromDB(t *testing.T) {
|
||||
// Open a SQLite in-memory database
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
// Create a connection from the existing database
|
||||
conn := NewConnectionFromDB("test-connection", DatabaseTypeSQLite, db)
|
||||
if conn == nil {
|
||||
t.Fatal("Expected connection to be created")
|
||||
}
|
||||
|
||||
// Verify connection properties
|
||||
if conn.Name() != "test-connection" {
|
||||
t.Errorf("Expected name 'test-connection', got '%s'", conn.Name())
|
||||
}
|
||||
|
||||
if conn.Type() != DatabaseTypeSQLite {
|
||||
t.Errorf("Expected type DatabaseTypeSQLite, got '%s'", conn.Type())
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewConnectionFromDB_Connect(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
conn := NewConnectionFromDB("test-connection", DatabaseTypeSQLite, db)
|
||||
ctx := context.Background()
|
||||
|
||||
// Connect should verify the existing connection works
|
||||
err = conn.Connect(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("Expected Connect to succeed, got error: %v", err)
|
||||
}
|
||||
|
||||
// Cleanup
|
||||
conn.Close()
|
||||
}
|
||||
|
||||
func TestNewConnectionFromDB_Native(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
conn := NewConnectionFromDB("test-connection", DatabaseTypeSQLite, db)
|
||||
ctx := context.Background()
|
||||
|
||||
err = conn.Connect(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// Get native DB
|
||||
nativeDB, err := conn.Native()
|
||||
if err != nil {
|
||||
t.Errorf("Expected Native to succeed, got error: %v", err)
|
||||
}
|
||||
|
||||
if nativeDB != db {
|
||||
t.Error("Expected Native to return the same database instance")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewConnectionFromDB_Bun(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
conn := NewConnectionFromDB("test-connection", DatabaseTypeSQLite, db)
|
||||
ctx := context.Background()
|
||||
|
||||
err = conn.Connect(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// Get Bun ORM
|
||||
bunDB, err := conn.Bun()
|
||||
if err != nil {
|
||||
t.Errorf("Expected Bun to succeed, got error: %v", err)
|
||||
}
|
||||
|
||||
if bunDB == nil {
|
||||
t.Error("Expected Bun to return a non-nil instance")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewConnectionFromDB_GORM(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
conn := NewConnectionFromDB("test-connection", DatabaseTypeSQLite, db)
|
||||
ctx := context.Background()
|
||||
|
||||
err = conn.Connect(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// Get GORM
|
||||
gormDB, err := conn.GORM()
|
||||
if err != nil {
|
||||
t.Errorf("Expected GORM to succeed, got error: %v", err)
|
||||
}
|
||||
|
||||
if gormDB == nil {
|
||||
t.Error("Expected GORM to return a non-nil instance")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewConnectionFromDB_HealthCheck(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
conn := NewConnectionFromDB("test-connection", DatabaseTypeSQLite, db)
|
||||
ctx := context.Background()
|
||||
|
||||
err = conn.Connect(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// Health check should succeed
|
||||
err = conn.HealthCheck(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("Expected HealthCheck to succeed, got error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewConnectionFromDB_Stats(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
conn := NewConnectionFromDB("test-connection", DatabaseTypeSQLite, db)
|
||||
ctx := context.Background()
|
||||
|
||||
err = conn.Connect(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
stats := conn.Stats()
|
||||
if stats == nil {
|
||||
t.Fatal("Expected stats to be returned")
|
||||
}
|
||||
|
||||
if stats.Name != "test-connection" {
|
||||
t.Errorf("Expected stats.Name to be 'test-connection', got '%s'", stats.Name)
|
||||
}
|
||||
|
||||
if stats.Type != DatabaseTypeSQLite {
|
||||
t.Errorf("Expected stats.Type to be DatabaseTypeSQLite, got '%s'", stats.Type)
|
||||
}
|
||||
|
||||
if !stats.Connected {
|
||||
t.Error("Expected stats.Connected to be true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewConnectionFromDB_PostgreSQL(t *testing.T) {
|
||||
// This test just verifies the factory works with PostgreSQL type
|
||||
// It won't actually connect since we're using SQLite
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
conn := NewConnectionFromDB("test-pg", DatabaseTypePostgreSQL, db)
|
||||
if conn == nil {
|
||||
t.Fatal("Expected connection to be created")
|
||||
}
|
||||
|
||||
if conn.Type() != DatabaseTypePostgreSQL {
|
||||
t.Errorf("Expected type DatabaseTypePostgreSQL, got '%s'", conn.Type())
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,415 @@
|
||||
package dbmanager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// Manager manages multiple named database connections
|
||||
type Manager interface {
|
||||
// Connection retrieval
|
||||
Get(name string) (Connection, error)
|
||||
GetDefault() (Connection, error)
|
||||
GetAll() map[string]Connection
|
||||
|
||||
// Default database management
|
||||
GetDefaultDatabase() (common.Database, error)
|
||||
SetDefaultDatabase(name string) error
|
||||
|
||||
// Lifecycle
|
||||
Connect(ctx context.Context) error
|
||||
Close() error
|
||||
HealthCheck(ctx context.Context) error
|
||||
|
||||
// Stats
|
||||
Stats() *ManagerStats
|
||||
}
|
||||
|
||||
// ManagerStats contains statistics about the connection manager
|
||||
type ManagerStats struct {
|
||||
TotalConnections int
|
||||
HealthyCount int
|
||||
UnhealthyCount int
|
||||
ConnectionStats map[string]*ConnectionStats
|
||||
}
|
||||
|
||||
// connectionManager implements Manager
|
||||
type connectionManager struct {
|
||||
connections map[string]Connection
|
||||
config ManagerConfig
|
||||
mu sync.RWMutex
|
||||
|
||||
// Background health check
|
||||
healthTicker *time.Ticker
|
||||
stopChan chan struct{}
|
||||
healthMu sync.Mutex // guards healthTicker and stopChan
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
var (
|
||||
// singleton instance of the manager
|
||||
instance Manager
|
||||
// instanceMu protects the singleton instance
|
||||
instanceMu sync.RWMutex
|
||||
)
|
||||
|
||||
// SetupManager initializes the singleton database manager with the provided configuration.
|
||||
// This function must be called before GetInstance().
|
||||
// Returns an error if the manager is already initialized or if configuration is invalid.
|
||||
func SetupManager(cfg ManagerConfig) error {
|
||||
instanceMu.Lock()
|
||||
defer instanceMu.Unlock()
|
||||
|
||||
if instance != nil {
|
||||
return fmt.Errorf("manager already initialized")
|
||||
}
|
||||
|
||||
mgr, err := NewManager(cfg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create manager: %w", err)
|
||||
}
|
||||
|
||||
instance = mgr
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetInstance returns the singleton instance of the database manager.
|
||||
// Returns an error if SetupManager has not been called yet.
|
||||
func GetInstance() (Manager, error) {
|
||||
instanceMu.RLock()
|
||||
defer instanceMu.RUnlock()
|
||||
|
||||
if instance == nil {
|
||||
return nil, fmt.Errorf("manager not initialized: call SetupManager first")
|
||||
}
|
||||
|
||||
return instance, nil
|
||||
}
|
||||
|
||||
// ResetInstance resets the singleton instance (primarily for testing purposes).
|
||||
// WARNING: This should only be used in tests. Calling this in production code
|
||||
// while the manager is in use can lead to undefined behavior.
|
||||
func ResetInstance() {
|
||||
instanceMu.Lock()
|
||||
defer instanceMu.Unlock()
|
||||
|
||||
if instance != nil {
|
||||
if err := instance.Close(); err != nil {
|
||||
logger.Error("Failed to close manager during reset: %v", err)
|
||||
}
|
||||
}
|
||||
instance = nil
|
||||
}
|
||||
|
||||
// NewManager creates a new database connection manager
|
||||
func NewManager(cfg ManagerConfig) (Manager, error) {
|
||||
// Apply defaults and validate configuration
|
||||
cfg.ApplyDefaults()
|
||||
if err := cfg.Validate(); err != nil {
|
||||
return nil, fmt.Errorf("invalid configuration: %w", err)
|
||||
}
|
||||
|
||||
mgr := &connectionManager{
|
||||
connections: make(map[string]Connection),
|
||||
config: cfg,
|
||||
}
|
||||
|
||||
return mgr, nil
|
||||
}
|
||||
|
||||
// Get retrieves a named connection
|
||||
func (m *connectionManager) Get(name string) (Connection, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
conn, ok := m.connections[name]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("%w: %s", ErrConnectionNotFound, name)
|
||||
}
|
||||
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// GetDefault retrieves the default connection
|
||||
func (m *connectionManager) GetDefault() (Connection, error) {
|
||||
m.mu.RLock()
|
||||
defaultName := m.config.DefaultConnection
|
||||
m.mu.RUnlock()
|
||||
|
||||
if defaultName == "" {
|
||||
return nil, ErrNoDefaultConnection
|
||||
}
|
||||
|
||||
return m.Get(defaultName)
|
||||
}
|
||||
|
||||
// GetAll returns all connections
|
||||
func (m *connectionManager) GetAll() map[string]Connection {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
// Create a copy to avoid concurrent access issues
|
||||
result := make(map[string]Connection, len(m.connections))
|
||||
for name, conn := range m.connections {
|
||||
result[name] = conn
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetDefaultDatabase returns the common.Database interface from the default connection
|
||||
func (m *connectionManager) GetDefaultDatabase() (common.Database, error) {
|
||||
conn, err := m.GetDefault()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
db, err := conn.Database()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get database from default connection: %w", err)
|
||||
}
|
||||
|
||||
return db, nil
|
||||
}
|
||||
|
||||
// SetDefaultDatabase sets the default database connection by name
|
||||
func (m *connectionManager) SetDefaultDatabase(name string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// Verify the connection exists
|
||||
if _, ok := m.connections[name]; !ok {
|
||||
return fmt.Errorf("%w: %s", ErrConnectionNotFound, name)
|
||||
}
|
||||
|
||||
m.config.DefaultConnection = name
|
||||
logger.Info("Default database connection changed: name=%s", name)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Connect establishes all configured database connections
|
||||
func (m *connectionManager) Connect(ctx context.Context) error {
|
||||
// Dial outside m.mu so a slow connect never blocks Get/Stats/HealthCheck.
|
||||
m.mu.RLock()
|
||||
names := make([]string, 0, len(m.config.Connections))
|
||||
for name := range m.config.Connections {
|
||||
if _, exists := m.connections[name]; !exists {
|
||||
names = append(names, name)
|
||||
}
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
opened := make(map[string]Connection, len(names))
|
||||
closeOpened := func() {
|
||||
for name, conn := range opened {
|
||||
if err := conn.Close(); err != nil {
|
||||
logger.Error("Failed to close connection after failed Connect: name=%s, error=%v", name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for _, name := range names {
|
||||
// Get a copy of the connection config
|
||||
connCfg := m.config.Connections[name]
|
||||
// Apply global defaults to connection config
|
||||
connCfg.ApplyDefaults(&m.config)
|
||||
connCfg.Name = name
|
||||
|
||||
// Create connection using factory
|
||||
conn, err := createConnection(connCfg)
|
||||
if err != nil {
|
||||
closeOpened()
|
||||
return fmt.Errorf("failed to create connection '%s': %w", name, err)
|
||||
}
|
||||
|
||||
// Connect
|
||||
if err := conn.Connect(ctx); err != nil {
|
||||
closeOpened()
|
||||
return fmt.Errorf("failed to connect '%s': %w", name, err)
|
||||
}
|
||||
|
||||
opened[name] = conn
|
||||
logger.Info("Database connection established: name=%s, type=%s", name, connCfg.Type)
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
for name, conn := range opened {
|
||||
if _, exists := m.connections[name]; exists {
|
||||
// Lost a race with a concurrent Connect; drop our duplicate.
|
||||
_ = conn.Close()
|
||||
continue
|
||||
}
|
||||
m.connections[name] = conn
|
||||
}
|
||||
total := len(m.connections)
|
||||
m.mu.Unlock()
|
||||
|
||||
// Always start background health checks
|
||||
if m.config.HealthCheckInterval > 0 {
|
||||
m.startHealthChecker()
|
||||
logger.Info("Background health checker started: interval=%v", m.config.HealthCheckInterval)
|
||||
}
|
||||
|
||||
logger.Info("Database manager initialized: connections=%d", total)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes all database connections
|
||||
func (m *connectionManager) Close() error {
|
||||
// Stop the health checker before taking mu. performHealthCheck acquires
|
||||
// a read lock, so waiting for the goroutine while holding the write lock
|
||||
// would deadlock.
|
||||
m.stopHealthChecker()
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// Close all connections
|
||||
var errors []error
|
||||
for name, conn := range m.connections {
|
||||
if err := conn.Close(); err != nil {
|
||||
errors = append(errors, fmt.Errorf("failed to close connection '%s': %w", name, err))
|
||||
logger.Error("Failed to close connection: name=%s, error=%v", name, err)
|
||||
} else {
|
||||
logger.Info("Connection closed: name=%s", name)
|
||||
}
|
||||
}
|
||||
|
||||
m.connections = make(map[string]Connection)
|
||||
|
||||
if len(errors) > 0 {
|
||||
return fmt.Errorf("errors closing connections: %v", errors)
|
||||
}
|
||||
|
||||
logger.Info("Database manager closed")
|
||||
return nil
|
||||
}
|
||||
|
||||
// HealthCheck performs health checks on all connections
|
||||
func (m *connectionManager) HealthCheck(ctx context.Context) error {
|
||||
m.mu.RLock()
|
||||
connections := make(map[string]Connection, len(m.connections))
|
||||
for name, conn := range m.connections {
|
||||
connections[name] = conn
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
var errors []error
|
||||
for name, conn := range connections {
|
||||
if err := conn.HealthCheck(ctx); err != nil {
|
||||
errors = append(errors, fmt.Errorf("connection '%s': %w", name, err))
|
||||
}
|
||||
}
|
||||
|
||||
if len(errors) > 0 {
|
||||
return fmt.Errorf("health check failed for %d connections: %v", len(errors), errors)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stats returns statistics for all connections
|
||||
func (m *connectionManager) Stats() *ManagerStats {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
stats := &ManagerStats{
|
||||
TotalConnections: len(m.connections),
|
||||
ConnectionStats: make(map[string]*ConnectionStats),
|
||||
}
|
||||
|
||||
for name, conn := range m.connections {
|
||||
connStats := conn.Stats()
|
||||
stats.ConnectionStats[name] = connStats
|
||||
|
||||
if connStats.Connected && connStats.HealthCheckStatus == "healthy" {
|
||||
stats.HealthyCount++
|
||||
} else {
|
||||
stats.UnhealthyCount++
|
||||
}
|
||||
}
|
||||
|
||||
return stats
|
||||
}
|
||||
|
||||
// startHealthChecker starts background health checking
|
||||
func (m *connectionManager) startHealthChecker() {
|
||||
m.healthMu.Lock()
|
||||
defer m.healthMu.Unlock()
|
||||
|
||||
if m.healthTicker != nil {
|
||||
return // Already running
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(m.config.HealthCheckInterval)
|
||||
stop := make(chan struct{})
|
||||
m.healthTicker = ticker
|
||||
m.stopChan = stop
|
||||
|
||||
m.wg.Add(1)
|
||||
go func() {
|
||||
defer m.wg.Done()
|
||||
logger.Info("Health checker started: interval=%v", m.config.HealthCheckInterval)
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
m.performHealthCheck()
|
||||
case <-stop:
|
||||
logger.Info("Health checker stopped")
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// stopHealthChecker stops background health checking. Safe to call repeatedly.
|
||||
func (m *connectionManager) stopHealthChecker() {
|
||||
m.healthMu.Lock()
|
||||
defer m.healthMu.Unlock()
|
||||
|
||||
if m.healthTicker == nil {
|
||||
return
|
||||
}
|
||||
m.healthTicker.Stop()
|
||||
close(m.stopChan)
|
||||
m.wg.Wait()
|
||||
m.healthTicker = nil
|
||||
m.stopChan = nil
|
||||
}
|
||||
|
||||
// performHealthCheck performs a health check on all connections
|
||||
func (m *connectionManager) performHealthCheck() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
m.mu.RLock()
|
||||
connections := make([]struct {
|
||||
name string
|
||||
conn Connection
|
||||
}, 0, len(m.connections))
|
||||
for name, conn := range m.connections {
|
||||
connections = append(connections, struct {
|
||||
name string
|
||||
conn Connection
|
||||
}{name, conn})
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
defer m.PublishMetrics()
|
||||
|
||||
for _, item := range connections {
|
||||
if err := item.conn.HealthCheck(ctx); err != nil {
|
||||
// Do not reconnect here: *sql.DB discards bad connections and dials
|
||||
// new ones by itself, while Reconnect closes the pool and breaks
|
||||
// every handle already handed out. Reconnect is operator-only.
|
||||
logger.Warn("Health check failed: connection=%s, error=%v", item.name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,281 @@
|
||||
package dbmanager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/uptrace/bun"
|
||||
"go.mongodb.org/mongo-driver/mongo"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
|
||||
_ "github.com/mattn/go-sqlite3"
|
||||
)
|
||||
|
||||
type healthCheckStubConnection struct {
|
||||
healthErr error
|
||||
reconnectCalls int
|
||||
}
|
||||
|
||||
func (c *healthCheckStubConnection) Name() string { return "stub" }
|
||||
func (c *healthCheckStubConnection) Type() DatabaseType { return DatabaseTypePostgreSQL }
|
||||
func (c *healthCheckStubConnection) Bun() (*bun.DB, error) { return nil, fmt.Errorf("not implemented") }
|
||||
func (c *healthCheckStubConnection) GORM() (*gorm.DB, error) {
|
||||
return nil, fmt.Errorf("not implemented")
|
||||
}
|
||||
func (c *healthCheckStubConnection) Native() (*sql.DB, error) {
|
||||
return nil, fmt.Errorf("not implemented")
|
||||
}
|
||||
func (c *healthCheckStubConnection) DB() (*sql.DB, error) { return nil, fmt.Errorf("not implemented") }
|
||||
func (c *healthCheckStubConnection) Database() (common.Database, error) {
|
||||
return nil, fmt.Errorf("not implemented")
|
||||
}
|
||||
func (c *healthCheckStubConnection) MongoDB() (*mongo.Client, error) {
|
||||
return nil, fmt.Errorf("not implemented")
|
||||
}
|
||||
func (c *healthCheckStubConnection) Connect(ctx context.Context) error { return nil }
|
||||
func (c *healthCheckStubConnection) Close() error { return nil }
|
||||
func (c *healthCheckStubConnection) HealthCheck(ctx context.Context) error { return c.healthErr }
|
||||
func (c *healthCheckStubConnection) Reconnect(ctx context.Context) error {
|
||||
c.reconnectCalls++
|
||||
return nil
|
||||
}
|
||||
func (c *healthCheckStubConnection) Stats() *ConnectionStats { return &ConnectionStats{} }
|
||||
|
||||
func TestBackgroundHealthChecker(t *testing.T) {
|
||||
// Create a SQLite in-memory database
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
// Create manager config with a short health check interval for testing
|
||||
cfg := ManagerConfig{
|
||||
DefaultConnection: "test",
|
||||
Connections: map[string]ConnectionConfig{
|
||||
"test": {
|
||||
Name: "test",
|
||||
Type: DatabaseTypeSQLite,
|
||||
FilePath: ":memory:",
|
||||
},
|
||||
},
|
||||
HealthCheckInterval: 1 * time.Second, // Short interval for testing
|
||||
EnableAutoReconnect: true,
|
||||
}
|
||||
|
||||
// Create manager
|
||||
mgr, err := NewManager(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create manager: %v", err)
|
||||
}
|
||||
|
||||
// Connect - this should start the background health checker
|
||||
ctx := context.Background()
|
||||
err = mgr.Connect(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer mgr.Close()
|
||||
|
||||
// Get the connection to verify it's healthy
|
||||
conn, err := mgr.Get("test")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get connection: %v", err)
|
||||
}
|
||||
|
||||
// Verify initial health check
|
||||
err = conn.HealthCheck(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("Initial health check failed: %v", err)
|
||||
}
|
||||
|
||||
// Wait for a few health check cycles
|
||||
time.Sleep(3 * time.Second)
|
||||
|
||||
// Get stats to verify the connection is still healthy
|
||||
stats := conn.Stats()
|
||||
if stats == nil {
|
||||
t.Fatal("Expected stats to be returned")
|
||||
}
|
||||
|
||||
if !stats.Connected {
|
||||
t.Error("Expected connection to still be connected")
|
||||
}
|
||||
|
||||
if stats.HealthCheckStatus == "" {
|
||||
t.Error("Expected health check status to be set")
|
||||
}
|
||||
|
||||
// Verify the manager has started the health checker
|
||||
if cm, ok := mgr.(*connectionManager); ok {
|
||||
if cm.healthTicker == nil {
|
||||
t.Error("Expected health ticker to be running")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultHealthCheckInterval(t *testing.T) {
|
||||
// Verify the default health check interval is 15 seconds
|
||||
defaults := DefaultManagerConfig()
|
||||
|
||||
expectedInterval := 15 * time.Second
|
||||
if defaults.HealthCheckInterval != expectedInterval {
|
||||
t.Errorf("Expected default health check interval to be %v, got %v",
|
||||
expectedInterval, defaults.HealthCheckInterval)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyDefaultsHealthCheckInterval(t *testing.T) {
|
||||
cfg := ManagerConfig{}
|
||||
cfg.ApplyDefaults()
|
||||
if cfg.HealthCheckInterval != 15*time.Second {
|
||||
t.Errorf("Expected health check interval to be 15s, got %v", cfg.HealthCheckInterval)
|
||||
}
|
||||
|
||||
// A negative interval disables the background checker and is preserved.
|
||||
cfg = ManagerConfig{HealthCheckInterval: -1}
|
||||
cfg.ApplyDefaults()
|
||||
if cfg.HealthCheckInterval >= 0 {
|
||||
t.Errorf("Expected negative interval to be preserved, got %v", cfg.HealthCheckInterval)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerHealthCheck(t *testing.T) {
|
||||
// Create a SQLite in-memory database
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
// Create manager config
|
||||
cfg := ManagerConfig{
|
||||
DefaultConnection: "test",
|
||||
Connections: map[string]ConnectionConfig{
|
||||
"test": {
|
||||
Name: "test",
|
||||
Type: DatabaseTypeSQLite,
|
||||
FilePath: ":memory:",
|
||||
},
|
||||
},
|
||||
HealthCheckInterval: 15 * time.Second,
|
||||
EnableAutoReconnect: true,
|
||||
}
|
||||
|
||||
// Create and connect manager
|
||||
mgr, err := NewManager(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create manager: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
err = mgr.Connect(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer mgr.Close()
|
||||
|
||||
// Perform health check on all connections
|
||||
err = mgr.HealthCheck(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("Health check failed: %v", err)
|
||||
}
|
||||
|
||||
// Get stats
|
||||
stats := mgr.Stats()
|
||||
if stats == nil {
|
||||
t.Fatal("Expected stats to be returned")
|
||||
}
|
||||
|
||||
if stats.TotalConnections != 1 {
|
||||
t.Errorf("Expected 1 total connection, got %d", stats.TotalConnections)
|
||||
}
|
||||
|
||||
if stats.HealthyCount != 1 {
|
||||
t.Errorf("Expected 1 healthy connection, got %d", stats.HealthyCount)
|
||||
}
|
||||
|
||||
if stats.UnhealthyCount != 0 {
|
||||
t.Errorf("Expected 0 unhealthy connections, got %d", stats.UnhealthyCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerStatsAfterClose(t *testing.T) {
|
||||
cfg := ManagerConfig{
|
||||
DefaultConnection: "test",
|
||||
Connections: map[string]ConnectionConfig{
|
||||
"test": {
|
||||
Name: "test",
|
||||
Type: DatabaseTypeSQLite,
|
||||
FilePath: ":memory:",
|
||||
},
|
||||
},
|
||||
HealthCheckInterval: 15 * time.Second,
|
||||
}
|
||||
|
||||
mgr, err := NewManager(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create manager: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
err = mgr.Connect(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
|
||||
// Close the manager
|
||||
err = mgr.Close()
|
||||
if err != nil {
|
||||
t.Errorf("Failed to close manager: %v", err)
|
||||
}
|
||||
|
||||
// Stats should show no connections
|
||||
stats := mgr.Stats()
|
||||
if stats.TotalConnections != 0 {
|
||||
t.Errorf("Expected 0 total connections after close, got %d", stats.TotalConnections)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerformHealthCheckSkipsReconnectForTransientFailures(t *testing.T) {
|
||||
conn := &healthCheckStubConnection{
|
||||
healthErr: fmt.Errorf("connection 'primary' health check: dial tcp 127.0.0.1:5432: connect: connection refused"),
|
||||
}
|
||||
|
||||
mgr := &connectionManager{
|
||||
connections: map[string]Connection{"primary": conn},
|
||||
config: ManagerConfig{
|
||||
EnableAutoReconnect: true,
|
||||
},
|
||||
}
|
||||
|
||||
mgr.performHealthCheck()
|
||||
|
||||
if conn.reconnectCalls != 0 {
|
||||
t.Fatalf("expected no reconnect attempts for transient health failure, got %d", conn.reconnectCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerformHealthCheckNeverReconnects(t *testing.T) {
|
||||
conn := &healthCheckStubConnection{
|
||||
healthErr: NewConnectionError("primary", "health check", fmt.Errorf("sql: database is closed")),
|
||||
}
|
||||
|
||||
mgr := &connectionManager{
|
||||
connections: map[string]Connection{"primary": conn},
|
||||
config: ManagerConfig{
|
||||
EnableAutoReconnect: true,
|
||||
},
|
||||
}
|
||||
|
||||
mgr.performHealthCheck()
|
||||
|
||||
if conn.reconnectCalls != 0 {
|
||||
t.Fatalf("health check must not close the shared pool via Reconnect, got %d", conn.reconnectCalls)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,160 @@
|
||||
package dbmanager
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/prometheus/client_golang/prometheus/promauto"
|
||||
)
|
||||
|
||||
var (
|
||||
// connectionsTotal tracks the total number of configured database connections
|
||||
connectionsTotal = promauto.NewGaugeVec(
|
||||
prometheus.GaugeOpts{
|
||||
Name: "dbmanager_connections_total",
|
||||
Help: "Total number of configured database connections",
|
||||
},
|
||||
[]string{"type"},
|
||||
)
|
||||
|
||||
// connectionStatus tracks connection health status (1=healthy, 0=unhealthy)
|
||||
connectionStatus = promauto.NewGaugeVec(
|
||||
prometheus.GaugeOpts{
|
||||
Name: "dbmanager_connection_status",
|
||||
Help: "Connection status (1=healthy, 0=unhealthy)",
|
||||
},
|
||||
[]string{"name", "type"},
|
||||
)
|
||||
|
||||
// connectionPoolSize tracks connection pool sizes
|
||||
connectionPoolSize = promauto.NewGaugeVec(
|
||||
prometheus.GaugeOpts{
|
||||
Name: "dbmanager_connection_pool_size",
|
||||
Help: "Current connection pool size",
|
||||
},
|
||||
[]string{"name", "type", "state"}, // state: open, idle, in_use
|
||||
)
|
||||
|
||||
// connectionWaitCount tracks how many times connections had to wait for availability
|
||||
connectionWaitCount = promauto.NewCounterVec(
|
||||
prometheus.CounterOpts{
|
||||
Name: "dbmanager_connection_wait_count",
|
||||
Help: "Number of times connections had to wait for availability",
|
||||
},
|
||||
[]string{"name", "type"},
|
||||
)
|
||||
|
||||
// connectionWaitDuration tracks total time connections spent waiting
|
||||
connectionWaitDuration = promauto.NewCounterVec(
|
||||
prometheus.CounterOpts{
|
||||
Name: "dbmanager_connection_wait_duration_seconds",
|
||||
Help: "Total time connections spent waiting for availability",
|
||||
},
|
||||
[]string{"name", "type"},
|
||||
)
|
||||
|
||||
// reconnectAttempts tracks reconnection attempts and their outcomes
|
||||
reconnectAttempts = promauto.NewCounterVec(
|
||||
prometheus.CounterOpts{
|
||||
Name: "dbmanager_reconnect_attempts_total",
|
||||
Help: "Total number of reconnection attempts",
|
||||
},
|
||||
[]string{"name", "type", "result"}, // result: success, failure
|
||||
)
|
||||
|
||||
// connectionLifetimeClosed tracks connections closed due to max lifetime
|
||||
connectionLifetimeClosed = promauto.NewCounterVec(
|
||||
prometheus.CounterOpts{
|
||||
Name: "dbmanager_connection_lifetime_closed_total",
|
||||
Help: "Total connections closed due to exceeding max lifetime",
|
||||
},
|
||||
[]string{"name", "type"},
|
||||
)
|
||||
|
||||
// connectionIdleClosed tracks connections closed due to max idle time
|
||||
connectionIdleClosed = promauto.NewCounterVec(
|
||||
prometheus.CounterOpts{
|
||||
Name: "dbmanager_connection_idle_closed_total",
|
||||
Help: "Total connections closed due to exceeding max idle time",
|
||||
},
|
||||
[]string{"name", "type"},
|
||||
)
|
||||
)
|
||||
|
||||
// PublishMetrics publishes current metrics for all connections
|
||||
func (m *connectionManager) PublishMetrics() {
|
||||
stats := m.Stats()
|
||||
|
||||
// Count connections by type
|
||||
typeCount := make(map[DatabaseType]int)
|
||||
for _, connStats := range stats.ConnectionStats {
|
||||
typeCount[connStats.Type]++
|
||||
}
|
||||
|
||||
// Update total connections gauge
|
||||
for dbType, count := range typeCount {
|
||||
connectionsTotal.WithLabelValues(string(dbType)).Set(float64(count))
|
||||
}
|
||||
|
||||
// Update per-connection metrics
|
||||
for name, connStats := range stats.ConnectionStats {
|
||||
labels := prometheus.Labels{
|
||||
"name": name,
|
||||
"type": string(connStats.Type),
|
||||
}
|
||||
|
||||
// Connection status
|
||||
status := float64(0)
|
||||
if connStats.Connected && connStats.HealthCheckStatus == "healthy" {
|
||||
status = 1
|
||||
}
|
||||
connectionStatus.With(labels).Set(status)
|
||||
|
||||
// Pool size metrics (SQL databases only)
|
||||
if connStats.Type != DatabaseTypeMongoDB {
|
||||
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "open").Set(float64(connStats.OpenConnections))
|
||||
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "idle").Set(float64(connStats.Idle))
|
||||
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "in_use").Set(float64(connStats.InUse))
|
||||
|
||||
// sql.DBStats values are cumulative, so add only the growth since
|
||||
// the last publish to keep these true counters.
|
||||
prev := lastPublished.swap(name, connStats)
|
||||
connectionWaitCount.With(labels).Add(float64(connStats.WaitCount - prev.WaitCount))
|
||||
connectionWaitDuration.With(labels).Add((connStats.WaitDuration - prev.WaitDuration).Seconds())
|
||||
connectionLifetimeClosed.With(labels).Add(float64(connStats.MaxLifetimeClosed - prev.MaxLifetimeClosed))
|
||||
connectionIdleClosed.With(labels).Add(float64(connStats.MaxIdleClosed - prev.MaxIdleClosed))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RecordReconnectAttempt records a reconnection attempt
|
||||
func RecordReconnectAttempt(name string, dbType DatabaseType, success bool) {
|
||||
result := "failure"
|
||||
if success {
|
||||
result = "success"
|
||||
}
|
||||
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,319 @@
|
||||
# PostgreSQL NOTIFY/LISTEN Support
|
||||
|
||||
The `dbmanager` package provides built-in support for PostgreSQL's NOTIFY/LISTEN functionality through the `PostgresListener` type.
|
||||
|
||||
## Overview
|
||||
|
||||
PostgreSQL NOTIFY/LISTEN is a simple pub/sub mechanism that allows database clients to:
|
||||
- **LISTEN** on named channels to receive notifications
|
||||
- **NOTIFY** channels to send messages to all listeners
|
||||
- Receive asynchronous notifications without polling
|
||||
|
||||
## Features
|
||||
|
||||
- ✅ Subscribe to multiple channels simultaneously
|
||||
- ✅ Callback-based notification handling
|
||||
- ✅ Automatic reconnection on connection loss
|
||||
- ✅ Automatic resubscription after reconnection
|
||||
- ✅ Thread-safe operations
|
||||
- ✅ Panic recovery in notification handlers
|
||||
- ✅ Dedicated connection for listening (doesn't interfere with queries)
|
||||
|
||||
## Usage
|
||||
|
||||
### Basic Usage
|
||||
|
||||
```go
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
|
||||
)
|
||||
|
||||
func main() {
|
||||
// Create PostgreSQL provider
|
||||
cfg := &providers.Config{
|
||||
Name: "primary",
|
||||
Type: "postgres",
|
||||
Host: "localhost",
|
||||
Port: 5432,
|
||||
User: "postgres",
|
||||
Password: "password",
|
||||
Database: "myapp",
|
||||
}
|
||||
|
||||
provider := providers.NewPostgresProvider()
|
||||
ctx := context.Background()
|
||||
|
||||
if err := provider.Connect(ctx, cfg); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
defer provider.Close()
|
||||
|
||||
// Get listener
|
||||
listener, err := provider.GetListener(ctx)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
// Subscribe to a channel
|
||||
err = listener.Listen("events", func(channel, payload string) {
|
||||
fmt.Printf("Received on %s: %s\n", channel, payload)
|
||||
})
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
// Send a notification
|
||||
err = listener.Notify(ctx, "events", "Hello, World!")
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
// Keep the program running
|
||||
time.Sleep(1 * time.Second)
|
||||
}
|
||||
```
|
||||
|
||||
### Multiple Channels
|
||||
|
||||
```go
|
||||
listener, _ := provider.GetListener(ctx)
|
||||
|
||||
// Listen to different channels with different handlers
|
||||
listener.Listen("user_events", func(channel, payload string) {
|
||||
fmt.Printf("User event: %s\n", payload)
|
||||
})
|
||||
|
||||
listener.Listen("order_events", func(channel, payload string) {
|
||||
fmt.Printf("Order event: %s\n", payload)
|
||||
})
|
||||
|
||||
listener.Listen("payment_events", func(channel, payload string) {
|
||||
fmt.Printf("Payment event: %s\n", payload)
|
||||
})
|
||||
```
|
||||
|
||||
### Unsubscribing
|
||||
|
||||
```go
|
||||
// Stop listening to a specific channel
|
||||
err := listener.Unlisten("user_events")
|
||||
if err != nil {
|
||||
fmt.Printf("Failed to unlisten: %v\n", err)
|
||||
}
|
||||
```
|
||||
|
||||
### Checking Active Channels
|
||||
|
||||
```go
|
||||
// Get list of channels currently being listened to
|
||||
channels := listener.Channels()
|
||||
fmt.Printf("Listening to: %v\n", channels)
|
||||
```
|
||||
|
||||
### Checking Connection Status
|
||||
|
||||
```go
|
||||
if listener.IsConnected() {
|
||||
fmt.Println("Listener is connected")
|
||||
} else {
|
||||
fmt.Println("Listener is disconnected")
|
||||
}
|
||||
```
|
||||
|
||||
## Integration with DBManager
|
||||
|
||||
When using the DBManager, you can access the listener through the PostgreSQL provider:
|
||||
|
||||
```go
|
||||
// Initialize DBManager
|
||||
mgr, err := dbmanager.NewManager(dbmanager.FromConfig(cfg.DBManager))
|
||||
mgr.Connect(ctx)
|
||||
defer mgr.Close()
|
||||
|
||||
// Get PostgreSQL connection
|
||||
conn, err := mgr.Get("primary")
|
||||
|
||||
// Note: You'll need to cast to the underlying provider type
|
||||
// This requires exposing the provider through the Connection interface
|
||||
// or providing a helper method
|
||||
```
|
||||
|
||||
## Use Cases
|
||||
|
||||
### Cache Invalidation
|
||||
|
||||
```go
|
||||
listener.Listen("cache_invalidation", func(channel, payload string) {
|
||||
// Parse the payload to determine what to invalidate
|
||||
cache.Invalidate(payload)
|
||||
})
|
||||
```
|
||||
|
||||
### Real-time Updates
|
||||
|
||||
```go
|
||||
listener.Listen("data_updates", func(channel, payload string) {
|
||||
// Broadcast update to WebSocket clients
|
||||
websocketBroadcast(payload)
|
||||
})
|
||||
```
|
||||
|
||||
### Configuration Reload
|
||||
|
||||
```go
|
||||
listener.Listen("config_reload", func(channel, payload string) {
|
||||
// Reload application configuration
|
||||
config.Reload()
|
||||
})
|
||||
```
|
||||
|
||||
### Distributed Locking
|
||||
|
||||
```go
|
||||
listener.Listen("lock_released", func(channel, payload string) {
|
||||
// Attempt to acquire the lock
|
||||
tryAcquireLock(payload)
|
||||
})
|
||||
```
|
||||
|
||||
## Automatic Reconnection
|
||||
|
||||
The listener automatically handles connection failures:
|
||||
|
||||
1. When a connection error is detected, the listener initiates reconnection
|
||||
2. Once reconnected, it automatically resubscribes to all previous channels
|
||||
3. Notification handlers remain active throughout the reconnection process
|
||||
|
||||
No manual intervention is required for reconnection.
|
||||
|
||||
## Error Handling
|
||||
|
||||
### Handler Panics
|
||||
|
||||
If a notification handler panics, the panic is recovered and logged. The listener continues to function normally:
|
||||
|
||||
```go
|
||||
listener.Listen("events", func(channel, payload string) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
log.Printf("Handler panic: %v", r)
|
||||
}
|
||||
}()
|
||||
|
||||
// Your event processing logic
|
||||
processEvent(payload)
|
||||
})
|
||||
```
|
||||
|
||||
### Connection Errors
|
||||
|
||||
Connection errors trigger automatic reconnection. Check logs for reconnection events when `EnableLogging` is true.
|
||||
|
||||
## Thread Safety
|
||||
|
||||
All `PostgresListener` methods are thread-safe and can be called concurrently from multiple goroutines.
|
||||
|
||||
## Performance Considerations
|
||||
|
||||
1. **Dedicated Connection**: The listener uses a dedicated PostgreSQL connection separate from the query connection pool
|
||||
2. **Asynchronous Handlers**: Notification handlers run in separate goroutines to avoid blocking
|
||||
3. **Lightweight**: NOTIFY/LISTEN has minimal overhead compared to polling
|
||||
|
||||
## Comparison with Polling
|
||||
|
||||
| Feature | NOTIFY/LISTEN | Polling |
|
||||
|---------|---------------|---------|
|
||||
| Latency | Low (near real-time) | High (depends on poll interval) |
|
||||
| Database Load | Minimal | High (constant queries) |
|
||||
| Scalability | Excellent | Poor |
|
||||
| Complexity | Simple | Moderate |
|
||||
|
||||
## Limitations
|
||||
|
||||
1. **PostgreSQL Only**: This feature is specific to PostgreSQL and not available for other databases
|
||||
2. **No Message Persistence**: Notifications are not stored; if no listener is connected, the message is lost
|
||||
3. **Payload Limit**: Notification payload is limited to 8000 bytes in PostgreSQL
|
||||
4. **No Guaranteed Delivery**: If a listener disconnects, in-flight notifications may be lost
|
||||
|
||||
## Best Practices
|
||||
|
||||
1. **Keep Handlers Fast**: Notification handlers should be quick; for heavy processing, send work to a queue
|
||||
2. **Use JSON Payloads**: Encode structured data as JSON for easy parsing
|
||||
3. **Handle Errors Gracefully**: Always recover from panics in handlers
|
||||
4. **Close Properly**: Always close the provider to ensure the listener is properly shut down
|
||||
5. **Monitor Connection Status**: Use `IsConnected()` for health checks
|
||||
|
||||
## Example: Real-World Application
|
||||
|
||||
```go
|
||||
// Subscribe to various application events
|
||||
listener, _ := provider.GetListener(ctx)
|
||||
|
||||
// User registration events
|
||||
listener.Listen("user_registered", func(channel, payload string) {
|
||||
var event UserRegisteredEvent
|
||||
json.Unmarshal([]byte(payload), &event)
|
||||
|
||||
// Send welcome email
|
||||
sendWelcomeEmail(event.UserID)
|
||||
|
||||
// Invalidate user count cache
|
||||
cache.Delete("user_count")
|
||||
})
|
||||
|
||||
// Order placement events
|
||||
listener.Listen("order_placed", func(channel, payload string) {
|
||||
var event OrderPlacedEvent
|
||||
json.Unmarshal([]byte(payload), &event)
|
||||
|
||||
// Notify warehouse system
|
||||
warehouse.ProcessOrder(event.OrderID)
|
||||
|
||||
// Update inventory cache
|
||||
cache.Invalidate("inventory:" + event.ProductID)
|
||||
})
|
||||
|
||||
// Configuration changes
|
||||
listener.Listen("config_updated", func(channel, payload string) {
|
||||
// Reload configuration from database
|
||||
appConfig.Reload()
|
||||
})
|
||||
```
|
||||
|
||||
## Triggering Notifications from SQL
|
||||
|
||||
You can trigger notifications directly from PostgreSQL triggers or functions:
|
||||
|
||||
```sql
|
||||
-- Example trigger to notify on new user
|
||||
CREATE OR REPLACE FUNCTION notify_user_registered()
|
||||
RETURNS TRIGGER AS $$
|
||||
BEGIN
|
||||
PERFORM pg_notify('user_registered',
|
||||
json_build_object(
|
||||
'user_id', NEW.id,
|
||||
'email', NEW.email,
|
||||
'timestamp', NOW()
|
||||
)::text
|
||||
);
|
||||
RETURN NEW;
|
||||
END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
|
||||
CREATE TRIGGER user_registered_trigger
|
||||
AFTER INSERT ON users
|
||||
FOR EACH ROW
|
||||
EXECUTE FUNCTION notify_user_registered();
|
||||
```
|
||||
|
||||
## Additional Resources
|
||||
|
||||
- [PostgreSQL NOTIFY Documentation](https://www.postgresql.org/docs/current/sql-notify.html)
|
||||
- [PostgreSQL LISTEN Documentation](https://www.postgresql.org/docs/current/sql-listen.html)
|
||||
- [pgx Driver Documentation](https://github.com/jackc/pgx)
|
||||
@@ -0,0 +1,124 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"sync"
|
||||
|
||||
"go.mongodb.org/mongo-driver/mongo"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// ExistingDBProvider wraps an existing *sql.DB connection
|
||||
// This allows using dbmanager features with a database connection
|
||||
// that was opened outside of the dbmanager package
|
||||
type ExistingDBProvider struct {
|
||||
db *sql.DB
|
||||
name string
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewExistingDBProvider creates a new provider wrapping an existing *sql.DB
|
||||
func NewExistingDBProvider(db *sql.DB, name string) *ExistingDBProvider {
|
||||
return &ExistingDBProvider{
|
||||
db: db,
|
||||
name: name,
|
||||
}
|
||||
}
|
||||
|
||||
// Connect verifies the existing database connection is valid
|
||||
// It does NOT create a new connection, but ensures the existing one works
|
||||
func (p *ExistingDBProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
|
||||
if p.db == nil {
|
||||
return fmt.Errorf("database connection is nil")
|
||||
}
|
||||
|
||||
// Verify the connection works
|
||||
if err := p.db.PingContext(ctx); err != nil {
|
||||
return fmt.Errorf("failed to ping existing database: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Refresh verifies the wrapped database is still reachable. The pool belongs to
|
||||
// the caller and cannot be re-dialed here, so it is never closed to "reconnect".
|
||||
func (p *ExistingDBProvider) Refresh(ctx context.Context) error {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
|
||||
if p.db == nil {
|
||||
return fmt.Errorf("database connection is nil")
|
||||
}
|
||||
return p.db.PingContext(ctx)
|
||||
}
|
||||
|
||||
// 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
|
||||
func (p *ExistingDBProvider) HealthCheck(ctx context.Context) error {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
|
||||
if p.db == nil {
|
||||
return fmt.Errorf("database connection is nil")
|
||||
}
|
||||
|
||||
return p.db.PingContext(ctx)
|
||||
}
|
||||
|
||||
// GetNative returns the wrapped *sql.DB
|
||||
func (p *ExistingDBProvider) GetNative() (*sql.DB, error) {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
|
||||
if p.db == nil {
|
||||
return nil, fmt.Errorf("database connection is nil")
|
||||
}
|
||||
|
||||
return p.db, nil
|
||||
}
|
||||
|
||||
// GetMongo returns an error since this is a SQL database
|
||||
func (p *ExistingDBProvider) GetMongo() (*mongo.Client, error) {
|
||||
return nil, ErrNotMongoDB
|
||||
}
|
||||
|
||||
// Stats returns connection statistics
|
||||
func (p *ExistingDBProvider) Stats() *ConnectionStats {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
|
||||
stats := &ConnectionStats{
|
||||
Name: p.name,
|
||||
Type: "sql", // Generic since we don't know the specific type
|
||||
Connected: p.db != nil,
|
||||
}
|
||||
|
||||
if p.db != nil {
|
||||
dbStats := p.db.Stats()
|
||||
stats.OpenConnections = dbStats.OpenConnections
|
||||
stats.InUse = dbStats.InUse
|
||||
stats.Idle = dbStats.Idle
|
||||
stats.WaitCount = dbStats.WaitCount
|
||||
stats.WaitDuration = dbStats.WaitDuration
|
||||
stats.MaxIdleClosed = dbStats.MaxIdleClosed
|
||||
stats.MaxLifetimeClosed = dbStats.MaxLifetimeClosed
|
||||
}
|
||||
|
||||
return stats
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "github.com/mattn/go-sqlite3"
|
||||
)
|
||||
|
||||
func TestNewExistingDBProvider(t *testing.T) {
|
||||
// Open a SQLite in-memory database
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
// Create provider
|
||||
provider := NewExistingDBProvider(db, "test-db")
|
||||
if provider == nil {
|
||||
t.Fatal("Expected provider to be created")
|
||||
}
|
||||
|
||||
if provider.name != "test-db" {
|
||||
t.Errorf("Expected name 'test-db', got '%s'", provider.name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingDBProvider_Connect(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
provider := NewExistingDBProvider(db, "test-db")
|
||||
ctx := context.Background()
|
||||
|
||||
// Connect should verify the connection works
|
||||
err = provider.Connect(ctx, nil)
|
||||
if err != nil {
|
||||
t.Errorf("Expected Connect to succeed, got error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingDBProvider_Connect_NilDB(t *testing.T) {
|
||||
provider := NewExistingDBProvider(nil, "test-db")
|
||||
ctx := context.Background()
|
||||
|
||||
err := provider.Connect(ctx, nil)
|
||||
if err == nil {
|
||||
t.Error("Expected Connect to fail with nil database")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingDBProvider_GetNative(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
provider := NewExistingDBProvider(db, "test-db")
|
||||
|
||||
nativeDB, err := provider.GetNative()
|
||||
if err != nil {
|
||||
t.Errorf("Expected GetNative to succeed, got error: %v", err)
|
||||
}
|
||||
|
||||
if nativeDB != db {
|
||||
t.Error("Expected GetNative to return the same database instance")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingDBProvider_GetNative_NilDB(t *testing.T) {
|
||||
provider := NewExistingDBProvider(nil, "test-db")
|
||||
|
||||
_, err := provider.GetNative()
|
||||
if err == nil {
|
||||
t.Error("Expected GetNative to fail with nil database")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingDBProvider_HealthCheck(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
provider := NewExistingDBProvider(db, "test-db")
|
||||
ctx := context.Background()
|
||||
|
||||
err = provider.HealthCheck(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("Expected HealthCheck to succeed, got error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingDBProvider_HealthCheck_ClosedDB(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
|
||||
provider := NewExistingDBProvider(db, "test-db")
|
||||
|
||||
// Close the database
|
||||
db.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
err = provider.HealthCheck(ctx)
|
||||
if err == nil {
|
||||
t.Error("Expected HealthCheck to fail with closed database")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingDBProvider_GetMongo(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
provider := NewExistingDBProvider(db, "test-db")
|
||||
|
||||
_, err = provider.GetMongo()
|
||||
if err != ErrNotMongoDB {
|
||||
t.Errorf("Expected ErrNotMongoDB, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingDBProvider_Stats(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
// Set some connection pool settings to test stats
|
||||
db.SetMaxOpenConns(10)
|
||||
db.SetMaxIdleConns(5)
|
||||
db.SetConnMaxLifetime(time.Hour)
|
||||
|
||||
provider := NewExistingDBProvider(db, "test-db")
|
||||
|
||||
stats := provider.Stats()
|
||||
if stats == nil {
|
||||
t.Fatal("Expected stats to be returned")
|
||||
}
|
||||
|
||||
if stats.Name != "test-db" {
|
||||
t.Errorf("Expected stats.Name to be 'test-db', got '%s'", stats.Name)
|
||||
}
|
||||
|
||||
if stats.Type != "sql" {
|
||||
t.Errorf("Expected stats.Type to be 'sql', got '%s'", stats.Type)
|
||||
}
|
||||
|
||||
if !stats.Connected {
|
||||
t.Error("Expected stats.Connected to be true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingDBProvider_Close_LeavesDBOpen(t *testing.T) {
|
||||
db, err := sql.Open("sqlite3", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to open database: %v", err)
|
||||
}
|
||||
|
||||
provider := NewExistingDBProvider(db, "test-db")
|
||||
|
||||
err = provider.Close()
|
||||
if err != nil {
|
||||
t.Errorf("Expected Close to succeed, got error: %v", err)
|
||||
}
|
||||
|
||||
// The caller owns the database, so Close must leave it open
|
||||
defer db.Close()
|
||||
if err := db.Ping(); err != nil {
|
||||
t.Errorf("Expected caller's database to stay open, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingDBProvider_Close_NilDB(t *testing.T) {
|
||||
provider := NewExistingDBProvider(nil, "test-db")
|
||||
|
||||
err := provider.Close()
|
||||
if err != nil {
|
||||
t.Errorf("Expected Close to succeed with nil database, got error: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,214 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"go.mongodb.org/mongo-driver/mongo"
|
||||
"go.mongodb.org/mongo-driver/mongo/options"
|
||||
"go.mongodb.org/mongo-driver/mongo/readpref"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// MongoProvider implements Provider for MongoDB databases
|
||||
type MongoProvider struct {
|
||||
client *mongo.Client
|
||||
config ConnectionConfig
|
||||
}
|
||||
|
||||
// NewMongoProvider creates a new MongoDB provider
|
||||
func NewMongoProvider() *MongoProvider {
|
||||
return &MongoProvider{}
|
||||
}
|
||||
|
||||
// Connect establishes a MongoDB connection
|
||||
func (p *MongoProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
||||
// Build DSN
|
||||
dsn, err := cfg.BuildDSN()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to build DSN: %w", err)
|
||||
}
|
||||
|
||||
// Create client options
|
||||
clientOpts := options.Client().ApplyURI(dsn)
|
||||
|
||||
// Set connection pool size
|
||||
if cfg.GetMaxOpenConns() != nil {
|
||||
maxPoolSize := uint64(*cfg.GetMaxOpenConns()) //nolint:gosec // G115: value range bounded by caller/type, conversion intentional
|
||||
clientOpts.SetMaxPoolSize(maxPoolSize)
|
||||
}
|
||||
|
||||
// MaxIdleConns is a ceiling on idle connections, not a pre-warmed minimum
|
||||
// (MinPoolSize), so only the idle-time limit maps onto the Mongo pool.
|
||||
if cfg.GetConnMaxIdleTime() != nil {
|
||||
clientOpts.SetMaxConnIdleTime(*cfg.GetConnMaxIdleTime())
|
||||
}
|
||||
|
||||
// Set timeouts
|
||||
clientOpts.SetConnectTimeout(cfg.GetConnectTimeout())
|
||||
if cfg.GetQueryTimeout() > 0 {
|
||||
clientOpts.SetTimeout(cfg.GetQueryTimeout())
|
||||
}
|
||||
|
||||
// Set read preference if specified
|
||||
if cfg.GetReadPreference() != "" {
|
||||
rp, err := parseReadPreference(cfg.GetReadPreference())
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid read preference: %w", err)
|
||||
}
|
||||
clientOpts.SetReadPreference(rp)
|
||||
}
|
||||
|
||||
// Connect with retry logic
|
||||
var client *mongo.Client
|
||||
var lastErr error
|
||||
|
||||
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
|
||||
|
||||
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||
if attempt > 0 {
|
||||
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Info("Retrying MongoDB connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-time.After(delay):
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
// Create MongoDB client
|
||||
client, err = mongo.Connect(ctx, clientOpts)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Warn("Failed to connect to MongoDB: %v", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Ping the database to verify connection
|
||||
pingCtx, cancel := context.WithTimeout(ctx, cfg.GetConnectTimeout())
|
||||
err = client.Ping(pingCtx, readpref.Primary())
|
||||
cancel()
|
||||
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
_ = client.Disconnect(ctx)
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Warn("Failed to ping MongoDB: %v", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Connection successful
|
||||
break
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to connect after %d attempts: %w", retryAttempts, lastErr)
|
||||
}
|
||||
|
||||
p.client = client
|
||||
p.config = cfg
|
||||
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Info("MongoDB connection established: name=%s, host=%s, database=%s", cfg.GetName(), cfg.GetHost(), cfg.GetDatabase())
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the MongoDB connection
|
||||
func (p *MongoProvider) Close() error {
|
||||
if p.client == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
err := p.client.Disconnect(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to close MongoDB connection: %w", err)
|
||||
}
|
||||
|
||||
if p.config.GetEnableLogging() {
|
||||
logger.Info("MongoDB connection closed: name=%s", p.config.GetName())
|
||||
}
|
||||
|
||||
p.client = nil
|
||||
return nil
|
||||
}
|
||||
|
||||
// HealthCheck verifies the MongoDB connection is alive
|
||||
func (p *MongoProvider) HealthCheck(ctx context.Context) error {
|
||||
if p.client == nil {
|
||||
return fmt.Errorf("MongoDB client is nil")
|
||||
}
|
||||
|
||||
// Use a short timeout for health checks
|
||||
healthCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if err := p.client.Ping(healthCtx, readpref.Primary()); err != nil {
|
||||
return fmt.Errorf("health check failed: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetNative returns an error for MongoDB (not a SQL database)
|
||||
func (p *MongoProvider) GetNative() (*sql.DB, error) {
|
||||
return nil, ErrNotSQLDatabase
|
||||
}
|
||||
|
||||
// GetMongo returns the MongoDB client
|
||||
func (p *MongoProvider) GetMongo() (*mongo.Client, error) {
|
||||
if p.client == nil {
|
||||
return nil, fmt.Errorf("MongoDB client is not initialized")
|
||||
}
|
||||
return p.client, nil
|
||||
}
|
||||
|
||||
// Stats returns connection statistics for MongoDB
|
||||
func (p *MongoProvider) Stats() *ConnectionStats {
|
||||
if p.client == nil {
|
||||
return &ConnectionStats{
|
||||
Name: p.config.GetName(),
|
||||
Type: "mongodb",
|
||||
Connected: false,
|
||||
}
|
||||
}
|
||||
|
||||
// MongoDB doesn't expose detailed connection pool stats like sql.DB
|
||||
// We return basic stats
|
||||
return &ConnectionStats{
|
||||
Name: p.config.GetName(),
|
||||
Type: "mongodb",
|
||||
Connected: true,
|
||||
}
|
||||
}
|
||||
|
||||
// parseReadPreference parses a read preference string into a readpref.ReadPref
|
||||
func parseReadPreference(rp string) (*readpref.ReadPref, error) {
|
||||
switch rp {
|
||||
case "primary":
|
||||
return readpref.Primary(), nil
|
||||
case "primaryPreferred":
|
||||
return readpref.PrimaryPreferred(), nil
|
||||
case "secondary":
|
||||
return readpref.Secondary(), nil
|
||||
case "secondaryPreferred":
|
||||
return readpref.SecondaryPreferred(), nil
|
||||
case "nearest":
|
||||
return readpref.Nearest(), nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unknown read preference: %s", rp)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
_ "github.com/microsoft/go-mssqldb" // MSSQL driver
|
||||
"go.mongodb.org/mongo-driver/mongo"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// MSSQLProvider implements Provider for Microsoft SQL Server databases
|
||||
type MSSQLProvider struct {
|
||||
db *sql.DB
|
||||
config ConnectionConfig
|
||||
}
|
||||
|
||||
// NewMSSQLProvider creates a new MSSQL provider
|
||||
func NewMSSQLProvider() *MSSQLProvider {
|
||||
return &MSSQLProvider{}
|
||||
}
|
||||
|
||||
// Connect establishes a MSSQL connection
|
||||
func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
||||
// Build DSN
|
||||
dsn, err := cfg.BuildDSN()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to build DSN: %w", err)
|
||||
}
|
||||
|
||||
// Connect with retry logic
|
||||
var db *sql.DB
|
||||
var lastErr error
|
||||
|
||||
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
|
||||
|
||||
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||
if attempt > 0 {
|
||||
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Info("Retrying MSSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-time.After(delay):
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
// Open database connection
|
||||
db, err = sql.Open("sqlserver", dsn)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Warn("Failed to open MSSQL connection: %v", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Test the connection with context timeout
|
||||
connectCtx, cancel := context.WithTimeout(ctx, cfg.GetConnectTimeout())
|
||||
err = db.PingContext(connectCtx)
|
||||
cancel()
|
||||
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
db.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Warn("Failed to ping MSSQL database: %v", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Connection successful
|
||||
break
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to connect after %d attempts: %w", retryAttempts, lastErr)
|
||||
}
|
||||
|
||||
// Configure connection pool
|
||||
if cfg.GetMaxOpenConns() != nil {
|
||||
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
|
||||
}
|
||||
if cfg.GetMaxIdleConns() != nil {
|
||||
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
|
||||
}
|
||||
if cfg.GetConnMaxLifetime() != nil {
|
||||
db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime())
|
||||
}
|
||||
if cfg.GetConnMaxIdleTime() != nil {
|
||||
db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime())
|
||||
}
|
||||
|
||||
p.db = db
|
||||
p.config = cfg
|
||||
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Info("MSSQL connection established: name=%s, host=%s, database=%s", cfg.GetName(), cfg.GetHost(), cfg.GetDatabase())
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the MSSQL connection
|
||||
func (p *MSSQLProvider) Close() error {
|
||||
if p.db == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
err := p.db.Close()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to close MSSQL connection: %w", err)
|
||||
}
|
||||
|
||||
if p.config.GetEnableLogging() {
|
||||
logger.Info("MSSQL connection closed: name=%s", p.config.GetName())
|
||||
}
|
||||
|
||||
p.db = nil
|
||||
return nil
|
||||
}
|
||||
|
||||
// HealthCheck verifies the MSSQL connection is alive
|
||||
func (p *MSSQLProvider) HealthCheck(ctx context.Context) error {
|
||||
if p.db == nil {
|
||||
return fmt.Errorf("database connection is nil")
|
||||
}
|
||||
|
||||
// Use a short timeout for health checks
|
||||
healthCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if err := p.db.PingContext(healthCtx); err != nil {
|
||||
return fmt.Errorf("health check failed: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetNative returns the native *sql.DB connection
|
||||
func (p *MSSQLProvider) GetNative() (*sql.DB, error) {
|
||||
if p.db == nil {
|
||||
return nil, fmt.Errorf("database connection is not initialized")
|
||||
}
|
||||
return p.db, nil
|
||||
}
|
||||
|
||||
// GetMongo returns an error for MSSQL (not a MongoDB connection)
|
||||
func (p *MSSQLProvider) GetMongo() (*mongo.Client, error) {
|
||||
return nil, ErrNotMongoDB
|
||||
}
|
||||
|
||||
// Stats returns connection pool statistics
|
||||
func (p *MSSQLProvider) Stats() *ConnectionStats {
|
||||
if p.db == nil {
|
||||
return &ConnectionStats{
|
||||
Name: p.config.GetName(),
|
||||
Type: "mssql",
|
||||
Connected: false,
|
||||
}
|
||||
}
|
||||
|
||||
stats := p.db.Stats()
|
||||
|
||||
return &ConnectionStats{
|
||||
Name: p.config.GetName(),
|
||||
Type: "mssql",
|
||||
Connected: true,
|
||||
OpenConnections: stats.OpenConnections,
|
||||
InUse: stats.InUse,
|
||||
Idle: stats.Idle,
|
||||
WaitCount: stats.WaitCount,
|
||||
WaitDuration: stats.WaitDuration,
|
||||
MaxIdleClosed: stats.MaxIdleClosed,
|
||||
MaxLifetimeClosed: stats.MaxLifetimeClosed,
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,247 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go.mongodb.org/mongo-driver/mongo"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// PostgresProvider implements Provider for PostgreSQL databases
|
||||
type PostgresProvider struct {
|
||||
db *sql.DB
|
||||
connector *pgConnector
|
||||
config ConnectionConfig
|
||||
listener *PostgresListener
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// NewPostgresProvider creates a new PostgreSQL provider
|
||||
func NewPostgresProvider() *PostgresProvider {
|
||||
return &PostgresProvider{}
|
||||
}
|
||||
|
||||
// Connect establishes a PostgreSQL connection
|
||||
func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
||||
connCfg, err := buildPGXConfig(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// The connector and *sql.DB are created once; the pool is never closed to
|
||||
// recover from errors (see Refresh).
|
||||
connector := newPGConnector(connCfg)
|
||||
db := sql.OpenDB(connector)
|
||||
|
||||
// Connect with retry logic
|
||||
var lastErr error
|
||||
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(cfg)
|
||||
|
||||
connected := false
|
||||
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||
if attempt > 0 {
|
||||
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Info("Retrying PostgreSQL connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-time.After(delay):
|
||||
case <-ctx.Done():
|
||||
db.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
// Test the connection with context timeout
|
||||
connectCtx, cancel := context.WithTimeout(ctx, cfg.GetConnectTimeout())
|
||||
err = db.PingContext(connectCtx)
|
||||
cancel()
|
||||
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Warn("Failed to ping PostgreSQL database: %v", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
connected = true
|
||||
break
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
// Configure connection pool
|
||||
if cfg.GetMaxOpenConns() != nil {
|
||||
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
|
||||
}
|
||||
if cfg.GetMaxIdleConns() != nil {
|
||||
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
|
||||
}
|
||||
if cfg.GetConnMaxLifetime() != nil {
|
||||
db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime())
|
||||
}
|
||||
if cfg.GetConnMaxIdleTime() != nil {
|
||||
db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime())
|
||||
}
|
||||
|
||||
p.db = db
|
||||
p.connector = connector
|
||||
p.config = cfg
|
||||
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Info("PostgreSQL connection established: name=%s, host=%s, database=%s", cfg.GetName(), cfg.GetHost(), cfg.GetDatabase())
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Refresh retires every pooled connection and dials fresh ones on demand,
|
||||
// without closing the *sql.DB. Handles already handed out keep working:
|
||||
// connections in use finish their current query and are then discarded.
|
||||
func (p *PostgresProvider) Refresh(ctx context.Context) error {
|
||||
if p.db == nil || p.connector == nil {
|
||||
return fmt.Errorf("database connection is not initialized")
|
||||
}
|
||||
|
||||
connCfg, err := buildPGXConfig(p.config)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.connector.swap(connCfg)
|
||||
|
||||
pingCtx, cancel := context.WithTimeout(ctx, p.config.GetConnectTimeout())
|
||||
defer cancel()
|
||||
if err := p.db.PingContext(pingCtx); err != nil {
|
||||
return fmt.Errorf("failed to ping after refresh: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the PostgreSQL connection. A listener failure does not stop the
|
||||
// pool from being closed.
|
||||
func (p *PostgresProvider) Close() error {
|
||||
var errs []error
|
||||
|
||||
p.mu.Lock()
|
||||
listener := p.listener
|
||||
p.listener = nil
|
||||
p.mu.Unlock()
|
||||
|
||||
if listener != nil {
|
||||
if err := listener.Close(); err != nil {
|
||||
errs = append(errs, fmt.Errorf("failed to close listener: %w", err))
|
||||
}
|
||||
}
|
||||
|
||||
if p.db != nil {
|
||||
if err := p.db.Close(); err != nil {
|
||||
errs = append(errs, fmt.Errorf("failed to close PostgreSQL connection: %w", err))
|
||||
} else if p.config.GetEnableLogging() {
|
||||
logger.Info("PostgreSQL connection closed: name=%s", p.config.GetName())
|
||||
}
|
||||
p.db = nil
|
||||
p.connector = nil
|
||||
}
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
// HealthCheck verifies the PostgreSQL connection is alive
|
||||
func (p *PostgresProvider) HealthCheck(ctx context.Context) error {
|
||||
if p.db == nil {
|
||||
return fmt.Errorf("database connection is nil")
|
||||
}
|
||||
|
||||
// Use a short timeout for health checks
|
||||
healthCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if err := p.db.PingContext(healthCtx); err != nil {
|
||||
return fmt.Errorf("health check failed: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetNative returns the native *sql.DB connection
|
||||
func (p *PostgresProvider) GetNative() (*sql.DB, error) {
|
||||
if p.db == nil {
|
||||
return nil, fmt.Errorf("database connection is not initialized")
|
||||
}
|
||||
return p.db, nil
|
||||
}
|
||||
|
||||
// GetMongo returns an error for PostgreSQL (not a MongoDB connection)
|
||||
func (p *PostgresProvider) GetMongo() (*mongo.Client, error) {
|
||||
return nil, ErrNotMongoDB
|
||||
}
|
||||
|
||||
// Stats returns connection pool statistics
|
||||
func (p *PostgresProvider) Stats() *ConnectionStats {
|
||||
if p.db == nil {
|
||||
return &ConnectionStats{
|
||||
Name: p.config.GetName(),
|
||||
Type: "postgres",
|
||||
Connected: false,
|
||||
}
|
||||
}
|
||||
|
||||
stats := p.db.Stats()
|
||||
|
||||
return &ConnectionStats{
|
||||
Name: p.config.GetName(),
|
||||
Type: "postgres",
|
||||
Connected: true,
|
||||
OpenConnections: stats.OpenConnections,
|
||||
InUse: stats.InUse,
|
||||
Idle: stats.Idle,
|
||||
WaitCount: stats.WaitCount,
|
||||
WaitDuration: stats.WaitDuration,
|
||||
MaxIdleClosed: stats.MaxIdleClosed,
|
||||
MaxLifetimeClosed: stats.MaxLifetimeClosed,
|
||||
}
|
||||
}
|
||||
|
||||
// GetListener returns a PostgreSQL listener for NOTIFY/LISTEN functionality
|
||||
// The listener is lazily initialized on first call and reused for subsequent calls
|
||||
func (p *PostgresProvider) GetListener(ctx context.Context) (*PostgresListener, error) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
|
||||
// Return existing listener if already created
|
||||
if p.listener != nil {
|
||||
return p.listener, nil
|
||||
}
|
||||
|
||||
// Create new listener
|
||||
listener := NewPostgresListener(p.config)
|
||||
|
||||
// Connect the listener
|
||||
if err := listener.Connect(ctx); err != nil {
|
||||
return nil, fmt.Errorf("failed to connect listener: %w", err)
|
||||
}
|
||||
|
||||
p.listener = listener
|
||||
return p.listener, nil
|
||||
}
|
||||
|
||||
// calculateBackoff calculates exponential backoff delay
|
||||
func calculateBackoff(attempt int, initial, maxDelay time.Duration) time.Duration {
|
||||
delay := initial * time.Duration(math.Pow(2, float64(attempt)))
|
||||
if delay > maxDelay {
|
||||
delay = maxDelay
|
||||
}
|
||||
return delay
|
||||
}
|
||||
@@ -0,0 +1,459 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// NotificationHandler is called when a notification is received
|
||||
type NotificationHandler func(channel string, payload string)
|
||||
|
||||
// PostgresListener manages PostgreSQL LISTEN/NOTIFY functionality
|
||||
type PostgresListener struct {
|
||||
config ConnectionConfig
|
||||
conn *pgx.Conn
|
||||
|
||||
// Channel subscriptions
|
||||
channels map[string]NotificationHandler
|
||||
mu sync.RWMutex
|
||||
// connMu serialises use of the single pgx.Conn: it is not safe for
|
||||
// concurrent use, so the notification wait and LISTEN/UNLISTEN/NOTIFY take
|
||||
// turns. Lock order: connMu before mu.
|
||||
connMu sync.Mutex
|
||||
|
||||
// Lifecycle management
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
closed bool
|
||||
closeMu sync.Mutex
|
||||
reconnectC chan struct{}
|
||||
startOnce sync.Once // background goroutines start exactly once
|
||||
}
|
||||
|
||||
// NewPostgresListener creates a new PostgreSQL listener
|
||||
func NewPostgresListener(cfg ConnectionConfig) *PostgresListener {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
return &PostgresListener{
|
||||
config: cfg,
|
||||
channels: make(map[string]NotificationHandler),
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
reconnectC: make(chan struct{}, 1),
|
||||
}
|
||||
}
|
||||
|
||||
// Connect establishes a dedicated connection for listening and starts the
|
||||
// background loops (once per listener).
|
||||
func (l *PostgresListener) Connect(ctx context.Context) error {
|
||||
conn, err := l.dial(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
l.swapConn(conn)
|
||||
|
||||
l.startOnce.Do(func() {
|
||||
go l.handleNotifications()
|
||||
go l.handleReconnection()
|
||||
})
|
||||
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Info("PostgreSQL listener connected: name=%s", l.config.GetName())
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// dial opens and verifies a new dedicated connection, with retries.
|
||||
func (l *PostgresListener) dial(ctx context.Context) (*pgx.Conn, error) {
|
||||
connConfig, err := buildPGXConfig(l.config)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var lastErr error
|
||||
|
||||
retryAttempts, retryDelay, retryMaxDelay := retryPolicy(l.config)
|
||||
|
||||
for attempt := 0; attempt < retryAttempts; attempt++ {
|
||||
if attempt > 0 {
|
||||
delay := calculateBackoff(attempt, retryDelay, retryMaxDelay)
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Info("Retrying PostgreSQL listener connection: attempt=%d/%d, delay=%v", attempt+1, retryAttempts, delay)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-time.After(delay):
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
conn, err := pgx.ConnectConfig(ctx, connConfig)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Warn("Failed to connect PostgreSQL listener: %v", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Test the connection
|
||||
if err = conn.Ping(ctx); err != nil {
|
||||
lastErr = err
|
||||
_ = closeConnBounded(conn)
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Warn("Failed to ping PostgreSQL listener: %v", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("failed to connect listener after %d attempts: %w", retryAttempts, lastErr)
|
||||
}
|
||||
|
||||
// closeConnBounded closes a pgx connection without ever waiting on a dead socket.
|
||||
func closeConnBounded(conn *pgx.Conn) error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), listenerCloseTimeout)
|
||||
defer cancel()
|
||||
return conn.Close(ctx)
|
||||
}
|
||||
|
||||
const (
|
||||
listenerCloseTimeout = 2 * time.Second
|
||||
notificationPollInterval = 500 * time.Millisecond
|
||||
)
|
||||
|
||||
// currentConn returns the live connection, or an error if the listener is
|
||||
// closed or not yet connected.
|
||||
func (l *PostgresListener) currentConn() (*pgx.Conn, error) {
|
||||
l.closeMu.Lock()
|
||||
closed := l.closed
|
||||
l.closeMu.Unlock()
|
||||
if closed {
|
||||
return nil, fmt.Errorf("listener is closed")
|
||||
}
|
||||
|
||||
l.mu.RLock()
|
||||
conn := l.conn
|
||||
l.mu.RUnlock()
|
||||
if conn == nil {
|
||||
return nil, fmt.Errorf("listener connection is not initialized")
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// Listen subscribes to a PostgreSQL notification channel
|
||||
func (l *PostgresListener) Listen(channel string, handler NotificationHandler) error {
|
||||
// Take the connection between notification waits (each wait is short).
|
||||
l.connMu.Lock()
|
||||
defer l.connMu.Unlock()
|
||||
|
||||
conn, err := l.currentConn()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := conn.Exec(l.ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{channel}.Sanitize())); err != nil {
|
||||
return fmt.Errorf("failed to listen on channel %s: %w", channel, err)
|
||||
}
|
||||
|
||||
l.mu.Lock()
|
||||
l.channels[channel] = handler
|
||||
l.mu.Unlock()
|
||||
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Info("Listening on channel: name=%s, channel=%s", l.config.GetName(), channel)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Unlisten unsubscribes from a PostgreSQL notification channel
|
||||
func (l *PostgresListener) Unlisten(channel string) error {
|
||||
l.connMu.Lock()
|
||||
defer l.connMu.Unlock()
|
||||
|
||||
conn, err := l.currentConn()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := conn.Exec(l.ctx, fmt.Sprintf("UNLISTEN %s", pgx.Identifier{channel}.Sanitize())); err != nil {
|
||||
return fmt.Errorf("failed to unlisten from channel %s: %w", channel, err)
|
||||
}
|
||||
|
||||
l.mu.Lock()
|
||||
delete(l.channels, channel)
|
||||
l.mu.Unlock()
|
||||
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Info("Unlistened from channel: name=%s, channel=%s", l.config.GetName(), channel)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Notify sends a notification to a PostgreSQL channel
|
||||
func (l *PostgresListener) Notify(ctx context.Context, channel string, payload string) error {
|
||||
l.connMu.Lock()
|
||||
defer l.connMu.Unlock()
|
||||
|
||||
conn, err := l.currentConn()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := conn.Exec(ctx, "SELECT pg_notify($1, $2)", channel, payload); err != nil {
|
||||
return fmt.Errorf("failed to notify channel %s: %w", channel, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the listener and all subscriptions. Closing the connection drops
|
||||
// every subscription server-side, so no UNLISTEN round trips are needed, and
|
||||
// the close itself is bounded so a dead socket cannot hang the caller.
|
||||
func (l *PostgresListener) Close() error {
|
||||
l.closeMu.Lock()
|
||||
if l.closed {
|
||||
l.closeMu.Unlock()
|
||||
return nil
|
||||
}
|
||||
l.closed = true
|
||||
l.closeMu.Unlock()
|
||||
|
||||
// Cancel context to stop background goroutines
|
||||
l.cancel()
|
||||
|
||||
// The cancelled ctx makes the notification wait return promptly, releasing
|
||||
// connMu; closing the conn while it is being read would race inside pgx.
|
||||
l.connMu.Lock()
|
||||
l.mu.Lock()
|
||||
conn := l.conn
|
||||
l.conn = nil
|
||||
l.channels = make(map[string]NotificationHandler)
|
||||
l.mu.Unlock()
|
||||
|
||||
if conn == nil {
|
||||
l.connMu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
err := closeConnBounded(conn)
|
||||
l.connMu.Unlock()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to close listener connection: %w", err)
|
||||
}
|
||||
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Info("PostgreSQL listener closed: name=%s", l.config.GetName())
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// handleNotifications processes incoming notifications
|
||||
func (l *PostgresListener) handleNotifications() {
|
||||
for {
|
||||
select {
|
||||
case <-l.ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
l.connMu.Lock()
|
||||
l.mu.RLock()
|
||||
conn := l.conn
|
||||
l.mu.RUnlock()
|
||||
|
||||
if conn == nil {
|
||||
l.connMu.Unlock()
|
||||
// Connection not available, wait for reconnection
|
||||
if !l.sleep(100 * time.Millisecond) {
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Wait for a notification with a short timeout, so Listen/Unlisten/Notify
|
||||
// waiting on connMu are served promptly.
|
||||
ctx, cancel := context.WithTimeout(l.ctx, notificationPollInterval)
|
||||
notification, err := conn.WaitForNotification(ctx)
|
||||
cancel()
|
||||
l.connMu.Unlock()
|
||||
|
||||
if err != nil {
|
||||
// Check if context was cancelled
|
||||
if l.ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
|
||||
// Check if it's a connection error
|
||||
if pgconn.Timeout(err) {
|
||||
// Timeout is normal, continue waiting
|
||||
continue
|
||||
}
|
||||
|
||||
// Connection error, trigger reconnection
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Warn("Notification error, triggering reconnection: %v", err)
|
||||
}
|
||||
select {
|
||||
case l.reconnectC <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
if !l.sleep(1 * time.Second) {
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Process notification
|
||||
l.mu.RLock()
|
||||
handler, exists := l.channels[notification.Channel]
|
||||
l.mu.RUnlock()
|
||||
|
||||
if exists && handler != nil {
|
||||
// Call handler in a goroutine to avoid blocking
|
||||
go func(ch, payload string) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Error("Notification handler panic: channel=%s, error=%v", ch, r)
|
||||
}
|
||||
}
|
||||
}()
|
||||
handler(ch, payload)
|
||||
}(notification.Channel, notification.Payload)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// sleep waits for d or until the listener is closed; it reports whether the
|
||||
// listener is still running.
|
||||
func (l *PostgresListener) sleep(d time.Duration) bool {
|
||||
t := time.NewTimer(d)
|
||||
defer t.Stop()
|
||||
select {
|
||||
case <-t.C:
|
||||
return true
|
||||
case <-l.ctx.Done():
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// handleReconnection manages automatic reconnection. It runs as a single
|
||||
// goroutine and dials replacement connections directly rather than through the
|
||||
// public Connect, so no extra loops are started.
|
||||
func (l *PostgresListener) handleReconnection() {
|
||||
for {
|
||||
select {
|
||||
case <-l.ctx.Done():
|
||||
return
|
||||
case <-l.reconnectC:
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Info("Attempting to reconnect listener: name=%s", l.config.GetName())
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(l.ctx, 30*time.Second)
|
||||
err := l.reconnect(ctx)
|
||||
cancel()
|
||||
|
||||
if err != nil {
|
||||
if l.ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Error("Failed to reconnect listener: name=%s, error=%v", l.config.GetName(), err)
|
||||
}
|
||||
// Retry after delay
|
||||
if !l.sleep(5 * time.Second) {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case l.reconnectC <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if l.config.GetEnableLogging() {
|
||||
logger.Info("Listener reconnected successfully: name=%s", l.config.GetName())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// reconnect replaces the connection and resubscribes every channel on the new
|
||||
// connection before publishing it, so the notification loop never touches a
|
||||
// half-initialised conn.
|
||||
func (l *PostgresListener) reconnect(ctx context.Context) error {
|
||||
conn, err := l.dial(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
l.mu.RLock()
|
||||
channels := make([]string, 0, len(l.channels))
|
||||
for ch := range l.channels {
|
||||
channels = append(channels, ch)
|
||||
}
|
||||
l.mu.RUnlock()
|
||||
|
||||
for _, ch := range channels {
|
||||
if _, err := conn.Exec(ctx, fmt.Sprintf("LISTEN %s", pgx.Identifier{ch}.Sanitize())); err != nil {
|
||||
_ = closeConnBounded(conn)
|
||||
return fmt.Errorf("failed to resubscribe to channel %s: %w", ch, err)
|
||||
}
|
||||
}
|
||||
|
||||
if l.ctx.Err() != nil {
|
||||
_ = closeConnBounded(conn)
|
||||
return l.ctx.Err()
|
||||
}
|
||||
l.swapConn(conn)
|
||||
return nil
|
||||
}
|
||||
|
||||
// swapConn installs conn and closes the previous one. The old connection is
|
||||
// closed under connMu so it is never closed while another goroutine is using it.
|
||||
func (l *PostgresListener) swapConn(conn *pgx.Conn) {
|
||||
l.connMu.Lock()
|
||||
l.mu.Lock()
|
||||
old := l.conn
|
||||
l.conn = conn
|
||||
l.mu.Unlock()
|
||||
if old != nil {
|
||||
_ = closeConnBounded(old)
|
||||
}
|
||||
l.connMu.Unlock()
|
||||
}
|
||||
|
||||
// IsConnected returns true if the listener is connected
|
||||
func (l *PostgresListener) IsConnected() bool {
|
||||
l.mu.RLock()
|
||||
defer l.mu.RUnlock()
|
||||
return l.conn != nil
|
||||
}
|
||||
|
||||
// Channels returns the list of channels currently being listened to
|
||||
func (l *PostgresListener) Channels() []string {
|
||||
l.mu.RLock()
|
||||
defer l.mu.RUnlock()
|
||||
|
||||
channels := make([]string, 0, len(l.channels))
|
||||
for ch := range l.channels {
|
||||
channels = append(channels, ch)
|
||||
}
|
||||
return channels
|
||||
}
|
||||
@@ -0,0 +1,229 @@
|
||||
package providers_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbmanager"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbmanager/providers"
|
||||
)
|
||||
|
||||
// ExamplePostgresListener_basic demonstrates basic LISTEN/NOTIFY usage
|
||||
func ExamplePostgresListener_basic() {
|
||||
// Create a connection config
|
||||
cfg := &dbmanager.ConnectionConfig{
|
||||
Name: "example",
|
||||
Type: dbmanager.DatabaseTypePostgreSQL,
|
||||
Host: "localhost",
|
||||
Port: 5432,
|
||||
User: "postgres",
|
||||
Password: "password",
|
||||
Database: "testdb",
|
||||
ConnectTimeout: 10 * time.Second,
|
||||
EnableLogging: true,
|
||||
}
|
||||
|
||||
// Create and connect PostgreSQL provider
|
||||
provider := providers.NewPostgresProvider()
|
||||
ctx := context.Background()
|
||||
|
||||
if err := provider.Connect(ctx, cfg); err != nil {
|
||||
log.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer provider.Close()
|
||||
|
||||
// Get listener
|
||||
listener, err := provider.GetListener(ctx)
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to get listener: %v", err)
|
||||
}
|
||||
|
||||
// Subscribe to a channel with a handler
|
||||
err = listener.Listen("user_events", func(channel, payload string) {
|
||||
fmt.Printf("Received notification on %s: %s\n", channel, payload)
|
||||
})
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to listen: %v", err)
|
||||
}
|
||||
|
||||
// Send a notification
|
||||
err = listener.Notify(ctx, "user_events", `{"event":"user_created","user_id":123}`)
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to notify: %v", err)
|
||||
}
|
||||
|
||||
// Wait for notification to be processed
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
// Unsubscribe from the channel
|
||||
if err := listener.Unlisten("user_events"); err != nil {
|
||||
log.Fatalf("Failed to unlisten: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ExamplePostgresListener_multipleChannels demonstrates listening to multiple channels
|
||||
func ExamplePostgresListener_multipleChannels() {
|
||||
cfg := &dbmanager.ConnectionConfig{
|
||||
Name: "example",
|
||||
Type: dbmanager.DatabaseTypePostgreSQL,
|
||||
Host: "localhost",
|
||||
Port: 5432,
|
||||
User: "postgres",
|
||||
Password: "password",
|
||||
Database: "testdb",
|
||||
ConnectTimeout: 10 * time.Second,
|
||||
EnableLogging: false,
|
||||
}
|
||||
|
||||
provider := providers.NewPostgresProvider()
|
||||
ctx := context.Background()
|
||||
|
||||
if err := provider.Connect(ctx, cfg); err != nil {
|
||||
log.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer provider.Close()
|
||||
|
||||
listener, err := provider.GetListener(ctx)
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to get listener: %v", err)
|
||||
}
|
||||
|
||||
// Listen to multiple channels
|
||||
channels := []string{"orders", "payments", "notifications"}
|
||||
for _, ch := range channels {
|
||||
channel := ch // Capture for closure
|
||||
err := listener.Listen(channel, func(ch, payload string) {
|
||||
fmt.Printf("[%s] %s\n", ch, payload)
|
||||
})
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to listen on %s: %v", channel, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Send notifications to different channels
|
||||
listener.Notify(ctx, "orders", "New order #12345")
|
||||
listener.Notify(ctx, "payments", "Payment received $99.99")
|
||||
listener.Notify(ctx, "notifications", "Welcome email sent")
|
||||
|
||||
// Wait for notifications
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
|
||||
// Check active channels
|
||||
activeChannels := listener.Channels()
|
||||
fmt.Printf("Listening to %d channels: %v\n", len(activeChannels), activeChannels)
|
||||
}
|
||||
|
||||
// ExamplePostgresListener_withDBManager demonstrates usage with DBManager
|
||||
func ExamplePostgresListener_withDBManager() {
|
||||
// This example shows how to use the listener with the full DBManager
|
||||
|
||||
// Assume we have a DBManager instance and get a connection
|
||||
// conn, _ := dbMgr.Get("primary")
|
||||
|
||||
// Get the underlying provider (this would need to be exposed via the Connection interface)
|
||||
// For now, this is a conceptual example
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Create provider directly for demonstration
|
||||
cfg := &dbmanager.ConnectionConfig{
|
||||
Name: "primary",
|
||||
Type: dbmanager.DatabaseTypePostgreSQL,
|
||||
Host: "localhost",
|
||||
Port: 5432,
|
||||
User: "postgres",
|
||||
Password: "password",
|
||||
Database: "myapp",
|
||||
ConnectTimeout: 10 * time.Second,
|
||||
}
|
||||
|
||||
provider := providers.NewPostgresProvider()
|
||||
if err := provider.Connect(ctx, cfg); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
defer provider.Close()
|
||||
|
||||
// Get listener
|
||||
listener, err := provider.GetListener(ctx)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
|
||||
// Subscribe to application events
|
||||
listener.Listen("cache_invalidation", func(channel, payload string) {
|
||||
fmt.Printf("Cache invalidation request: %s\n", payload)
|
||||
// Handle cache invalidation logic here
|
||||
})
|
||||
|
||||
listener.Listen("config_reload", func(channel, payload string) {
|
||||
fmt.Printf("Configuration reload request: %s\n", payload)
|
||||
// Handle configuration reload logic here
|
||||
})
|
||||
|
||||
// Simulate receiving notifications
|
||||
listener.Notify(ctx, "cache_invalidation", "user:123")
|
||||
listener.Notify(ctx, "config_reload", "database")
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
|
||||
// ExamplePostgresListener_errorHandling demonstrates error handling and reconnection
|
||||
func ExamplePostgresListener_errorHandling() {
|
||||
cfg := &dbmanager.ConnectionConfig{
|
||||
Name: "example",
|
||||
Type: dbmanager.DatabaseTypePostgreSQL,
|
||||
Host: "localhost",
|
||||
Port: 5432,
|
||||
User: "postgres",
|
||||
Password: "password",
|
||||
Database: "testdb",
|
||||
ConnectTimeout: 10 * time.Second,
|
||||
EnableLogging: true,
|
||||
}
|
||||
|
||||
provider := providers.NewPostgresProvider()
|
||||
ctx := context.Background()
|
||||
|
||||
if err := provider.Connect(ctx, cfg); err != nil {
|
||||
log.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer provider.Close()
|
||||
|
||||
listener, err := provider.GetListener(ctx)
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to get listener: %v", err)
|
||||
}
|
||||
|
||||
// The listener automatically reconnects if the connection is lost
|
||||
// Subscribe with error handling in the callback
|
||||
err = listener.Listen("critical_events", func(channel, payload string) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
fmt.Printf("Handler panic recovered: %v\n", r)
|
||||
}
|
||||
}()
|
||||
|
||||
// Process the event
|
||||
fmt.Printf("Processing critical event: %s\n", payload)
|
||||
|
||||
// If processing fails, the panic will be caught by the defer above
|
||||
// The listener will continue to function normally
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
fmt.Printf("Failed to listen: %v\n", err)
|
||||
return
|
||||
}
|
||||
|
||||
// Check if listener is connected
|
||||
if listener.IsConnected() {
|
||||
fmt.Println("Listener is connected and ready")
|
||||
}
|
||||
|
||||
// Send a notification
|
||||
listener.Notify(ctx, "critical_events", "system_alert")
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
@@ -0,0 +1,108 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"go.mongodb.org/mongo-driver/mongo"
|
||||
)
|
||||
|
||||
// Common errors
|
||||
var (
|
||||
// ErrNotSQLDatabase is returned when attempting SQL operations on a non-SQL database
|
||||
ErrNotSQLDatabase = errors.New("not a SQL database")
|
||||
|
||||
// ErrNotMongoDB is returned when attempting MongoDB operations on a non-MongoDB connection
|
||||
ErrNotMongoDB = errors.New("not a MongoDB connection")
|
||||
)
|
||||
|
||||
// ConnectionStats contains statistics about a database connection
|
||||
type ConnectionStats struct {
|
||||
Name string
|
||||
Type string // Database type as string to avoid circular dependency
|
||||
Connected bool
|
||||
LastHealthCheck time.Time
|
||||
HealthCheckStatus string
|
||||
|
||||
// SQL connection pool stats
|
||||
OpenConnections int
|
||||
InUse int
|
||||
Idle int
|
||||
WaitCount int64
|
||||
WaitDuration time.Duration
|
||||
MaxIdleClosed int64
|
||||
MaxLifetimeClosed int64
|
||||
}
|
||||
|
||||
// ConnectionConfig is a minimal interface for configuration
|
||||
// The actual implementation is in dbmanager package
|
||||
type ConnectionConfig interface {
|
||||
BuildDSN() (string, error)
|
||||
GetName() string
|
||||
GetType() string
|
||||
GetHost() string
|
||||
GetPort() int
|
||||
GetUser() string
|
||||
GetPassword() string
|
||||
GetDatabase() string
|
||||
GetFilePath() string
|
||||
GetConnectTimeout() time.Duration
|
||||
GetQueryTimeout() time.Duration
|
||||
GetEnableLogging() bool
|
||||
GetEnableMetrics() bool
|
||||
GetMaxOpenConns() *int
|
||||
GetMaxIdleConns() *int
|
||||
GetConnMaxLifetime() *time.Duration
|
||||
GetConnMaxIdleTime() *time.Duration
|
||||
GetReadPreference() string
|
||||
GetRetryAttempts() int
|
||||
GetRetryDelay() time.Duration
|
||||
GetRetryMaxDelay() time.Duration
|
||||
}
|
||||
|
||||
// retryPolicy returns the configured retry settings, falling back to defaults.
|
||||
func retryPolicy(cfg ConnectionConfig) (attempts int, delay, maxDelay time.Duration) {
|
||||
attempts, delay, maxDelay = cfg.GetRetryAttempts(), cfg.GetRetryDelay(), cfg.GetRetryMaxDelay()
|
||||
if attempts <= 0 {
|
||||
attempts = 3
|
||||
}
|
||||
if delay <= 0 {
|
||||
delay = time.Second
|
||||
}
|
||||
if maxDelay <= 0 {
|
||||
maxDelay = 10 * time.Second
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Refresher is implemented by providers that can retire their pooled
|
||||
// connections and dial fresh ones without closing the shared *sql.DB, so
|
||||
// handles already handed out keep working.
|
||||
type Refresher interface {
|
||||
Refresh(ctx context.Context) error
|
||||
}
|
||||
|
||||
// Provider creates and manages the underlying database connection
|
||||
type Provider interface {
|
||||
// Connect establishes the database connection
|
||||
Connect(ctx context.Context, cfg ConnectionConfig) error
|
||||
|
||||
// Close closes the connection
|
||||
Close() error
|
||||
|
||||
// HealthCheck verifies the connection is alive
|
||||
HealthCheck(ctx context.Context) error
|
||||
|
||||
// GetNative returns the native *sql.DB (SQL databases only)
|
||||
// Returns an error for non-SQL databases
|
||||
GetNative() (*sql.DB, error)
|
||||
|
||||
// GetMongo returns the MongoDB client (MongoDB only)
|
||||
// Returns an error for non-MongoDB databases
|
||||
GetMongo() (*mongo.Client, error)
|
||||
|
||||
// Stats returns connection statistics
|
||||
Stats() *ConnectionStats
|
||||
}
|
||||
@@ -0,0 +1,192 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
_ "github.com/glebarez/sqlite" // Pure Go SQLite driver
|
||||
"go.mongodb.org/mongo-driver/mongo"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// SQLiteProvider implements Provider for SQLite databases
|
||||
type SQLiteProvider struct {
|
||||
db *sql.DB
|
||||
dbMu sync.RWMutex
|
||||
config ConnectionConfig
|
||||
}
|
||||
|
||||
// NewSQLiteProvider creates a new SQLite provider
|
||||
func NewSQLiteProvider() *SQLiteProvider {
|
||||
return &SQLiteProvider{}
|
||||
}
|
||||
|
||||
// isMemoryDSN reports whether the SQLite DSN refers to a private in-memory
|
||||
// database (each pooled connection would get its own empty database).
|
||||
func isMemoryDSN(dsn string) bool {
|
||||
path := dsn
|
||||
if i := strings.IndexByte(path, '?'); i >= 0 {
|
||||
path = path[:i]
|
||||
}
|
||||
if path == ":memory:" || path == "" {
|
||||
return true
|
||||
}
|
||||
if strings.Contains(dsn, "mode=memory") && !strings.Contains(dsn, "cache=shared") {
|
||||
return true
|
||||
}
|
||||
return path == "file::memory:" && !strings.Contains(dsn, "cache=shared")
|
||||
}
|
||||
|
||||
// Connect establishes a SQLite connection
|
||||
func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) error {
|
||||
// Build DSN
|
||||
dsn, err := cfg.BuildDSN()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to build DSN: %w", err)
|
||||
}
|
||||
|
||||
// Open database connection
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to open SQLite connection: %w", err)
|
||||
}
|
||||
|
||||
// Test the connection with context timeout
|
||||
connectCtx, cancel := context.WithTimeout(ctx, cfg.GetConnectTimeout())
|
||||
err = db.PingContext(connectCtx)
|
||||
cancel()
|
||||
|
||||
if err != nil {
|
||||
db.Close() //nolint:gosec // G104: best-effort call, error intentionally ignored
|
||||
return fmt.Errorf("failed to ping SQLite database: %w", err)
|
||||
}
|
||||
|
||||
if isMemoryDSN(dsn) {
|
||||
// A private in-memory database exists per connection and disappears when
|
||||
// that connection closes, so pin the pool to one connection that is
|
||||
// never recycled.
|
||||
db.SetMaxOpenConns(1)
|
||||
db.SetMaxIdleConns(1)
|
||||
db.SetConnMaxLifetime(0)
|
||||
db.SetConnMaxIdleTime(0)
|
||||
} else {
|
||||
// SQLite works best with few writers; default to 1 unless configured.
|
||||
if cfg.GetMaxOpenConns() != nil {
|
||||
db.SetMaxOpenConns(*cfg.GetMaxOpenConns())
|
||||
} else {
|
||||
db.SetMaxOpenConns(1)
|
||||
}
|
||||
if cfg.GetMaxIdleConns() != nil {
|
||||
db.SetMaxIdleConns(*cfg.GetMaxIdleConns())
|
||||
}
|
||||
if cfg.GetConnMaxLifetime() != nil {
|
||||
db.SetConnMaxLifetime(*cfg.GetConnMaxLifetime())
|
||||
}
|
||||
if cfg.GetConnMaxIdleTime() != nil {
|
||||
db.SetConnMaxIdleTime(*cfg.GetConnMaxIdleTime())
|
||||
}
|
||||
}
|
||||
|
||||
p.dbMu.Lock()
|
||||
p.db = db
|
||||
p.dbMu.Unlock()
|
||||
p.config = cfg
|
||||
|
||||
if cfg.GetEnableLogging() {
|
||||
logger.Info("SQLite connection established: name=%s, filepath=%s", cfg.GetName(), cfg.GetFilePath())
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the SQLite connection
|
||||
func (p *SQLiteProvider) Close() error {
|
||||
if p.db == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
err := p.db.Close()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to close SQLite connection: %w", err)
|
||||
}
|
||||
|
||||
if p.config.GetEnableLogging() {
|
||||
logger.Info("SQLite connection closed: name=%s", p.config.GetName())
|
||||
}
|
||||
|
||||
p.db = nil
|
||||
return nil
|
||||
}
|
||||
|
||||
// HealthCheck verifies the SQLite connection is alive
|
||||
func (p *SQLiteProvider) HealthCheck(ctx context.Context) error {
|
||||
if p.db == nil {
|
||||
return fmt.Errorf("database connection is nil")
|
||||
}
|
||||
|
||||
// Use a short timeout for health checks
|
||||
healthCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// Execute a simple query to verify the database is accessible
|
||||
var result int
|
||||
if err := p.getDB().QueryRowContext(healthCtx, "SELECT 1").Scan(&result); err != nil {
|
||||
return fmt.Errorf("health check failed: %w", err)
|
||||
}
|
||||
|
||||
if result != 1 {
|
||||
return fmt.Errorf("health check returned unexpected result: %d", result)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *SQLiteProvider) getDB() *sql.DB {
|
||||
p.dbMu.RLock()
|
||||
defer p.dbMu.RUnlock()
|
||||
return p.db
|
||||
}
|
||||
|
||||
// GetNative returns the native *sql.DB connection
|
||||
func (p *SQLiteProvider) GetNative() (*sql.DB, error) {
|
||||
if p.db == nil {
|
||||
return nil, fmt.Errorf("database connection is not initialized")
|
||||
}
|
||||
return p.db, nil
|
||||
}
|
||||
|
||||
// GetMongo returns an error for SQLite (not a MongoDB connection)
|
||||
func (p *SQLiteProvider) GetMongo() (*mongo.Client, error) {
|
||||
return nil, ErrNotMongoDB
|
||||
}
|
||||
|
||||
// Stats returns connection pool statistics
|
||||
func (p *SQLiteProvider) Stats() *ConnectionStats {
|
||||
if p.db == nil {
|
||||
return &ConnectionStats{
|
||||
Name: p.config.GetName(),
|
||||
Type: "sqlite",
|
||||
Connected: false,
|
||||
}
|
||||
}
|
||||
|
||||
stats := p.db.Stats()
|
||||
|
||||
return &ConnectionStats{
|
||||
Name: p.config.GetName(),
|
||||
Type: "sqlite",
|
||||
Connected: true,
|
||||
OpenConnections: stats.OpenConnections,
|
||||
InUse: stats.InUse,
|
||||
Idle: stats.Idle,
|
||||
WaitCount: stats.WaitCount,
|
||||
WaitDuration: stats.WaitDuration,
|
||||
MaxIdleClosed: stats.MaxIdleClosed,
|
||||
MaxLifetimeClosed: stats.MaxLifetimeClosed,
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user