mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-09-28 19:12:00 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
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 |
@@ -1,5 +1,7 @@
|
||||
.PHONY: test test-unit 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..."
|
||||
@@ -49,7 +51,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"; \
|
||||
@@ -58,7 +62,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"; \
|
||||
|
||||
@@ -135,6 +135,73 @@ For complete documentation including setup, headers, lifecycle hooks, cursor pag
|
||||
|
||||
For detailed examples of reading data, cursor pagination, recursive CRUD operations, filtering, sorting, and more, see [pkg/resolvespec/README.md](pkg/resolvespec/README.md).
|
||||
|
||||
## PostGIS & Vector (PostgreSQL only)
|
||||
|
||||
First-class support for PostGIS geometry/geography and pgvector columns in `resolvespec` + `restheadspec`. No extra dependencies. On non-Postgres databases the spatial/vector operators simply don't match.
|
||||
|
||||
### Column types (`pkg/spectypes`)
|
||||
|
||||
| Go type | SQL type | Wire / JSON |
|
||||
|--------------------|--------------|--------------------------------------------------------|
|
||||
| `SqlGeometry` | `geometry` | JSON in/out = **GeoJSON**; also accepts EWKT / hex-EWKB |
|
||||
| `SqlGeography` | `geography` | same as `SqlGeometry` |
|
||||
| `SqlVector` | `vector` | `[]float32` ⇄ `[1,2,3]` |
|
||||
| `SqlHalfVector` | `halfvec` | `[]float32` ⇄ `[1,2,3]` |
|
||||
| `SqlSparseVector` | `sparsevec` | `{"dim":8,"indices":[1,4],"values":[0.5,0.2]}` |
|
||||
| `SqlBitVector` | `bit`/`varbit` | bool array or `"1011"` string |
|
||||
|
||||
- Geometry `Value()` emits `SRID=<n>;<WKT>` (PostGIS implicit text→geometry cast; no wrapper function needed).
|
||||
- Declare dimensioned types with a tag: `gorm:"type:vector(1536)"` — the tag wins over the canonical name in metadata/OpenAPI.
|
||||
- Metadata endpoint and OpenAPI schema report `geometry`/`vector`/`halfvec`/`sparsevec`/`bit`.
|
||||
|
||||
### Spatial filter operators
|
||||
|
||||
`value` is a geometry (GeoJSON object, EWKT string, or hex-EWKB) unless noted.
|
||||
|
||||
| Operator | Value shape |
|
||||
|----------|-------------|
|
||||
| `st_intersects`, `st_contains`, `st_within`, `st_covers`, `st_coveredby`, `st_overlaps`, `st_touches`, `st_crosses`, `st_equals`, `st_disjoint` | geometry |
|
||||
| `st_dwithin` | `{"geom": <geometry>, "distance": <meters>}` |
|
||||
| `bbox` (alias `&&`) | geometry, or `{"bbox":[minx,miny,maxx,maxy],"srid":4326}` |
|
||||
|
||||
### Vector similarity filter operators
|
||||
|
||||
| Operator | pgvector op | Value shape |
|
||||
|----------|-------------|-------------|
|
||||
| `l2_within` / `euclidean_within` | `<->` | `{"vector":[...], "distance": <n>}` |
|
||||
| `cosine_within` | `<=>` | same (also `"lt"`/`"lte"`/`"gt"`/`"gte"` instead of `"distance"`) |
|
||||
| `ip_within` / `inner_within` | `<#>` | same |
|
||||
|
||||
### KNN search (ordering + distance column)
|
||||
|
||||
**resolvespec** — `options.vector_search`:
|
||||
|
||||
```json
|
||||
{ "options": { "vector_search": {
|
||||
"column": "embedding",
|
||||
"vector": [0.1, 0.2, 0.3],
|
||||
"metric": "cosine",
|
||||
"as": "_distance",
|
||||
"direction": "asc"
|
||||
}}}
|
||||
```
|
||||
|
||||
Orders rows by distance; when `as` is set, returns the distance as an extra column (all model columns are auto-selected).
|
||||
`metric`: `l2` (default) | `cosine` | `ip`.
|
||||
|
||||
**restheadspec** — headers:
|
||||
|
||||
```HTTP
|
||||
X-Vector-Search-embedding: cosine
|
||||
X-Vector-Search-Vector: [0.1,0.2,0.3]
|
||||
X-Vector-Search-As: _distance
|
||||
X-Vector-Search-Dir: asc
|
||||
```
|
||||
|
||||
Spatial/vector filters via headers: `X-SpatialFilter-<col>` / `X-VectorFilter-<col>` with a JSON operator object, e.g.
|
||||
`X-SpatialFilter-geom: {"op":"st_dwithin","geom":"SRID=4326;POINT(0 0)","distance":1000}`
|
||||
(optional `"logic":"or"`).
|
||||
|
||||
## Installation
|
||||
|
||||
```Shell
|
||||
|
||||
@@ -8,6 +8,7 @@ require (
|
||||
github.com/eclipse/paho.mqtt.golang v1.5.1
|
||||
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
|
||||
@@ -32,13 +33,13 @@ require (
|
||||
github.com/uptrace/bun/driver/sqliteshim v1.2.16
|
||||
github.com/uptrace/bunrouter v1.0.23
|
||||
go.mongodb.org/mongo-driver v1.17.9
|
||||
go.opentelemetry.io/otel v1.43.0
|
||||
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.43.0
|
||||
go.opentelemetry.io/otel/trace 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.51.0
|
||||
golang.org/x/crypto v0.55.0
|
||||
golang.org/x/oauth2 v0.36.0
|
||||
golang.org/x/time v0.15.0
|
||||
gorm.io/driver/postgres v1.6.0
|
||||
@@ -61,7 +62,6 @@ require (
|
||||
github.com/containerd/platforms v0.2.1 // indirect
|
||||
github.com/cpuguy83/dockercfg v0.3.2 // indirect
|
||||
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
|
||||
github.com/distribution/reference v0.6.0 // indirect
|
||||
github.com/docker/docker v28.5.1+incompatible // indirect
|
||||
github.com/docker/go-connections v0.6.0 // indirect
|
||||
@@ -137,22 +137,21 @@ require (
|
||||
github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 // indirect
|
||||
github.com/yusufpapurcu/wmi v1.2.4 // indirect
|
||||
go.opentelemetry.io/auto/sdk v1.2.1 // indirect
|
||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 // indirect
|
||||
go.opentelemetry.io/otel/metric v1.43.0 // 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-20260508232706-74f9aab9d74a // indirect
|
||||
golang.org/x/mod v0.36.0 // indirect
|
||||
golang.org/x/net v0.54.0 // indirect
|
||||
golang.org/x/sync v0.20.0 // indirect
|
||||
golang.org/x/sys v0.44.0 // indirect
|
||||
golang.org/x/text v0.37.0 // indirect
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20260519071638-aa98bba5eb94 // indirect
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260519071638-aa98bba5eb94 // indirect
|
||||
google.golang.org/grpc v1.81.1 // 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/sys v0.47.0 // indirect
|
||||
golang.org/x/text v0.41.0 // indirect
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||
google.golang.org/grpc v1.83.2 // indirect
|
||||
google.golang.org/protobuf v1.36.11 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
modernc.org/libc v1.72.3 // indirect
|
||||
|
||||
@@ -5,43 +5,35 @@ github.com/AdaLogics/go-fuzz-headers v0.0.0-20240806141605-e8a1dd7889d6/go.mod h
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.0/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.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.18.0 h1:Gt0j3wceWMwPmiazCa8MzMA0MfhmPIz0Qp0FJ6qcM0U=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.18.0/go.mod h1:Ot/6aikWnKWi4l9QB7qVSwa8iMphQNqkWALMoNT3rzM=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1 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.10.1 h1:B+blDbyVIG3WaikNxPnhPiJ1MThR03b3vKGtER95TP4=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.10.1/go.mod h1:JdM5psgjfBf5fo2uWOZhflPWyDBZ/O/CNAH9CtsuZE4=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1 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.11.1 h1:FPKJS1T+clwv+OLGt13a8UjqeRuh0O4SJ3lUriThc+4=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.11.1/go.mod h1:j2chePtV91HrC22tGoRX3sGY42uF13WzmmV80/OdVAA=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.12.0 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.3.1 h1:Wgf5rZba3YZqeTNJPtvqZoBu1sBN/L4sry+u2U3Y75w=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.3.1/go.mod h1:xxCBG/f/4Vbmh2XQJBsOmNdxWUY5j/s27jujKPbQf14=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.4.0 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.1.1 h1:bFWuoEKg+gImo7pvkiQEFAc8ocibADgXeiLAxWhWmkI=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.1.1/go.mod h1:Vih/3yc6yac2JzU4hzpaDupBJP0Flaia9rXXrU8xyww=
|
||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.2.0 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.4.2 h1:oygO0locgZJe7PpYPXT5A29ZkwJaPqcva7BVeemZOZs=
|
||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.4.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.6.0 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=
|
||||
@@ -68,14 +60,13 @@ github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ3
|
||||
github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY=
|
||||
github.com/creack/pty v1.1.18/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/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/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/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=
|
||||
@@ -94,12 +85,8 @@ 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/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.40.0 h1:VTJMN9zbTvqDqPwheRVLcp0qcUcM+8eFivvGocAaSbo=
|
||||
github.com/getsentry/sentry-go v0.40.0/go.mod h1:eRXCoh3uvmjQLY6qu63BjUZnaBu5L5WhMV1RwYO8W5s=
|
||||
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=
|
||||
@@ -115,16 +102,13 @@ github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
|
||||
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
|
||||
github.com/go-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.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo=
|
||||
github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
||||
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=
|
||||
@@ -137,8 +121,6 @@ github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/
|
||||
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
||||
github.com/google/jsonschema-go v0.4.2 h1:tmrUohrwoLZZS/P3x7ex0WAVknEkBZM46iALbcqoRA8=
|
||||
github.com/google/jsonschema-go v0.4.2/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE=
|
||||
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=
|
||||
@@ -152,8 +134,6 @@ github.com/gorilla/securecookie v1.1.1/go.mod h1:ra0sb63/xPlUeL+yeDciTfxMRAA+MP+
|
||||
github.com/gorilla/sessions v1.2.1/go.mod h1:dk2InVEVJ0sfLlnXv9EAgkf6ecYs/i80K/zI+bUmuGM=
|
||||
github.com/gorilla/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=
|
||||
@@ -164,8 +144,6 @@ github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsI
|
||||
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
||||
github.com/jackc/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.8.0 h1:TYPDoleBBme0xGSAX3/+NujXXtpZn9HBONkQC7IEZSo=
|
||||
github.com/jackc/pgx/v5 v5.8.0/go.mod h1:QVeDInX2m9VyzvNeiCJVjCkNFqzsNb43204HshNSZKw=
|
||||
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=
|
||||
@@ -183,10 +161,10 @@ github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkr
|
||||
github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
|
||||
github.com/jinzhu/now v1.1.5/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.2 h1:iiPHWW0YrcFgpBYhsA6D1+fqHssJscY/Tm/y2Uqnapk=
|
||||
github.com/klauspost/compress v1.18.2/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4=
|
||||
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=
|
||||
@@ -200,21 +178,13 @@ 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/mark3labs/mcp-go v0.46.0 h1:8KRibF4wcKejbLsHxCA/QBVUr5fQ9nwz/n8lGqmaALo=
|
||||
github.com/mark3labs/mcp-go v0.46.0/go.mod h1:JKTC7R2LLVagkEWK7Kwu7DbmA6iIvnNAod6yrHiQMag=
|
||||
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.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
||||
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
||||
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.33 h1:A5blZ5ulQo2AtayQ9/limgHEkFreKj1Dv226a1K73s0=
|
||||
github.com/mattn/go-sqlite3 v1.14.33/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y=
|
||||
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.9.5 h1:orwya0X/5bsL1o+KasupTkk2eNTNFkTQG0BEe/HxCn0=
|
||||
github.com/microsoft/go-mssqldb v1.9.5/go.mod h1:VCP2a0KEZZtGLRHd1PsLavLFYy/3xX2yJUPycv3Sr2Q=
|
||||
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=
|
||||
@@ -237,20 +207,14 @@ github.com/mochi-mqtt/server/v2 v2.7.9 h1:y0g4vrSLAag7T07l2oCzOa/+nKVLoazKEWAArw
|
||||
github.com/mochi-mqtt/server/v2 v2.7.9/go.mod h1:lZD3j35AVNqJL5cezlnSkuG05c0FCHSsfAKSPBOSbqc=
|
||||
github.com/modocache/gover v0.0.0-20171022184752-b58185e213c5/go.mod h1:caMODM3PzxT8aQXRPkAt8xlV/e7d7w8GM5g0fa5F0D8=
|
||||
github.com/montanaflynn/stats v0.7.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
||||
github.com/montanaflynn/stats v0.7.1 h1:etflOAAHORrCC44V+aR6Ftzort912ZU+YLiSTuV8eaE=
|
||||
github.com/montanaflynn/stats v0.7.1/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/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.11 h1:q44qGV008kYd9W1b1nEBkNzvnWxtRSQ7A8BoqRrcfa0=
|
||||
github.com/nats-io/nkeys v0.4.11/go.mod h1:szDimtgmfOi9n25JpfIdGw12tZFYXqhGxjhVxsatHVE=
|
||||
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=
|
||||
@@ -261,8 +225,6 @@ 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=
|
||||
@@ -281,29 +243,20 @@ 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.67.4 h1:yR3NqWO1/UyO1w2PhUvXlGQs/PtFmoveVO0KZ4+Lvsc=
|
||||
github.com/prometheus/common v0.67.4/go.mod h1:gP0fq6YjjNCLssJCQp0yk4M8W6ikLURwkdd/YKtTbyI=
|
||||
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.19.2 h1:zUMhqEW66Ex7OXIiDkll3tl9a1ZdilUOd/F6ZXw4Vws=
|
||||
github.com/prometheus/procfs v0.19.2/go.mod h1:M0aotyiemPhBCM0z5w87kL22CxfcH05ZpYlu+b4J7mw=
|
||||
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.2 h1:P2EGsA4qVIM3Pp+aPocCJ7DguDHhqrXNhVcEp4ViluI=
|
||||
github.com/redis/go-redis/v9 v9.17.2/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.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.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII=
|
||||
github.com/rogpeppe/go-internal v1.13.1/go.mod h1:uMEvuHeurkdAXX61udpOXGD/AzZDWNMNyH2VO9fmH0o=
|
||||
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
|
||||
github.com/rs/xid v1.4.0 h1:qd7wPTDkN6KQx2VmMBLrpHkiyQwgFXRnkOLacUiaSNY=
|
||||
github.com/rs/xid v1.4.0/go.mod h1:trrq9SKmegXys3aeAKXMUTdJsYXVwGY3RLcfgqegfbg=
|
||||
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=
|
||||
@@ -344,8 +297,6 @@ 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/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=
|
||||
@@ -364,22 +315,10 @@ github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc h1:9lRDQMhESg+zvGYm
|
||||
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/mssqldialect v1.2.17 h1:xEUH4WamuY9rXT9d8wHVZanhmLJCrc4s4v7frDH/PMc=
|
||||
github.com/uptrace/bun/dialect/mssqldialect v1.2.17/go.mod h1:i1NRx/5cz1nivwtV7FEb/gP3CIbRTj4AQC9/Q0lNVno=
|
||||
github.com/uptrace/bun/dialect/mssqldialect v1.2.18 h1:nYzHoyJKJlIyl5i95Exi8ZTK8ooKWG+o3z3f404d/yQ=
|
||||
github.com/uptrace/bun/dialect/mssqldialect v1.2.18/go.mod h1:Su45Je7z66sfeZ3d1ZsnOQEK8xfzGgaMzBvtoE8yFhk=
|
||||
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/pgdialect v1.2.17 h1:DFmhOollvbYHvooxoS8ZIbiGC0wXIzstKeFUmWs+TP4=
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.17/go.mod h1:ej8ZDsvLETvyELlRDfUtIoA57sWnATv1GhOEVsuVG/k=
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.18 h1:IZ6nM2+OYrL8lkEAy7UkSEZvoa3vluTAUlZfPtlRB2k=
|
||||
github.com/uptrace/bun/dialect/pgdialect v1.2.18/go.mod h1:Tqdf4QP1okrGYpXfodXvCOK6Ob1OOTwSaoAzCgBB3IU=
|
||||
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/dialect/sqlitedialect v1.2.17 h1:ZipEoNr+wQJQleGy2poKSSoaQDavzc+nXTDp3ZzkA0E=
|
||||
github.com/uptrace/bun/dialect/sqlitedialect v1.2.17/go.mod h1:phXmrxxeYqUhMU09FgazbfNxq9LlArdqjZqHc1ILy9U=
|
||||
github.com/uptrace/bun/dialect/sqlitedialect v1.2.18 h1:Z33SY/U++XK9uGWqS4h8OZVxfCXguIG+sU9cYq2PGFQ=
|
||||
github.com/uptrace/bun/dialect/sqlitedialect v1.2.18/go.mod h1:1MVOS/Ncy4FZbkJcgUFH6OqYoQinYNjkEwsmNQEXz2A=
|
||||
github.com/uptrace/bun/driver/sqliteshim v1.2.16 h1:M6Dh5kkDWFbUWBrOsIE1g1zdZ5JbSytTD4piFRBOUAI=
|
||||
github.com/uptrace/bun/driver/sqliteshim v1.2.16/go.mod h1:iKdJ06P3XS+pwKcONjSIK07bbhksH3lWsw3mpfr0+bY=
|
||||
github.com/uptrace/bunrouter v1.0.23 h1:Bi7NKw3uCQkcA/GUCtDNPq5LE5UdR9pe+UyWbjHB/wU=
|
||||
@@ -403,47 +342,30 @@ github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78/go.mod h1:aL8wCCfTfS
|
||||
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
|
||||
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
|
||||
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
||||
go.mongodb.org/mongo-driver v1.17.6 h1:87JUG1wZfWsr6rIz3ZmpH90rL5tea7O3IHuSwHUpsss=
|
||||
go.mongodb.org/mongo-driver v1.17.6/go.mod h1:Hy04i7O2kC4RS06ZrhPRqj/u4DTYkFDAAccj+rVKqgQ=
|
||||
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.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA=
|
||||
go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A=
|
||||
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.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 v1.43.0 h1:mYIM03dnh5zfN7HautFE4ieIig9amkNANT+xcVxAj9I=
|
||||
go.opentelemetry.io/otel v1.43.0/go.mod h1:JuG+u74mvjvcm8vj8pI5XiHy1zDeoCS2LB1spIq7Ay0=
|
||||
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/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.38.0 h1:lwI4Dc5leUqENgGuQImwLo4WnuXFPetmPpkLi2IrX54=
|
||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.38.0/go.mod h1:Kz/oCE7z5wuyhPxsXDuaPteSWqjSBD5YaSdbxZYGbGk=
|
||||
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/metric v1.43.0 h1:d7638QeInOnuwOONPp4JAOGfbCEpYb+K6DVWvdxGzgM=
|
||||
go.opentelemetry.io/otel/metric v1.43.0/go.mod h1:RDnPtIxvqlgO8GRW18W6Z/4P462ldprJtfxHxyKd2PY=
|
||||
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 v1.43.0 h1:pi5mE86i5rTeLXqoF/hhiBtUNcrAGHLKQdhg4h4V9Dg=
|
||||
go.opentelemetry.io/otel/sdk v1.43.0/go.mod h1:P+IkVU3iWukmiit/Yf9AWvpyRDlUeBaRg6Y+C58QHzg=
|
||||
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/sdk/metric v1.43.0 h1:S88dyqXjJkuBNLeMcVPRFXpRw2fuwdvfCGLEo89fDkw=
|
||||
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/otel/trace v1.43.0 h1:BkNrHpup+4k4w+ZZ86CZoHHEkohws8AY+WTX09nk+3A=
|
||||
go.opentelemetry.io/otel/trace v1.43.0/go.mod h1:/QJhyVBUUswCphDVxq+8mld+AvhXZLhe+8WVFxiFff0=
|
||||
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=
|
||||
@@ -452,12 +374,8 @@ 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.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0=
|
||||
go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
||||
go.uber.org/zap v1.27.1 h1:08RqriUEv8+ArZRYSTXy1LeBScaMpVSTBhCeaZYfMYc=
|
||||
go.uber.org/zap v1.27.1/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E=
|
||||
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.3 h1:6gvOSjQoTB3vt1l+CU+tSyi/HOjfOjRLJ4YwYZGwRO0=
|
||||
go.yaml.in/yaml/v2 v2.4.3/go.mod h1:zSxWcmIDjOzPXpjlTTbAsKokqkDNAVtZO0WOMiT90s8=
|
||||
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=
|
||||
@@ -474,24 +392,16 @@ golang.org/x/crypto v0.21.0/go.mod h1:0BP7YvVV9gBbVKyeTG0Gyn+gZm94bibOW5BjDEYAOM
|
||||
golang.org/x/crypto v0.22.0/go.mod h1:vr6Su+7cTlO45qkww3VDJlzDn0ctJvRgYbC2NvXHt+M=
|
||||
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
||||
golang.org/x/crypto v0.24.0/go.mod h1:Z1PMYSOR5nyMcyAVAIQSKCDwalqy85Aqn1x3Ws4L5DM=
|
||||
golang.org/x/crypto v0.46.0 h1:cKRW/pmt1pKAfetfu+RCEvjvZkA9RimPbh7bhFjGVBU=
|
||||
golang.org/x/crypto v0.46.0/go.mod h1:Evb/oLKmMraqjZ2iQTwDwvCtJkczlDuTmdJXoZVzqU0=
|
||||
golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI=
|
||||
golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8=
|
||||
golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93 h1:fQsdNF2N+/YewlRZiricy4P1iimyPKZ/xwniHj8Q2a0=
|
||||
golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93/go.mod h1:EPRbTFwzwjXj9NpYyyrvenVh9Y+GFeEvMNh7Xuz7xgU=
|
||||
golang.org/x/exp v0.0.0-20260508232706-74f9aab9d74a h1:+3jdDGGB8NGb1Zktc737jlt3/A5f6UlwSzmvqUuufxw=
|
||||
golang.org/x/exp v0.0.0-20260508232706-74f9aab9d74a/go.mod h1:d2fgXJLVs4dYDHUk5lwMIfzRzSrWCfGZb0ZqeLa/Vcw=
|
||||
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.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
|
||||
golang.org/x/mod v0.31.0/go.mod h1:43JraMp9cGx1Rx3AqioxrbrhNsLl2l/iNAvuBkrezpg=
|
||||
golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4=
|
||||
golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ=
|
||||
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=
|
||||
@@ -509,12 +419,8 @@ 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.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU=
|
||||
golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY=
|
||||
golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w=
|
||||
golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ=
|
||||
golang.org/x/oauth2 v0.34.0 h1:hqK/t4AKgbqWkdkcAeI8XLmbK+4m4G5YeQRrmiotGlw=
|
||||
golang.org/x/oauth2 v0.34.0/go.mod h1:lzm5WQJQwKZ3nwavOZ3IS5Aulzxi68dUSgRHujetwEA=
|
||||
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=
|
||||
@@ -524,10 +430,8 @@ golang.org/x/sync v0.3.0/go.mod h1:FU7BRWz2tNW+3quACPkgCx/L+uEAv1htQ0V83Z9Rj+Y=
|
||||
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.9.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
|
||||
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
||||
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
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=
|
||||
@@ -551,10 +455,8 @@ golang.org/x/sys v0.18.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.19.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.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.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk=
|
||||
golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
||||
golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ=
|
||||
golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
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=
|
||||
@@ -570,9 +472,8 @@ golang.org/x/term v0.18.0/go.mod h1:ILwASektA3OnRv7amZ1xhE/KTR+u50pbXfZ03+6Nx58=
|
||||
golang.org/x/term v0.19.0/go.mod h1:2CuTdWZ7KHSQwUzKva0cbMg6q2DMI3Mmxp+gKJbskEk=
|
||||
golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
|
||||
golang.org/x/term v0.21.0/go.mod h1:ooXLefLobQVslOqselCNF4SxFAaoS6KujMbsGzSDmX0=
|
||||
golang.org/x/term v0.38.0 h1:PQ5pkm/rLO6HnxFR7N2lJHOZX6Kez5Y1gDSJla6jo7Q=
|
||||
golang.org/x/term v0.38.0/go.mod h1:bSEAKrOT1W+VSu9TSCMtoGEOUcKxOKgl3LE5QEF/xVg=
|
||||
golang.org/x/term v0.43.0 h1:S4RLU2sB31O/NCl+zFN9Aru9A/Cq2aqKpTZJ6B+DwT4=
|
||||
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=
|
||||
@@ -587,12 +488,8 @@ 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.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU=
|
||||
golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY=
|
||||
golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc=
|
||||
golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38=
|
||||
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/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=
|
||||
@@ -601,26 +498,18 @@ 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.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA=
|
||||
golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc=
|
||||
golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8=
|
||||
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=
|
||||
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
|
||||
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/api v0.0.0-20260519071638-aa98bba5eb94 h1:DddG61lE5LkX6144z22i0gma9BMBs5aZ9B8lZLobxyw=
|
||||
google.golang.org/genproto/googleapis/api v0.0.0-20260519071638-aa98bba5eb94/go.mod h1:1dCETSCY2YKZNXQE3h4fun3TYwF5p8jejRKZgfWAgAY=
|
||||
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/genproto/googleapis/rpc v0.0.0-20260519071638-aa98bba5eb94 h1:eZCjr/aAF8c5ccm5pb6T4EXgIei5MlAAPWPJk+5ArfY=
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260519071638-aa98bba5eb94/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
|
||||
google.golang.org/grpc v1.75.0 h1:+TW+dqTd2Biwe6KKfhE5JpiYIBWq865PhKGSXiivqt4=
|
||||
google.golang.org/grpc v1.75.0/go.mod h1:JtPAzKiq4v1xcAB2hydNlWI2RnF85XXcV0mhKXr2ecQ=
|
||||
google.golang.org/grpc v1.81.1 h1:VnnIIZ88UzOOKLukQi+ImGz8O1Wdp8nAGGnvOfEIWQQ=
|
||||
google.golang.org/grpc v1.81.1/go.mod h1:xGH9GfzOyMTGIOXBJmXt+BX/V0kcdQbdcuwQ/zNw42I=
|
||||
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=
|
||||
@@ -644,37 +533,28 @@ 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/cc/v4 v4.28.2 h1:3tQ0lf2ADtoby2EtSP+J7IE2SHwEJdP8ioR59wx7XpY=
|
||||
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
|
||||
modernc.org/ccgo/v4 v4.30.1/go.mod h1:bIOeI1JL54Utlxn+LwrFyjCx2n2RDiYEaJVSrgdrRfM=
|
||||
modernc.org/cc/v4 v4.28.2/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI=
|
||||
modernc.org/ccgo/v4 v4.34.0 h1:yRLPFZieg532OT4rp4JFNIVcquwalMX26G95WQDqwCQ=
|
||||
modernc.org/fileutil v1.3.40 h1:ZGMswMNc9JOCrcrakF1HrvmergNLAmxOPjizirpfqBA=
|
||||
modernc.org/fileutil v1.3.40/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc=
|
||||
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.4 h1:zZGmCMUVPORtKv95c2ReQN5VDjvkoRm9GWPTEPuvlWg=
|
||||
modernc.org/libc v1.67.4/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.42.2 h1:7hkZUNJvJFN2PgfUdjni9Kbvd4ef4mNLOu0B9FGxM74=
|
||||
modernc.org/sqlite v1.42.2/go.mod h1:+VkC6v3pLOAE0A0uVucQEcbVW0I5nHCeDaBf+DpsQT8=
|
||||
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=
|
||||
|
||||
@@ -339,7 +339,7 @@ func (b *BunSelectQuery) Column(columns ...string) common.SelectQuery {
|
||||
|
||||
func (b *BunSelectQuery) ColumnExpr(query string, args ...interface{}) common.SelectQuery {
|
||||
if len(args) > 0 {
|
||||
b.query = b.query.ColumnExpr(query, args)
|
||||
b.query = b.query.ColumnExpr(query, args...)
|
||||
} else {
|
||||
b.query = b.query.ColumnExpr(query)
|
||||
}
|
||||
@@ -669,10 +669,7 @@ func (b *BunSelectQuery) PreloadRelation(relation string, apply ...func(common.S
|
||||
b.query = b.query.Relation(relation, func(sq *bun.SelectQuery) *bun.SelectQuery {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
err := logger.HandlePanic("BunSelectQuery.PreloadRelation", r)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_ = logger.HandlePanic("BunSelectQuery.PreloadRelation", r)
|
||||
}
|
||||
}()
|
||||
if len(apply) == 0 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -82,6 +82,20 @@ type queryMetricsBunUser struct {
|
||||
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)
|
||||
@@ -346,3 +360,35 @@ func TestBunAdapterRecordsEntityAndTableMetrics(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -1040,3 +1040,21 @@ func BuildInCondition(column string, v interface{}) (query string, args []interf
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
+147
-89
@@ -542,8 +542,8 @@ func TestSplitByAND(t *testing.T) {
|
||||
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'",
|
||||
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.
|
||||
@@ -889,93 +889,151 @@ func TestSanitizeWhereClause_PreservesParenthesesWithOR(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')",
|
||||
},
|
||||
// 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')`,
|
||||
},
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -42,6 +42,10 @@ type RequestOptions struct {
|
||||
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:"-"`
|
||||
@@ -90,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"`
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -109,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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -197,7 +197,7 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
||||
hookCtx.Tx = tx
|
||||
|
||||
// Execute BeforeQueryList hook (inside transaction)
|
||||
if err := h.hooks.Execute(BeforeQueryList, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeQueryList, hookCtx); err != nil {
|
||||
logger.Error("BeforeQueryList hook failed: %v", err)
|
||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||
return err
|
||||
@@ -261,7 +261,7 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
||||
|
||||
// Execute BeforeSQLExec hook
|
||||
hookCtx.SQLQuery = sqlquery
|
||||
if err := h.hooks.Execute(BeforeSQLExec, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
|
||||
logger.Error("BeforeSQLExec hook failed: %v", err)
|
||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||
return err
|
||||
@@ -331,7 +331,10 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
||||
w.Header().Set("Content-Range", fmt.Sprintf("items %d-%d/%d", respOffset, respOffset+len(dbobjlist), total))
|
||||
logger.Info("Serving: Records %d of %d", len(dbobjlist), total)
|
||||
|
||||
// Execute BeforeResponse hook
|
||||
// Execute BeforeResponse hook. The transaction has already committed by
|
||||
// this point, so hooks must use the pooled connection rather than the
|
||||
// now-dead tx.
|
||||
hookCtx.Tx = h.db
|
||||
hookCtx.Result = dbobjlist
|
||||
hookCtx.Total = total
|
||||
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
||||
@@ -560,7 +563,7 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
||||
hookCtx.Tx = tx
|
||||
|
||||
// Execute BeforeQuery hook (inside transaction)
|
||||
if err := h.hooks.Execute(BeforeQuery, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeQuery, hookCtx); err != nil {
|
||||
logger.Error("BeforeQuery hook failed: %v", err)
|
||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||
return err
|
||||
@@ -579,7 +582,7 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
||||
sqlquery = hookCtx.SQLQuery
|
||||
|
||||
// Execute BeforeSQLExec hook
|
||||
if err := h.hooks.Execute(BeforeSQLExec, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
|
||||
logger.Error("BeforeSQLExec hook failed: %v", err)
|
||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||
return err
|
||||
@@ -631,7 +634,10 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
||||
return
|
||||
}
|
||||
|
||||
// Execute BeforeResponse hook
|
||||
// Execute BeforeResponse hook. The transaction has already committed by
|
||||
// this point, so hooks must use the pooled connection rather than the
|
||||
// now-dead tx.
|
||||
hookCtx.Tx = h.db
|
||||
hookCtx.Result = dbobj
|
||||
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
||||
logger.Error("BeforeResponse hook failed: %v", err)
|
||||
|
||||
@@ -28,6 +28,10 @@ const (
|
||||
|
||||
// Response hooks (before response is sent)
|
||||
BeforeResponse HookType = "before_response"
|
||||
|
||||
// BeforeOp fires immediately before every SQL operation (query, query list, SQL exec).
|
||||
// It fires at each individual SQL-operation hook point, so it runs once per statement executed.
|
||||
BeforeOp HookType = "before_op"
|
||||
)
|
||||
|
||||
// HookContext contains all the data available to a hook
|
||||
@@ -130,6 +134,16 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
|
||||
// (BeforeQuery, BeforeQueryList, or BeforeSQLExec). BeforeOp always runs first so it can
|
||||
// observe/veto every SQL operation regardless of type.
|
||||
func (r *HookRegistry) ExecuteBeforeOp(hookType HookType, ctx *HookContext) error {
|
||||
if err := r.Execute(BeforeOp, ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
return r.Execute(hookType, ctx)
|
||||
}
|
||||
|
||||
// Clear removes all hooks for the specified type
|
||||
func (r *HookRegistry) Clear(hookType HookType) {
|
||||
delete(r.hooks, hookType)
|
||||
|
||||
@@ -51,7 +51,7 @@ func (h *Handler) ParseParameters(r *http.Request) *RequestParameters {
|
||||
FieldFilters: make(map[string]string),
|
||||
SearchFilters: make(map[string]string),
|
||||
SearchOps: make(map[string]FilterOperator),
|
||||
Limit: 20, // Default limit
|
||||
Limit: 100000, // Default limit
|
||||
Offset: 0, // Default offset
|
||||
ResponseFormat: "simple", // Default format
|
||||
ComplexAPI: false, // Default to simple API
|
||||
|
||||
@@ -71,6 +71,16 @@ func (f *funcSpecSecurityContext) GetUserID() (int, bool) {
|
||||
return int(f.ctx.UserContext.UserID), true
|
||||
}
|
||||
|
||||
// GetUserRef returns an opaque user identifier for row security lookups.
|
||||
// It returns the full *security.UserContext so providers can read JWT claims
|
||||
// (e.g. a UUID subject) instead of relying on the int user ID.
|
||||
func (f *funcSpecSecurityContext) GetUserRef() (any, bool) {
|
||||
if f.ctx.UserContext == nil {
|
||||
return nil, false
|
||||
}
|
||||
return f.ctx.UserContext, true
|
||||
}
|
||||
|
||||
func (f *funcSpecSecurityContext) GetSchema() string {
|
||||
// funcspec doesn't have a schema concept, extract from SQL query or use default
|
||||
return "public"
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ModelRules defines the permissions and security settings for a model
|
||||
@@ -59,12 +60,44 @@ func NewModelRegistry() *DefaultModelRegistry {
|
||||
}
|
||||
}
|
||||
|
||||
// lockRetryAttempts/lockRetryDelay bound how long the try-lock helpers below
|
||||
// will spin before giving up, so a contended registriesMutex can never hang
|
||||
// a caller of GetDefaultRegistry/SetDefaultRegistry.
|
||||
const (
|
||||
lockRetryAttempts = 20
|
||||
lockRetryDelay = 1 * time.Millisecond
|
||||
)
|
||||
|
||||
// GetDefaultRegistry returns the current default registry. It uses a
|
||||
// bounded TryRLock instead of a blocking RLock so it can never hang;
|
||||
// if the lock can't be acquired in time it falls back to the last known
|
||||
// value without synchronization.
|
||||
func GetDefaultRegistry() *DefaultModelRegistry {
|
||||
for i := 0; i < lockRetryAttempts; i++ {
|
||||
if registriesMutex.TryRLock() {
|
||||
defer registriesMutex.RUnlock()
|
||||
return defaultRegistry
|
||||
}
|
||||
time.Sleep(lockRetryDelay)
|
||||
}
|
||||
return defaultRegistry
|
||||
}
|
||||
|
||||
// SetDefaultRegistry replaces the default registry. It uses a bounded
|
||||
// TryLock instead of a blocking Lock so it can never hang; if the lock
|
||||
// can't be acquired in time the call is a no-op.
|
||||
func SetDefaultRegistry(registry *DefaultModelRegistry) {
|
||||
registriesMutex.Lock()
|
||||
acquired := false
|
||||
for i := 0; i < lockRetryAttempts; i++ {
|
||||
if registriesMutex.TryLock() {
|
||||
acquired = true
|
||||
break
|
||||
}
|
||||
time.Sleep(lockRetryDelay)
|
||||
}
|
||||
if !acquired {
|
||||
return
|
||||
}
|
||||
defer registriesMutex.Unlock()
|
||||
|
||||
foundAt := -1
|
||||
@@ -90,8 +123,34 @@ func AddRegistry(registry *DefaultModelRegistry) {
|
||||
registries = append(registries, registry)
|
||||
}
|
||||
|
||||
// tryLock attempts to acquire the registry's write lock, retrying briefly.
|
||||
// Returns false if it could not be acquired within the bound.
|
||||
func (r *DefaultModelRegistry) tryLock() bool {
|
||||
for i := 0; i < lockRetryAttempts; i++ {
|
||||
if r.mutex.TryLock() {
|
||||
return true
|
||||
}
|
||||
time.Sleep(lockRetryDelay)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// tryRLock attempts to acquire the registry's read lock, retrying briefly.
|
||||
// Returns false if it could not be acquired within the bound.
|
||||
func (r *DefaultModelRegistry) tryRLock() bool {
|
||||
for i := 0; i < lockRetryAttempts; i++ {
|
||||
if r.mutex.TryRLock() {
|
||||
return true
|
||||
}
|
||||
time.Sleep(lockRetryDelay)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (r *DefaultModelRegistry) RegisterModel(name string, model interface{}) error {
|
||||
r.mutex.Lock()
|
||||
if !r.tryLock() {
|
||||
return fmt.Errorf("failed to register model %s: registry locked", name)
|
||||
}
|
||||
defer r.mutex.Unlock()
|
||||
|
||||
if _, exists := r.models[name]; exists {
|
||||
@@ -137,7 +196,9 @@ func (r *DefaultModelRegistry) RegisterModel(name string, model interface{}) err
|
||||
}
|
||||
|
||||
func (r *DefaultModelRegistry) GetModel(name string) (interface{}, error) {
|
||||
r.mutex.RLock()
|
||||
if !r.tryRLock() {
|
||||
return nil, fmt.Errorf("failed to get model %s: registry locked", name)
|
||||
}
|
||||
defer r.mutex.RUnlock()
|
||||
|
||||
model, exists := r.models[name]
|
||||
@@ -149,7 +210,9 @@ func (r *DefaultModelRegistry) GetModel(name string) (interface{}, error) {
|
||||
}
|
||||
|
||||
func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
|
||||
r.mutex.RLock()
|
||||
if !r.tryRLock() {
|
||||
return make(map[string]interface{})
|
||||
}
|
||||
defer r.mutex.RUnlock()
|
||||
|
||||
result := make(map[string]interface{})
|
||||
@@ -253,14 +316,26 @@ func IterateModels(fn func(name string, model interface{})) {
|
||||
// GetModels returns a list of all models from all registries
|
||||
// Models are collected in registry order, with duplicates included
|
||||
func GetModels() []interface{} {
|
||||
registriesMutex.RLock()
|
||||
acquired := false
|
||||
for i := 0; i < lockRetryAttempts; i++ {
|
||||
if registriesMutex.TryRLock() {
|
||||
acquired = true
|
||||
break
|
||||
}
|
||||
time.Sleep(lockRetryDelay)
|
||||
}
|
||||
if !acquired {
|
||||
return nil
|
||||
}
|
||||
defer registriesMutex.RUnlock()
|
||||
|
||||
var models []interface{}
|
||||
seen := make(map[string]bool)
|
||||
|
||||
for _, registry := range registries {
|
||||
registry.mutex.RLock()
|
||||
if !registry.tryRLock() {
|
||||
continue
|
||||
}
|
||||
for name, model := range registry.models {
|
||||
// Only add the first occurrence of each model name
|
||||
if !seen[name] {
|
||||
|
||||
+56
-8
@@ -313,7 +313,7 @@ func (h *Handler) handleRequest(client *Client, msg *Message) {
|
||||
// handleRead processes a read operation
|
||||
func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext) {
|
||||
// Execute before hook
|
||||
if err := h.hooks.Execute(BeforeRead, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil {
|
||||
logger.Error("[MQTTSpec] BeforeRead hook failed: %v", err)
|
||||
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
||||
return
|
||||
@@ -356,7 +356,7 @@ func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext)
|
||||
// handleCreate processes a create operation
|
||||
func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContext) {
|
||||
// Execute before hook
|
||||
if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
|
||||
logger.Error("[MQTTSpec] BeforeCreate hook failed: %v", err)
|
||||
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
||||
return
|
||||
@@ -390,7 +390,7 @@ func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContex
|
||||
// handleUpdate processes an update operation
|
||||
func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContext) {
|
||||
// Execute before hook
|
||||
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil {
|
||||
logger.Error("[MQTTSpec] BeforeUpdate hook failed: %v", err)
|
||||
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
||||
return
|
||||
@@ -424,7 +424,7 @@ func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContex
|
||||
// handleDelete processes a delete operation
|
||||
func (h *Handler) handleDelete(client *Client, msg *Message, hookCtx *HookContext) {
|
||||
// Execute before hook
|
||||
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil {
|
||||
logger.Error("[MQTTSpec] BeforeDelete hook failed: %v", err)
|
||||
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
||||
return
|
||||
@@ -676,7 +676,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) {
|
||||
|
||||
// Apply columns
|
||||
if hookCtx.Options != nil && len(hookCtx.Options.Columns) > 0 {
|
||||
query = query.Column(hookCtx.Options.Columns...)
|
||||
query = common.ApplySelectColumns(query, hookCtx.Model, "", hookCtx.Options.Columns)
|
||||
}
|
||||
|
||||
// Apply preloads (simplified)
|
||||
@@ -686,6 +686,18 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// Execute BeforeScan hooks - pass query so hooks (e.g. row-level security) can modify it
|
||||
if hookCtx.Metadata == nil {
|
||||
hookCtx.Metadata = make(map[string]interface{})
|
||||
}
|
||||
hookCtx.Metadata["query"] = query
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
|
||||
return nil, fmt.Errorf("BeforeScan hook failed: %w", err)
|
||||
}
|
||||
if modifiedQuery, ok := hookCtx.Metadata["query"].(common.SelectQuery); ok {
|
||||
query = modifiedQuery
|
||||
}
|
||||
|
||||
// Execute query
|
||||
if err := query.ScanModel(hookCtx.Context); err != nil {
|
||||
return nil, fmt.Errorf("failed to read record: %w", err)
|
||||
@@ -702,9 +714,19 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
if hookCtx.Options != nil {
|
||||
// Apply filters
|
||||
for _, filter := range hookCtx.Options.Filters {
|
||||
if cond, jargs, ok := common.BuildJSONFilterCondition(hookCtx.Model, "", filter.Column, filter.Operator, filter.Value); ok {
|
||||
query = query.Where(cond, jargs...)
|
||||
continue
|
||||
}
|
||||
op := strings.ToLower(filter.Operator)
|
||||
if op == "like" || op == "ilike" {
|
||||
query = query.Where(fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||
// citext columns are already case-insensitive; casting to TEXT would
|
||||
// switch to case-sensitive matching and defeat a citext index.
|
||||
if reflection.IsCitextColumn(hookCtx.Model, filter.Column) {
|
||||
query = query.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||
} else {
|
||||
query = query.Where(fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||
}
|
||||
} else {
|
||||
query = query.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||
}
|
||||
@@ -716,6 +738,10 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
if sort.Direction == "desc" {
|
||||
direction = "DESC"
|
||||
}
|
||||
if expr, jargs, _, ok := common.ResolveJSONColumnExpr(hookCtx.Model, "", sort.Column); ok {
|
||||
query = query.OrderExpr(fmt.Sprintf("%s %s", expr, direction), jargs...)
|
||||
continue
|
||||
}
|
||||
query = query.Order(fmt.Sprintf("%s %s", sort.Column, direction))
|
||||
}
|
||||
|
||||
@@ -734,10 +760,22 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
|
||||
// Apply columns
|
||||
if len(hookCtx.Options.Columns) > 0 {
|
||||
query = query.Column(hookCtx.Options.Columns...)
|
||||
query = common.ApplySelectColumns(query, hookCtx.Model, "", hookCtx.Options.Columns)
|
||||
}
|
||||
}
|
||||
|
||||
// Execute BeforeScan hooks - pass query so hooks (e.g. row-level security) can modify it
|
||||
if hookCtx.Metadata == nil {
|
||||
hookCtx.Metadata = make(map[string]interface{})
|
||||
}
|
||||
hookCtx.Metadata["query"] = query
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
|
||||
return nil, nil, fmt.Errorf("BeforeScan hook failed: %w", err)
|
||||
}
|
||||
if modifiedQuery, ok := hookCtx.Metadata["query"].(common.SelectQuery); ok {
|
||||
query = modifiedQuery
|
||||
}
|
||||
|
||||
// Execute query
|
||||
if err := query.ScanModel(hookCtx.Context); err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to read records: %w", err)
|
||||
@@ -748,9 +786,19 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
countQuery := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
||||
if hookCtx.Options != nil {
|
||||
for _, filter := range hookCtx.Options.Filters {
|
||||
if cond, jargs, ok := common.BuildJSONFilterCondition(hookCtx.Model, "", filter.Column, filter.Operator, filter.Value); ok {
|
||||
countQuery = countQuery.Where(cond, jargs...)
|
||||
continue
|
||||
}
|
||||
op := strings.ToLower(filter.Operator)
|
||||
if op == "like" || op == "ilike" {
|
||||
countQuery = countQuery.Where(fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||
// citext columns are already case-insensitive; casting to TEXT would
|
||||
// switch to case-sensitive matching and defeat a citext index.
|
||||
if reflection.IsCitextColumn(hookCtx.Model, filter.Column) {
|
||||
countQuery = countQuery.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||
} else {
|
||||
countQuery = countQuery.Where(fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||
}
|
||||
} else {
|
||||
countQuery = countQuery.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||
}
|
||||
|
||||
@@ -34,6 +34,7 @@ const (
|
||||
AfterUpdate = websocketspec.AfterUpdate
|
||||
BeforeDelete = websocketspec.BeforeDelete
|
||||
AfterDelete = websocketspec.AfterDelete
|
||||
BeforeScan = websocketspec.BeforeScan
|
||||
|
||||
// Subscription hooks
|
||||
BeforeSubscribe = websocketspec.BeforeSubscribe
|
||||
@@ -46,6 +47,9 @@ const (
|
||||
AfterConnect = websocketspec.AfterConnect
|
||||
BeforeDisconnect = websocketspec.BeforeDisconnect
|
||||
AfterDisconnect = websocketspec.AfterDisconnect
|
||||
|
||||
// BeforeOp fires immediately before every SQL operation (read, create, update, delete)
|
||||
BeforeOp = websocketspec.BeforeOp
|
||||
)
|
||||
|
||||
// NewHookRegistry creates a new hook registry
|
||||
|
||||
@@ -27,25 +27,31 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
||||
return security.LoadSecurityRules(secCtx, securityList)
|
||||
})
|
||||
|
||||
// Hook 2: AfterRead - Apply column-level security (masking)
|
||||
// Hook 2: BeforeScan - Apply row-level security filters
|
||||
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
return security.ApplyRowSecurity(secCtx, securityList)
|
||||
})
|
||||
|
||||
// Hook 3: AfterRead - Apply column-level security (masking)
|
||||
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
return security.ApplyColumnSecurity(secCtx, securityList)
|
||||
})
|
||||
|
||||
// Hook 3 (Optional): Audit logging
|
||||
// Hook 4 (Optional): Audit logging
|
||||
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
return security.LogDataAccess(secCtx)
|
||||
})
|
||||
|
||||
// Hook 4: BeforeUpdate - enforce CanUpdate rule from context/registry
|
||||
// Hook 5: BeforeUpdate - enforce CanUpdate rule from context/registry
|
||||
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
return security.CheckModelUpdateAllowed(secCtx)
|
||||
})
|
||||
|
||||
// Hook 5: BeforeDelete - enforce CanDelete rule from context/registry
|
||||
// Hook 6: BeforeDelete - enforce CanDelete rule from context/registry
|
||||
handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
return security.CheckModelDeleteAllowed(secCtx)
|
||||
@@ -71,6 +77,17 @@ func (s *securityContext) GetUserID() (int, bool) {
|
||||
return security.GetUserID(s.ctx.Context)
|
||||
}
|
||||
|
||||
// GetUserRef returns an opaque user identifier for row security lookups.
|
||||
// It prefers the full *security.UserContext (so providers can read JWT claims,
|
||||
// e.g. a UUID subject) and falls back to the int user ID.
|
||||
func (s *securityContext) GetUserRef() (any, bool) {
|
||||
if userCtx, ok := security.GetUserContext(s.ctx.Context); ok {
|
||||
return userCtx, true
|
||||
}
|
||||
userID, ok := security.GetUserID(s.ctx.Context)
|
||||
return userID, ok
|
||||
}
|
||||
|
||||
func (s *securityContext) GetSchema() string {
|
||||
return s.ctx.Schema
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// OpenAPISpec represents the OpenAPI 3.0 specification structure
|
||||
@@ -440,6 +441,28 @@ func (g *Generator) generatePropertySchema(field reflect.StructField) *Schema {
|
||||
schema.Description = desc
|
||||
}
|
||||
|
||||
// spectypes PostGIS / pgvector wrappers get dedicated schemas.
|
||||
if n, ok := spectypes.SQLTypeName(field.Type); ok {
|
||||
switch n {
|
||||
case "geometry", "geography":
|
||||
schema.Type = "object"
|
||||
schema.Format = "geojson"
|
||||
return schema
|
||||
case "vector", "halfvec":
|
||||
schema.Type = "array"
|
||||
schema.Items = &Schema{Type: "number"}
|
||||
schema.Format = "vector"
|
||||
return schema
|
||||
case "sparsevec":
|
||||
schema.Type = "object"
|
||||
schema.Format = "sparsevec"
|
||||
return schema
|
||||
case "bit":
|
||||
schema.Type = "string"
|
||||
return schema
|
||||
}
|
||||
}
|
||||
|
||||
switch fieldType.Kind() {
|
||||
case reflect.String:
|
||||
schema.Type = "string"
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
type PrimaryKeyNameProvider interface {
|
||||
@@ -438,6 +439,63 @@ func GetSQLModelColumns(model any) []string {
|
||||
return columns
|
||||
}
|
||||
|
||||
// HasColumn reports whether the model has a struct field that bun/gorm would
|
||||
// scan a column named columnName into. Unlike GetSQLModelColumns, this
|
||||
// includes scanonly fields (e.g. a `bun:"jsonvalue_product_cost,scanonly"`
|
||||
// field added specifically to receive a computed/JSON-path SELECT expression)
|
||||
// since those are legitimate scan targets even though they are not writable.
|
||||
// Matching is case-insensitive against the resolved bun/gorm/json column name
|
||||
// and against the bare Go field name.
|
||||
func HasColumn(model any, columnName string) bool {
|
||||
if columnName == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
modelType := reflect.TypeOf(model)
|
||||
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||
modelType = modelType.Elem()
|
||||
}
|
||||
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||
return false
|
||||
}
|
||||
|
||||
return hasColumnInType(modelType, columnName)
|
||||
}
|
||||
|
||||
func hasColumnInType(typ reflect.Type, columnName string) bool {
|
||||
for i := 0; i < typ.NumField(); i++ {
|
||||
field := typ.Field(i)
|
||||
if !field.IsExported() {
|
||||
continue
|
||||
}
|
||||
|
||||
bunTag := field.Tag.Get("bun")
|
||||
gormTag := field.Tag.Get("gorm")
|
||||
|
||||
if field.Anonymous {
|
||||
fieldType := field.Type
|
||||
if fieldType.Kind() == reflect.Pointer {
|
||||
fieldType = fieldType.Elem()
|
||||
}
|
||||
if fieldType.Kind() == reflect.Struct {
|
||||
if hasColumnInType(fieldType, columnName) {
|
||||
return true
|
||||
}
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
if bunTag == "-" || gormTag == "-" {
|
||||
continue
|
||||
}
|
||||
|
||||
if strings.EqualFold(getColumnNameFromField(field), columnName) || strings.EqualFold(field.Name, columnName) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// collectSQLColumnsFromType recursively collects SQL column names from a struct type
|
||||
// scanOnlyEmbedded indicates if we're inside a scan-only embedded struct
|
||||
func collectSQLColumnsFromType(typ reflect.Type, columns *[]string, scanOnlyEmbedded bool) {
|
||||
@@ -728,19 +786,19 @@ func GetColumnTypeFromModel(model interface{}, colName string) reflect.Kind {
|
||||
// Parse JSON tag (format: "name,omitempty")
|
||||
parts := strings.Split(jsonTag, ",")
|
||||
if parts[0] == sourceColName {
|
||||
return field.Type.Kind()
|
||||
return spectypes.UnwrapKind(field.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// Check field name (case-insensitive)
|
||||
if strings.EqualFold(field.Name, sourceColName) {
|
||||
return field.Type.Kind()
|
||||
return spectypes.UnwrapKind(field.Type)
|
||||
}
|
||||
|
||||
// Check snake_case conversion
|
||||
snakeCaseName := ToSnakeCase(field.Name)
|
||||
if snakeCaseName == sourceColName {
|
||||
return field.Type.Kind()
|
||||
return spectypes.UnwrapKind(field.Type)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
package reflection
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
type geoModel struct {
|
||||
ID int64 `json:"id"`
|
||||
Location spectypes.SqlGeometry `json:"location"`
|
||||
Area spectypes.SqlGeography `json:"area"`
|
||||
Embedding spectypes.SqlVector `json:"embedding"`
|
||||
HalfEmb spectypes.SqlHalfVector `json:"half_emb"`
|
||||
Name spectypes.SqlString `json:"name"`
|
||||
}
|
||||
|
||||
func TestGetColumnSQLTypeName(t *testing.T) {
|
||||
m := geoModel{}
|
||||
cases := map[string]string{
|
||||
"location": "geometry",
|
||||
"area": "geography",
|
||||
"embedding": "vector",
|
||||
"half_emb": "halfvec",
|
||||
"name": "text",
|
||||
}
|
||||
for col, want := range cases {
|
||||
got, ok := GetColumnSQLTypeName(m, col)
|
||||
if !ok || got != want {
|
||||
t.Errorf("GetColumnSQLTypeName(%q) = %q, %v; want %q", col, got, ok, want)
|
||||
}
|
||||
}
|
||||
if _, ok := GetColumnSQLTypeName(m, "id"); ok {
|
||||
t.Error("id is not a spectypes column")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsSpatialColumn(t *testing.T) {
|
||||
m := geoModel{}
|
||||
if !IsSpatialColumn(m, "location") || !IsSpatialColumn(m, "area") {
|
||||
t.Error("location/area should be spatial")
|
||||
}
|
||||
if IsSpatialColumn(m, "embedding") || IsSpatialColumn(m, "name") {
|
||||
t.Error("embedding/name should not be spatial")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsVectorColumn(t *testing.T) {
|
||||
m := geoModel{}
|
||||
if !IsVectorColumn(m, "embedding") || !IsVectorColumn(m, "half_emb") {
|
||||
t.Error("embedding/half_emb should be vector")
|
||||
}
|
||||
if IsVectorColumn(m, "location") || IsVectorColumn(m, "name") {
|
||||
t.Error("location/name should not be vector")
|
||||
}
|
||||
}
|
||||
@@ -3,6 +3,8 @@ package reflection
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// Test models for GORM
|
||||
@@ -409,6 +411,23 @@ func TestGetModelColumnsWithEmbedded(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestHasColumn(t *testing.T) {
|
||||
m := ModelWithEmbedded{}
|
||||
|
||||
for _, col := range []string{"name", "description", "rid_base", "created_at", "cql1", "cql2"} {
|
||||
if !HasColumn(m, col) {
|
||||
t.Errorf("HasColumn(%q) = false, want true", col)
|
||||
}
|
||||
}
|
||||
|
||||
if HasColumn(m, "nonexistent_column") {
|
||||
t.Error("HasColumn(nonexistent_column) = true, want false")
|
||||
}
|
||||
if HasColumn(m, "") {
|
||||
t.Error("HasColumn(\"\") = true, want false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsColumnWritableWithEmbedded(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -1047,6 +1066,22 @@ func TestGetColumnTypeFromModel(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// SqlNull-wrapped columns (e.g. nullable bigint foreign keys) must report the
|
||||
// wrapped value's Kind, not reflect.Struct, so numeric eq/gt/lt filters don't
|
||||
// get an unnecessary CAST(... AS TEXT) that defeats the column's index.
|
||||
type SqlNullFKModel struct {
|
||||
RidParent spectypes.SqlInt64 `bun:"rid_parent" json:"rid_parent"`
|
||||
}
|
||||
|
||||
func TestGetColumnTypeFromModel_SqlNullWrapper(t *testing.T) {
|
||||
model := SqlNullFKModel{RidParent: spectypes.NewSqlInt64(90446096)}
|
||||
|
||||
result := GetColumnTypeFromModel(model, "rid_parent")
|
||||
if result != reflect.Int64 {
|
||||
t.Errorf("GetColumnTypeFromModel(rid_parent) = %v, want %v (SqlInt64 must unwrap to its numeric Kind)", result, reflect.Int64)
|
||||
}
|
||||
}
|
||||
|
||||
// ============= Tests for relation functions =============
|
||||
|
||||
// Models for relation testing
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
package reflection
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// getColumnStructField resolves the struct field that backs colName (matched by
|
||||
// json tag, field name or snake_case), following the same rules as
|
||||
// GetColumnTypeFromModel.
|
||||
func getColumnStructField(model interface{}, colName string) (reflect.StructField, bool) {
|
||||
if model == nil {
|
||||
return reflect.StructField{}, false
|
||||
}
|
||||
sourceColName := ExtractSourceColumn(colName)
|
||||
|
||||
modelType := reflect.TypeOf(model)
|
||||
for modelType != nil && modelType.Kind() == reflect.Pointer {
|
||||
modelType = modelType.Elem()
|
||||
}
|
||||
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||
return reflect.StructField{}, false
|
||||
}
|
||||
|
||||
for i := 0; i < modelType.NumField(); i++ {
|
||||
field := modelType.Field(i)
|
||||
|
||||
if jsonTag := field.Tag.Get("json"); jsonTag != "" {
|
||||
if name := jsonTagName(jsonTag); name == sourceColName {
|
||||
return field, true
|
||||
}
|
||||
}
|
||||
if equalFold(field.Name, sourceColName) {
|
||||
return field, true
|
||||
}
|
||||
if ToSnakeCase(field.Name) == sourceColName {
|
||||
return field, true
|
||||
}
|
||||
}
|
||||
return reflect.StructField{}, false
|
||||
}
|
||||
|
||||
// getColumnFieldType resolves the reflect.Type of the struct field that backs
|
||||
// colName (matched by json tag, field name or snake_case), following the same
|
||||
// rules as GetColumnTypeFromModel.
|
||||
func getColumnFieldType(model interface{}, colName string) (reflect.Type, bool) {
|
||||
f, ok := getColumnStructField(model, colName)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
return f.Type, true
|
||||
}
|
||||
|
||||
func jsonTagName(tag string) string {
|
||||
for i := 0; i < len(tag); i++ {
|
||||
if tag[i] == ',' {
|
||||
return tag[:i]
|
||||
}
|
||||
}
|
||||
return tag
|
||||
}
|
||||
|
||||
func equalFold(a, b string) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
for i := 0; i < len(a); i++ {
|
||||
ca, cb := a[i], b[i]
|
||||
if 'A' <= ca && ca <= 'Z' {
|
||||
ca += 'a' - 'A'
|
||||
}
|
||||
if 'A' <= cb && cb <= 'Z' {
|
||||
cb += 'a' - 'A'
|
||||
}
|
||||
if ca != cb {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// GetColumnSQLTypeName returns the canonical PostgreSQL type name for a column
|
||||
// backed by a spectypes wrapper (e.g. "geometry", "vector", "jsonb"), or
|
||||
// ("", false) if the column is not found or not a spectypes type.
|
||||
func GetColumnSQLTypeName(model interface{}, colName string) (string, bool) {
|
||||
t, ok := getColumnFieldType(model, colName)
|
||||
if !ok {
|
||||
return "", false
|
||||
}
|
||||
return spectypes.SQLTypeName(t)
|
||||
}
|
||||
|
||||
// IsSpatialColumn reports whether colName is backed by a PostGIS
|
||||
// geometry/geography wrapper.
|
||||
func IsSpatialColumn(model interface{}, colName string) bool {
|
||||
t, ok := getColumnFieldType(model, colName)
|
||||
return ok && spectypes.IsSpatialType(t)
|
||||
}
|
||||
|
||||
// IsVectorColumn reports whether colName is backed by a pgvector wrapper.
|
||||
func IsVectorColumn(model interface{}, colName string) bool {
|
||||
t, ok := getColumnFieldType(model, colName)
|
||||
return ok && spectypes.IsVectorType(t)
|
||||
}
|
||||
|
||||
var rawMessageType = reflect.TypeOf(json.RawMessage(nil))
|
||||
|
||||
// IsJSONColumn reports whether colName is backed by a JSON/JSONB column on the
|
||||
// model. It recognises the spectypes SqlJSONB wrapper, encoding/json.RawMessage,
|
||||
// map-typed fields, and fields carrying a bun/gorm `type:json` / `type:jsonb`
|
||||
// tag. colName should be a bare column name (callers pass the parsed base column
|
||||
// of a JSON path, not the full "col->>'x'" expression).
|
||||
func IsJSONColumn(model interface{}, colName string) bool {
|
||||
f, ok := getColumnStructField(model, colName)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
|
||||
ft := f.Type
|
||||
for ft != nil && ft.Kind() == reflect.Pointer {
|
||||
ft = ft.Elem()
|
||||
}
|
||||
if ft == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
if spectypes.IsJSONType(ft) {
|
||||
return true
|
||||
}
|
||||
if ft == rawMessageType {
|
||||
return true
|
||||
}
|
||||
if ft.Kind() == reflect.Map {
|
||||
return true
|
||||
}
|
||||
if tagDeclaresJSON(f.Tag.Get("bun")) || tagDeclaresJSON(f.Tag.Get("gorm")) {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// tagDeclaresJSON reports whether an ORM struct tag declares a json/jsonb column
|
||||
// type, e.g. `bun:"meta,type:jsonb"` or `gorm:"column:meta;type:json"`.
|
||||
func tagDeclaresJSON(tag string) bool {
|
||||
return columnTypeTagValue(tag) == "json" || strings.HasPrefix(columnTypeTagValue(tag), "json(") ||
|
||||
columnTypeTagValue(tag) == "jsonb" || strings.HasPrefix(columnTypeTagValue(tag), "jsonb(")
|
||||
}
|
||||
|
||||
// columnTypeTagValue extracts the lower-cased value of a `type:` entry from a
|
||||
// bun or gorm struct tag, e.g. `bun:"name,type:citext"` -> "citext". Returns ""
|
||||
// if the tag carries no `type:` entry.
|
||||
func columnTypeTagValue(tag string) string {
|
||||
if tag == "" {
|
||||
return ""
|
||||
}
|
||||
for _, part := range strings.FieldsFunc(tag, func(r rune) bool {
|
||||
return r == ',' || r == ';' || r == ' '
|
||||
}) {
|
||||
value, found := strings.CutPrefix(strings.TrimSpace(part), "type:")
|
||||
if !found {
|
||||
continue
|
||||
}
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// IsCitextColumn reports whether colName carries an explicit `type:citext`
|
||||
// bun/gorm tag. citext columns must never be CAST(... AS TEXT) for comparisons:
|
||||
// that swaps in case-sensitive text semantics and defeats any citext index.
|
||||
func IsCitextColumn(model interface{}, colName string) bool {
|
||||
f, ok := getColumnStructField(model, colName)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
tagVal := columnTypeTagValue(f.Tag.Get("bun"))
|
||||
if tagVal == "" {
|
||||
tagVal = columnTypeTagValue(f.Tag.Get("gorm"))
|
||||
}
|
||||
return tagVal == "citext"
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package reflection
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
type jsonColModel struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Meta spectypes.SqlJSONB `json:"meta"`
|
||||
Raw json.RawMessage `json:"raw"`
|
||||
Attrs map[string]interface{} `json:"attrs"`
|
||||
Config []byte `json:"config" bun:"config,type:jsonb"`
|
||||
Settings string `json:"settings" gorm:"column:settings;type:json"`
|
||||
Blob []byte `json:"blob"`
|
||||
}
|
||||
|
||||
func TestIsJSONColumn(t *testing.T) {
|
||||
m := jsonColModel{}
|
||||
|
||||
jsonCols := []string{"meta", "raw", "attrs", "config", "settings"}
|
||||
for _, c := range jsonCols {
|
||||
if !IsJSONColumn(m, c) {
|
||||
t.Errorf("expected %q to be a JSON column", c)
|
||||
}
|
||||
}
|
||||
|
||||
notJSON := []string{"id", "name", "blob", "missing"}
|
||||
for _, c := range notJSON {
|
||||
if IsJSONColumn(m, c) {
|
||||
t.Errorf("expected %q NOT to be a JSON column", c)
|
||||
}
|
||||
}
|
||||
|
||||
if IsJSONColumn(nil, "meta") {
|
||||
t.Error("nil model must not report JSON columns")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTagDeclaresJSON(t *testing.T) {
|
||||
cases := map[string]bool{
|
||||
"config,type:jsonb": true,
|
||||
"column:settings;type:json": true,
|
||||
"col,type:text": false,
|
||||
"column:name": false,
|
||||
"": false,
|
||||
"col,type:jsonb,notnull": true,
|
||||
"column:x;type:varchar(255)": false,
|
||||
"col , type:json": true,
|
||||
}
|
||||
for tag, want := range cases {
|
||||
if got := tagDeclaresJSON(tag); got != want {
|
||||
t.Errorf("tagDeclaresJSON(%q) = %v; want %v", tag, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
+34
-15
@@ -269,7 +269,7 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
|
||||
}
|
||||
|
||||
// Filters
|
||||
query = h.applyFilters(query, options.Filters)
|
||||
query = h.applyFilters(query, options.Filters, model)
|
||||
|
||||
// Custom operators
|
||||
for _, customOp := range options.CustomOperators {
|
||||
@@ -340,18 +340,28 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
|
||||
|
||||
var data interface{}
|
||||
if id != "" {
|
||||
singleResult := reflect.New(modelType).Interface()
|
||||
pkName := reflection.GetPrimaryKeyName(singleResult)
|
||||
pkName := reflection.GetPrimaryKeyName(model)
|
||||
query = query.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
|
||||
if err := query.Scan(ctx, singleResult); err != nil {
|
||||
// Scan through the model configured on the query. Bun rejects Scan with
|
||||
// a destination when the query preloads a has-many relation.
|
||||
if err := query.ScanModel(ctx); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, nil, fmt.Errorf("record not found")
|
||||
}
|
||||
return nil, nil, fmt.Errorf("query error: %w", err)
|
||||
}
|
||||
data = singleResult
|
||||
|
||||
// The configured model is a slice so the same query construction works
|
||||
// for both collection and single-record reads. Extract its one result.
|
||||
scannedResults := reflect.ValueOf(modelPtr).Elem()
|
||||
if scannedResults.Len() == 0 {
|
||||
return nil, nil, fmt.Errorf("record not found")
|
||||
}
|
||||
data = scannedResults.Index(0).Interface()
|
||||
} else {
|
||||
if err := query.Scan(ctx, modelPtr); err != nil {
|
||||
// Use the model already configured on the query. This is required by
|
||||
// Bun whenever the query includes a has-many preload.
|
||||
if err := query.ScanModel(ctx); err != nil {
|
||||
return nil, nil, fmt.Errorf("query error: %w", err)
|
||||
}
|
||||
data = reflect.ValueOf(modelPtr).Elem().Interface()
|
||||
@@ -741,8 +751,10 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
|
||||
return recordToDelete, nil
|
||||
}
|
||||
|
||||
// applyFilters applies all filters with OR grouping logic.
|
||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
|
||||
// applyFilters applies all filters with OR grouping logic. model, when
|
||||
// non-nil, lets citext columns be recognised so LIKE/ILIKE compares them
|
||||
// natively instead of casting to TEXT (which would defeat a citext index).
|
||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
||||
if len(filters) == 0 {
|
||||
return query
|
||||
}
|
||||
@@ -758,10 +770,10 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
||||
orGroup = append(orGroup, filters[j])
|
||||
j++
|
||||
}
|
||||
query = h.applyFilterGroup(query, orGroup)
|
||||
query = h.applyFilterGroup(query, orGroup, model)
|
||||
i = j
|
||||
} else {
|
||||
condition, args := h.buildFilterCondition(filters[i])
|
||||
condition, args := h.buildFilterCondition(filters[i], model)
|
||||
if condition != "" {
|
||||
query = query.Where(condition, args...)
|
||||
}
|
||||
@@ -772,12 +784,12 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
||||
return query
|
||||
}
|
||||
|
||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
|
||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
||||
var conditions []string
|
||||
var args []interface{}
|
||||
|
||||
for _, filter := range filters {
|
||||
condition, filterArgs := h.buildFilterCondition(filter)
|
||||
condition, filterArgs := h.buildFilterCondition(filter, model)
|
||||
if condition != "" {
|
||||
conditions = append(conditions, condition)
|
||||
args = append(args, filterArgs...)
|
||||
@@ -793,7 +805,14 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
||||
return query.Where("("+strings.Join(conditions, " OR ")+")", args...)
|
||||
}
|
||||
|
||||
func (h *Handler) buildFilterCondition(filter common.FilterOption) (condition string, args []interface{}) {
|
||||
func (h *Handler) buildFilterCondition(filter common.FilterOption, model interface{}) (condition string, args []interface{}) {
|
||||
// citext columns are already case-insensitive; casting to TEXT would
|
||||
// switch to case-sensitive matching and defeat a citext index.
|
||||
likeColumn := filter.Column
|
||||
if !reflection.IsCitextColumn(model, filter.Column) {
|
||||
likeColumn = fmt.Sprintf("CAST(%s AS TEXT)", filter.Column)
|
||||
}
|
||||
|
||||
switch filter.Operator {
|
||||
case "eq", "=":
|
||||
return fmt.Sprintf("%s = ?", filter.Column), []interface{}{filter.Value}
|
||||
@@ -808,9 +827,9 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption) (condition st
|
||||
case "lte", "<=":
|
||||
return fmt.Sprintf("%s <= ?", filter.Column), []interface{}{filter.Value}
|
||||
case "like":
|
||||
return fmt.Sprintf("CAST(%s AS TEXT) LIKE ?", filter.Column), []interface{}{filter.Value}
|
||||
return fmt.Sprintf("%s LIKE ?", likeColumn), []interface{}{filter.Value}
|
||||
case "ilike":
|
||||
return fmt.Sprintf("CAST(%s AS TEXT) ILIKE ?", filter.Column), []interface{}{filter.Value}
|
||||
return fmt.Sprintf("%s ILIKE ?", likeColumn), []interface{}{filter.Value}
|
||||
case "in":
|
||||
condition, args := common.BuildInCondition(filter.Column, filter.Value)
|
||||
return condition, args
|
||||
|
||||
@@ -17,6 +17,7 @@ package resolvemcp
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"runtime/debug"
|
||||
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/uptrace/bun"
|
||||
@@ -25,6 +26,7 @@ import (
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
)
|
||||
|
||||
@@ -82,11 +84,20 @@ func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler) {
|
||||
// - GET {basePath}/sse — SSE connection endpoint
|
||||
// - POST {basePath}/message — JSON-RPC message endpoint
|
||||
func SetupBunRouterRoutes(router *bunrouter.Router, handler *Handler) {
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
logger.Error("panic in resolvemcp.SetupBunRouterRoutes: %v\n%s", rec, debug.Stack())
|
||||
}
|
||||
}()
|
||||
|
||||
basePath := handler.config.BasePath
|
||||
h := handler.SSEServer()
|
||||
|
||||
router.GET(basePath+"/sse", bunrouter.HTTPHandler(h))
|
||||
logger.Info("Registered resolvemcp bunrouter route GET %s/sse", basePath)
|
||||
|
||||
router.POST(basePath+"/message", bunrouter.HTTPHandler(h))
|
||||
logger.Info("Registered resolvemcp bunrouter route POST %s/message", basePath)
|
||||
}
|
||||
|
||||
// NewSSEServer returns an http.Handler that serves MCP over SSE.
|
||||
|
||||
@@ -84,6 +84,17 @@ func (s *securityContext) GetUserID() (int, bool) {
|
||||
return security.GetUserID(s.ctx.Context)
|
||||
}
|
||||
|
||||
// GetUserRef returns an opaque user identifier for row security lookups.
|
||||
// It prefers the full *security.UserContext (so providers can read JWT claims,
|
||||
// e.g. a UUID subject) and falls back to the int user ID.
|
||||
func (s *securityContext) GetUserRef() (any, bool) {
|
||||
if userCtx, ok := security.GetUserContext(s.ctx.Context); ok {
|
||||
return userCtx, true
|
||||
}
|
||||
userID, ok := security.GetUserID(s.ctx.Context)
|
||||
return userID, ok
|
||||
}
|
||||
|
||||
func (s *securityContext) GetSchema() string {
|
||||
return s.ctx.Schema
|
||||
}
|
||||
|
||||
+31
-3
@@ -85,13 +85,19 @@ func buildModelInfo(schema, entity string, model interface{}) modelInfo {
|
||||
|
||||
// Skip relation fields (slice or user-defined struct that isn't time.Time).
|
||||
fieldType, found := modelType.FieldByName(d.Name)
|
||||
var unwrappedType reflect.Type
|
||||
isSQLType := false
|
||||
if found {
|
||||
ft := fieldType.Type
|
||||
if ft.Kind() == reflect.Pointer {
|
||||
if sqlType, ok := unwrapSQLType(ft); ok {
|
||||
unwrappedType = sqlType
|
||||
ft = sqlType
|
||||
isSQLType = true
|
||||
} else if ft.Kind() == reflect.Pointer {
|
||||
ft = ft.Elem()
|
||||
}
|
||||
isUserStruct := ft.Kind() == reflect.Struct && ft.Name() != "Time" && ft.PkgPath() != ""
|
||||
if ft.Kind() == reflect.Slice || isUserStruct {
|
||||
if !isSQLType && (ft.Kind() == reflect.Slice || isUserStruct) {
|
||||
info.relationNames = append(info.relationNames, jsonName)
|
||||
continue
|
||||
}
|
||||
@@ -104,6 +110,9 @@ func buildModelInfo(schema, entity string, model interface{}) modelInfo {
|
||||
|
||||
// Derive Go type name, unwrapping pointer if needed.
|
||||
goType := d.DataType
|
||||
if isSQLType {
|
||||
goType = unwrappedType.Name()
|
||||
}
|
||||
if goType == "" && found {
|
||||
ft := fieldType.Type
|
||||
for ft.Kind() == reflect.Pointer {
|
||||
@@ -125,7 +134,7 @@ func buildModelInfo(schema, entity string, model interface{}) modelInfo {
|
||||
isPrimary: isPrimary,
|
||||
isUnique: d.SQLKey == "unique" || d.SQLKey == "uniqueindex",
|
||||
isFK: d.SQLKey == "foreign_key",
|
||||
nullable: d.Nullable,
|
||||
nullable: isSQLType || d.Nullable,
|
||||
}
|
||||
info.columns = append(info.columns, ci)
|
||||
}
|
||||
@@ -134,6 +143,25 @@ func buildModelInfo(schema, entity string, model interface{}) modelInfo {
|
||||
return info
|
||||
}
|
||||
|
||||
// unwrapSQLType returns the value type wrapped by a spectypes SQL value. These
|
||||
// types are scalar columns even when their Go representation is a struct or a
|
||||
// slice (for example, SqlNull[string] and SqlJSONB).
|
||||
func unwrapSQLType(t reflect.Type) (reflect.Type, bool) {
|
||||
for t.Kind() == reflect.Pointer {
|
||||
t = t.Elem()
|
||||
}
|
||||
|
||||
if t.PkgPath() != "github.com/bitechdev/ResolveSpec/pkg/spectypes" {
|
||||
return nil, false
|
||||
}
|
||||
if t.Kind() == reflect.Struct {
|
||||
if value, ok := t.FieldByName("Val"); ok {
|
||||
return value.Type, true
|
||||
}
|
||||
}
|
||||
return t, true
|
||||
}
|
||||
|
||||
// fieldJSONName returns the JSON tag name for a struct field, falling back to the field name.
|
||||
func fieldJSONName(modelType reflect.Type, fieldName string) string {
|
||||
field, ok := modelType.FieldByName(fieldName)
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
func TestBuildModelInfo_UnwrapsSQLTypes(t *testing.T) {
|
||||
type related struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
type model struct {
|
||||
Name spectypes.SqlString `gorm:"column:name" json:"name"`
|
||||
Metadata spectypes.SqlJSONB `gorm:"column:metadata;type:jsonb" json:"metadata"`
|
||||
Related related `json:"related"`
|
||||
}
|
||||
|
||||
info := buildModelInfo("public", "models", model{})
|
||||
columns := make(map[string]columnInfo, len(info.columns))
|
||||
for _, column := range info.columns {
|
||||
columns[column.jsonName] = column
|
||||
}
|
||||
|
||||
if column, ok := columns["name"]; !ok || column.goType != "string" || !column.nullable {
|
||||
t.Errorf("expected name SQL wrapper column, got %+v", column)
|
||||
}
|
||||
if column, ok := columns["metadata"]; !ok || !column.nullable {
|
||||
t.Errorf("expected metadata SQL wrapper column, got %+v", column)
|
||||
}
|
||||
if len(info.relationNames) != 1 || info.relationNames[0] != "related" {
|
||||
t.Errorf("expected only related to be a relation, got %v", info.relationNames)
|
||||
}
|
||||
}
|
||||
@@ -57,11 +57,78 @@ func TestBuildFilterCondition(t *testing.T) {
|
||||
expectedCondition: "CAST(email AS TEXT) LIKE ?",
|
||||
expectedArgsCount: 1,
|
||||
},
|
||||
{
|
||||
name: "CONTAINS operator with single value",
|
||||
filter: common.FilterOption{
|
||||
Column: "tags",
|
||||
Operator: "contains",
|
||||
Value: "urgent",
|
||||
},
|
||||
expectedCondition: "tags && ARRAY[?]",
|
||||
expectedArgsCount: 1,
|
||||
},
|
||||
{
|
||||
name: "CONTAINS operator with multiple values",
|
||||
filter: common.FilterOption{
|
||||
Column: "tags",
|
||||
Operator: "contains",
|
||||
Value: []string{"urgent", "billing"},
|
||||
},
|
||||
expectedCondition: "tags && ARRAY[?,?]",
|
||||
expectedArgsCount: 2,
|
||||
},
|
||||
{
|
||||
name: "CONTAINS operator with empty value",
|
||||
filter: common.FilterOption{
|
||||
Column: "tags",
|
||||
Operator: "contains",
|
||||
Value: nil,
|
||||
},
|
||||
expectedCondition: "",
|
||||
expectedArgsCount: 0,
|
||||
},
|
||||
{
|
||||
name: "st_dwithin spatial operator",
|
||||
filter: common.FilterOption{
|
||||
Column: "geom",
|
||||
Operator: "st_dwithin",
|
||||
Value: map[string]any{
|
||||
"geom": "SRID=4326;POINT(0 0)",
|
||||
"distance": 1000.0,
|
||||
},
|
||||
},
|
||||
expectedCondition: "ST_DWithin(geom, ST_GeomFromEWKT(?), ?)",
|
||||
expectedArgsCount: 2,
|
||||
},
|
||||
{
|
||||
name: "st_intersects spatial operator",
|
||||
filter: common.FilterOption{
|
||||
Column: "geom",
|
||||
Operator: "st_intersects",
|
||||
Value: "SRID=4326;POLYGON((0 0,1 0,1 1,0 1,0 0))",
|
||||
LogicOperator: "",
|
||||
},
|
||||
expectedCondition: "ST_Intersects(geom, ST_GeomFromEWKT(?))",
|
||||
expectedArgsCount: 1,
|
||||
},
|
||||
{
|
||||
name: "l2_within vector operator",
|
||||
filter: common.FilterOption{
|
||||
Column: "embedding",
|
||||
Operator: "l2_within",
|
||||
Value: map[string]any{
|
||||
"vector": []any{1.0, 2.0, 3.0},
|
||||
"distance": 0.5,
|
||||
},
|
||||
},
|
||||
expectedCondition: "embedding <-> ? < ?",
|
||||
expectedArgsCount: 2,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
condition, args := h.buildFilterCondition(tt.filter)
|
||||
condition, args := h.buildFilterCondition(tt.filter, nil)
|
||||
|
||||
if condition != tt.expectedCondition {
|
||||
t.Errorf("Expected condition '%s', got '%s'", tt.expectedCondition, condition)
|
||||
|
||||
+764
-398
File diff suppressed because it is too large
Load Diff
@@ -5,6 +5,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
func TestNewHandler(t *testing.T) {
|
||||
@@ -41,6 +42,36 @@ func TestSetFallbackHandler(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetDefaultSort(t *testing.T) {
|
||||
handler := NewHandler(nil, nil)
|
||||
|
||||
// No default configured yet
|
||||
if got := handler.getDefaultSort("public", "users"); got != nil {
|
||||
t.Errorf("Expected no default sort, got %v", got)
|
||||
}
|
||||
|
||||
// Global default
|
||||
global := []common.SortOption{{Column: "created_at", Direction: "desc"}}
|
||||
handler.SetDefaultSort("", "", global...)
|
||||
|
||||
if got := handler.getDefaultSort("public", "users"); !reflect.DeepEqual(got, global) {
|
||||
t.Errorf("Expected global default sort %v, got %v", global, got)
|
||||
}
|
||||
|
||||
// Per-model default overrides the global default
|
||||
perModel := []common.SortOption{{Column: "name", Direction: "asc"}}
|
||||
handler.SetDefaultSort("public", "users", perModel...)
|
||||
|
||||
if got := handler.getDefaultSort("public", "users"); !reflect.DeepEqual(got, perModel) {
|
||||
t.Errorf("Expected per-model default sort %v, got %v", perModel, got)
|
||||
}
|
||||
|
||||
// Other models still fall back to the global default
|
||||
if got := handler.getDefaultSort("public", "orders"); !reflect.DeepEqual(got, global) {
|
||||
t.Errorf("Expected global default sort %v for unrelated model, got %v", global, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDatabase(t *testing.T) {
|
||||
handler := NewHandler(nil, nil)
|
||||
db := handler.GetDatabase()
|
||||
@@ -172,6 +203,48 @@ func TestGetColumnType(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateMetadata_UnwrapsSQLTypes(t *testing.T) {
|
||||
type related struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
type model struct {
|
||||
Name spectypes.SqlString `json:"name"`
|
||||
Count spectypes.SqlInt64 `json:"count"`
|
||||
CreatedAt spectypes.SqlTimeStamp `json:"created_at"`
|
||||
Metadata spectypes.SqlJSONB `json:"metadata" gorm:"type:jsonb"`
|
||||
Related related `json:"related"`
|
||||
}
|
||||
|
||||
metadata := NewHandler(nil, nil).generateMetadata("public", "models", model{})
|
||||
columns := make(map[string]common.Column, len(metadata.Columns))
|
||||
for _, column := range metadata.Columns {
|
||||
columns[column.Name] = column
|
||||
}
|
||||
|
||||
for name, wantType := range map[string]string{
|
||||
"name": "string",
|
||||
"count": "bigint",
|
||||
"created_at": "timestamp",
|
||||
"metadata": "jsonb",
|
||||
} {
|
||||
column, ok := columns[name]
|
||||
if !ok {
|
||||
t.Errorf("expected %q to be a metadata column", name)
|
||||
continue
|
||||
}
|
||||
if column.Type != wantType {
|
||||
t.Errorf("%q: expected type %q, got %q", name, wantType, column.Type)
|
||||
}
|
||||
if !column.IsNullable {
|
||||
t.Errorf("%q: expected SQL wrapper to be nullable", name)
|
||||
}
|
||||
}
|
||||
|
||||
if len(metadata.Relations) != 1 || metadata.Relations[0] != "related" {
|
||||
t.Errorf("expected only related to be a relation, got %v", metadata.Relations)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsNullable(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
@@ -3,6 +3,8 @@ package resolvespec
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
@@ -34,6 +36,11 @@ const (
|
||||
|
||||
// Scan/Execute operation hooks (for query building)
|
||||
BeforeScan HookType = "before_scan"
|
||||
|
||||
// BeforeOp fires immediately before every SQL operation (read, create, update, delete, scan).
|
||||
// Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each
|
||||
// individual SQL-operation hook point, so it runs once per statement executed.
|
||||
BeforeOp HookType = "before_op"
|
||||
)
|
||||
|
||||
// HookContext contains all the data available to a hook
|
||||
@@ -77,6 +84,7 @@ type HookFunc func(*HookContext) error
|
||||
// HookRegistry manages all registered hooks
|
||||
type HookRegistry struct {
|
||||
hooks map[HookType][]HookFunc
|
||||
mutex sync.RWMutex
|
||||
}
|
||||
|
||||
// NewHookRegistry creates a new hook registry
|
||||
@@ -86,8 +94,46 @@ func NewHookRegistry() *HookRegistry {
|
||||
}
|
||||
}
|
||||
|
||||
// hookLockRetryAttempts/hookLockRetryDelay bound how long the try-lock
|
||||
// helpers below will spin before giving up, so a contended mutex can
|
||||
// never hang a caller.
|
||||
const (
|
||||
hookLockRetryAttempts = 20
|
||||
hookLockRetryDelay = 1 * time.Millisecond
|
||||
)
|
||||
|
||||
// tryLock attempts to acquire the write lock, retrying briefly. Returns
|
||||
// false if it could not be acquired within the bound.
|
||||
func (r *HookRegistry) tryLock() bool {
|
||||
for i := 0; i < hookLockRetryAttempts; i++ {
|
||||
if r.mutex.TryLock() {
|
||||
return true
|
||||
}
|
||||
time.Sleep(hookLockRetryDelay)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// tryRLock attempts to acquire the read lock, retrying briefly. Returns
|
||||
// false if it could not be acquired within the bound.
|
||||
func (r *HookRegistry) tryRLock() bool {
|
||||
for i := 0; i < hookLockRetryAttempts; i++ {
|
||||
if r.mutex.TryRLock() {
|
||||
return true
|
||||
}
|
||||
time.Sleep(hookLockRetryDelay)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Register adds a new hook for the specified hook type
|
||||
func (r *HookRegistry) Register(hookType HookType, hook HookFunc) {
|
||||
if !r.tryLock() {
|
||||
logger.Error("Failed to register resolvespec hook for %s: registry locked", hookType)
|
||||
return
|
||||
}
|
||||
defer r.mutex.Unlock()
|
||||
|
||||
if r.hooks == nil {
|
||||
r.hooks = make(map[HookType][]HookFunc)
|
||||
}
|
||||
@@ -105,8 +151,13 @@ func (r *HookRegistry) RegisterMultiple(hookTypes []HookType, hook HookFunc) {
|
||||
// Execute runs all hooks for the specified type in order
|
||||
// If any hook returns an error, execution stops and the error is returned
|
||||
func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
||||
hooks, exists := r.hooks[hookType]
|
||||
if !exists || len(hooks) == 0 {
|
||||
if !r.tryRLock() {
|
||||
return fmt.Errorf("hook execution failed: registry locked")
|
||||
}
|
||||
hooks := append([]HookFunc(nil), r.hooks[hookType]...)
|
||||
r.mutex.RUnlock()
|
||||
|
||||
if len(hooks) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -128,20 +179,47 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
|
||||
// (BeforeRead, BeforeCreate, BeforeUpdate, BeforeDelete, or BeforeScan). BeforeOp always
|
||||
// runs first so it can observe/veto every SQL operation regardless of type.
|
||||
func (r *HookRegistry) ExecuteBeforeOp(hookType HookType, ctx *HookContext) error {
|
||||
if err := r.Execute(BeforeOp, ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
return r.Execute(hookType, ctx)
|
||||
}
|
||||
|
||||
// Clear removes all hooks for the specified type
|
||||
func (r *HookRegistry) Clear(hookType HookType) {
|
||||
if !r.tryLock() {
|
||||
logger.Error("Failed to clear resolvespec hooks for %s: registry locked", hookType)
|
||||
return
|
||||
}
|
||||
defer r.mutex.Unlock()
|
||||
|
||||
delete(r.hooks, hookType)
|
||||
logger.Info("Cleared all resolvespec hooks for %s", hookType)
|
||||
}
|
||||
|
||||
// ClearAll removes all registered hooks
|
||||
func (r *HookRegistry) ClearAll() {
|
||||
if !r.tryLock() {
|
||||
logger.Error("Failed to clear all resolvespec hooks: registry locked")
|
||||
return
|
||||
}
|
||||
defer r.mutex.Unlock()
|
||||
|
||||
r.hooks = make(map[HookType][]HookFunc)
|
||||
logger.Info("Cleared all resolvespec hooks")
|
||||
}
|
||||
|
||||
// Count returns the number of hooks registered for a specific type
|
||||
func (r *HookRegistry) Count(hookType HookType) int {
|
||||
if !r.tryRLock() {
|
||||
return 0
|
||||
}
|
||||
defer r.mutex.RUnlock()
|
||||
|
||||
if hooks, exists := r.hooks[hookType]; exists {
|
||||
return len(hooks)
|
||||
}
|
||||
@@ -155,6 +233,11 @@ func (r *HookRegistry) HasHooks(hookType HookType) bool {
|
||||
|
||||
// GetAllHookTypes returns all hook types that have registered hooks
|
||||
func (r *HookRegistry) GetAllHookTypes() []HookType {
|
||||
if !r.tryRLock() {
|
||||
return nil
|
||||
}
|
||||
defer r.mutex.RUnlock()
|
||||
|
||||
types := make([]HookType, 0, len(r.hooks))
|
||||
for hookType := range r.hooks {
|
||||
types = append(types, hookType)
|
||||
|
||||
@@ -0,0 +1,153 @@
|
||||
package resolvespec
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// jsonColModel has a real JSONB column so the dotted "data.x" shorthand is
|
||||
// recognised as JSON access.
|
||||
type jsonColModel struct {
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Name string `json:"name" bun:"name"`
|
||||
Data spectypes.SqlJSONB `json:"data" bun:"data"`
|
||||
}
|
||||
|
||||
type jsonCapCall struct {
|
||||
method string
|
||||
query string
|
||||
args []interface{}
|
||||
}
|
||||
|
||||
// jsonCapQuery records the string + args of the calls the handler makes.
|
||||
type jsonCapQuery struct {
|
||||
calls []jsonCapCall
|
||||
}
|
||||
|
||||
func (m *jsonCapQuery) rec(method, query string, args []interface{}) common.SelectQuery {
|
||||
m.calls = append(m.calls, jsonCapCall{method: method, query: query, args: args})
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *jsonCapQuery) Model(interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Table(string) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Column(cols ...string) common.SelectQuery {
|
||||
for _, c := range cols {
|
||||
m.rec("Column", c, nil)
|
||||
}
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) ColumnExpr(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("ColumnExpr", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) Where(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("Where", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) WhereOr(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("WhereOr", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) Join(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) LeftJoin(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Preload(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) PreloadRelation(string, ...func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) JoinRelation(string, ...func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) Order(o string) common.SelectQuery { return m.rec("Order", o, nil) }
|
||||
func (m *jsonCapQuery) OrderExpr(o string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("OrderExpr", o, args)
|
||||
}
|
||||
func (m *jsonCapQuery) Limit(int) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Offset(int) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Group(string) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Having(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Scan(context.Context, interface{}) error { return nil }
|
||||
func (m *jsonCapQuery) ScanModel(context.Context) error { return nil }
|
||||
func (m *jsonCapQuery) Count(context.Context) (int, error) { return 0, nil }
|
||||
func (m *jsonCapQuery) Exists(context.Context) (bool, error) { return false, nil }
|
||||
|
||||
func (m *jsonCapQuery) only(t *testing.T) jsonCapCall {
|
||||
t.Helper()
|
||||
if len(m.calls) != 1 {
|
||||
t.Fatalf("expected exactly 1 recorded call, got %d: %+v", len(m.calls), m.calls)
|
||||
}
|
||||
return m.calls[0]
|
||||
}
|
||||
|
||||
func TestBuildFilterCondition_JSONColumn(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := jsonColModel{}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
filter common.FilterOption
|
||||
wantCond string
|
||||
wantArgs []interface{}
|
||||
}{
|
||||
{
|
||||
name: "arrow syntax eq stays text",
|
||||
filter: common.FilterOption{Column: "data->>'city'", Operator: "eq", Value: "LA"},
|
||||
wantCond: `("data" #>> ?::text[]) = ?`,
|
||||
wantArgs: []interface{}{"{city}", "LA"},
|
||||
},
|
||||
{
|
||||
name: "dotted shorthand numeric cast inference",
|
||||
filter: common.FilterOption{Column: "data.age", Operator: "gt", Value: 18},
|
||||
wantCond: `(("data" #>> ?::text[]))::numeric > ?`,
|
||||
wantArgs: []interface{}{"{age}", 18},
|
||||
},
|
||||
{
|
||||
name: "hash path with explicit cast",
|
||||
filter: common.FilterOption{Column: "data#>>'{a,b}'::int", Operator: "lte", Value: "5"},
|
||||
wantCond: `(("data" #>> ?::text[]))::integer <= ?`,
|
||||
wantArgs: []interface{}{"{a,b}", "5"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cond, args := h.buildFilterCondition(tc.filter, model)
|
||||
if cond != tc.wantCond {
|
||||
t.Fatalf("cond = %q, want %q", cond, tc.wantCond)
|
||||
}
|
||||
if !reflect.DeepEqual(args, tc.wantArgs) {
|
||||
t.Fatalf("args = %#v, want %#v", args, tc.wantArgs)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Non-JSON column falls through to ordinary handling.
|
||||
cond, _ := h.buildFilterCondition(common.FilterOption{Column: "name", Operator: "eq", Value: "x"}, model)
|
||||
if cond != "name = ?" {
|
||||
t.Fatalf("non-JSON cond = %q", cond)
|
||||
}
|
||||
|
||||
// Without a model the dotted shorthand must NOT be treated as JSON.
|
||||
cond, _ = h.buildFilterCondition(common.FilterOption{Column: "data.age", Operator: "eq", Value: "x"}, nil)
|
||||
if cond != "data.age = ?" {
|
||||
t.Fatalf("nil-model dotted cond = %q, want ordinary handling", cond)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyFilter_JSONColumn(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := jsonColModel{}
|
||||
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "data->>'tier'", Operator: "in", Value: []string{"a", "b"}, LogicOperator: "OR",
|
||||
}, model)
|
||||
c := q.only(t)
|
||||
if c.method != "WhereOr" || c.query != `("data" #>> ?::text[]) IN (?,?)` {
|
||||
t.Fatalf("call = %+v", c)
|
||||
}
|
||||
if !reflect.DeepEqual(c.args, []interface{}{"{tier}", "a", "b"}) {
|
||||
t.Fatalf("args = %#v", c.args)
|
||||
}
|
||||
}
|
||||
+105
-88
@@ -2,6 +2,7 @@ package resolvespec
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"runtime/debug"
|
||||
"strings"
|
||||
|
||||
"github.com/gorilla/mux"
|
||||
@@ -12,6 +13,7 @@ import (
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/router"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
)
|
||||
|
||||
@@ -244,6 +246,11 @@ func wrapBunRouterHandler(handler bunrouter.HandlerFunc, authMiddleware Middlewa
|
||||
// Accepts bunrouter.Router or bunrouter.Group
|
||||
// authMiddleware is optional - if provided, routes will be protected with the middleware
|
||||
func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware MiddlewareFunc) {
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
logger.Error("panic in resolvespec.SetupBunRouterRoutes: %v\n%s", rec, debug.Stack())
|
||||
}
|
||||
}()
|
||||
|
||||
// CORS config
|
||||
corsConfig := common.DefaultCORSConfig()
|
||||
@@ -269,112 +276,122 @@ func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware M
|
||||
|
||||
// Loop through each registered model and create explicit routes
|
||||
for fullName := range allModels {
|
||||
// Parse the full name (e.g., "public.users" or just "users")
|
||||
schema, entity := parseModelName(fullName)
|
||||
func() {
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
logger.Error("panic registering resolvespec routes for model %s: %v\n%s", fullName, rec, debug.Stack())
|
||||
}
|
||||
}()
|
||||
|
||||
// Build the route paths
|
||||
entityPath := buildRoutePath(schema, entity)
|
||||
entityWithIDPath := entityPath + "/:id"
|
||||
// Parse the full name (e.g., "public.users" or just "users")
|
||||
schema, entity := parseModelName(fullName)
|
||||
|
||||
// Create closure variables to capture current schema and entity
|
||||
currentSchema := schema
|
||||
currentEntity := entity
|
||||
// Build the route paths
|
||||
entityPath := buildRoutePath(schema, entity)
|
||||
entityWithIDPath := entityPath + "/:id"
|
||||
|
||||
// POST route without ID
|
||||
postEntityHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
// Create closure variables to capture current schema and entity
|
||||
currentSchema := schema
|
||||
currentEntity := entity
|
||||
|
||||
// POST route without ID
|
||||
postEntityHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("POST", entityPath, wrapBunRouterHandler(postEntityHandler, authMiddleware))
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("POST", entityPath, wrapBunRouterHandler(postEntityHandler, authMiddleware))
|
||||
// POST route with ID
|
||||
postEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
}
|
||||
|
||||
// POST route with ID
|
||||
postEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("POST", entityWithIDPath, wrapBunRouterHandler(postEntityWithIDHandler, authMiddleware))
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("POST", entityWithIDPath, wrapBunRouterHandler(postEntityWithIDHandler, authMiddleware))
|
||||
// GET route without ID
|
||||
getEntityHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
|
||||
// GET route without ID
|
||||
getEntityHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("GET", entityPath, wrapBunRouterHandler(getEntityHandler, authMiddleware))
|
||||
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("GET", entityPath, wrapBunRouterHandler(getEntityHandler, authMiddleware))
|
||||
// GET route with ID
|
||||
getEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
}
|
||||
|
||||
// GET route with ID
|
||||
getEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("GET", entityWithIDPath, wrapBunRouterHandler(getEntityWithIDHandler, authMiddleware))
|
||||
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("GET", entityWithIDPath, wrapBunRouterHandler(getEntityWithIDHandler, authMiddleware))
|
||||
// OPTIONS route without ID (returns metadata)
|
||||
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||
r.Handle("OPTIONS", entityPath, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
optionsCorsConfig := corsConfig
|
||||
optionsCorsConfig.AllowedMethods = []string{"GET", "POST", "OPTIONS"}
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, optionsCorsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
|
||||
// OPTIONS route without ID (returns metadata)
|
||||
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||
r.Handle("OPTIONS", entityPath, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
optionsCorsConfig := corsConfig
|
||||
optionsCorsConfig.AllowedMethods = []string{"GET", "POST", "OPTIONS"}
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, optionsCorsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
})
|
||||
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
})
|
||||
// OPTIONS route with ID (returns metadata)
|
||||
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||
r.Handle("OPTIONS", entityWithIDPath, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
optionsCorsConfig := corsConfig
|
||||
optionsCorsConfig.AllowedMethods = []string{"POST", "OPTIONS"}
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, optionsCorsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
|
||||
// OPTIONS route with ID (returns metadata)
|
||||
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||
r.Handle("OPTIONS", entityWithIDPath, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||
optionsCorsConfig := corsConfig
|
||||
optionsCorsConfig.AllowedMethods = []string{"POST", "OPTIONS"}
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, optionsCorsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
})
|
||||
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
})
|
||||
logger.Info("Registered resolvespec bunrouter routes for model %s at %s", fullName, entityPath)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -25,12 +25,18 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
||||
// Hook 1: BeforeRead - Load security rules
|
||||
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
if security.IsModelSecurityDisabled(secCtx) {
|
||||
return nil
|
||||
}
|
||||
return security.LoadSecurityRules(secCtx, securityList)
|
||||
})
|
||||
|
||||
// Hook 2: BeforeScan - Apply row-level security filters
|
||||
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
if security.ShouldSkipRowSecurity(secCtx, hookCtx.Operation) {
|
||||
return nil
|
||||
}
|
||||
return security.ApplyRowSecurity(secCtx, securityList)
|
||||
})
|
||||
|
||||
@@ -78,6 +84,17 @@ func (s *securityContext) GetUserID() (int, bool) {
|
||||
return security.GetUserID(s.ctx.Context)
|
||||
}
|
||||
|
||||
// GetUserRef returns an opaque user identifier for row security lookups.
|
||||
// It prefers the full *security.UserContext (so providers can read JWT claims,
|
||||
// e.g. a UUID subject) and falls back to the int user ID.
|
||||
func (s *securityContext) GetUserRef() (any, bool) {
|
||||
if userCtx, ok := security.GetUserContext(s.ctx.Context); ok {
|
||||
return userCtx, true
|
||||
}
|
||||
userID, ok := security.GetUserID(s.ctx.Context)
|
||||
return userID, ok
|
||||
}
|
||||
|
||||
func (s *securityContext) GetSchema() string {
|
||||
return s.ctx.Schema
|
||||
}
|
||||
@@ -86,6 +103,10 @@ func (s *securityContext) GetEntity() string {
|
||||
return s.ctx.Entity
|
||||
}
|
||||
|
||||
func (s *securityContext) GetOperation() string {
|
||||
return s.ctx.Operation
|
||||
}
|
||||
|
||||
func (s *securityContext) GetModel() interface{} {
|
||||
return s.ctx.Model
|
||||
}
|
||||
|
||||
@@ -88,7 +88,7 @@ This will match any records where the column contains the search term (case-inse
|
||||
Search with specific operators (AND logic).
|
||||
|
||||
**Supported Operators:**
|
||||
- `contains` - Contains substring (case-insensitive)
|
||||
- `contains` - Contains substring (case-insensitive). Implemented as `CAST(col AS TEXT) ILIKE '%value%'` for every column type, including arrays (stringifies the array, then substring-matches). **Not** array containment — no GIN index use, and can false-positive on partial matches within array elements. resolvespec (a different spec package in this repo) defines `contains` differently: real PostgreSQL array-overlap (`&&`). Don't assume the two behave the same.
|
||||
- `beginswith` / `startswith` - Starts with (case-insensitive)
|
||||
- `endswith` - Ends with (case-insensitive)
|
||||
- `equals` / `eq` - Exact match
|
||||
|
||||
@@ -96,6 +96,8 @@ X-Limit: 50
|
||||
|
||||
**Available Operators**: `eq`, `neq`, `gt`, `gte`, `lt`, `lte`, `contains`, `startswith`, `endswith`, `between`, `betweeninclusive`, `in`, `empty`, `notempty`
|
||||
|
||||
> Note: `contains` here is a text-cast ILIKE substring match (works on any column type, including arrays, by stringifying first) — not array containment. resolvespec's `contains` operator has different semantics (real array overlap). See [HEADERS.md](HEADERS.md) for details.
|
||||
|
||||
For complete header documentation, see [HEADERS.md](HEADERS.md).
|
||||
|
||||
## Lifecycle Hooks
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
package restheadspec
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
)
|
||||
|
||||
func TestSetDefaultSort(t *testing.T) {
|
||||
handler := NewHandler(nil, nil)
|
||||
|
||||
// No default configured yet
|
||||
if got := handler.getDefaultSort("public", "users"); got != nil {
|
||||
t.Errorf("Expected no default sort, got %v", got)
|
||||
}
|
||||
|
||||
// Global default
|
||||
global := []common.SortOption{{Column: "created_at", Direction: "desc"}}
|
||||
handler.SetDefaultSort("", "", global...)
|
||||
|
||||
if got := handler.getDefaultSort("public", "users"); !reflect.DeepEqual(got, global) {
|
||||
t.Errorf("Expected global default sort %v, got %v", global, got)
|
||||
}
|
||||
|
||||
// Per-model default overrides the global default
|
||||
perModel := []common.SortOption{{Column: "name", Direction: "asc"}}
|
||||
handler.SetDefaultSort("public", "users", perModel...)
|
||||
|
||||
if got := handler.getDefaultSort("public", "users"); !reflect.DeepEqual(got, perModel) {
|
||||
t.Errorf("Expected per-model default sort %v, got %v", perModel, got)
|
||||
}
|
||||
|
||||
// Other models still fall back to the global default
|
||||
if got := handler.getDefaultSort("public", "orders"); !reflect.DeepEqual(got, global) {
|
||||
t.Errorf("Expected global default sort %v for unrelated model, got %v", global, got)
|
||||
}
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// detailTestModel is a simple model with gorm column/type tags for detail format tests.
|
||||
@@ -207,3 +208,30 @@ func TestBuildDetailFields_SkipsRelations(t *testing.T) {
|
||||
t.Errorf("expected 2 scalar fields (id, name), got %d", len(fields))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildDetailFields_UnwrapsSQLTypes(t *testing.T) {
|
||||
type model struct {
|
||||
Name spectypes.SqlString `gorm:"column:name" json:"name"`
|
||||
CreatedAt spectypes.SqlTimeStamp `gorm:"column:created_at" json:"created_at"`
|
||||
Metadata spectypes.SqlJSONB `gorm:"column:metadata;type:jsonb" json:"metadata"`
|
||||
}
|
||||
|
||||
fields := (&Handler{}).buildDetailFields(model{})
|
||||
byName := make(map[string]string, len(fields))
|
||||
for _, field := range fields {
|
||||
byName[field.Name] = field.DataType
|
||||
if !field.Nullable {
|
||||
t.Errorf("%q: expected SQL wrapper to be nullable", field.Name)
|
||||
}
|
||||
}
|
||||
|
||||
for name, want := range map[string]string{
|
||||
"name": "string",
|
||||
"created_at": "unknown",
|
||||
"metadata": "unknown",
|
||||
} {
|
||||
if got := byName[name]; got != want {
|
||||
t.Errorf("%q: expected type %q, got %q", name, want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,175 @@
|
||||
package restheadspec
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// atdetailModel mirrors the real-world model that triggered this regression:
|
||||
// rid_parent is a nullable bigint foreign key, backed by spectypes.SqlInt64
|
||||
// (a SqlNull[int64] alias). An eq filter on it was being rendered as
|
||||
// CAST(atdetail.rid_parent AS TEXT) = '90446096', which can't use the index
|
||||
// on rid_parent. Name is a citext column, which must never be cast to TEXT
|
||||
// either (that would switch to case-sensitive matching and lose its index).
|
||||
type atdetailModel struct {
|
||||
RidParent spectypes.SqlInt64 `json:"rid_parent" bun:"rid_parent"`
|
||||
Name string `json:"name" bun:"name,type:citext"`
|
||||
}
|
||||
|
||||
func TestValidateAndAdjustFilterForColumnType_SqlNullNumeric(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := atdetailModel{}
|
||||
|
||||
filter := &common.FilterOption{Column: "rid_parent", Operator: "eq", Value: "90446096"}
|
||||
info := h.ValidateAndAdjustFilterForColumnType(filter, model)
|
||||
|
||||
if info.NeedsCast {
|
||||
t.Fatalf("expected NeedsCast=false for a numeric SqlInt64 column with a numeric value, got true")
|
||||
}
|
||||
if !info.IsNumericType {
|
||||
t.Fatalf("expected IsNumericType=true for a SqlInt64 column")
|
||||
}
|
||||
if v, ok := filter.Value.(int64); !ok || v != 90446096 {
|
||||
t.Fatalf("expected filter.Value to be converted to int64(90446096), got %#v", filter.Value)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyFilter_SqlNullNumeric_NoCastKeepsIndexUsable(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := atdetailModel{}
|
||||
|
||||
filter := common.FilterOption{Column: "rid_parent", Operator: "eq", Value: "90446096"}
|
||||
castInfo := h.ValidateAndAdjustFilterForColumnType(&filter, model)
|
||||
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, filter, "public.atdetail", castInfo.NeedsCast, "AND", model)
|
||||
|
||||
c := q.only(t)
|
||||
const want = "atdetail.rid_parent = ?"
|
||||
if c.query != want {
|
||||
t.Fatalf("query = %q, want %q (must not CAST a numeric column to TEXT)", c.query, want)
|
||||
}
|
||||
if !reflect.DeepEqual(c.args, []interface{}{int64(90446096)}) {
|
||||
t.Fatalf("args = %#v", c.args)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFieldFilterHeader_SqlNullNumeric_EndToEnd reproduces the exact reported
|
||||
// regression: a request carrying the header
|
||||
//
|
||||
// x-fieldfilter-rid_parent: 90446096
|
||||
//
|
||||
// against a model whose rid_parent field is a nullable bigint (spectypes.SqlInt64).
|
||||
// Before the fix, this parsed to a filter that got CAST(atdetail.rid_parent AS TEXT) = '90446096',
|
||||
// making the query unable to use the index on rid_parent. It must now parse to
|
||||
// a native "atdetail.rid_parent = ?" comparison with an int64 argument.
|
||||
func TestFieldFilterHeader_SqlNullNumeric_EndToEnd(t *testing.T) {
|
||||
h := NewHandler(nil, nil)
|
||||
model := atdetailModel{}
|
||||
|
||||
req := &MockRequest{
|
||||
headers: map[string]string{
|
||||
"x-fieldfilter-rid_parent": "90446096",
|
||||
},
|
||||
queryParams: map[string]string{},
|
||||
}
|
||||
|
||||
options := h.parseOptionsFromHeaders(req, model)
|
||||
if len(options.Filters) != 1 {
|
||||
t.Fatalf("expected 1 filter parsed from x-fieldfilter-rid_parent, got %d: %+v", len(options.Filters), options.Filters)
|
||||
}
|
||||
|
||||
filter := options.Filters[0]
|
||||
if filter.Column != "rid_parent" || filter.Operator != "eq" {
|
||||
t.Fatalf("unexpected parsed filter: %+v", filter)
|
||||
}
|
||||
if filter.Value != "90446096" {
|
||||
t.Fatalf("expected raw header string value before type validation, got %#v", filter.Value)
|
||||
}
|
||||
|
||||
// This is the exact step that decided whether to CAST: ValidateAndAdjustFilterForColumnType
|
||||
// used to see reflect.Struct for the SqlInt64-wrapped column and cast to TEXT.
|
||||
castInfo := h.ValidateAndAdjustFilterForColumnType(&filter, model)
|
||||
if castInfo.NeedsCast {
|
||||
t.Fatalf("regression: numeric SqlInt64 column x-fieldfilter-rid_parent got NeedsCast=true, " +
|
||||
"which renders CAST(atdetail.rid_parent AS TEXT) = '90446096' and defeats the column's index")
|
||||
}
|
||||
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, filter, "public.atdetail", castInfo.NeedsCast, filter.LogicOperator, model)
|
||||
|
||||
c := q.only(t)
|
||||
const want = "atdetail.rid_parent = ?"
|
||||
if c.query != want {
|
||||
t.Fatalf("SQL condition = %q, want %q (no CAST, so the rid_parent index can still be used)", c.query, want)
|
||||
}
|
||||
if !reflect.DeepEqual(c.args, []interface{}{int64(90446096)}) {
|
||||
t.Fatalf("args = %#v, want [int64(90446096)]", c.args)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyFilter_Citext_NeverCastForEqOrIlike(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := atdetailModel{}
|
||||
|
||||
t.Run("eq", func(t *testing.T) {
|
||||
filter := common.FilterOption{Column: "name", Operator: "eq", Value: "Acme"}
|
||||
castInfo := h.ValidateAndAdjustFilterForColumnType(&filter, model)
|
||||
if castInfo.NeedsCast {
|
||||
t.Fatalf("citext column must never need a CAST")
|
||||
}
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, filter, "public.atdetail", castInfo.NeedsCast, "AND", model)
|
||||
if c := q.only(t); c.query != "atdetail.name = ?" {
|
||||
t.Fatalf("query = %q", c.query)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("ilike", func(t *testing.T) {
|
||||
filter := common.FilterOption{Column: "name", Operator: "ilike", Value: "%acme%"}
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, filter, "public.atdetail", false, "AND", model)
|
||||
if c := q.only(t); c.query != "atdetail.name ILIKE ?" {
|
||||
t.Fatalf("query = %q, want no CAST for a citext column", c.query)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestValidateAndAdjustFilterForColumnType_NumericColumn_Ilike reproduces a
|
||||
// global "search all columns" request (x-searchor-contains-<col> per column,
|
||||
// e.g. the X-Filter-All style OR group) landing an ILIKE filter with a
|
||||
// '%...%'-wrapped numeric-looking value on a numeric column such as
|
||||
// rid_parent. Before the fix, ValidateAndAdjustFilterForColumnType trimmed
|
||||
// the '%' wildcards, saw a numeric string, and rewrote filter.Value to an
|
||||
// int64 -- so applyFilter's CAST(col AS TEXT) ILIKE ? bound an integer
|
||||
// argument instead of the wildcard string, and Postgres rejected it with
|
||||
// "operator does not exist: text ~~* integer".
|
||||
func TestValidateAndAdjustFilterForColumnType_NumericColumn_Ilike(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := atdetailModel{}
|
||||
|
||||
filter := &common.FilterOption{Column: "rid_parent", Operator: "ilike", Value: "%345346346%"}
|
||||
info := h.ValidateAndAdjustFilterForColumnType(filter, model)
|
||||
|
||||
if !info.NeedsCast {
|
||||
t.Fatalf("expected NeedsCast=true so the numeric column is cast to TEXT for ILIKE")
|
||||
}
|
||||
if filter.Value != "%345346346%" {
|
||||
t.Fatalf("ILIKE must keep the wildcard-wrapped string value untouched, got %#v", filter.Value)
|
||||
}
|
||||
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, *filter, "public.atdetail", info.NeedsCast, "OR", model)
|
||||
|
||||
c := q.only(t)
|
||||
const want = "CAST(atdetail.rid_parent AS TEXT) ILIKE ?"
|
||||
if c.query != want {
|
||||
t.Fatalf("query = %q, want %q", c.query, want)
|
||||
}
|
||||
if !reflect.DeepEqual(c.args, []interface{}{"%345346346%"}) {
|
||||
t.Fatalf("args = %#v, want [\"%%345346346%%\"]", c.args)
|
||||
}
|
||||
}
|
||||
+789
-489
File diff suppressed because it is too large
Load Diff
@@ -182,6 +182,22 @@ func (h *Handler) parseOptionsFromHeaders(r common.Request, model interface{}) E
|
||||
h.parseSearchOp(&options, key, decodedValue, "AND")
|
||||
case strings.HasPrefix(key, "x-searchcols"):
|
||||
options.SearchColumns = h.parseCommaSeparated(decodedValue)
|
||||
case strings.HasPrefix(key, "x-spatialfilter-"):
|
||||
h.parseGeoFilter(&options, key, "x-spatialfilter-", decodedValue)
|
||||
case strings.HasPrefix(key, "x-vectorfilter-"):
|
||||
h.parseGeoFilter(&options, key, "x-vectorfilter-", decodedValue)
|
||||
|
||||
// pgvector KNN search
|
||||
case key == "x-vector-search-vector":
|
||||
h.ensureVectorSearch(&options).Vector = parseFloat32List(decodedValue)
|
||||
case key == "x-vector-search-as":
|
||||
h.ensureVectorSearch(&options).As = decodedValue
|
||||
case key == "x-vector-search-dir":
|
||||
h.ensureVectorSearch(&options).Direction = decodedValue
|
||||
case strings.HasPrefix(key, "x-vector-search-"):
|
||||
vs := h.ensureVectorSearch(&options)
|
||||
vs.Column = strings.TrimPrefix(key, "x-vector-search-")
|
||||
vs.Metric = decodedValue
|
||||
case strings.HasPrefix(key, "x-custom-sql-w"):
|
||||
if options.CustomSQLWhere != "" {
|
||||
options.CustomSQLWhere = fmt.Sprintf("%s AND (%s)", options.CustomSQLWhere, decodedValue)
|
||||
@@ -309,6 +325,83 @@ func (h *Handler) parseOptionsFromHeaders(r common.Request, model interface{}) E
|
||||
return options
|
||||
}
|
||||
|
||||
// ensureVectorSearch returns the options' VectorSearchOption, allocating it on
|
||||
// first use.
|
||||
func (h *Handler) ensureVectorSearch(options *ExtendedRequestOptions) *common.VectorSearchOption {
|
||||
if options.VectorSearch == nil {
|
||||
options.VectorSearch = &common.VectorSearchOption{}
|
||||
}
|
||||
return options.VectorSearch
|
||||
}
|
||||
|
||||
// parseFloat32List parses a JSON array ("[1,2,3]") or comma-separated list into
|
||||
// a []float32.
|
||||
func parseFloat32List(value string) []float32 {
|
||||
value = strings.TrimSpace(value)
|
||||
if value == "" {
|
||||
return nil
|
||||
}
|
||||
var raw []float64
|
||||
if err := json.Unmarshal([]byte(value), &raw); err == nil {
|
||||
out := make([]float32, len(raw))
|
||||
for i, f := range raw {
|
||||
out[i] = float32(f)
|
||||
}
|
||||
return out
|
||||
}
|
||||
parts := strings.Split(strings.Trim(value, "[]"), ",")
|
||||
out := make([]float32, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
f, err := strconv.ParseFloat(strings.TrimSpace(p), 32)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
out = append(out, float32(f))
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// parseGeoFilter parses an x-spatialfilter-<col> / x-vectorfilter-<col> header.
|
||||
// The value is a JSON object: {"op":"st_dwithin","geom":...,"distance":...} or
|
||||
// {"op":"st_intersects","value":<geojson>}. An optional "logic":"or" controls
|
||||
// how the filter combines with the previous one.
|
||||
func (h *Handler) parseGeoFilter(options *ExtendedRequestOptions, key, prefix, value string) {
|
||||
col := strings.TrimPrefix(key, prefix)
|
||||
if col == "" || strings.TrimSpace(value) == "" {
|
||||
return
|
||||
}
|
||||
var raw map[string]interface{}
|
||||
if err := json.Unmarshal([]byte(value), &raw); err != nil {
|
||||
logger.Warn("Invalid %s%s filter JSON: %v", prefix, col, err)
|
||||
return
|
||||
}
|
||||
op, _ := raw["op"].(string)
|
||||
if op == "" {
|
||||
logger.Warn("%s%s filter missing \"op\"", prefix, col)
|
||||
return
|
||||
}
|
||||
logicOp := "AND"
|
||||
if lo, ok := raw["logic"].(string); ok && strings.EqualFold(lo, "or") {
|
||||
logicOp = "OR"
|
||||
}
|
||||
|
||||
var fv interface{}
|
||||
if v, ok := raw["value"]; ok {
|
||||
fv = v
|
||||
} else {
|
||||
delete(raw, "op")
|
||||
delete(raw, "logic")
|
||||
fv = raw
|
||||
}
|
||||
|
||||
options.Filters = append(options.Filters, common.FilterOption{
|
||||
Column: col,
|
||||
Operator: op,
|
||||
Value: fv,
|
||||
LogicOperator: logicOp,
|
||||
})
|
||||
}
|
||||
|
||||
// parseSelectFields parses x-select-fields header
|
||||
func (h *Handler) parseSelectFields(options *ExtendedRequestOptions, value string) {
|
||||
if value == "" {
|
||||
@@ -1365,6 +1458,20 @@ func (h *Handler) ValidateAndAdjustFilterForColumnType(filter *common.FilterOpti
|
||||
return ColumnCastInfo{NeedsCast: false, IsNumericType: false}
|
||||
}
|
||||
|
||||
// Never cast geometry/geography or pgvector columns to TEXT — spatial and
|
||||
// vector operators need the native column type. Also bypass when the
|
||||
// operator itself is spatial/vector (e.g. st_dwithin, l2_within).
|
||||
if common.IsSpatialOperator(filter.Operator) || common.IsVectorOperator(filter.Operator) ||
|
||||
reflection.IsSpatialColumn(model, filter.Column) || reflection.IsVectorColumn(model, filter.Column) {
|
||||
return ColumnCastInfo{NeedsCast: false, IsNumericType: false}
|
||||
}
|
||||
|
||||
// Never cast citext columns to TEXT: CAST(col AS TEXT) swaps in case-sensitive
|
||||
// comparison semantics and prevents PostgreSQL from using a citext index.
|
||||
if reflection.IsCitextColumn(model, filter.Column) {
|
||||
return ColumnCastInfo{NeedsCast: false, IsNumericType: false}
|
||||
}
|
||||
|
||||
colType := reflection.GetColumnTypeFromModel(model, filter.Column)
|
||||
if colType == reflect.Invalid {
|
||||
// Column not found in model, no casting needed
|
||||
@@ -1372,6 +1479,18 @@ func (h *Handler) ValidateAndAdjustFilterForColumnType(filter *common.FilterOpti
|
||||
return ColumnCastInfo{NeedsCast: false, IsNumericType: false}
|
||||
}
|
||||
|
||||
// LIKE/ILIKE always compare against text, wildcards and all. Never coerce
|
||||
// the value to the column's native numeric/bool/time type here: doing so
|
||||
// strips the '%' wildcards and hands the driver a non-string argument,
|
||||
// which fails with "operator does not exist: text ~~* integer" once the
|
||||
// column is cast to TEXT below.
|
||||
if op := strings.ToLower(filter.Operator); op == "like" || op == "ilike" {
|
||||
if reflection.IsStringType(colType) {
|
||||
return ColumnCastInfo{NeedsCast: false, IsNumericType: false}
|
||||
}
|
||||
return ColumnCastInfo{NeedsCast: true, IsNumericType: reflection.IsNumericType(colType)}
|
||||
}
|
||||
|
||||
// Check if the input value is numeric
|
||||
valueIsNumeric := false
|
||||
if strVal, ok := filter.Value.(string); ok {
|
||||
|
||||
@@ -3,6 +3,8 @@ package restheadspec
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
@@ -34,6 +36,11 @@ const (
|
||||
|
||||
// Scan/Execute operation hooks
|
||||
BeforeScan HookType = "before_scan"
|
||||
|
||||
// BeforeOp fires immediately before every SQL operation (read, create, update, delete, scan).
|
||||
// Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each
|
||||
// individual SQL-operation hook point, so it runs once per statement executed.
|
||||
BeforeOp HookType = "before_op"
|
||||
)
|
||||
|
||||
// HookContext contains all the data available to a hook
|
||||
@@ -84,6 +91,7 @@ type HookFunc func(*HookContext) error
|
||||
// HookRegistry manages all registered hooks
|
||||
type HookRegistry struct {
|
||||
hooks map[HookType][]HookFunc
|
||||
mutex sync.RWMutex
|
||||
}
|
||||
|
||||
// NewHookRegistry creates a new hook registry
|
||||
@@ -93,8 +101,46 @@ func NewHookRegistry() *HookRegistry {
|
||||
}
|
||||
}
|
||||
|
||||
// hookLockRetryAttempts/hookLockRetryDelay bound how long the try-lock
|
||||
// helpers below will spin before giving up, so a contended mutex can
|
||||
// never hang a caller.
|
||||
const (
|
||||
hookLockRetryAttempts = 20
|
||||
hookLockRetryDelay = 1 * time.Millisecond
|
||||
)
|
||||
|
||||
// tryLock attempts to acquire the write lock, retrying briefly. Returns
|
||||
// false if it could not be acquired within the bound.
|
||||
func (r *HookRegistry) tryLock() bool {
|
||||
for i := 0; i < hookLockRetryAttempts; i++ {
|
||||
if r.mutex.TryLock() {
|
||||
return true
|
||||
}
|
||||
time.Sleep(hookLockRetryDelay)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// tryRLock attempts to acquire the read lock, retrying briefly. Returns
|
||||
// false if it could not be acquired within the bound.
|
||||
func (r *HookRegistry) tryRLock() bool {
|
||||
for i := 0; i < hookLockRetryAttempts; i++ {
|
||||
if r.mutex.TryRLock() {
|
||||
return true
|
||||
}
|
||||
time.Sleep(hookLockRetryDelay)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Register adds a new hook for the specified hook type
|
||||
func (r *HookRegistry) Register(hookType HookType, hook HookFunc) {
|
||||
if !r.tryLock() {
|
||||
logger.Error("Failed to register hook for %s: registry locked", hookType)
|
||||
return
|
||||
}
|
||||
defer r.mutex.Unlock()
|
||||
|
||||
if r.hooks == nil {
|
||||
r.hooks = make(map[HookType][]HookFunc)
|
||||
}
|
||||
@@ -112,8 +158,13 @@ func (r *HookRegistry) RegisterMultiple(hookTypes []HookType, hook HookFunc) {
|
||||
// Execute runs all hooks for the specified type in order
|
||||
// If any hook returns an error, execution stops and the error is returned
|
||||
func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
||||
hooks, exists := r.hooks[hookType]
|
||||
if !exists || len(hooks) == 0 {
|
||||
if !r.tryRLock() {
|
||||
return fmt.Errorf("hook execution failed: registry locked")
|
||||
}
|
||||
hooks := append([]HookFunc(nil), r.hooks[hookType]...)
|
||||
r.mutex.RUnlock()
|
||||
|
||||
if len(hooks) == 0 {
|
||||
// logger.Debug("No hooks registered for %s", hookType)
|
||||
return nil
|
||||
}
|
||||
@@ -137,20 +188,47 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
|
||||
// (BeforeRead, BeforeCreate, BeforeUpdate, BeforeDelete, or BeforeScan). BeforeOp always
|
||||
// runs first so it can observe/veto every SQL operation regardless of type.
|
||||
func (r *HookRegistry) ExecuteBeforeOp(hookType HookType, ctx *HookContext) error {
|
||||
if err := r.Execute(BeforeOp, ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
return r.Execute(hookType, ctx)
|
||||
}
|
||||
|
||||
// Clear removes all hooks for the specified type
|
||||
func (r *HookRegistry) Clear(hookType HookType) {
|
||||
if !r.tryLock() {
|
||||
logger.Error("Failed to clear hooks for %s: registry locked", hookType)
|
||||
return
|
||||
}
|
||||
defer r.mutex.Unlock()
|
||||
|
||||
delete(r.hooks, hookType)
|
||||
logger.Info("Cleared all hooks for %s", hookType)
|
||||
}
|
||||
|
||||
// ClearAll removes all registered hooks
|
||||
func (r *HookRegistry) ClearAll() {
|
||||
if !r.tryLock() {
|
||||
logger.Error("Failed to clear all hooks: registry locked")
|
||||
return
|
||||
}
|
||||
defer r.mutex.Unlock()
|
||||
|
||||
r.hooks = make(map[HookType][]HookFunc)
|
||||
logger.Info("Cleared all hooks")
|
||||
}
|
||||
|
||||
// Count returns the number of hooks registered for a specific type
|
||||
func (r *HookRegistry) Count(hookType HookType) int {
|
||||
if !r.tryRLock() {
|
||||
return 0
|
||||
}
|
||||
defer r.mutex.RUnlock()
|
||||
|
||||
if hooks, exists := r.hooks[hookType]; exists {
|
||||
return len(hooks)
|
||||
}
|
||||
@@ -164,6 +242,11 @@ func (r *HookRegistry) HasHooks(hookType HookType) bool {
|
||||
|
||||
// GetAllHookTypes returns all hook types that have registered hooks
|
||||
func (r *HookRegistry) GetAllHookTypes() []HookType {
|
||||
if !r.tryRLock() {
|
||||
return nil
|
||||
}
|
||||
defer r.mutex.RUnlock()
|
||||
|
||||
types := make([]HookType, 0, len(r.hooks))
|
||||
for hookType := range r.hooks {
|
||||
types = append(types, hookType)
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
package restheadspec
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
)
|
||||
|
||||
// jsonColModel exercises the JSON-column wiring: Data is a real JSONB column so
|
||||
// the dotted "data.x" shorthand is recognised as JSON access.
|
||||
type jsonColModel struct {
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Name string `json:"name" bun:"name"`
|
||||
Data spectypes.SqlJSONB `json:"data" bun:"data"`
|
||||
}
|
||||
|
||||
// jsonCapQuery is a minimal common.SelectQuery that records the string + args of
|
||||
// the calls the handler makes so a test can assert on them.
|
||||
type jsonCapQuery struct {
|
||||
calls []jsonCapCall
|
||||
}
|
||||
|
||||
type jsonCapCall struct {
|
||||
method string
|
||||
query string
|
||||
args []interface{}
|
||||
}
|
||||
|
||||
func (m *jsonCapQuery) rec(method, query string, args []interface{}) common.SelectQuery {
|
||||
m.calls = append(m.calls, jsonCapCall{method: method, query: query, args: args})
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *jsonCapQuery) Model(interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Table(string) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Column(cols ...string) common.SelectQuery {
|
||||
for _, c := range cols {
|
||||
m.rec("Column", c, nil)
|
||||
}
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) ColumnExpr(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("ColumnExpr", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) Where(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("Where", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) WhereOr(q string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("WhereOr", q, args)
|
||||
}
|
||||
func (m *jsonCapQuery) WhereIn(col string, values interface{}) common.SelectQuery {
|
||||
return m.rec("WhereIn", col, []interface{}{values})
|
||||
}
|
||||
func (m *jsonCapQuery) Order(o string) common.SelectQuery { return m.rec("Order", o, nil) }
|
||||
func (m *jsonCapQuery) OrderExpr(o string, args ...interface{}) common.SelectQuery {
|
||||
return m.rec("OrderExpr", o, args)
|
||||
}
|
||||
func (m *jsonCapQuery) Limit(int) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Offset(int) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Join(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) LeftJoin(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Group(string) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Having(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) Preload(string, ...interface{}) common.SelectQuery { return m }
|
||||
func (m *jsonCapQuery) PreloadRelation(string, ...func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) JoinRelation(string, ...func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||
return m
|
||||
}
|
||||
func (m *jsonCapQuery) Scan(context.Context, interface{}) error { return nil }
|
||||
func (m *jsonCapQuery) ScanModel(context.Context) error { return nil }
|
||||
func (m *jsonCapQuery) Count(context.Context) (int, error) { return 0, nil }
|
||||
func (m *jsonCapQuery) Exists(context.Context) (bool, error) { return false, nil }
|
||||
func (m *jsonCapQuery) GetUnderlyingQuery() interface{} { return nil }
|
||||
func (m *jsonCapQuery) GetModel() interface{} { return nil }
|
||||
|
||||
func (m *jsonCapQuery) only(t *testing.T) jsonCapCall {
|
||||
t.Helper()
|
||||
if len(m.calls) != 1 {
|
||||
t.Fatalf("expected exactly 1 recorded call, got %d: %+v", len(m.calls), m.calls)
|
||||
}
|
||||
return m.calls[0]
|
||||
}
|
||||
|
||||
func TestApplyFilter_JSONColumn(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := jsonColModel{}
|
||||
|
||||
t.Run("arrow syntax eq", func(t *testing.T) {
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "data->>'city'", Operator: "eq", Value: "LA",
|
||||
}, "public.things", false, "AND", model)
|
||||
c := q.only(t)
|
||||
if c.method != "Where" || c.query != `("things"."data" #>> ?::text[]) = ?` {
|
||||
t.Fatalf("call = %+v", c)
|
||||
}
|
||||
if !reflect.DeepEqual(c.args, []interface{}{"{city}", "LA"}) {
|
||||
t.Fatalf("args = %#v", c.args)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("dotted shorthand with numeric cast inference, OR logic", func(t *testing.T) {
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "data.age", Operator: "gt", Value: 18,
|
||||
}, "public.things", false, "OR", model)
|
||||
c := q.only(t)
|
||||
if c.method != "WhereOr" || c.query != `(("things"."data" #>> ?::text[]))::numeric > ?` {
|
||||
t.Fatalf("call = %+v", c)
|
||||
}
|
||||
if !reflect.DeepEqual(c.args, []interface{}{"{age}", 18}) {
|
||||
t.Fatalf("args = %#v", c.args)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-JSON column is untouched", func(t *testing.T) {
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "name", Operator: "eq", Value: "x",
|
||||
}, "public.things", false, "AND", model)
|
||||
c := q.only(t)
|
||||
if c.query != "things.name = ?" {
|
||||
t.Fatalf("call = %+v", c)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("nil model: explicit syntax still works, dotted does not", func(t *testing.T) {
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "data->>'city'", Operator: "eq", Value: "LA",
|
||||
}, "public.things", false, "AND", nil)
|
||||
if c := q.only(t); c.query != `("things"."data" #>> ?::text[]) = ?` {
|
||||
t.Fatalf("explicit call = %+v", c)
|
||||
}
|
||||
|
||||
q2 := &jsonCapQuery{}
|
||||
h.applyFilter(q2, common.FilterOption{
|
||||
Column: "data.city", Operator: "eq", Value: "LA",
|
||||
}, "public.things", false, "AND", nil)
|
||||
if c := q2.only(t); c.query == `("things"."data" #>> ?::text[]) = ?` {
|
||||
t.Fatalf("dotted shorthand should not resolve without a model: %+v", c)
|
||||
}
|
||||
})
|
||||
}
|
||||
+151
-135
@@ -55,6 +55,7 @@ package restheadspec
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"runtime/debug"
|
||||
"strings"
|
||||
|
||||
"github.com/gorilla/mux"
|
||||
@@ -308,6 +309,11 @@ func wrapBunRouterHandler(handler bunrouter.HandlerFunc, authMiddleware Middlewa
|
||||
// Accepts bunrouter.Router or bunrouter.Group
|
||||
// authMiddleware is optional - if provided, routes will be protected with the middleware
|
||||
func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware MiddlewareFunc) {
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
logger.Error("panic in restheadspec.SetupBunRouterRoutes: %v\n%s", rec, debug.Stack())
|
||||
}
|
||||
}()
|
||||
|
||||
// CORS config
|
||||
corsConfig := common.DefaultCORSConfig()
|
||||
@@ -333,171 +339,181 @@ func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware M
|
||||
|
||||
// Loop through each registered model and create explicit routes
|
||||
for fullName := range allModels {
|
||||
// Parse the full name (e.g., "public.users" or just "users")
|
||||
schema, entity := parseModelName(fullName)
|
||||
func() {
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
logger.Error("panic registering restheadspec routes for model %s: %v\n%s", fullName, rec, debug.Stack())
|
||||
}
|
||||
}()
|
||||
|
||||
// Build the route paths
|
||||
entityPath := buildRoutePath(schema, entity)
|
||||
entityWithIDPath := entityPath + "/:id"
|
||||
metadataPath := entityPath + "/metadata"
|
||||
// Parse the full name (e.g., "public.users" or just "users")
|
||||
schema, entity := parseModelName(fullName)
|
||||
|
||||
// Create closure variables to capture current schema and entity
|
||||
currentSchema := schema
|
||||
currentEntity := entity
|
||||
// Build the route paths
|
||||
entityPath := buildRoutePath(schema, entity)
|
||||
entityWithIDPath := entityPath + "/:id"
|
||||
metadataPath := entityPath + "/metadata"
|
||||
|
||||
// GET and POST for /{schema}/{entity}
|
||||
getEntityHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
// Create closure variables to capture current schema and entity
|
||||
currentSchema := schema
|
||||
currentEntity := entity
|
||||
|
||||
// GET and POST for /{schema}/{entity}
|
||||
getEntityHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("GET", entityPath, wrapBunRouterHandler(getEntityHandler, authMiddleware))
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("GET", entityPath, wrapBunRouterHandler(getEntityHandler, authMiddleware))
|
||||
postEntityHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
|
||||
postEntityHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("POST", entityPath, wrapBunRouterHandler(postEntityHandler, authMiddleware))
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("POST", entityPath, wrapBunRouterHandler(postEntityHandler, authMiddleware))
|
||||
// GET, POST, PUT, PATCH, DELETE for /{schema}/{entity}/:id
|
||||
getEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
}
|
||||
|
||||
// GET, POST, PUT, PATCH, DELETE for /{schema}/{entity}/:id
|
||||
getEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("GET", entityWithIDPath, wrapBunRouterHandler(getEntityWithIDHandler, authMiddleware))
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("GET", entityWithIDPath, wrapBunRouterHandler(getEntityWithIDHandler, authMiddleware))
|
||||
postEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
}
|
||||
|
||||
postEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("POST", entityWithIDPath, wrapBunRouterHandler(postEntityWithIDHandler, authMiddleware))
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("POST", entityWithIDPath, wrapBunRouterHandler(postEntityWithIDHandler, authMiddleware))
|
||||
putEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
}
|
||||
|
||||
putEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("PUT", entityWithIDPath, wrapBunRouterHandler(putEntityWithIDHandler, authMiddleware))
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("PUT", entityWithIDPath, wrapBunRouterHandler(putEntityWithIDHandler, authMiddleware))
|
||||
patchEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
}
|
||||
|
||||
patchEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("PATCH", entityWithIDPath, wrapBunRouterHandler(patchEntityWithIDHandler, authMiddleware))
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("PATCH", entityWithIDPath, wrapBunRouterHandler(patchEntityWithIDHandler, authMiddleware))
|
||||
deleteEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
}
|
||||
|
||||
deleteEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
"id": req.Param("id"),
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("DELETE", entityWithIDPath, wrapBunRouterHandler(deleteEntityWithIDHandler, authMiddleware))
|
||||
|
||||
handler.Handle(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("DELETE", entityWithIDPath, wrapBunRouterHandler(deleteEntityWithIDHandler, authMiddleware))
|
||||
// Metadata endpoint
|
||||
metadataHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
|
||||
// Metadata endpoint
|
||||
metadataHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("GET", metadataPath, wrapBunRouterHandler(metadataHandler, authMiddleware))
|
||||
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
}
|
||||
r.Handle("GET", metadataPath, wrapBunRouterHandler(metadataHandler, authMiddleware))
|
||||
// OPTIONS route without ID (returns metadata)
|
||||
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||
r.Handle("OPTIONS", entityPath, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
optionsCorsConfig := corsConfig
|
||||
optionsCorsConfig.AllowedMethods = []string{"GET", "POST", "OPTIONS"}
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, optionsCorsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
|
||||
// OPTIONS route without ID (returns metadata)
|
||||
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||
r.Handle("OPTIONS", entityPath, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
optionsCorsConfig := corsConfig
|
||||
optionsCorsConfig.AllowedMethods = []string{"GET", "POST", "OPTIONS"}
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, optionsCorsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
})
|
||||
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
})
|
||||
// OPTIONS route with ID (returns metadata)
|
||||
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||
r.Handle("OPTIONS", entityWithIDPath, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
optionsCorsConfig := corsConfig
|
||||
optionsCorsConfig.AllowedMethods = []string{"GET", "PUT", "PATCH", "DELETE", "POST", "OPTIONS"}
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, optionsCorsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
|
||||
// OPTIONS route with ID (returns metadata)
|
||||
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||
r.Handle("OPTIONS", entityWithIDPath, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
respAdapter := router.NewHTTPResponseWriter(w)
|
||||
reqAdapter := router.NewBunRouterRequest(req)
|
||||
optionsCorsConfig := corsConfig
|
||||
optionsCorsConfig.AllowedMethods = []string{"GET", "PUT", "PATCH", "DELETE", "POST", "OPTIONS"}
|
||||
common.SetCORSHeaders(respAdapter, reqAdapter, optionsCorsConfig)
|
||||
params := map[string]string{
|
||||
"schema": currentSchema,
|
||||
"entity": currentEntity,
|
||||
}
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
})
|
||||
|
||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||
return nil
|
||||
})
|
||||
logger.Info("Registered restheadspec bunrouter routes for model %s at %s", fullName, entityPath)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -77,6 +77,17 @@ func (s *securityContext) GetUserID() (int, bool) {
|
||||
return security.GetUserID(s.ctx.Context)
|
||||
}
|
||||
|
||||
// GetUserRef returns an opaque user identifier for row security lookups.
|
||||
// It prefers the full *security.UserContext (so providers can read JWT claims,
|
||||
// e.g. a UUID subject) and falls back to the int user ID.
|
||||
func (s *securityContext) GetUserRef() (any, bool) {
|
||||
if userCtx, ok := security.GetUserContext(s.ctx.Context); ok {
|
||||
return userCtx, true
|
||||
}
|
||||
userID, ok := security.GetUserID(s.ctx.Context)
|
||||
return userID, ok
|
||||
}
|
||||
|
||||
func (s *securityContext) GetSchema() string {
|
||||
return s.ctx.Schema
|
||||
}
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
package restheadspec
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
)
|
||||
|
||||
func TestBuildFilterCondition_Spatial(t *testing.T) {
|
||||
h := &Handler{}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
filter common.FilterOption
|
||||
wantCond string
|
||||
wantCount int
|
||||
}{
|
||||
{
|
||||
name: "st_dwithin",
|
||||
filter: common.FilterOption{
|
||||
Column: "geom",
|
||||
Operator: "st_dwithin",
|
||||
Value: map[string]interface{}{
|
||||
"geom": "SRID=4326;POINT(0 0)",
|
||||
"distance": 1000.0,
|
||||
},
|
||||
},
|
||||
wantCond: "ST_DWithin(geom, ST_GeomFromEWKT(?), ?)",
|
||||
wantCount: 2,
|
||||
},
|
||||
{
|
||||
name: "st_intersects",
|
||||
filter: common.FilterOption{
|
||||
Column: "geom",
|
||||
Operator: "st_intersects",
|
||||
Value: "SRID=4326;POINT(0 0)",
|
||||
},
|
||||
wantCond: "ST_Intersects(geom, ST_GeomFromEWKT(?))",
|
||||
wantCount: 1,
|
||||
},
|
||||
{
|
||||
name: "l2_within",
|
||||
filter: common.FilterOption{
|
||||
Column: "embedding",
|
||||
Operator: "l2_within",
|
||||
Value: map[string]interface{}{
|
||||
"vector": []interface{}{1.0, 2.0},
|
||||
"distance": 0.3,
|
||||
},
|
||||
},
|
||||
wantCond: "embedding <-> ? < ?",
|
||||
wantCount: 2,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
f := tt.filter
|
||||
cond, args := h.buildFilterCondition(f.Column, &f, "")
|
||||
if cond != tt.wantCond {
|
||||
t.Errorf("cond = %q, want %q", cond, tt.wantCond)
|
||||
}
|
||||
if len(args) != tt.wantCount {
|
||||
t.Errorf("args = %d, want %d", len(args), tt.wantCount)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseGeoFilter(t *testing.T) {
|
||||
h := &Handler{}
|
||||
options := &ExtendedRequestOptions{}
|
||||
|
||||
h.parseGeoFilter(options, "x-spatialfilter-geom", "x-spatialfilter-",
|
||||
`{"op":"st_dwithin","geom":"SRID=4326;POINT(0 0)","distance":500}`)
|
||||
|
||||
if len(options.Filters) != 1 {
|
||||
t.Fatalf("expected 1 filter, got %d", len(options.Filters))
|
||||
}
|
||||
f := options.Filters[0]
|
||||
if f.Column != "geom" || f.Operator != "st_dwithin" {
|
||||
t.Errorf("filter = %+v", f)
|
||||
}
|
||||
m, ok := f.Value.(map[string]interface{})
|
||||
if !ok || m["distance"] != float64(500) {
|
||||
t.Errorf("value = %v", f.Value)
|
||||
}
|
||||
if _, has := m["op"]; has {
|
||||
t.Error("op should be stripped from value map")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseGeoFilter_ExplicitValue(t *testing.T) {
|
||||
h := &Handler{}
|
||||
options := &ExtendedRequestOptions{}
|
||||
|
||||
h.parseGeoFilter(options, "x-vectorfilter-embedding", "x-vectorfilter-",
|
||||
`{"op":"cosine_within","logic":"or","value":{"vector":[1,2,3],"distance":0.2}}`)
|
||||
|
||||
if len(options.Filters) != 1 {
|
||||
t.Fatalf("expected 1 filter, got %d", len(options.Filters))
|
||||
}
|
||||
f := options.Filters[0]
|
||||
if f.Operator != "cosine_within" || f.LogicOperator != "OR" {
|
||||
t.Errorf("filter = %+v", f)
|
||||
}
|
||||
if _, ok := f.Value.(map[string]interface{}); !ok {
|
||||
t.Errorf("value type = %T", f.Value)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFloat32List(t *testing.T) {
|
||||
got := parseFloat32List("[1,2.5,3]")
|
||||
if len(got) != 3 || got[1] != 2.5 {
|
||||
t.Errorf("json array = %v", got)
|
||||
}
|
||||
got = parseFloat32List("1, 2, 3")
|
||||
if len(got) != 3 || got[2] != 3 {
|
||||
t.Errorf("csv = %v", got)
|
||||
}
|
||||
if parseFloat32List("") != nil {
|
||||
t.Error("empty should be nil")
|
||||
}
|
||||
}
|
||||
@@ -74,8 +74,8 @@ func (c *CompositeSecurityProvider) GetColumnSecurity(ctx context.Context, userI
|
||||
}
|
||||
|
||||
// GetRowSecurity delegates to the row security provider
|
||||
func (c *CompositeSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
|
||||
return c.rowSec.GetRowSecurity(ctx, userID, schema, table)
|
||||
func (c *CompositeSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
|
||||
return c.rowSec.GetRowSecurity(ctx, userRef, schema, table)
|
||||
}
|
||||
|
||||
// Optional interface implementations (if wrapped providers support them)
|
||||
|
||||
@@ -79,7 +79,7 @@ type mockRowSec struct {
|
||||
supportsCache bool
|
||||
}
|
||||
|
||||
func (m *mockRowSec) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
|
||||
func (m *mockRowSec) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
|
||||
return m.rowSec, m.err
|
||||
}
|
||||
|
||||
|
||||
@@ -1597,6 +1597,8 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||
client_name VARCHAR(255),
|
||||
grant_types TEXT[] DEFAULT ARRAY['authorization_code'],
|
||||
allowed_scopes TEXT[] DEFAULT ARRAY['openid','profile','email'],
|
||||
client_secret_hash TEXT, -- sha256 hex of the confidential-client secret; NULL for public clients
|
||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||
is_active BOOLEAN DEFAULT true,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
@@ -1634,13 +1636,15 @@ DECLARE
|
||||
BEGIN
|
||||
v_client_id := p_data->>'client_id';
|
||||
|
||||
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes)
|
||||
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method)
|
||||
VALUES (
|
||||
v_client_id,
|
||||
ARRAY(SELECT jsonb_array_elements_text(p_data->'redirect_uris')),
|
||||
p_data->>'client_name',
|
||||
COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'grant_types')), ARRAY['authorization_code']),
|
||||
COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'allowed_scopes')), ARRAY['openid','profile','email'])
|
||||
COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'allowed_scopes')), ARRAY['openid','profile','email']),
|
||||
NULLIF(p_data->>'client_secret_hash', ''),
|
||||
COALESCE(NULLIF(p_data->>'token_endpoint_auth_method', ''), 'none')
|
||||
)
|
||||
RETURNING to_jsonb(oauth_clients.*) INTO v_row;
|
||||
|
||||
|
||||
@@ -100,6 +100,8 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||
client_name VARCHAR(255),
|
||||
grant_types TEXT, -- JSON-encoded []string
|
||||
allowed_scopes TEXT, -- JSON-encoded []string
|
||||
client_secret_hash TEXT, -- sha256 hex of the confidential-client secret; NULL for public clients
|
||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||
is_active BOOLEAN DEFAULT 1,
|
||||
created_at TIMESTAMP
|
||||
);
|
||||
|
||||
+82
-26
@@ -14,6 +14,11 @@ import (
|
||||
type SecurityContext interface {
|
||||
GetContext() context.Context
|
||||
GetUserID() (int, bool)
|
||||
// GetUserRef returns an opaque user identifier for row security lookups.
|
||||
// Unlike GetUserID, it is not required to be an integer: implementations backed by
|
||||
// non-integer identifiers (e.g. UUIDs) can return a string, or the full
|
||||
// *security.UserContext so a RowSecurityProvider can read JWT claims directly.
|
||||
GetUserRef() (any, bool)
|
||||
GetSchema() string
|
||||
GetEntity() string
|
||||
GetModel() interface{}
|
||||
@@ -45,8 +50,13 @@ func loadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error
|
||||
// return err
|
||||
}
|
||||
|
||||
// Load row security rules using the provider
|
||||
_, err = securityList.LoadRowSecurity(secCtx.GetContext(), userID, schema, tablename, false)
|
||||
// Load row security rules using the provider. Row security uses the opaque
|
||||
// user ref (not the int-only user ID) so non-integer user identifiers work.
|
||||
userRef, refOK := secCtx.GetUserRef()
|
||||
if !refOK {
|
||||
userRef = userID
|
||||
}
|
||||
_, err = securityList.LoadRowSecurity(secCtx.GetContext(), userRef, schema, tablename, false)
|
||||
if err != nil {
|
||||
logger.Warn("Failed to load row security: %v", err)
|
||||
// Don't fail the request if no security rules exist
|
||||
@@ -58,25 +68,29 @@ func loadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error
|
||||
|
||||
// applyRowSecurity applies row-level security filters to the query (generic version)
|
||||
func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
||||
userID, ok := secCtx.GetUserID()
|
||||
userRef, ok := secCtx.GetUserRef()
|
||||
if !ok {
|
||||
return nil // No user context, skip
|
||||
userID, idOK := secCtx.GetUserID()
|
||||
if !idOK {
|
||||
return nil // No user context, skip
|
||||
}
|
||||
userRef = userID
|
||||
}
|
||||
|
||||
schema := secCtx.GetSchema()
|
||||
tablename := secCtx.GetEntity()
|
||||
|
||||
// Get row security template
|
||||
rowSec, err := securityList.GetRowSecurityTemplate(userID, schema, tablename)
|
||||
rowSec, err := securityList.GetRowSecurityTemplate(userRef, schema, tablename)
|
||||
if err != nil {
|
||||
// No row security defined, allow query to proceed
|
||||
logger.Debug("No row security for %s.%s@%d: %v", schema, tablename, userID, err)
|
||||
logger.Debug("No row security for %s.%s@%v: %v", schema, tablename, userRef, err)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Check if user has a blocking rule
|
||||
if rowSec.HasBlock {
|
||||
logger.Warn("User %d blocked from accessing %s.%s", userID, schema, tablename)
|
||||
logger.Warn("User %v blocked from accessing %s.%s", userRef, schema, tablename)
|
||||
return fmt.Errorf("access denied to %s", tablename)
|
||||
}
|
||||
|
||||
@@ -112,8 +126,8 @@ func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error
|
||||
// Generate the WHERE clause from template
|
||||
whereClause := rowSec.GetTemplate(pkName, modelType)
|
||||
|
||||
logger.Info("Applying row security filter for user %d on %s.%s: %s",
|
||||
userID, schema, tablename, whereClause)
|
||||
logger.Info("Applying row security filter for user %v on %s.%s: %s",
|
||||
userRef, schema, tablename, whereClause)
|
||||
|
||||
// Apply the WHERE clause to the query
|
||||
query := secCtx.GetQuery()
|
||||
@@ -218,9 +232,37 @@ func LoadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error
|
||||
// ApplyRowSecurity is a public wrapper for applyRowSecurity that accepts a SecurityContext
|
||||
// This allows other packages to apply row-level security using the generic interface
|
||||
func ApplyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
||||
// Spec adapters that expose the dispatched operation can enforce the same
|
||||
// model-rule bypass even when ApplyRowSecurity is called directly.
|
||||
if operationCtx, ok := secCtx.(interface{ GetOperation() string }); ok &&
|
||||
ShouldSkipRowSecurity(secCtx, operationCtx.GetOperation()) {
|
||||
return nil
|
||||
}
|
||||
return applyRowSecurity(secCtx, securityList)
|
||||
}
|
||||
|
||||
// ShouldSkipRowSecurity reports whether row-security enforcement should be
|
||||
// skipped for the operation. It uses the same model-rule resolution as
|
||||
// CheckModelAuthAllowed so the model registry remains the single source of
|
||||
// truth for security behavior.
|
||||
func ShouldSkipRowSecurity(secCtx SecurityContext, operation string) bool {
|
||||
rules, ok := resolveModelRules(secCtx)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
|
||||
return rules.SecurityDisabled || (operation == "read" && rules.CanPublicRead)
|
||||
}
|
||||
|
||||
// IsModelSecurityDisabled reports whether all model-level security processing
|
||||
// is disabled for the model. This is distinct from ShouldSkipRowSecurity:
|
||||
// CanPublicRead skips row filtering for reads but must still allow other read
|
||||
// security, such as column masking, to be loaded.
|
||||
func IsModelSecurityDisabled(secCtx SecurityContext) bool {
|
||||
rules, ok := resolveModelRules(secCtx)
|
||||
return ok && rules.SecurityDisabled
|
||||
}
|
||||
|
||||
// ApplyColumnSecurity is a public wrapper for applyColumnSecurity that accepts a SecurityContext
|
||||
// This allows other packages to apply column-level security using the generic interface
|
||||
func ApplyColumnSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
||||
@@ -289,25 +331,14 @@ func checkModelDeleteAllowed(secCtx SecurityContext) error {
|
||||
// 7. Guest (UserID == 0) → return "authentication required".
|
||||
// 8. Authenticated user → allow (operation-specific checks remain in BeforeUpdate/BeforeDelete).
|
||||
func CheckModelAuthAllowed(secCtx SecurityContext, operation string) error {
|
||||
rules, ok := GetModelRulesFromContext(secCtx.GetContext())
|
||||
rules, ok := resolveModelRules(secCtx)
|
||||
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 {
|
||||
// Model not registered - fall through to auth check
|
||||
userID, _ := secCtx.GetUserID()
|
||||
if userID == 0 {
|
||||
return fmt.Errorf("authentication required")
|
||||
}
|
||||
return nil
|
||||
// Model not registered - fall through to auth check
|
||||
userID, _ := secCtx.GetUserID()
|
||||
if userID == 0 {
|
||||
return fmt.Errorf("authentication required")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
if rules.SecurityDisabled {
|
||||
@@ -333,6 +364,31 @@ func CheckModelAuthAllowed(secCtx SecurityContext, operation string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// resolveModelRules returns model rules from the request context first, then
|
||||
// falls back to the schema-qualified and unqualified registry names.
|
||||
func resolveModelRules(secCtx SecurityContext) (modelregistry.ModelRules, bool) {
|
||||
if rules, ok := GetModelRulesFromContext(secCtx.GetContext()); ok {
|
||||
return rules, true
|
||||
}
|
||||
|
||||
schema := secCtx.GetSchema()
|
||||
entity := secCtx.GetEntity()
|
||||
var err error
|
||||
if schema != "" {
|
||||
var rules modelregistry.ModelRules
|
||||
rules, err = modelregistry.GetModelRulesByName(fmt.Sprintf("%s.%s", schema, entity))
|
||||
if err == nil {
|
||||
return rules, true
|
||||
}
|
||||
}
|
||||
|
||||
rules, err := modelregistry.GetModelRulesByName(entity)
|
||||
if err != nil {
|
||||
return modelregistry.ModelRules{}, false
|
||||
}
|
||||
return rules, true
|
||||
}
|
||||
|
||||
// CheckModelUpdateAllowed is the public wrapper for checkModelUpdateAllowed.
|
||||
func CheckModelUpdateAllowed(secCtx SecurityContext) error {
|
||||
return checkModelUpdateAllowed(secCtx)
|
||||
|
||||
@@ -26,6 +26,10 @@ func (m *mockSecurityContext) GetUserID() (int, bool) {
|
||||
return m.userID, m.hasUser
|
||||
}
|
||||
|
||||
func (m *mockSecurityContext) GetUserRef() (any, bool) {
|
||||
return m.userID, m.hasUser
|
||||
}
|
||||
|
||||
func (m *mockSecurityContext) GetSchema() string {
|
||||
return m.schema
|
||||
}
|
||||
|
||||
@@ -121,8 +121,12 @@ type ColumnSecurityProvider interface {
|
||||
|
||||
// RowSecurityProvider handles row-level security (filtering)
|
||||
type RowSecurityProvider interface {
|
||||
// GetRowSecurity loads row security rules for a user and entity
|
||||
GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error)
|
||||
// GetRowSecurity loads row security rules for a user and entity.
|
||||
// userRef identifies the user and is opaque to the caller: it may be an int ID,
|
||||
// a string/UUID, or the full *security.UserContext (see SecurityContext.GetUserRef),
|
||||
// so providers backed by non-integer user identifiers (e.g. UUIDs) or that need
|
||||
// access to JWT claims can implement row security without relying on a numeric ID.
|
||||
GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error)
|
||||
}
|
||||
|
||||
// SecurityProvider is the main interface combining all security concerns
|
||||
|
||||
+424
-35
@@ -3,8 +3,11 @@ package security
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
@@ -12,6 +15,9 @@ import (
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// OAuthServerConfig configures the MCP-standard OAuth2 authorization server.
|
||||
@@ -44,15 +50,32 @@ type OAuthServerConfig struct {
|
||||
|
||||
// AuthCodeTTL is the auth code lifetime. Defaults to 2 minutes.
|
||||
AuthCodeTTL time.Duration
|
||||
|
||||
// ResourceIdentifier is this server's protected-resource identifier, advertised in
|
||||
// RFC 9728 metadata. Defaults to Issuer.
|
||||
ResourceIdentifier string
|
||||
|
||||
// SigningKey signs id_tokens (RS256) and is exposed via the JWKS endpoint. If nil, an
|
||||
// RSA-2048 key is generated in memory when the server starts. Supply a persistent key
|
||||
// for multi-instance deployments so id_tokens remain verifiable across restarts/instances.
|
||||
SigningKey *rsa.PrivateKey
|
||||
}
|
||||
|
||||
// oauthClient is a dynamically registered OAuth2 client (RFC 7591).
|
||||
type oauthClient struct {
|
||||
ClientID string `json:"client_id"`
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
ClientName string `json:"client_name,omitempty"`
|
||||
GrantTypes []string `json:"grant_types"`
|
||||
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
||||
ClientID string `json:"client_id"`
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
ClientName string `json:"client_name,omitempty"`
|
||||
GrantTypes []string `json:"grant_types"`
|
||||
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
||||
ClientSecretHash string `json:"client_secret_hash,omitempty"`
|
||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
|
||||
}
|
||||
|
||||
// isConfidential reports whether the client has a registered secret and must
|
||||
// authenticate itself at the token endpoint.
|
||||
func (c *oauthClient) isConfidential() bool {
|
||||
return c.ClientSecretHash != ""
|
||||
}
|
||||
|
||||
// pendingAuth tracks an in-progress authorization code exchange.
|
||||
@@ -85,13 +108,25 @@ type externalProvider struct {
|
||||
// The server exposes these RFC-compliant endpoints:
|
||||
//
|
||||
// GET /.well-known/oauth-authorization-server RFC 8414 — server metadata discovery
|
||||
// GET /.well-known/openid-configuration OIDC discovery (superset of the above)
|
||||
// GET /.well-known/oauth-protected-resource RFC 9728 — protected resource metadata
|
||||
// POST /oauth/register RFC 7591 — dynamic client registration
|
||||
// GET /oauth/authorize OAuth 2.1 + PKCE — start authorization
|
||||
// POST /oauth/authorize Direct login form submission
|
||||
// POST /oauth/token Token exchange and refresh
|
||||
// POST /oauth/token Token exchange: authorization_code,
|
||||
// refresh_token, client_credentials (RFC 6749 §4.4)
|
||||
// POST /oauth/revoke RFC 7009 — token revocation
|
||||
// POST /oauth/introspect RFC 7662 — token introspection
|
||||
// GET /oauth/userinfo OIDC UserInfo endpoint
|
||||
// GET /oauth/jwks.json JWKS — id_token verification keys
|
||||
// GET {ProviderCallbackPath} Internal — external provider callback
|
||||
//
|
||||
// Confidential clients (registered with token_endpoint_auth_method other than "none", or
|
||||
// any grant_types including client_credentials) authenticate at /oauth/token via
|
||||
// client_secret_basic or client_secret_post. Public clients keep relying on PKCE alone.
|
||||
//
|
||||
// When the granted scope includes "openid", authorization_code and refresh_token responses
|
||||
// include an RS256-signed id_token (see OAuthServerConfig.SigningKey).
|
||||
type OAuthServer struct {
|
||||
cfg OAuthServerConfig
|
||||
auth *DatabaseAuthenticator // nil = only external providers
|
||||
@@ -102,6 +137,9 @@ type OAuthServer struct {
|
||||
pending map[string]*pendingAuth // provider_state → pending (external flow)
|
||||
codes map[string]*pendingAuth // auth_code → pending (post-auth)
|
||||
|
||||
signingKey *rsa.PrivateKey
|
||||
signingKeyID string
|
||||
|
||||
done chan struct{} // closed by Close() to stop background goroutines
|
||||
}
|
||||
|
||||
@@ -130,13 +168,32 @@ func NewOAuthServer(cfg OAuthServerConfig, auth *DatabaseAuthenticator) *OAuthSe
|
||||
}
|
||||
// Normalize issuer: remove trailing slash to ensure consistent endpoint URL construction.
|
||||
cfg.Issuer = strings.TrimSuffix(cfg.Issuer, "/")
|
||||
if cfg.ResourceIdentifier == "" {
|
||||
cfg.ResourceIdentifier = cfg.Issuer
|
||||
}
|
||||
|
||||
signingKey := cfg.SigningKey
|
||||
if signingKey == nil {
|
||||
var err error
|
||||
signingKey, err = rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
// Signing keys are only required for id_token issuance (OIDC "openid" scope);
|
||||
// leaving signingKey nil degrades gracefully by omitting id_token/JWKS support.
|
||||
signingKey = nil
|
||||
}
|
||||
}
|
||||
|
||||
s := &OAuthServer{
|
||||
cfg: cfg,
|
||||
auth: auth,
|
||||
clients: make(map[string]*oauthClient),
|
||||
pending: make(map[string]*pendingAuth),
|
||||
codes: make(map[string]*pendingAuth),
|
||||
done: make(chan struct{}),
|
||||
cfg: cfg,
|
||||
auth: auth,
|
||||
clients: make(map[string]*oauthClient),
|
||||
pending: make(map[string]*pendingAuth),
|
||||
codes: make(map[string]*pendingAuth),
|
||||
signingKey: signingKey,
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
if signingKey != nil {
|
||||
s.signingKeyID = rsaKeyID(&signingKey.PublicKey)
|
||||
}
|
||||
go s.cleanupExpired()
|
||||
return s
|
||||
@@ -178,11 +235,15 @@ func (s *OAuthServer) ProviderCallbackPath() string {
|
||||
func (s *OAuthServer) HTTPHandler() http.Handler {
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/.well-known/oauth-authorization-server", s.metadataHandler)
|
||||
mux.HandleFunc("/.well-known/openid-configuration", s.openIDConfigurationHandler)
|
||||
mux.HandleFunc("/.well-known/oauth-protected-resource", s.protectedResourceHandler)
|
||||
mux.HandleFunc("/oauth/register", s.registerHandler)
|
||||
mux.HandleFunc("/oauth/authorize", s.authorizeHandler)
|
||||
mux.HandleFunc("/oauth/token", s.tokenHandler)
|
||||
mux.HandleFunc("/oauth/revoke", s.revokeHandler)
|
||||
mux.HandleFunc("/oauth/introspect", s.introspectHandler)
|
||||
mux.HandleFunc("/oauth/userinfo", s.userinfoHandler)
|
||||
mux.HandleFunc("/oauth/jwks.json", s.jwksHandler)
|
||||
mux.HandleFunc(s.cfg.ProviderCallbackPath, s.providerCallbackHandler)
|
||||
return mux
|
||||
}
|
||||
@@ -217,25 +278,127 @@ func (s *OAuthServer) cleanupExpired() {
|
||||
// RFC 8414 — Server metadata
|
||||
// --------------------------------------------------------------------------
|
||||
|
||||
func (s *OAuthServer) metadataHandler(w http.ResponseWriter, r *http.Request) {
|
||||
// serverMetadata builds the fields shared by RFC 8414 authorization-server
|
||||
// metadata and OIDC discovery metadata.
|
||||
func (s *OAuthServer) serverMetadata() map[string]interface{} {
|
||||
issuer := s.cfg.Issuer
|
||||
meta := map[string]interface{}{
|
||||
grantTypes := []string{"authorization_code", "refresh_token"}
|
||||
if s.auth != nil {
|
||||
grantTypes = append(grantTypes, "client_credentials")
|
||||
}
|
||||
return map[string]interface{}{
|
||||
"issuer": issuer,
|
||||
"authorization_endpoint": issuer + "/oauth/authorize",
|
||||
"token_endpoint": issuer + "/oauth/token",
|
||||
"registration_endpoint": issuer + "/oauth/register",
|
||||
"revocation_endpoint": issuer + "/oauth/revoke",
|
||||
"introspection_endpoint": issuer + "/oauth/introspect",
|
||||
"userinfo_endpoint": issuer + "/oauth/userinfo",
|
||||
"jwks_uri": issuer + "/oauth/jwks.json",
|
||||
"scopes_supported": s.cfg.DefaultScopes,
|
||||
"response_types_supported": []string{"code"},
|
||||
"grant_types_supported": []string{"authorization_code", "refresh_token"},
|
||||
"grant_types_supported": grantTypes,
|
||||
"code_challenge_methods_supported": []string{"S256"},
|
||||
"token_endpoint_auth_methods_supported": []string{"none"},
|
||||
"token_endpoint_auth_methods_supported": []string{"none", "client_secret_basic", "client_secret_post"},
|
||||
}
|
||||
}
|
||||
|
||||
func (s *OAuthServer) metadataHandler(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(s.serverMetadata()) //nolint:errcheck
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
// OIDC discovery — GET /.well-known/openid-configuration
|
||||
// --------------------------------------------------------------------------
|
||||
|
||||
func (s *OAuthServer) openIDConfigurationHandler(w http.ResponseWriter, r *http.Request) {
|
||||
meta := s.serverMetadata()
|
||||
meta["subject_types_supported"] = []string{"public"}
|
||||
meta["id_token_signing_alg_values_supported"] = []string{"RS256"}
|
||||
meta["claims_supported"] = []string{"sub", "iss", "aud", "exp", "iat", "email", "preferred_username"}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(meta) //nolint:errcheck
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
// RFC 9728 — Protected Resource Metadata
|
||||
// --------------------------------------------------------------------------
|
||||
|
||||
func (s *OAuthServer) protectedResourceHandler(w http.ResponseWriter, r *http.Request) {
|
||||
meta := map[string]interface{}{
|
||||
"resource": s.cfg.ResourceIdentifier,
|
||||
"authorization_servers": []string{s.cfg.Issuer},
|
||||
"scopes_supported": s.cfg.DefaultScopes,
|
||||
"bearer_methods_supported": []string{"header"},
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(meta) //nolint:errcheck
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
// JWKS — GET /oauth/jwks.json
|
||||
// --------------------------------------------------------------------------
|
||||
|
||||
func (s *OAuthServer) jwksHandler(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if s.signingKey == nil {
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{"keys": []interface{}{}}) //nolint:errcheck
|
||||
return
|
||||
}
|
||||
pub := s.signingKey.PublicKey
|
||||
jwk := map[string]interface{}{
|
||||
"kty": "RSA",
|
||||
"use": "sig",
|
||||
"alg": "RS256",
|
||||
"kid": s.signingKeyID,
|
||||
"n": base64.RawURLEncoding.EncodeToString(pub.N.Bytes()),
|
||||
"e": base64.RawURLEncoding.EncodeToString(bigEndianBytes(pub.E)),
|
||||
}
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{"keys": []interface{}{jwk}}) //nolint:errcheck
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
// Userinfo — GET/POST /oauth/userinfo
|
||||
// --------------------------------------------------------------------------
|
||||
|
||||
func (s *OAuthServer) userinfoHandler(w http.ResponseWriter, r *http.Request) {
|
||||
auth := r.Header.Get("Authorization")
|
||||
token := strings.TrimPrefix(auth, "Bearer ")
|
||||
if token == "" || token == auth {
|
||||
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`)
|
||||
writeOAuthError(w, "invalid_token", "missing bearer token", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
authToUse := s.auth
|
||||
if authToUse == nil {
|
||||
s.mu.RLock()
|
||||
if len(s.providers) > 0 {
|
||||
authToUse = s.providers[0].auth
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
}
|
||||
if authToUse == nil {
|
||||
writeOAuthError(w, "invalid_token", "no authenticator configured", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
info, err := authToUse.OAuthIntrospectToken(r.Context(), token)
|
||||
if err != nil || !info.Active {
|
||||
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`)
|
||||
writeOAuthError(w, "invalid_token", "token is inactive or invalid", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{ //nolint:errcheck
|
||||
"sub": info.Sub,
|
||||
"preferred_username": info.Username,
|
||||
"email": info.Email,
|
||||
})
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
// RFC 7591 — Dynamic client registration
|
||||
// --------------------------------------------------------------------------
|
||||
@@ -246,10 +409,11 @@ func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
var req struct {
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
ClientName string `json:"client_name"`
|
||||
GrantTypes []string `json:"grant_types"`
|
||||
AllowedScopes []string `json:"allowed_scopes"`
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
ClientName string `json:"client_name"`
|
||||
GrantTypes []string `json:"grant_types"`
|
||||
AllowedScopes []string `json:"allowed_scopes"`
|
||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeOAuthError(w, "invalid_request", "malformed JSON", http.StatusBadRequest)
|
||||
@@ -272,21 +436,48 @@ func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
|
||||
http.Error(w, "server error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// client_credentials is a machine-to-machine grant and requires a confidential
|
||||
// client (RFC 6749 §4.4), so it always forces secret issuance regardless of the
|
||||
// requested auth method.
|
||||
authMethod := req.TokenEndpointAuthMethod
|
||||
if authMethod == "" {
|
||||
authMethod = "none"
|
||||
}
|
||||
if oauthSliceContains(grantTypes, "client_credentials") && authMethod == "none" {
|
||||
authMethod = "client_secret_basic"
|
||||
}
|
||||
|
||||
var plaintextSecret string
|
||||
var secretHash string
|
||||
if authMethod != "none" {
|
||||
plaintextSecret, err = randomOAuthToken()
|
||||
if err != nil {
|
||||
http.Error(w, "server error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
secretHash = hashClientSecret(plaintextSecret)
|
||||
}
|
||||
|
||||
client := &oauthClient{
|
||||
ClientID: clientID,
|
||||
RedirectURIs: req.RedirectURIs,
|
||||
ClientName: req.ClientName,
|
||||
GrantTypes: grantTypes,
|
||||
AllowedScopes: allowedScopes,
|
||||
ClientID: clientID,
|
||||
RedirectURIs: req.RedirectURIs,
|
||||
ClientName: req.ClientName,
|
||||
GrantTypes: grantTypes,
|
||||
AllowedScopes: allowedScopes,
|
||||
ClientSecretHash: secretHash,
|
||||
TokenEndpointAuthMethod: authMethod,
|
||||
}
|
||||
|
||||
if s.cfg.PersistClients && s.auth != nil {
|
||||
dbClient := &OAuthServerClient{
|
||||
ClientID: client.ClientID,
|
||||
RedirectURIs: client.RedirectURIs,
|
||||
ClientName: client.ClientName,
|
||||
GrantTypes: client.GrantTypes,
|
||||
AllowedScopes: client.AllowedScopes,
|
||||
ClientID: client.ClientID,
|
||||
RedirectURIs: client.RedirectURIs,
|
||||
ClientName: client.ClientName,
|
||||
GrantTypes: client.GrantTypes,
|
||||
AllowedScopes: client.AllowedScopes,
|
||||
ClientSecretHash: client.ClientSecretHash,
|
||||
TokenEndpointAuthMethod: client.TokenEndpointAuthMethod,
|
||||
}
|
||||
if _, err := s.auth.OAuthRegisterClient(r.Context(), dbClient); err != nil {
|
||||
http.Error(w, "server error", http.StatusInternalServerError)
|
||||
@@ -298,9 +489,25 @@ func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
|
||||
s.clients[clientID] = client
|
||||
s.mu.Unlock()
|
||||
|
||||
// RFC 7591 registration response: the plaintext secret is returned exactly once here
|
||||
// and never persisted or served again — only its hash (client.ClientSecretHash) is
|
||||
// stored, and that hash is deliberately excluded from this response.
|
||||
resp := map[string]interface{}{
|
||||
"client_id": client.ClientID,
|
||||
"redirect_uris": client.RedirectURIs,
|
||||
"client_name": client.ClientName,
|
||||
"grant_types": client.GrantTypes,
|
||||
"allowed_scopes": client.AllowedScopes,
|
||||
"token_endpoint_auth_method": client.TokenEndpointAuthMethod,
|
||||
}
|
||||
if plaintextSecret != "" {
|
||||
resp["client_secret"] = plaintextSecret
|
||||
resp["client_secret_expires_at"] = 0
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
json.NewEncoder(w).Encode(client) //nolint:errcheck
|
||||
json.NewEncoder(w).Encode(resp) //nolint:errcheck
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
@@ -577,6 +784,8 @@ func (s *OAuthServer) tokenHandler(w http.ResponseWriter, r *http.Request) {
|
||||
s.handleAuthCodeGrant(w, r)
|
||||
case "refresh_token":
|
||||
s.handleRefreshGrant(w, r)
|
||||
case "client_credentials":
|
||||
s.handleClientCredentialsGrant(w, r)
|
||||
default:
|
||||
writeOAuthError(w, "unsupported_grant_type", "", http.StatusBadRequest)
|
||||
}
|
||||
@@ -593,6 +802,15 @@ func (s *OAuthServer) handleAuthCodeGrant(w http.ResponseWriter, r *http.Request
|
||||
return
|
||||
}
|
||||
|
||||
// Confidential clients (those registered with a client_secret) must authenticate;
|
||||
// public clients keep relying on PKCE alone, unchanged from prior behavior.
|
||||
if client, ok := s.lookupOrFetchClient(r.Context(), clientID); ok && client.isConfidential() {
|
||||
if _, err := s.authenticateClient(r); err != nil {
|
||||
writeOAuthError(w, "invalid_client", err.Error(), http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
var sessionToken string
|
||||
var refreshToken string
|
||||
var scopes []string
|
||||
@@ -647,12 +865,13 @@ func (s *OAuthServer) handleAuthCodeGrant(w http.ResponseWriter, r *http.Request
|
||||
scopes = pending.Scopes
|
||||
}
|
||||
|
||||
s.writeOAuthToken(w, sessionToken, refreshToken, scopes)
|
||||
s.writeOAuthToken(w, r, sessionToken, refreshToken, clientID, scopes, true)
|
||||
}
|
||||
|
||||
func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request) {
|
||||
refreshToken := r.FormValue("refresh_token")
|
||||
providerName := r.FormValue("provider")
|
||||
clientID := r.FormValue("client_id")
|
||||
if refreshToken == "" {
|
||||
writeOAuthError(w, "invalid_request", "refresh_token required", http.StatusBadRequest)
|
||||
return
|
||||
@@ -666,7 +885,7 @@ func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request)
|
||||
writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
s.writeOAuthToken(w, loginResp.Token, loginResp.RefreshToken, nil)
|
||||
s.writeOAuthToken(w, r, loginResp.Token, loginResp.RefreshToken, clientID, nil, false)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -676,13 +895,86 @@ func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request)
|
||||
writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
s.writeOAuthToken(w, loginResp.Token, loginResp.RefreshToken, nil)
|
||||
s.writeOAuthToken(w, r, loginResp.Token, loginResp.RefreshToken, clientID, nil, false)
|
||||
return
|
||||
}
|
||||
|
||||
writeOAuthError(w, "invalid_grant", "no provider available for refresh", http.StatusBadRequest)
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
// RFC 6749 §4.4 — Client credentials grant
|
||||
// --------------------------------------------------------------------------
|
||||
|
||||
func (s *OAuthServer) handleClientCredentialsGrant(w http.ResponseWriter, r *http.Request) {
|
||||
if s.auth == nil {
|
||||
writeOAuthError(w, "unsupported_grant_type", "client_credentials requires a local user store", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
client, err := s.authenticateClient(r)
|
||||
if err != nil {
|
||||
w.Header().Set("WWW-Authenticate", `Basic realm="oauth"`)
|
||||
writeOAuthError(w, "invalid_client", err.Error(), http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
if !oauthSliceContains(client.GrantTypes, "client_credentials") {
|
||||
writeOAuthError(w, "unauthorized_client", "client is not authorized for client_credentials", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
requested := strings.Fields(r.FormValue("scope"))
|
||||
effectiveScopes := client.AllowedScopes
|
||||
if len(requested) > 0 {
|
||||
effectiveScopes = nil
|
||||
for _, sc := range requested {
|
||||
if oauthSliceContains(client.AllowedScopes, sc) {
|
||||
effectiveScopes = append(effectiveScopes, sc)
|
||||
}
|
||||
}
|
||||
if len(effectiveScopes) == 0 {
|
||||
writeOAuthError(w, "invalid_scope", "no requested scope is allowed for this client", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// client_credentials tokens have no end user, but the rest of the stack (RLS-scoping
|
||||
// hooks, introspection) expects every access token to resolve to a user_sessions row
|
||||
// with a user_id. Represent the client as a deterministic synthetic "service account"
|
||||
// user so the existing get-or-create/create-session/introspection pipeline handles it
|
||||
// unchanged — no new tables or code paths required.
|
||||
userCtx := &UserContext{
|
||||
UserName: "client:" + client.ClientID,
|
||||
Email: "oauth-client-" + client.ClientID + "@service.internal",
|
||||
RemoteID: client.ClientID,
|
||||
Roles: effectiveScopes,
|
||||
}
|
||||
userID, err := s.auth.oauth2GetOrCreateUser(r.Context(), userCtx, "oauth2_client")
|
||||
if err != nil {
|
||||
http.Error(w, "server error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
sessionToken, err := randomOAuthToken()
|
||||
if err != nil {
|
||||
http.Error(w, "server error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
expiresAt := time.Now().Add(s.cfg.AccessTokenTTL)
|
||||
err = s.auth.oauth2CreateSession(r.Context(), sessionToken, userID, &oauth2.Token{
|
||||
AccessToken: sessionToken,
|
||||
TokenType: "Bearer",
|
||||
}, expiresAt, "oauth2_client")
|
||||
if err != nil {
|
||||
http.Error(w, "server error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// No refresh token per RFC 6749 §4.4.3, and no id_token — client_credentials has no
|
||||
// end-user subject to represent in OIDC terms.
|
||||
s.writeOAuthToken(w, r, sessionToken, "", client.ClientID, effectiveScopes, false)
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
// RFC 7009 — Token revocation
|
||||
// --------------------------------------------------------------------------
|
||||
@@ -879,7 +1171,10 @@ func oauthSliceContains(slice []string, s string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *OAuthServer) writeOAuthToken(w http.ResponseWriter, accessToken, refreshToken string, scopes []string) {
|
||||
// writeOAuthToken writes the token response. When issueIDToken is true and the granted
|
||||
// scopes include "openid", an RS256-signed id_token is included (OIDC); client_credentials
|
||||
// responses always pass issueIDToken=false since that grant has no end-user subject.
|
||||
func (s *OAuthServer) writeOAuthToken(w http.ResponseWriter, r *http.Request, accessToken, refreshToken, clientID string, scopes []string, issueIDToken bool) {
|
||||
expiresIn := int64(s.cfg.AccessTokenTTL.Seconds())
|
||||
resp := map[string]interface{}{
|
||||
"access_token": accessToken,
|
||||
@@ -892,12 +1187,106 @@ func (s *OAuthServer) writeOAuthToken(w http.ResponseWriter, accessToken, refres
|
||||
if len(scopes) > 0 {
|
||||
resp["scope"] = strings.Join(scopes, " ")
|
||||
}
|
||||
if issueIDToken && oauthSliceContains(scopes, "openid") {
|
||||
if idToken, err := s.buildIDToken(r.Context(), accessToken, clientID, scopes); err == nil {
|
||||
resp["id_token"] = idToken
|
||||
}
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.Header().Set("Pragma", "no-cache")
|
||||
json.NewEncoder(w).Encode(resp) //nolint:errcheck
|
||||
}
|
||||
|
||||
// buildIDToken issues an OIDC id_token for the just-issued access token by reusing the
|
||||
// existing introspection pipeline to resolve the subject's claims.
|
||||
func (s *OAuthServer) buildIDToken(ctx context.Context, accessToken, clientID string, scopes []string) (string, error) {
|
||||
if s.signingKey == nil {
|
||||
return "", fmt.Errorf("no signing key configured")
|
||||
}
|
||||
authToUse := s.auth
|
||||
if authToUse == nil {
|
||||
s.mu.RLock()
|
||||
if len(s.providers) > 0 {
|
||||
authToUse = s.providers[0].auth
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
}
|
||||
if authToUse == nil {
|
||||
return "", fmt.Errorf("no authenticator configured")
|
||||
}
|
||||
info, err := authToUse.OAuthIntrospectToken(ctx, accessToken)
|
||||
if err != nil || !info.Active {
|
||||
return "", fmt.Errorf("token not active")
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
claims := jwt.MapClaims{
|
||||
"iss": s.cfg.Issuer,
|
||||
"sub": info.Sub,
|
||||
"aud": clientID,
|
||||
"exp": now.Add(s.cfg.AccessTokenTTL).Unix(),
|
||||
"iat": now.Unix(),
|
||||
}
|
||||
if oauthSliceContains(scopes, "profile") && info.Username != "" {
|
||||
claims["preferred_username"] = info.Username
|
||||
}
|
||||
if oauthSliceContains(scopes, "email") && info.Email != "" {
|
||||
claims["email"] = info.Email
|
||||
}
|
||||
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims)
|
||||
token.Header["kid"] = s.signingKeyID
|
||||
return token.SignedString(s.signingKey)
|
||||
}
|
||||
|
||||
// authenticateClient validates client_secret_basic (Authorization: Basic) or
|
||||
// client_secret_post (client_id/client_secret form fields) credentials against a
|
||||
// registered confidential client's stored secret hash.
|
||||
func (s *OAuthServer) authenticateClient(r *http.Request) (*oauthClient, error) {
|
||||
clientID, clientSecret, ok := r.BasicAuth()
|
||||
if !ok {
|
||||
clientID = r.FormValue("client_id")
|
||||
clientSecret = r.FormValue("client_secret")
|
||||
}
|
||||
if clientID == "" || clientSecret == "" {
|
||||
return nil, fmt.Errorf("client authentication required")
|
||||
}
|
||||
client, ok := s.lookupOrFetchClient(r.Context(), clientID)
|
||||
if !ok || !client.isConfidential() {
|
||||
return nil, fmt.Errorf("invalid client credentials")
|
||||
}
|
||||
if subtle.ConstantTimeCompare([]byte(hashClientSecret(clientSecret)), []byte(client.ClientSecretHash)) != 1 {
|
||||
return nil, fmt.Errorf("invalid client credentials")
|
||||
}
|
||||
return client, nil
|
||||
}
|
||||
|
||||
func hashClientSecret(secret string) string {
|
||||
sum := sha256.Sum256([]byte(secret))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// rsaKeyID derives a stable JWKS "kid" from an RSA public key's modulus.
|
||||
func rsaKeyID(pub *rsa.PublicKey) string {
|
||||
sum := sha256.Sum256(pub.N.Bytes())
|
||||
return base64.RawURLEncoding.EncodeToString(sum[:8])
|
||||
}
|
||||
|
||||
// bigEndianBytes encodes a small positive int (e.g. an RSA public exponent) as
|
||||
// minimal big-endian bytes for JWK "e" encoding.
|
||||
func bigEndianBytes(n int) []byte {
|
||||
if n == 0 {
|
||||
return []byte{0}
|
||||
}
|
||||
var b []byte
|
||||
for n > 0 {
|
||||
b = append([]byte{byte(n & 0xff)}, b...)
|
||||
n >>= 8
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func writeOAuthError(w http.ResponseWriter, errCode, description string, status int) {
|
||||
resp := map[string]string{"error": errCode}
|
||||
if description != "" {
|
||||
|
||||
@@ -9,11 +9,13 @@ import (
|
||||
|
||||
// OAuthServerClient is a persisted RFC 7591 registered OAuth2 client.
|
||||
type OAuthServerClient struct {
|
||||
ClientID string `json:"client_id"`
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
ClientName string `json:"client_name,omitempty"`
|
||||
GrantTypes []string `json:"grant_types"`
|
||||
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
||||
ClientID string `json:"client_id"`
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
ClientName string `json:"client_name,omitempty"`
|
||||
GrantTypes []string `json:"grant_types"`
|
||||
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
||||
ClientSecretHash string `json:"client_secret_hash,omitempty"`
|
||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
|
||||
}
|
||||
|
||||
// OAuthCode is a short-lived authorization code.
|
||||
|
||||
@@ -15,6 +15,15 @@ import (
|
||||
// Array columns (redirect_uris, grant_types, allowed_scopes, scopes) are
|
||||
// JSON-encoded TEXT instead of native Postgres arrays.
|
||||
|
||||
// nullIfEmpty converts an empty string to a SQL NULL so optional TEXT columns
|
||||
// (e.g. client_secret_hash for public clients) stay unset rather than "".
|
||||
func nullIfEmpty(s string) any {
|
||||
if s == "" {
|
||||
return nil
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func (a *DatabaseAuthenticator) oauthRegisterClientDirect(ctx context.Context, client *OAuthServerClient) (*OAuthServerClient, error) {
|
||||
grantTypes := client.GrantTypes
|
||||
if len(grantTypes) == 0 {
|
||||
@@ -38,11 +47,16 @@ func (a *DatabaseAuthenticator) oauthRegisterClientDirect(ctx context.Context, c
|
||||
return nil, fmt.Errorf("failed to marshal allowed_scopes: %w", err)
|
||||
}
|
||||
|
||||
authMethod := client.TokenEndpointAuthMethod
|
||||
if authMethod == "" {
|
||||
authMethod = "none"
|
||||
}
|
||||
|
||||
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
query := rewritePlaceholders(db, fmt.Sprintf(
|
||||
`INSERT INTO %s (client_id, redirect_uris, client_name, grant_types, allowed_scopes, is_active, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`,
|
||||
`INSERT INTO %s (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method, is_active, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
a.tableNames.OAuthClients))
|
||||
_, err := db.ExecContext(ctx, query, client.ClientID, string(redirectURIsJSON), client.ClientName, string(grantTypesJSON), string(allowedScopesJSON), true, time.Now())
|
||||
_, err := db.ExecContext(ctx, query, client.ClientID, string(redirectURIsJSON), client.ClientName, string(grantTypesJSON), string(allowedScopesJSON), nullIfEmpty(client.ClientSecretHash), authMethod, true, time.Now())
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
@@ -50,23 +64,25 @@ func (a *DatabaseAuthenticator) oauthRegisterClientDirect(ctx context.Context, c
|
||||
}
|
||||
|
||||
return &OAuthServerClient{
|
||||
ClientID: client.ClientID,
|
||||
RedirectURIs: client.RedirectURIs,
|
||||
ClientName: client.ClientName,
|
||||
GrantTypes: grantTypes,
|
||||
AllowedScopes: allowedScopes,
|
||||
ClientID: client.ClientID,
|
||||
RedirectURIs: client.RedirectURIs,
|
||||
ClientName: client.ClientName,
|
||||
GrantTypes: grantTypes,
|
||||
AllowedScopes: allowedScopes,
|
||||
ClientSecretHash: client.ClientSecretHash,
|
||||
TokenEndpointAuthMethod: authMethod,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (a *DatabaseAuthenticator) oauthGetClientDirect(ctx context.Context, clientID string) (*OAuthServerClient, error) {
|
||||
var redirectURIsJSON, grantTypesJSON, allowedScopesJSON sql.NullString
|
||||
var clientName sql.NullString
|
||||
var clientName, clientSecretHash, authMethod sql.NullString
|
||||
|
||||
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
query := rewritePlaceholders(db, fmt.Sprintf(
|
||||
`SELECT redirect_uris, client_name, grant_types, allowed_scopes FROM %s WHERE client_id = ? AND is_active = ?`,
|
||||
`SELECT redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method FROM %s WHERE client_id = ? AND is_active = ?`,
|
||||
a.tableNames.OAuthClients))
|
||||
return db.QueryRowContext(ctx, query, clientID, true).Scan(&redirectURIsJSON, &clientName, &grantTypesJSON, &allowedScopesJSON)
|
||||
return db.QueryRowContext(ctx, query, clientID, true).Scan(&redirectURIsJSON, &clientName, &grantTypesJSON, &allowedScopesJSON, &clientSecretHash, &authMethod)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
@@ -75,7 +91,12 @@ func (a *DatabaseAuthenticator) oauthGetClientDirect(ctx context.Context, client
|
||||
return nil, fmt.Errorf("failed to get client: %w", err)
|
||||
}
|
||||
|
||||
result := &OAuthServerClient{ClientID: clientID, ClientName: clientName.String}
|
||||
result := &OAuthServerClient{
|
||||
ClientID: clientID,
|
||||
ClientName: clientName.String,
|
||||
ClientSecretHash: clientSecretHash.String,
|
||||
TokenEndpointAuthMethod: authMethod.String,
|
||||
}
|
||||
if redirectURIsJSON.Valid {
|
||||
_ = json.Unmarshal([]byte(redirectURIsJSON.String), &result.RedirectURIs)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,483 @@
|
||||
package security
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rsa"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"math/big"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
func newTestOAuthServer(t *testing.T) (*OAuthServer, *DatabaseAuthenticator) {
|
||||
t.Helper()
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
||||
srv := NewOAuthServer(OAuthServerConfig{Issuer: "https://auth.example.com", PersistCodes: true}, auth)
|
||||
t.Cleanup(srv.Close)
|
||||
return srv, auth
|
||||
}
|
||||
|
||||
// s256Challenge computes the PKCE S256 code_challenge for a given verifier,
|
||||
// matching validatePKCESHA256 in oauth_server.go.
|
||||
func s256Challenge(verifier string) string {
|
||||
h := sha256.Sum256([]byte(verifier))
|
||||
return base64.RawURLEncoding.EncodeToString(h[:])
|
||||
}
|
||||
|
||||
func doJSON(t *testing.T, mux http.Handler, method, path string, body map[string]interface{}) (*httptest.ResponseRecorder, map[string]interface{}) {
|
||||
t.Helper()
|
||||
var reqBody *strings.Reader
|
||||
if body != nil {
|
||||
b, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal body: %v", err)
|
||||
}
|
||||
reqBody = strings.NewReader(string(b))
|
||||
} else {
|
||||
reqBody = strings.NewReader("")
|
||||
}
|
||||
req := httptest.NewRequest(method, path, reqBody)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
mux.ServeHTTP(rec, req)
|
||||
|
||||
var parsed map[string]interface{}
|
||||
if rec.Body.Len() > 0 {
|
||||
_ = json.Unmarshal(rec.Body.Bytes(), &parsed)
|
||||
}
|
||||
return rec, parsed
|
||||
}
|
||||
|
||||
func doForm(t *testing.T, mux http.Handler, path string, form url.Values, basicUser, basicPass string) (*httptest.ResponseRecorder, map[string]interface{}) {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
if basicUser != "" {
|
||||
req.SetBasicAuth(basicUser, basicPass)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
mux.ServeHTTP(rec, req)
|
||||
|
||||
var parsed map[string]interface{}
|
||||
if rec.Body.Len() > 0 {
|
||||
_ = json.Unmarshal(rec.Body.Bytes(), &parsed)
|
||||
}
|
||||
return rec, parsed
|
||||
}
|
||||
|
||||
func doGet(mux http.Handler, path, bearer string) (*httptest.ResponseRecorder, map[string]interface{}) {
|
||||
req := httptest.NewRequest(http.MethodGet, path, nil)
|
||||
if bearer != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+bearer)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
mux.ServeHTTP(rec, req)
|
||||
var parsed map[string]interface{}
|
||||
if rec.Body.Len() > 0 {
|
||||
_ = json.Unmarshal(rec.Body.Bytes(), &parsed)
|
||||
}
|
||||
return rec, parsed
|
||||
}
|
||||
|
||||
func TestOAuthServer_RegisterConfidentialClient_IssuesSecretOnce(t *testing.T) {
|
||||
srv, _ := newTestOAuthServer(t)
|
||||
mux := srv.HTTPHandler()
|
||||
|
||||
rec, resp := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
|
||||
"redirect_uris": []string{"https://app.example.com/callback"},
|
||||
"grant_types": []string{"client_credentials"},
|
||||
})
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("register status = %d, body = %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
secret, _ := resp["client_secret"].(string)
|
||||
if secret == "" {
|
||||
t.Fatal("expected client_secret to be present in registration response")
|
||||
}
|
||||
if _, ok := resp["client_secret_hash"]; ok {
|
||||
t.Error("client_secret_hash must never be returned in the registration response")
|
||||
}
|
||||
if resp["token_endpoint_auth_method"] != "client_secret_basic" {
|
||||
t.Errorf("token_endpoint_auth_method = %v, want client_secret_basic", resp["token_endpoint_auth_method"])
|
||||
}
|
||||
if resp["client_id"] == "" || resp["client_id"] == nil {
|
||||
t.Fatal("expected non-empty client_id")
|
||||
}
|
||||
|
||||
// Fetching the client back (e.g. via a later authorize/token call) must never leak the hash.
|
||||
clientID := resp["client_id"].(string)
|
||||
fetched, ok := srv.lookupOrFetchClient(context.Background(), clientID)
|
||||
if !ok {
|
||||
t.Fatal("expected client to be found")
|
||||
}
|
||||
if fetched.ClientSecretHash == "" {
|
||||
t.Error("expected ClientSecretHash to be stored internally")
|
||||
}
|
||||
if fetched.ClientSecretHash == secret {
|
||||
t.Error("stored hash must not equal the plaintext secret")
|
||||
}
|
||||
|
||||
// Public registration (no client_credentials) should stay a public client with no secret.
|
||||
recPublic, respPublic := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
|
||||
"redirect_uris": []string{"https://app.example.com/callback"},
|
||||
})
|
||||
if recPublic.Code != http.StatusCreated {
|
||||
t.Fatalf("public register status = %d", recPublic.Code)
|
||||
}
|
||||
if _, ok := respPublic["client_secret"]; ok {
|
||||
t.Error("public client registration should not receive a client_secret")
|
||||
}
|
||||
if respPublic["token_endpoint_auth_method"] != "none" {
|
||||
t.Errorf("public client token_endpoint_auth_method = %v, want none", respPublic["token_endpoint_auth_method"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthServer_ClientCredentialsGrant(t *testing.T) {
|
||||
srv, _ := newTestOAuthServer(t)
|
||||
mux := srv.HTTPHandler()
|
||||
|
||||
_, reg := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
|
||||
"redirect_uris": []string{"https://app.example.com/callback"},
|
||||
"grant_types": []string{"client_credentials"},
|
||||
"allowed_scopes": []string{"read", "write"},
|
||||
})
|
||||
clientID := reg["client_id"].(string)
|
||||
clientSecret := reg["client_secret"].(string)
|
||||
|
||||
t.Run("valid credentials", func(t *testing.T) {
|
||||
rec, resp := doForm(t, mux, "/oauth/token", url.Values{
|
||||
"grant_type": {"client_credentials"},
|
||||
"scope": {"read"},
|
||||
}, clientID, clientSecret)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("token status = %d, body = %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
if resp["access_token"] == "" || resp["access_token"] == nil {
|
||||
t.Error("expected non-empty access_token")
|
||||
}
|
||||
if _, ok := resp["refresh_token"]; ok {
|
||||
t.Error("client_credentials must not issue a refresh_token (RFC 6749 §4.4.3)")
|
||||
}
|
||||
if resp["scope"] != "read" {
|
||||
t.Errorf("scope = %v, want read", resp["scope"])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("scope not allowed for client", func(t *testing.T) {
|
||||
rec, resp := doForm(t, mux, "/oauth/token", url.Values{
|
||||
"grant_type": {"client_credentials"},
|
||||
"scope": {"admin"},
|
||||
}, clientID, clientSecret)
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want 400", rec.Code)
|
||||
}
|
||||
if resp["error"] != "invalid_scope" {
|
||||
t.Errorf("error = %v, want invalid_scope", resp["error"])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("wrong secret", func(t *testing.T) {
|
||||
rec, resp := doForm(t, mux, "/oauth/token", url.Values{
|
||||
"grant_type": {"client_credentials"},
|
||||
}, clientID, "wrong-secret")
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401", rec.Code)
|
||||
}
|
||||
if resp["error"] != "invalid_client" {
|
||||
t.Errorf("error = %v, want invalid_client", resp["error"])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("public client cannot use client_credentials", func(t *testing.T) {
|
||||
_, pubReg := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
|
||||
"redirect_uris": []string{"https://app.example.com/callback"},
|
||||
})
|
||||
pubClientID := pubReg["client_id"].(string)
|
||||
|
||||
rec, _ := doForm(t, mux, "/oauth/token", url.Values{
|
||||
"grant_type": {"client_credentials"},
|
||||
"client_id": {pubClientID},
|
||||
}, "", "")
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 for public client attempting client_credentials", rec.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestOAuthServer_ConfidentialClient_AuthCodeGrantRequiresSecret(t *testing.T) {
|
||||
srv, auth := newTestOAuthServer(t)
|
||||
mux := srv.HTTPHandler()
|
||||
|
||||
regResp, err := auth.Register(context.Background(), RegisterRequest{
|
||||
Username: "nadia", Password: "p", Email: "nadia@example.com",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Register() error = %v", err)
|
||||
}
|
||||
|
||||
_, reg := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
|
||||
"redirect_uris": []string{"https://app.example.com/callback"},
|
||||
"token_endpoint_auth_method": "client_secret_basic",
|
||||
})
|
||||
clientID := reg["client_id"].(string)
|
||||
clientSecret := reg["client_secret"].(string)
|
||||
|
||||
verifier := "verifier-nadia-1234567890"
|
||||
challenge := s256Challenge(verifier)
|
||||
if err := auth.OAuthSaveCode(context.Background(), &OAuthCode{
|
||||
Code: "code-no-auth",
|
||||
ClientID: clientID,
|
||||
RedirectURI: "https://app.example.com/callback",
|
||||
CodeChallenge: challenge,
|
||||
SessionToken: regResp.Token,
|
||||
Scopes: []string{"profile"},
|
||||
ExpiresAt: futureTime(),
|
||||
}); err != nil {
|
||||
t.Fatalf("OAuthSaveCode() error = %v", err)
|
||||
}
|
||||
|
||||
// No client credentials supplied -> must be rejected for a confidential client.
|
||||
rec, resp := doForm(t, mux, "/oauth/token", url.Values{
|
||||
"grant_type": {"authorization_code"},
|
||||
"code": {"code-no-auth"},
|
||||
"redirect_uri": {"https://app.example.com/callback"},
|
||||
"client_id": {clientID},
|
||||
"code_verifier": {verifier},
|
||||
}, "", "")
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 without client auth, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
if resp["error"] != "invalid_client" {
|
||||
t.Errorf("error = %v, want invalid_client", resp["error"])
|
||||
}
|
||||
|
||||
// Same code, now with correct client credentials -> succeeds.
|
||||
if err := auth.OAuthSaveCode(context.Background(), &OAuthCode{
|
||||
Code: "code-with-auth",
|
||||
ClientID: clientID,
|
||||
RedirectURI: "https://app.example.com/callback",
|
||||
CodeChallenge: challenge,
|
||||
SessionToken: regResp.Token,
|
||||
Scopes: []string{"profile"},
|
||||
ExpiresAt: futureTime(),
|
||||
}); err != nil {
|
||||
t.Fatalf("OAuthSaveCode() error = %v", err)
|
||||
}
|
||||
rec2, _ := doForm(t, mux, "/oauth/token", url.Values{
|
||||
"grant_type": {"authorization_code"},
|
||||
"code": {"code-with-auth"},
|
||||
"redirect_uri": {"https://app.example.com/callback"},
|
||||
"client_id": {clientID},
|
||||
"code_verifier": {verifier},
|
||||
}, clientID, clientSecret)
|
||||
if rec2.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200 with client auth, body=%s", rec2.Code, rec2.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthServer_PublicClient_AuthCodeGrantUnaffected(t *testing.T) {
|
||||
srv, auth := newTestOAuthServer(t)
|
||||
mux := srv.HTTPHandler()
|
||||
|
||||
regResp, err := auth.Register(context.Background(), RegisterRequest{
|
||||
Username: "oscar", Password: "p", Email: "oscar@example.com",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Register() error = %v", err)
|
||||
}
|
||||
|
||||
_, reg := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
|
||||
"redirect_uris": []string{"https://app.example.com/callback"},
|
||||
})
|
||||
clientID := reg["client_id"].(string)
|
||||
if _, ok := reg["client_secret"]; ok {
|
||||
t.Fatal("expected no client_secret for a default (public) registration")
|
||||
}
|
||||
|
||||
verifier := "verifier-oscar-1234567890"
|
||||
challenge := s256Challenge(verifier)
|
||||
if err := auth.OAuthSaveCode(context.Background(), &OAuthCode{
|
||||
Code: "public-code",
|
||||
ClientID: clientID,
|
||||
RedirectURI: "https://app.example.com/callback",
|
||||
CodeChallenge: challenge,
|
||||
SessionToken: regResp.Token,
|
||||
Scopes: []string{"profile"},
|
||||
ExpiresAt: futureTime(),
|
||||
}); err != nil {
|
||||
t.Fatalf("OAuthSaveCode() error = %v", err)
|
||||
}
|
||||
|
||||
// No credentials needed for a public client — PKCE alone is sufficient, unchanged.
|
||||
rec, _ := doForm(t, mux, "/oauth/token", url.Values{
|
||||
"grant_type": {"authorization_code"},
|
||||
"code": {"public-code"},
|
||||
"redirect_uri": {"https://app.example.com/callback"},
|
||||
"client_id": {clientID},
|
||||
"code_verifier": {verifier},
|
||||
}, "", "")
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200 for public client without credentials, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthServer_ProtectedResourceMetadata(t *testing.T) {
|
||||
srv, _ := newTestOAuthServer(t)
|
||||
mux := srv.HTTPHandler()
|
||||
|
||||
rec, resp := doGet(mux, "/.well-known/oauth-protected-resource", "")
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d", rec.Code)
|
||||
}
|
||||
if resp["resource"] != "https://auth.example.com" {
|
||||
t.Errorf("resource = %v", resp["resource"])
|
||||
}
|
||||
servers, ok := resp["authorization_servers"].([]interface{})
|
||||
if !ok || len(servers) != 1 || servers[0] != "https://auth.example.com" {
|
||||
t.Errorf("authorization_servers = %v", resp["authorization_servers"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthServer_OIDCDiscoveryAndIDToken(t *testing.T) {
|
||||
srv, auth := newTestOAuthServer(t)
|
||||
mux := srv.HTTPHandler()
|
||||
|
||||
// Discovery document.
|
||||
rec, disc := doGet(mux, "/.well-known/openid-configuration", "")
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("discovery status = %d", rec.Code)
|
||||
}
|
||||
if disc["jwks_uri"] != "https://auth.example.com/oauth/jwks.json" {
|
||||
t.Errorf("jwks_uri = %v", disc["jwks_uri"])
|
||||
}
|
||||
grantTypes, _ := disc["grant_types_supported"].([]interface{})
|
||||
found := false
|
||||
for _, g := range grantTypes {
|
||||
if g == "client_credentials" {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Errorf("expected client_credentials in grant_types_supported, got %v", grantTypes)
|
||||
}
|
||||
|
||||
// JWKS.
|
||||
rec, _ = doGet(mux, "/oauth/jwks.json", "")
|
||||
var jwks struct {
|
||||
Keys []struct {
|
||||
Kid string `json:"kid"`
|
||||
N string `json:"n"`
|
||||
E string `json:"e"`
|
||||
} `json:"keys"`
|
||||
}
|
||||
_ = json.Unmarshal(rec.Body.Bytes(), &jwks)
|
||||
if len(jwks.Keys) != 1 {
|
||||
t.Fatalf("expected 1 JWKS key, got %d", len(jwks.Keys))
|
||||
}
|
||||
|
||||
// End-to-end authorization_code flow with scope=openid, verifying the id_token
|
||||
// signature against the published JWKS key.
|
||||
regResp, err := auth.Register(context.Background(), RegisterRequest{
|
||||
Username: "olivia", Password: "p", Email: "olivia@example.com",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Register() error = %v", err)
|
||||
}
|
||||
|
||||
_, clientReg := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
|
||||
"redirect_uris": []string{"https://app.example.com/callback"},
|
||||
})
|
||||
clientID := clientReg["client_id"].(string)
|
||||
|
||||
verifier := "test-code-verifier-1234567890"
|
||||
challenge := s256Challenge(verifier)
|
||||
if err := auth.OAuthSaveCode(context.Background(), &OAuthCode{
|
||||
Code: "oidc-code",
|
||||
ClientID: clientID,
|
||||
RedirectURI: "https://app.example.com/callback",
|
||||
CodeChallenge: challenge,
|
||||
SessionToken: regResp.Token,
|
||||
Scopes: []string{"openid", "profile", "email"},
|
||||
ExpiresAt: futureTime(),
|
||||
}); err != nil {
|
||||
t.Fatalf("OAuthSaveCode() error = %v", err)
|
||||
}
|
||||
|
||||
tokRec, tokenResp := doForm(t, mux, "/oauth/token", url.Values{
|
||||
"grant_type": {"authorization_code"},
|
||||
"code": {"oidc-code"},
|
||||
"redirect_uri": {"https://app.example.com/callback"},
|
||||
"client_id": {clientID},
|
||||
"code_verifier": {verifier},
|
||||
}, "", "")
|
||||
if tokRec.Code != http.StatusOK {
|
||||
t.Fatalf("token exchange status = %d, body = %s", tokRec.Code, tokRec.Body.String())
|
||||
}
|
||||
idTokenStr, _ := tokenResp["id_token"].(string)
|
||||
if idTokenStr == "" {
|
||||
t.Fatal("expected id_token in response for scope containing openid")
|
||||
}
|
||||
|
||||
nBytes, err := base64.RawURLEncoding.DecodeString(jwks.Keys[0].N)
|
||||
if err != nil {
|
||||
t.Fatalf("decode n: %v", err)
|
||||
}
|
||||
eBytes, err := base64.RawURLEncoding.DecodeString(jwks.Keys[0].E)
|
||||
if err != nil {
|
||||
t.Fatalf("decode e: %v", err)
|
||||
}
|
||||
pub := &rsa.PublicKey{N: new(big.Int).SetBytes(nBytes), E: int(new(big.Int).SetBytes(eBytes).Int64())}
|
||||
|
||||
parsed, err := jwt.Parse(idTokenStr, func(token *jwt.Token) (interface{}, error) {
|
||||
return pub, nil
|
||||
}, jwt.WithValidMethods([]string{"RS256"}))
|
||||
if err != nil || !parsed.Valid {
|
||||
t.Fatalf("id_token did not verify against JWKS key: %v", err)
|
||||
}
|
||||
claims := parsed.Claims.(jwt.MapClaims)
|
||||
if claims["iss"] != "https://auth.example.com" {
|
||||
t.Errorf("iss claim = %v", claims["iss"])
|
||||
}
|
||||
if claims["aud"] != clientID {
|
||||
t.Errorf("aud claim = %v, want %v", claims["aud"], clientID)
|
||||
}
|
||||
if claims["preferred_username"] != "olivia" {
|
||||
t.Errorf("preferred_username claim = %v", claims["preferred_username"])
|
||||
}
|
||||
if claims["email"] != "olivia@example.com" {
|
||||
t.Errorf("email claim = %v", claims["email"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthServer_Userinfo(t *testing.T) {
|
||||
srv, auth := newTestOAuthServer(t)
|
||||
mux := srv.HTTPHandler()
|
||||
|
||||
regResp, err := auth.Register(context.Background(), RegisterRequest{
|
||||
Username: "pete", Password: "p", Email: "pete@example.com",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Register() error = %v", err)
|
||||
}
|
||||
|
||||
rec, info := doGet(mux, "/oauth/userinfo", regResp.Token)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("userinfo status = %d, body = %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
if info["preferred_username"] != "pete" {
|
||||
t.Errorf("preferred_username = %v", info["preferred_username"])
|
||||
}
|
||||
|
||||
rec2, _ := doGet(mux, "/oauth/userinfo", "not-a-real-token")
|
||||
if rec2.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 for invalid token", rec2.Code)
|
||||
}
|
||||
}
|
||||
@@ -260,10 +260,7 @@ func (p *DatabasePasskeyProvider) BeginAuthentication(ctx context.Context, usern
|
||||
// If username is provided, get user's credentials
|
||||
var allowCredentials []PasskeyCredentialDescriptor
|
||||
if username != "" {
|
||||
var creds []struct {
|
||||
ID string `json:"credential_id"`
|
||||
Transports []string `json:"transports"`
|
||||
}
|
||||
var creds []passkeyCredential
|
||||
|
||||
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyGetCredsByUsername) {
|
||||
_, directCreds, err := p.getCredsByUsernameDirect(ctx, username)
|
||||
|
||||
@@ -39,7 +39,7 @@ func (p *DatabasePasskeyProvider) storeCredentialDirect(ctx context.Context, par
|
||||
var exists int
|
||||
checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT 1 FROM %s WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials))
|
||||
if err := db.QueryRowContext(ctx, checkQuery, params.CredentialID).Scan(&exists); err == nil {
|
||||
return fmt.Errorf("Credential already exists")
|
||||
return fmt.Errorf("credential already exists")
|
||||
} else if !errors.Is(err, sql.ErrNoRows) {
|
||||
return err
|
||||
}
|
||||
@@ -48,7 +48,7 @@ func (p *DatabasePasskeyProvider) storeCredentialDirect(ctx context.Context, par
|
||||
userCheckQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT 1 FROM %s WHERE id = ?`, p.tableNames.Users))
|
||||
if err := db.QueryRowContext(ctx, userCheckQuery, params.UserID).Scan(&userExists); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return fmt.Errorf("User not found")
|
||||
return fmt.Errorf("user not found")
|
||||
}
|
||||
return err
|
||||
}
|
||||
@@ -78,7 +78,7 @@ func (p *DatabasePasskeyProvider) getCredentialDirect(ctx context.Context, crede
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return 0, 0, fmt.Errorf("Credential not found")
|
||||
return 0, 0, fmt.Errorf("credential not found")
|
||||
}
|
||||
return 0, 0, fmt.Errorf("failed to get credential: %w", err)
|
||||
}
|
||||
@@ -91,7 +91,7 @@ func (p *DatabasePasskeyProvider) updateCounterDirect(ctx context.Context, crede
|
||||
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT sign_count FROM %s WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials))
|
||||
if err := db.QueryRowContext(ctx, query, credentialIDB64).Scan(&oldCounter); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return fmt.Errorf("Credential not found")
|
||||
return fmt.Errorf("credential not found")
|
||||
}
|
||||
return err
|
||||
}
|
||||
@@ -188,7 +188,7 @@ func (p *DatabasePasskeyProvider) deleteCredentialDirect(ctx context.Context, us
|
||||
return err
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("Credential not found")
|
||||
return fmt.Errorf("credential not found")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
@@ -206,24 +206,19 @@ func (p *DatabasePasskeyProvider) updateNameDirect(ctx context.Context, userID i
|
||||
return err
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("Credential not found")
|
||||
return fmt.Errorf("credential not found")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func (p *DatabasePasskeyProvider) getCredsByUsernameDirect(ctx context.Context, username string) (int, []struct {
|
||||
type passkeyCredential struct {
|
||||
ID string `json:"credential_id"`
|
||||
Transports []string `json:"transports"`
|
||||
}, error) {
|
||||
type credT = struct {
|
||||
ID string `json:"credential_id"`
|
||||
Transports []string `json:"transports"`
|
||||
}
|
||||
var userID int
|
||||
var creds []credT
|
||||
}
|
||||
|
||||
err := p.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
func (p *DatabasePasskeyProvider) getCredsByUsernameDirect(ctx context.Context, username string) (userID int, creds []passkeyCredential, err error) {
|
||||
err = p.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
userQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT id FROM %s WHERE username = ? AND is_active = ?`, p.tableNames.Users))
|
||||
if err := db.QueryRowContext(ctx, userQuery, username, true).Scan(&userID); err != nil {
|
||||
return err
|
||||
@@ -236,7 +231,7 @@ func (p *DatabasePasskeyProvider) getCredsByUsernameDirect(ctx context.Context,
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
creds = make([]credT, 0)
|
||||
creds = make([]passkeyCredential, 0)
|
||||
for rows.Next() {
|
||||
var credID string
|
||||
var transportsJSON sql.NullString
|
||||
@@ -247,13 +242,13 @@ func (p *DatabasePasskeyProvider) getCredsByUsernameDirect(ctx context.Context,
|
||||
if transportsJSON.Valid && transportsJSON.String != "" {
|
||||
_ = json.Unmarshal([]byte(transportsJSON.String), &transports)
|
||||
}
|
||||
creds = append(creds, credT{ID: credID, Transports: transports})
|
||||
creds = append(creds, passkeyCredential{ID: credID, Transports: transports})
|
||||
}
|
||||
return rows.Err()
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return 0, nil, fmt.Errorf("User not found")
|
||||
return 0, nil, fmt.Errorf("user not found")
|
||||
}
|
||||
return 0, nil, fmt.Errorf("failed to get credentials: %w", err)
|
||||
}
|
||||
|
||||
@@ -34,7 +34,10 @@ type RowSecurity struct {
|
||||
Tablename string `json:"tablename"`
|
||||
Template string `json:"template"`
|
||||
HasBlock bool `json:"has_block"`
|
||||
UserID int `json:"user_id"`
|
||||
// UserID is the opaque user reference the security rules were loaded for.
|
||||
// It may be an int, a string/UUID, or a *UserContext, depending on what the
|
||||
// RowSecurityProvider/SecurityContext.GetUserRef implementation returns.
|
||||
UserID any `json:"user_id"`
|
||||
}
|
||||
|
||||
func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Type) string {
|
||||
@@ -42,7 +45,7 @@ func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Typ
|
||||
str = strings.ReplaceAll(str, "{PrimaryKeyName}", pPrimaryKeyName)
|
||||
str = strings.ReplaceAll(str, "{TableName}", m.Tablename)
|
||||
str = strings.ReplaceAll(str, "{SchemaName}", m.Schema)
|
||||
str = strings.ReplaceAll(str, "{UserID}", fmt.Sprintf("%d", m.UserID))
|
||||
str = strings.ReplaceAll(str, "{UserID}", fmt.Sprintf("%v", m.UserID))
|
||||
return str
|
||||
}
|
||||
|
||||
@@ -413,7 +416,7 @@ func (m *SecurityList) ClearSecurity(pUserID int, pSchema, pTablename string) er
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserID int, pSchema, pTablename string, pOverwrite bool) (RowSecurity, error) {
|
||||
func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserRef any, pSchema, pTablename string, pOverwrite bool) (RowSecurity, error) {
|
||||
if m.provider == nil {
|
||||
return RowSecurity{}, fmt.Errorf("security provider not set")
|
||||
}
|
||||
@@ -424,10 +427,10 @@ func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserID int, pSchema
|
||||
if m.RowSecurity == nil {
|
||||
m.RowSecurity = make(map[string]RowSecurity, 0)
|
||||
}
|
||||
secKey := fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)
|
||||
secKey := fmt.Sprintf("%s.%s@%v", pSchema, pTablename, pUserRef)
|
||||
|
||||
// Call the provider to load security rules
|
||||
record, err := m.provider.GetRowSecurity(ctx, pUserID, pSchema, pTablename)
|
||||
record, err := m.provider.GetRowSecurity(ctx, pUserRef, pSchema, pTablename)
|
||||
if err != nil {
|
||||
return RowSecurity{}, fmt.Errorf("GetRowSecurity failed: %v", err)
|
||||
}
|
||||
@@ -436,7 +439,7 @@ func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserID int, pSchema
|
||||
return record, nil
|
||||
}
|
||||
|
||||
func (m *SecurityList) GetRowSecurityTemplate(pUserID int, pSchema, pTablename string) (RowSecurity, error) {
|
||||
func (m *SecurityList) GetRowSecurityTemplate(pUserRef any, pSchema, pTablename string) (RowSecurity, error) {
|
||||
defer logger.CatchPanic("GetRowSecurityTemplate")()
|
||||
|
||||
if m.RowSecurity == nil {
|
||||
@@ -446,7 +449,7 @@ func (m *SecurityList) GetRowSecurityTemplate(pUserID int, pSchema, pTablename s
|
||||
m.RowSecurityMutex.RLock()
|
||||
defer m.RowSecurityMutex.RUnlock()
|
||||
|
||||
rowSec, ok := m.RowSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
|
||||
rowSec, ok := m.RowSecurity[fmt.Sprintf("%s.%s@%v", pSchema, pTablename, pUserRef)]
|
||||
if !ok {
|
||||
return RowSecurity{}, fmt.Errorf("no row security data")
|
||||
}
|
||||
|
||||
@@ -44,7 +44,7 @@ func (m *mockSecurityProvider) GetColumnSecurity(ctx context.Context, userID int
|
||||
return m.columnSecurity, nil
|
||||
}
|
||||
|
||||
func (m *mockSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
|
||||
func (m *mockSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
|
||||
return m.rowSecurity, nil
|
||||
}
|
||||
|
||||
|
||||
+40
-12
@@ -551,11 +551,27 @@ func (a *DatabaseAuthenticator) RefreshToken(ctx context.Context, refreshToken s
|
||||
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
||||
}
|
||||
|
||||
return &LoginResponse{
|
||||
// A resolvespec_refresh_token implementation that issues its own rotating
|
||||
// refresh token (independent of the access/session token) returns it
|
||||
// under claims.refresh_token, since UserContext has no dedicated field
|
||||
// for it. Surface that into LoginResponse.RefreshToken so callers don't
|
||||
// need to reach into User.Claims themselves. claims.expires_in
|
||||
// (seconds) similarly overrides the default access-token ExpiresIn when
|
||||
// the procedure provides a real value. Implementations that don't set
|
||||
// these claims keep today's behavior unchanged (empty RefreshToken,
|
||||
// 24h ExpiresIn default).
|
||||
resp := &LoginResponse{
|
||||
Token: userCtx.SessionID, // New session token from stored procedure
|
||||
User: &userCtx,
|
||||
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
||||
}, nil
|
||||
}
|
||||
if refreshToken, ok := userCtx.Claims["refresh_token"].(string); ok && refreshToken != "" {
|
||||
resp.RefreshToken = refreshToken
|
||||
}
|
||||
if expiresIn, ok := userCtx.Claims["expires_in"].(float64); ok && expiresIn > 0 {
|
||||
resp.ExpiresIn = int64(expiresIn)
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// JWTAuthenticator provides JWT token-based authentication
|
||||
@@ -907,17 +923,29 @@ func (p *DatabaseRowSecurityProvider) reconnectDB() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *DatabaseRowSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
|
||||
func (p *DatabaseRowSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
|
||||
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.RowSecurity) {
|
||||
return RowSecurity{}, ErrDirectModeUnsupported
|
||||
}
|
||||
|
||||
var template string
|
||||
var hasBlock bool
|
||||
// resolvespec_row_security's p_user_id is a scalar integer. GetUserRef() may
|
||||
// hand back the full *UserContext so non-DB providers can inspect claims;
|
||||
// unwrap it here before it reaches the SQL args.
|
||||
switch v := userRef.(type) {
|
||||
case *UserContext:
|
||||
if v != nil {
|
||||
userRef = v.UserID
|
||||
}
|
||||
case UserContext:
|
||||
userRef = v.UserID
|
||||
}
|
||||
|
||||
var template sql.NullString
|
||||
var hasBlock sql.NullBool
|
||||
|
||||
runQuery := func() error {
|
||||
query := fmt.Sprintf(`SELECT p_template, p_block FROM %s($1, $2, $3)`, p.sqlNames.RowSecurity)
|
||||
return p.getDB().QueryRowContext(ctx, query, schema, table, userID).Scan(&template, &hasBlock)
|
||||
return p.getDB().QueryRowContext(ctx, query, schema, table, userRef).Scan(&template, &hasBlock)
|
||||
}
|
||||
err := runQuery()
|
||||
if isDBClosed(err) {
|
||||
@@ -932,9 +960,9 @@ func (p *DatabaseRowSecurityProvider) GetRowSecurity(ctx context.Context, userID
|
||||
return RowSecurity{
|
||||
Schema: schema,
|
||||
Tablename: table,
|
||||
UserID: userID,
|
||||
Template: template,
|
||||
HasBlock: hasBlock,
|
||||
UserID: userRef,
|
||||
Template: template.String,
|
||||
HasBlock: hasBlock.Bool,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -969,14 +997,14 @@ func NewConfigRowSecurityProvider(templates map[string]string, blocked map[strin
|
||||
}
|
||||
}
|
||||
|
||||
func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
|
||||
func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
|
||||
key := fmt.Sprintf("%s.%s", schema, table)
|
||||
|
||||
if p.blocked[key] {
|
||||
return RowSecurity{
|
||||
Schema: schema,
|
||||
Tablename: table,
|
||||
UserID: userID,
|
||||
UserID: userRef,
|
||||
HasBlock: true,
|
||||
}, nil
|
||||
}
|
||||
@@ -985,7 +1013,7 @@ func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userID i
|
||||
return RowSecurity{
|
||||
Schema: schema,
|
||||
Tablename: table,
|
||||
UserID: userID,
|
||||
UserID: userRef,
|
||||
Template: template,
|
||||
HasBlock: false,
|
||||
}, nil
|
||||
|
||||
@@ -23,8 +23,8 @@ import (
|
||||
// than introducing a mismatch between modes.
|
||||
|
||||
var (
|
||||
errUsernameExists = errors.New("Username already exists")
|
||||
errEmailExists = errors.New("Email already exists")
|
||||
errUsernameExists = errors.New("username already exists")
|
||||
errEmailExists = errors.New("email already exists")
|
||||
)
|
||||
|
||||
func (a *DatabaseAuthenticator) loginDirect(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
||||
@@ -40,7 +40,7 @@ func (a *DatabaseAuthenticator) loginDirect(ctx context.Context, req LoginReques
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("Invalid credentials")
|
||||
return nil, fmt.Errorf("invalid credentials")
|
||||
}
|
||||
return nil, fmt.Errorf("login query failed: %w", err)
|
||||
}
|
||||
@@ -88,13 +88,13 @@ func (a *DatabaseAuthenticator) loginDirect(ctx context.Context, req LoginReques
|
||||
|
||||
func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req RegisterRequest) (*LoginResponse, error) {
|
||||
if req.Username == "" {
|
||||
return nil, fmt.Errorf("Username is required")
|
||||
return nil, fmt.Errorf("username is required")
|
||||
}
|
||||
if req.Email == "" {
|
||||
return nil, fmt.Errorf("Email is required")
|
||||
return nil, fmt.Errorf("email is required")
|
||||
}
|
||||
if req.Password == "" {
|
||||
return nil, fmt.Errorf("Password is required")
|
||||
return nil, fmt.Errorf("password is required")
|
||||
}
|
||||
|
||||
rolesStr := strings.Join(req.Roles, ",")
|
||||
@@ -197,7 +197,7 @@ func (a *DatabaseAuthenticator) logoutDirect(ctx context.Context, req LogoutRequ
|
||||
return fmt.Errorf("logout query failed: %w", err)
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("Session not found")
|
||||
return fmt.Errorf("session not found")
|
||||
}
|
||||
|
||||
if req.Token != "" {
|
||||
@@ -222,7 +222,7 @@ func (a *DatabaseAuthenticator) sessionDirect(ctx context.Context, token string)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("Invalid or expired session")
|
||||
return nil, fmt.Errorf("invalid or expired session")
|
||||
}
|
||||
return nil, fmt.Errorf("session query failed: %w", err)
|
||||
}
|
||||
@@ -262,7 +262,7 @@ func (a *DatabaseAuthenticator) refreshTokenDirect(ctx context.Context, oldToken
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("Invalid or expired refresh token")
|
||||
return nil, fmt.Errorf("invalid or expired refresh token")
|
||||
}
|
||||
return nil, fmt.Errorf("refresh token query failed: %w", err)
|
||||
}
|
||||
|
||||
@@ -793,6 +793,49 @@ func TestDatabaseAuthenticatorRefreshToken(t *testing.T) {
|
||||
t.Errorf("unfulfilled expectations: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
// A resolvespec_refresh_token implementation that rotates its own
|
||||
// independent refresh token (not just reusing the session/access token)
|
||||
// has nowhere else to put the new refresh token and real access-token
|
||||
// expiry than under UserContext.Claims, since UserContext has no
|
||||
// dedicated fields for either. RefreshToken must surface those claims
|
||||
// keys into LoginResponse.RefreshToken/ExpiresIn rather than silently
|
||||
// dropping them (see the "successful token refresh" case above, which
|
||||
// covers an implementation that has no independent refresh token at all
|
||||
// and gets the 24h default instead).
|
||||
t.Run("surfaces rotated refresh token and expiry from claims", func(t *testing.T) {
|
||||
refreshToken := "refresh-token-abc"
|
||||
|
||||
sessionRows := sqlmock.NewRows([]string{"p_success", "p_error", "p_user"}).
|
||||
AddRow(true, nil, `{"user_id":1,"user_name":"testuser"}`)
|
||||
mock.ExpectQuery(`SELECT p_success, p_error, p_user::text FROM resolvespec_session`).
|
||||
WithArgs(refreshToken, "refresh").
|
||||
WillReturnRows(sessionRows)
|
||||
|
||||
refreshRows := sqlmock.NewRows([]string{"p_success", "p_error", "p_user"}).
|
||||
AddRow(true, nil, `{"user_id":1,"user_name":"testuser","session_id":"new-access-789","claims":{"refresh_token":"new-refresh-def","expires_in":900}}`)
|
||||
mock.ExpectQuery(`SELECT p_success, p_error, p_user::text FROM resolvespec_refresh_token`).
|
||||
WithArgs(refreshToken, sqlmock.AnyArg()).
|
||||
WillReturnRows(refreshRows)
|
||||
|
||||
resp, err := auth.RefreshToken(ctx, refreshToken)
|
||||
if err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
if resp.Token != "new-access-789" {
|
||||
t.Errorf("expected token new-access-789, got %s", resp.Token)
|
||||
}
|
||||
if resp.RefreshToken != "new-refresh-def" {
|
||||
t.Errorf("expected rotated refresh token new-refresh-def, got %q", resp.RefreshToken)
|
||||
}
|
||||
if resp.ExpiresIn != 900 {
|
||||
t.Errorf("expected ExpiresIn 900 from claims, got %d", resp.ExpiresIn)
|
||||
}
|
||||
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Errorf("unfulfilled expectations: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestDatabaseAuthenticatorReconnectsClosedDBPaths(t *testing.T) {
|
||||
|
||||
@@ -22,7 +22,7 @@ func (p *DatabaseTwoFactorProvider) enable2FADirect(ctx context.Context, userID
|
||||
if rows, err := res.RowsAffected(); err != nil {
|
||||
return err
|
||||
} else if rows == 0 {
|
||||
return fmt.Errorf("User not found")
|
||||
return fmt.Errorf("user not found")
|
||||
}
|
||||
|
||||
delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, p.tableNames.UserTOTPBackupCodes))
|
||||
@@ -50,7 +50,7 @@ func (p *DatabaseTwoFactorProvider) disable2FADirect(ctx context.Context, userID
|
||||
if rows, err := res.RowsAffected(); err != nil {
|
||||
return err
|
||||
} else if rows == 0 {
|
||||
return fmt.Errorf("User not found")
|
||||
return fmt.Errorf("user not found")
|
||||
}
|
||||
|
||||
delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, p.tableNames.UserTOTPBackupCodes))
|
||||
@@ -67,7 +67,7 @@ func (p *DatabaseTwoFactorProvider) get2FAStatusDirect(ctx context.Context, user
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return false, fmt.Errorf("User not found")
|
||||
return false, fmt.Errorf("user not found")
|
||||
}
|
||||
return false, fmt.Errorf("get 2FA status query failed: %w", err)
|
||||
}
|
||||
@@ -83,7 +83,7 @@ func (p *DatabaseTwoFactorProvider) get2FASecretDirect(ctx context.Context, user
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return "", fmt.Errorf("User not found")
|
||||
return "", fmt.Errorf("user not found")
|
||||
}
|
||||
return "", fmt.Errorf("get 2FA secret query failed: %w", err)
|
||||
}
|
||||
@@ -101,7 +101,7 @@ func (p *DatabaseTwoFactorProvider) regenerateBackupCodesDirect(ctx context.Cont
|
||||
return err
|
||||
}
|
||||
if count == 0 {
|
||||
return fmt.Errorf("User not found or TOTP not enabled")
|
||||
return fmt.Errorf("user not found or TOTP not enabled")
|
||||
}
|
||||
|
||||
delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, p.tableNames.UserTOTPBackupCodes))
|
||||
@@ -134,7 +134,7 @@ func (p *DatabaseTwoFactorProvider) validateBackupCodeDirect(ctx context.Context
|
||||
return err
|
||||
}
|
||||
if used {
|
||||
return fmt.Errorf("Backup code already used")
|
||||
return fmt.Errorf("backup code already used")
|
||||
}
|
||||
|
||||
updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET used = ?, used_at = ? WHERE id = ?`, p.tableNames.UserTOTPBackupCodes))
|
||||
|
||||
@@ -0,0 +1,249 @@
|
||||
// Package quickproxy provides a small reverse-proxy layer that tries a set
|
||||
// of configured upstream targets first, and falls back to a caller-supplied
|
||||
// http.Handler (typically static file serving) when the upstream is
|
||||
// unreachable or returns 404.
|
||||
package quickproxy
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httputil"
|
||||
"net/url"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Rule maps a URL path prefix to an upstream target.
|
||||
// A Rule with URLPrefix "/" acts as a catch-all passthrough.
|
||||
type Rule struct {
|
||||
// URLPrefix is the URL path prefix this rule matches. Must start with "/".
|
||||
URLPrefix string
|
||||
|
||||
// Target is the upstream base URL, e.g. "http://localhost:3000".
|
||||
// The incoming request path and query are forwarded unchanged; only the
|
||||
// scheme and host are rewritten to Target's.
|
||||
Target string
|
||||
|
||||
// Exclude is a list of URL path prefixes that this rule should not
|
||||
// proxy, even though they fall under URLPrefix. Each entry is a full
|
||||
// path from root and must itself start with URLPrefix (e.g. rule
|
||||
// URLPrefix "/api" excluding a subpath must use "/api/health", not
|
||||
// "/health"). A request matching an Exclude prefix is treated as if
|
||||
// this rule didn't match at all: matching continues against any other
|
||||
// configured rule, falling back if none match. This is typically used
|
||||
// to carve out paths (e.g. "/health") from a catch-all "/" rule so
|
||||
// they're served by the fallback handler instead of being proxied.
|
||||
Exclude []string
|
||||
}
|
||||
|
||||
// DefaultTimeout is the dial and response-header timeout applied to
|
||||
// upstream requests when no WithTimeout option is given. It does not limit
|
||||
// response body streaming.
|
||||
const DefaultTimeout = 10 * time.Second
|
||||
|
||||
// Option configures a Service.
|
||||
type Option func(*options)
|
||||
|
||||
type options struct {
|
||||
timeout time.Duration
|
||||
}
|
||||
|
||||
// WithTimeout sets the dial and response-header timeout used when
|
||||
// connecting to upstream targets. It does not limit response body
|
||||
// streaming, so it won't interrupt long-lived downloads or SSE/WebSocket
|
||||
// connections once established.
|
||||
func WithTimeout(d time.Duration) Option {
|
||||
return func(o *options) { o.timeout = d }
|
||||
}
|
||||
|
||||
// compiledRule pairs a Rule with its ready-to-use reverse proxy.
|
||||
type compiledRule struct {
|
||||
prefix string
|
||||
excludes []string
|
||||
proxy *httputil.ReverseProxy
|
||||
}
|
||||
|
||||
// excluded reports whether path falls under one of the rule's Exclude prefixes.
|
||||
func (r *compiledRule) excluded(path string) bool {
|
||||
for _, ex := range r.excludes {
|
||||
if strings.HasPrefix(path, ex) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Service holds a compiled set of proxy rules and performs longest-prefix
|
||||
// matching against them. A Service is safe for concurrent use once
|
||||
// returned from NewService; Handler must be called once per Service to
|
||||
// wire up the fallback handler before the returned http.Handler is served.
|
||||
type Service struct {
|
||||
rules []compiledRule // sorted by descending prefix length
|
||||
}
|
||||
|
||||
// errUpstreamNotFound is a sentinel error returned from ModifyResponse to
|
||||
// make ReverseProxy invoke ErrorHandler (our fallback path) instead of
|
||||
// writing the upstream's 404 to the client. Nothing has been written to
|
||||
// the ResponseWriter yet when this happens.
|
||||
var errUpstreamNotFound = errors.New("quickproxy: upstream returned 404")
|
||||
|
||||
// NewService compiles the given rules into a Service. Rules are matched by
|
||||
// longest URLPrefix, so a catch-all "/" rule can coexist with more specific
|
||||
// rules such as "/api".
|
||||
func NewService(rules []Rule, opts ...Option) (*Service, error) {
|
||||
if len(rules) == 0 {
|
||||
return nil, fmt.Errorf("quickproxy: no rules configured")
|
||||
}
|
||||
|
||||
cfg := options{timeout: DefaultTimeout}
|
||||
for _, opt := range opts {
|
||||
opt(&cfg)
|
||||
}
|
||||
|
||||
seen := make(map[string]bool, len(rules))
|
||||
compiled := make([]compiledRule, 0, len(rules))
|
||||
|
||||
for _, r := range rules {
|
||||
if !strings.HasPrefix(r.URLPrefix, "/") {
|
||||
return nil, fmt.Errorf("quickproxy: rule prefix %q must start with /", r.URLPrefix)
|
||||
}
|
||||
if seen[r.URLPrefix] {
|
||||
return nil, fmt.Errorf("quickproxy: duplicate rule prefix %q", r.URLPrefix)
|
||||
}
|
||||
seen[r.URLPrefix] = true
|
||||
|
||||
target, err := url.Parse(r.Target)
|
||||
if err != nil || target.Scheme == "" || target.Host == "" {
|
||||
return nil, fmt.Errorf("quickproxy: invalid target %q for prefix %q", r.Target, r.URLPrefix)
|
||||
}
|
||||
|
||||
for _, ex := range r.Exclude {
|
||||
if !strings.HasPrefix(ex, "/") {
|
||||
return nil, fmt.Errorf("quickproxy: exclude prefix %q for rule %q must start with /", ex, r.URLPrefix)
|
||||
}
|
||||
if !strings.HasPrefix(ex, r.URLPrefix) {
|
||||
return nil, fmt.Errorf("quickproxy: exclude prefix %q for rule %q must itself start with the rule's URLPrefix", ex, r.URLPrefix)
|
||||
}
|
||||
}
|
||||
|
||||
compiled = append(compiled, compiledRule{
|
||||
prefix: r.URLPrefix,
|
||||
excludes: r.Exclude,
|
||||
proxy: newReverseProxy(target, cfg.timeout),
|
||||
})
|
||||
}
|
||||
|
||||
// Longest prefix first, so the first match in Handler is always the
|
||||
// most specific one.
|
||||
sort.Slice(compiled, func(i, j int) bool {
|
||||
return len(compiled[i].prefix) > len(compiled[j].prefix)
|
||||
})
|
||||
|
||||
return &Service{rules: compiled}, nil
|
||||
}
|
||||
|
||||
func newReverseProxy(target *url.URL, timeout time.Duration) *httputil.ReverseProxy {
|
||||
transport := &http.Transport{
|
||||
DialContext: (&net.Dialer{
|
||||
Timeout: timeout,
|
||||
}).DialContext,
|
||||
ResponseHeaderTimeout: timeout,
|
||||
}
|
||||
|
||||
return &httputil.ReverseProxy{
|
||||
Transport: transport,
|
||||
Director: func(req *http.Request) {
|
||||
originalHost := req.Host
|
||||
|
||||
req.URL.Scheme = target.Scheme
|
||||
req.URL.Host = target.Host
|
||||
req.Host = target.Host
|
||||
|
||||
if originalHost != "" {
|
||||
req.Header.Set("X-Forwarded-Host", originalHost)
|
||||
}
|
||||
},
|
||||
ModifyResponse: func(resp *http.Response) error {
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return errUpstreamNotFound
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// Handler returns an http.Handler that tries the configured proxy rules
|
||||
// first (longest-prefix match), and calls fallback when no rule matches,
|
||||
// the upstream is unreachable, or the upstream returns 404. Any other
|
||||
// upstream response (2xx, other 4xx, 5xx) is streamed through to the
|
||||
// client unchanged.
|
||||
//
|
||||
// Handler wires up ErrorHandler on the Service's compiled rules, so it
|
||||
// should be called once per Service, before the returned http.Handler
|
||||
// starts serving requests.
|
||||
func (s *Service) Handler(fallback http.Handler) http.Handler {
|
||||
if fallback == nil {
|
||||
fallback = http.NotFoundHandler()
|
||||
}
|
||||
|
||||
for i := range s.rules {
|
||||
s.rules[i].proxy.ErrorHandler = func(w http.ResponseWriter, r *http.Request, _ error) {
|
||||
// ReverseProxy consumes and closes r.Body while attempting the
|
||||
// upstream request, even when that attempt fails (per the
|
||||
// http.RoundTripper contract). Restore a fresh copy from
|
||||
// r.GetBody, set below, before handing the request to fallback.
|
||||
if r.GetBody != nil {
|
||||
if body, err := r.GetBody(); err == nil {
|
||||
r.Body = body
|
||||
}
|
||||
}
|
||||
fallback.ServeHTTP(w, r)
|
||||
}
|
||||
}
|
||||
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
rule := s.match(r.URL.Path)
|
||||
if rule == nil {
|
||||
fallback.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
// Buffer the body so it can be replayed to fallback if the upstream
|
||||
// attempt fails; see ErrorHandler above.
|
||||
if r.Body != nil && r.Body != http.NoBody {
|
||||
bodyBytes, err := io.ReadAll(r.Body)
|
||||
r.Body.Close()
|
||||
if err != nil {
|
||||
http.Error(w, "failed to read request body", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
r.Body = io.NopCloser(bytes.NewReader(bodyBytes))
|
||||
r.GetBody = func() (io.ReadCloser, error) {
|
||||
return io.NopCloser(bytes.NewReader(bodyBytes)), nil
|
||||
}
|
||||
}
|
||||
|
||||
rule.proxy.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
// match returns the longest-prefix rule matching path, or nil if none match.
|
||||
// A rule whose Exclude covers path is skipped, and matching continues
|
||||
// against the next-longest-prefix rule.
|
||||
func (s *Service) match(path string) *compiledRule {
|
||||
for i := range s.rules {
|
||||
if !strings.HasPrefix(path, s.rules[i].prefix) {
|
||||
continue
|
||||
}
|
||||
if s.rules[i].excluded(path) {
|
||||
continue
|
||||
}
|
||||
return &s.rules[i]
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,392 @@
|
||||
package quickproxy
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestNewService_Validation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
rules []Rule
|
||||
wantErr bool
|
||||
}{
|
||||
{"no rules", nil, true},
|
||||
{"empty rules", []Rule{}, true},
|
||||
{"bad prefix", []Rule{{URLPrefix: "api", Target: "http://localhost:1"}}, true},
|
||||
{"bad target", []Rule{{URLPrefix: "/api", Target: "not-a-url"}}, true},
|
||||
{"missing host", []Rule{{URLPrefix: "/api", Target: "http://"}}, true},
|
||||
{"duplicate prefix", []Rule{
|
||||
{URLPrefix: "/api", Target: "http://localhost:1"},
|
||||
{URLPrefix: "/api", Target: "http://localhost:2"},
|
||||
}, true},
|
||||
{"bad exclude prefix", []Rule{
|
||||
{URLPrefix: "/", Target: "http://localhost:1", Exclude: []string{"health"}},
|
||||
}, true},
|
||||
{"exclude outside rule's URLPrefix", []Rule{
|
||||
{URLPrefix: "/api", Target: "http://localhost:1", Exclude: []string{"/health"}},
|
||||
}, true},
|
||||
{"valid", []Rule{{URLPrefix: "/api", Target: "http://localhost:1"}}, false},
|
||||
{"valid with exclude", []Rule{
|
||||
{URLPrefix: "/", Target: "http://localhost:1", Exclude: []string{"/health"}},
|
||||
}, false},
|
||||
{"valid with nested exclude", []Rule{
|
||||
{URLPrefix: "/api", Target: "http://localhost:1", Exclude: []string{"/api/health"}},
|
||||
}, false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, err := NewService(tt.rules)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Fatalf("NewService() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func fallbackHandler(body string) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(body))
|
||||
})
|
||||
}
|
||||
|
||||
func TestHandler_ProxiesSuccessResponse(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte("upstream:" + r.URL.Path))
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
svc, err := NewService([]Rule{{URLPrefix: "/api", Target: upstream.URL}})
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
handler := svc.Handler(fallbackHandler("fallback"))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/widgets", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
if got := rr.Body.String(); got != "upstream:/api/widgets" {
|
||||
t.Fatalf("body = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_404FallsBack(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
_, _ = w.Write([]byte("upstream not found"))
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
svc, err := NewService([]Rule{{URLPrefix: "/", Target: upstream.URL}})
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
handler := svc.Handler(fallbackHandler("fallback-content"))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/missing.html", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
if got := rr.Body.String(); got != "fallback-content" {
|
||||
t.Fatalf("body = %q, want fallback-content", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_UnreachableUpstreamFallsBack(t *testing.T) {
|
||||
// A closed listener address: nothing is listening, so dialing fails.
|
||||
unreachable := "http://127.0.0.1:1"
|
||||
|
||||
svc, err := NewService([]Rule{{URLPrefix: "/", Target: unreachable}}, WithTimeout(500*time.Millisecond))
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
handler := svc.Handler(fallbackHandler("fallback-content"))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/anything", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
if got := rr.Body.String(); got != "fallback-content" {
|
||||
t.Fatalf("body = %q, want fallback-content", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_UnreachableUpstreamFallsBackWithBody(t *testing.T) {
|
||||
// A closed listener address: nothing is listening, so dialing fails and
|
||||
// ReverseProxy invokes ErrorHandler. The fallback handler must still see
|
||||
// the original request body, even though ReverseProxy consumed and
|
||||
// closed it while attempting (and failing) the upstream request.
|
||||
unreachable := "http://127.0.0.1:1"
|
||||
|
||||
svc, err := NewService([]Rule{{URLPrefix: "/", Target: unreachable}}, WithTimeout(500*time.Millisecond))
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
echoBody := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
body, err := io.ReadAll(r.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("fallback reading body: %v", err)
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(body)
|
||||
})
|
||||
|
||||
handler := svc.Handler(echoBody)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/submit", strings.NewReader("payload=1"))
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
if got := rr.Body.String(); got != "payload=1" {
|
||||
t.Fatalf("body = %q, want payload=1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_404FallsBackWithBody(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
svc, err := NewService([]Rule{{URLPrefix: "/", Target: upstream.URL}})
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
echoBody := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
body, err := io.ReadAll(r.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("fallback reading body: %v", err)
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(body)
|
||||
})
|
||||
|
||||
handler := svc.Handler(echoBody)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPut, "/missing", strings.NewReader("payload=2"))
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
if got := rr.Body.String(); got != "payload=2" {
|
||||
t.Fatalf("body = %q, want payload=2", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_NonNotFoundErrorsPassThrough(t *testing.T) {
|
||||
codes := []int{http.StatusOK, http.StatusForbidden, http.StatusBadRequest, http.StatusInternalServerError}
|
||||
|
||||
for _, code := range codes {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(code)
|
||||
_, _ = w.Write([]byte("upstream response"))
|
||||
}))
|
||||
|
||||
svc, err := NewService([]Rule{{URLPrefix: "/", Target: upstream.URL}})
|
||||
if err != nil {
|
||||
upstream.Close()
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
handler := svc.Handler(fallbackHandler("fallback-content"))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/x", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != code {
|
||||
t.Errorf("status for upstream code %d = %d, want %d", code, rr.Code, code)
|
||||
}
|
||||
if got := rr.Body.String(); got != "upstream response" {
|
||||
t.Errorf("body for upstream code %d = %q, want passthrough", code, got)
|
||||
}
|
||||
|
||||
upstream.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_LongestPrefixMatch(t *testing.T) {
|
||||
specific := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
_, _ = w.Write([]byte("specific"))
|
||||
}))
|
||||
defer specific.Close()
|
||||
|
||||
general := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
_, _ = w.Write([]byte("general"))
|
||||
}))
|
||||
defer general.Close()
|
||||
|
||||
svc, err := NewService([]Rule{
|
||||
{URLPrefix: "/", Target: general.URL},
|
||||
{URLPrefix: "/api/v1", Target: specific.URL},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
handler := svc.Handler(fallbackHandler("fallback"))
|
||||
|
||||
for path, want := range map[string]string{
|
||||
"/api/v1/thing": "specific",
|
||||
"/api/other": "general",
|
||||
"/anything": "general",
|
||||
} {
|
||||
req := httptest.NewRequest(http.MethodGet, path, nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if got := rr.Body.String(); got != want {
|
||||
t.Errorf("path %s: body = %q, want %q", path, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_ExcludeFallsBackToFallback(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
_, _ = w.Write([]byte("upstream:" + r.URL.Path))
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
svc, err := NewService([]Rule{
|
||||
{URLPrefix: "/", Target: upstream.URL, Exclude: []string{"/health"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
handler := svc.Handler(fallbackHandler("fallback-content"))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/health", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if got := rr.Body.String(); got != "fallback-content" {
|
||||
t.Fatalf("body = %q, want fallback-content", got)
|
||||
}
|
||||
|
||||
req = httptest.NewRequest(http.MethodGet, "/health/live", nil)
|
||||
rr = httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if got := rr.Body.String(); got != "fallback-content" {
|
||||
t.Fatalf("body = %q, want fallback-content", got)
|
||||
}
|
||||
|
||||
req = httptest.NewRequest(http.MethodGet, "/other", nil)
|
||||
rr = httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if got := rr.Body.String(); got != "upstream:/other" {
|
||||
t.Fatalf("body = %q, want upstream:/other", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_ExcludeFallsThroughToNextRule(t *testing.T) {
|
||||
specific := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
_, _ = w.Write([]byte("specific"))
|
||||
}))
|
||||
defer specific.Close()
|
||||
|
||||
general := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
_, _ = w.Write([]byte("general"))
|
||||
}))
|
||||
defer general.Close()
|
||||
|
||||
svc, err := NewService([]Rule{
|
||||
{URLPrefix: "/api", Target: general.URL},
|
||||
{URLPrefix: "/api/v1", Target: specific.URL, Exclude: []string{"/api/v1/health"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
handler := svc.Handler(fallbackHandler("fallback"))
|
||||
|
||||
for path, want := range map[string]string{
|
||||
"/api/v1/thing": "specific",
|
||||
"/api/v1/health": "general",
|
||||
} {
|
||||
req := httptest.NewRequest(http.MethodGet, path, nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if got := rr.Body.String(); got != want {
|
||||
t.Errorf("path %s: body = %q, want %q", path, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_NoMatchFallsBack(t *testing.T) {
|
||||
svc, err := NewService([]Rule{{URLPrefix: "/api", Target: "http://127.0.0.1:1"}})
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
handler := svc.Handler(fallbackHandler("fallback-content"))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/other", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
if got := rr.Body.String(); got != "fallback-content" {
|
||||
t.Fatalf("body = %q, want fallback-content", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandler_AllMethodsProxied(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(r.Method + ":" + string(body)))
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
svc, err := NewService([]Rule{{URLPrefix: "/api", Target: upstream.URL}})
|
||||
if err != nil {
|
||||
t.Fatalf("NewService: %v", err)
|
||||
}
|
||||
|
||||
handler := svc.Handler(fallbackHandler("fallback"))
|
||||
|
||||
methods := []string{http.MethodGet, http.MethodPost, http.MethodPut, http.MethodPatch, http.MethodDelete}
|
||||
for _, method := range methods {
|
||||
req := httptest.NewRequest(method, "/api/widgets", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
|
||||
want := method + ":"
|
||||
if got := rr.Body.String(); got != want {
|
||||
t.Errorf("method %s: body = %q, want %q", method, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
package spectypes
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// CIString is a string that stores, scans, and returns its value exactly as
|
||||
// given (no case normalization), but compares case-insensitively via Equal
|
||||
// and EqualString. Use it as a bun model field type for columns (e.g.
|
||||
// citext, or codes matched case-insensitively) where you want Go-side
|
||||
// case-insensitive comparisons without forcing the stored/returned value to
|
||||
// a particular case.
|
||||
type CIString string
|
||||
|
||||
// Value implements driver.Valuer. The value is passed through unchanged.
|
||||
func (s CIString) Value() (driver.Value, error) {
|
||||
return string(s), nil
|
||||
}
|
||||
|
||||
// Scan implements sql.Scanner. The value is stored unchanged.
|
||||
func (s *CIString) Scan(value any) error {
|
||||
switch v := value.(type) {
|
||||
case string:
|
||||
*s = CIString(v)
|
||||
case []byte:
|
||||
*s = CIString(v)
|
||||
case nil:
|
||||
*s = ""
|
||||
default:
|
||||
return fmt.Errorf("cannot scan %T into CIString", value)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// String implements fmt.Stringer.
|
||||
func (s CIString) String() string { return string(s) }
|
||||
|
||||
// Equal reports whether s and other are equal, ignoring case.
|
||||
func (s CIString) Equal(other CIString) bool {
|
||||
return strings.EqualFold(string(s), string(other))
|
||||
}
|
||||
|
||||
// EqualString reports whether s equals other, ignoring case.
|
||||
func (s CIString) EqualString(other string) bool {
|
||||
return strings.EqualFold(string(s), other)
|
||||
}
|
||||
|
||||
// Compare returns -1, 0, or +1 if s is less than, equal to, or greater than
|
||||
// other, ignoring case. Useful with slices.SortFunc or similar.
|
||||
func (s CIString) Compare(other CIString) int {
|
||||
return strings.Compare(strings.ToLower(string(s)), strings.ToLower(string(other)))
|
||||
}
|
||||
|
||||
// Less reports whether s sorts before other, ignoring case. Suitable for
|
||||
// sort.Slice or slices.SortFunc comparisons.
|
||||
func (s CIString) Less(other CIString) bool {
|
||||
return s.Compare(other) < 0
|
||||
}
|
||||
|
||||
// LCString is a string that always stores, scans, and returns as lowercase.
|
||||
// Use it as a bun model field type for columns that must be normalized to
|
||||
// lowercase (e.g. codes, slugs, emails) rather than merely compared
|
||||
// case-insensitively; see CIString if the original case must be preserved.
|
||||
type LCString string
|
||||
|
||||
// Value implements driver.Valuer, always lowercase.
|
||||
func (s LCString) Value() (driver.Value, error) {
|
||||
return strings.ToLower(string(s)), nil
|
||||
}
|
||||
|
||||
// Scan implements sql.Scanner, always lowercase.
|
||||
func (s *LCString) Scan(value any) error {
|
||||
switch v := value.(type) {
|
||||
case string:
|
||||
*s = LCString(strings.ToLower(v))
|
||||
case []byte:
|
||||
*s = LCString(strings.ToLower(string(v)))
|
||||
case nil:
|
||||
*s = ""
|
||||
default:
|
||||
return fmt.Errorf("cannot scan %T into LCString", value)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// String implements fmt.Stringer, always lowercase.
|
||||
func (s LCString) String() string { return strings.ToLower(string(s)) }
|
||||
|
||||
// Equal reports whether s and other are equal (case-insensitively, since
|
||||
// both normalize to lowercase).
|
||||
func (s LCString) Equal(other LCString) bool {
|
||||
return s.String() == other.String()
|
||||
}
|
||||
|
||||
// EqualString reports whether s equals other, ignoring case.
|
||||
func (s LCString) EqualString(other string) bool {
|
||||
return s.String() == strings.ToLower(other)
|
||||
}
|
||||
|
||||
// MarshalJSON implements json.Marshaler, always lowercase. Needed because
|
||||
// encoding/json marshals a bare string-kind type as-is and does not call
|
||||
// Value/String, so a value constructed directly (not scanned from the DB)
|
||||
// would otherwise serialize with its original case.
|
||||
func (s LCString) MarshalJSON() ([]byte, error) {
|
||||
return json.Marshal(strings.ToLower(string(s)))
|
||||
}
|
||||
|
||||
// UnmarshalJSON implements json.Unmarshaler, always lowercase.
|
||||
func (s *LCString) UnmarshalJSON(b []byte) error {
|
||||
var str string
|
||||
if err := json.Unmarshal(b, &str); err != nil {
|
||||
return err
|
||||
}
|
||||
*s = LCString(strings.ToLower(str))
|
||||
return nil
|
||||
}
|
||||
|
||||
// UCString is a string that always stores, scans, and returns as uppercase.
|
||||
// Use it as a bun model field type for columns that must be normalized to
|
||||
// uppercase (e.g. table prefix codes) rather than merely compared
|
||||
// case-insensitively; see CIString if the original case must be preserved.
|
||||
type UCString string
|
||||
|
||||
// Value implements driver.Valuer, always uppercase.
|
||||
func (s UCString) Value() (driver.Value, error) {
|
||||
return strings.ToUpper(string(s)), nil
|
||||
}
|
||||
|
||||
// Scan implements sql.Scanner, always uppercase.
|
||||
func (s *UCString) Scan(value any) error {
|
||||
switch v := value.(type) {
|
||||
case string:
|
||||
*s = UCString(strings.ToUpper(v))
|
||||
case []byte:
|
||||
*s = UCString(strings.ToUpper(string(v)))
|
||||
case nil:
|
||||
*s = ""
|
||||
default:
|
||||
return fmt.Errorf("cannot scan %T into UCString", value)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// String implements fmt.Stringer, always uppercase.
|
||||
func (s UCString) String() string { return strings.ToUpper(string(s)) }
|
||||
|
||||
// Equal reports whether s and other are equal (case-insensitively, since
|
||||
// both normalize to uppercase).
|
||||
func (s UCString) Equal(other UCString) bool {
|
||||
return s.String() == other.String()
|
||||
}
|
||||
|
||||
// EqualString reports whether s equals other, ignoring case.
|
||||
func (s UCString) EqualString(other string) bool {
|
||||
return s.String() == strings.ToUpper(other)
|
||||
}
|
||||
|
||||
// MarshalJSON implements json.Marshaler, always uppercase. Needed because
|
||||
// encoding/json marshals a bare string-kind type as-is and does not call
|
||||
// Value/String, so a value constructed directly (not scanned from the DB)
|
||||
// would otherwise serialize with its original case.
|
||||
func (s UCString) MarshalJSON() ([]byte, error) {
|
||||
return json.Marshal(strings.ToUpper(string(s)))
|
||||
}
|
||||
|
||||
// UnmarshalJSON implements json.Unmarshaler, always uppercase.
|
||||
func (s *UCString) UnmarshalJSON(b []byte) error {
|
||||
var str string
|
||||
if err := json.Unmarshal(b, &str); err != nil {
|
||||
return err
|
||||
}
|
||||
*s = UCString(strings.ToUpper(str))
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,379 @@
|
||||
package spectypes
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"sort"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCIString_Scan(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input interface{}
|
||||
expected CIString
|
||||
}{
|
||||
{name: "plain string", input: "MixedCase", expected: "MixedCase"},
|
||||
{name: "bytes as string", input: []byte("FromBytes"), expected: "FromBytes"},
|
||||
{name: "nil value", input: nil, expected: ""},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var s CIString
|
||||
if err := s.Scan(tt.input); err != nil {
|
||||
t.Fatalf("Scan failed: %v", err)
|
||||
}
|
||||
if s != tt.expected {
|
||||
t.Errorf("expected %q, got %q", tt.expected, s)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIString_Scan_InvalidType(t *testing.T) {
|
||||
var s CIString
|
||||
if err := s.Scan(123); err == nil {
|
||||
t.Fatal("expected error scanning int into CIString, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIString_Value(t *testing.T) {
|
||||
s := CIString("MixedCase")
|
||||
v, err := s.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if v != "MixedCase" {
|
||||
t.Errorf("expected %q, got %q (case must be preserved)", "MixedCase", v)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIString_String(t *testing.T) {
|
||||
s := CIString("MixedCase")
|
||||
if s.String() != "MixedCase" {
|
||||
t.Errorf("expected %q, got %q", "MixedCase", s.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIString_Equal(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
a, b CIString
|
||||
expected bool
|
||||
}{
|
||||
{name: "same case", a: "ABC", b: "ABC", expected: true},
|
||||
{name: "different case", a: "ABC", b: "abc", expected: true},
|
||||
{name: "mixed case", a: "AbC", b: "aBc", expected: true},
|
||||
{name: "not equal", a: "ABC", b: "XYZ", expected: false},
|
||||
{name: "both empty", a: "", b: "", expected: true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := tt.a.Equal(tt.b); got != tt.expected {
|
||||
t.Errorf("Equal(%q, %q) = %v, want %v", tt.a, tt.b, got, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIString_EqualString(t *testing.T) {
|
||||
s := CIString("ABC")
|
||||
if !s.EqualString("abc") {
|
||||
t.Error("expected EqualString to match case-insensitively")
|
||||
}
|
||||
if s.EqualString("xyz") {
|
||||
t.Error("expected EqualString to not match different strings")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIString_Compare(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
a, b CIString
|
||||
expected int
|
||||
}{
|
||||
{name: "equal same case", a: "abc", b: "abc", expected: 0},
|
||||
{name: "equal different case", a: "ABC", b: "abc", expected: 0},
|
||||
{name: "less", a: "abc", b: "xyz", expected: -1},
|
||||
{name: "less different case", a: "ABC", b: "xyz", expected: -1},
|
||||
{name: "greater", a: "xyz", b: "abc", expected: 1},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := tt.a.Compare(tt.b); got != tt.expected {
|
||||
t.Errorf("Compare(%q, %q) = %v, want %v", tt.a, tt.b, got, tt.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIString_Less(t *testing.T) {
|
||||
if !CIString("abc").Less("xyz") {
|
||||
t.Error("expected abc < xyz")
|
||||
}
|
||||
if CIString("xyz").Less("abc") {
|
||||
t.Error("expected xyz not < abc")
|
||||
}
|
||||
if CIString("ABC").Less("abc") {
|
||||
t.Error("expected ABC not < abc (equal ignoring case)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIString_Sort(t *testing.T) {
|
||||
vals := []CIString{"banana", "Apple", "cherry", "apple"}
|
||||
sort.Slice(vals, func(i, j int) bool { return vals[i].Less(vals[j]) })
|
||||
|
||||
// After a case-insensitive sort, "Apple"/"apple" must be adjacent and first,
|
||||
// followed by banana then cherry.
|
||||
if !vals[0].EqualString("apple") || !vals[1].EqualString("apple") {
|
||||
t.Errorf("expected the two apple variants first, got %v", vals)
|
||||
}
|
||||
if !vals[2].EqualString("banana") {
|
||||
t.Errorf("expected banana third, got %v", vals)
|
||||
}
|
||||
if !vals[3].EqualString("cherry") {
|
||||
t.Errorf("expected cherry fourth, got %v", vals)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLCString_Scan(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input interface{}
|
||||
expected LCString
|
||||
}{
|
||||
{name: "mixed case string", input: "MixedCase", expected: "mixedcase"},
|
||||
{name: "bytes mixed case", input: []byte("FromBytes"), expected: "frombytes"},
|
||||
{name: "nil value", input: nil, expected: ""},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var s LCString
|
||||
if err := s.Scan(tt.input); err != nil {
|
||||
t.Fatalf("Scan failed: %v", err)
|
||||
}
|
||||
if s != tt.expected {
|
||||
t.Errorf("expected %q, got %q", tt.expected, s)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLCString_Scan_InvalidType(t *testing.T) {
|
||||
var s LCString
|
||||
if err := s.Scan(123); err == nil {
|
||||
t.Fatal("expected error scanning int into LCString, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLCString_Value(t *testing.T) {
|
||||
s := LCString("MixedCase")
|
||||
v, err := s.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if v != "mixedcase" {
|
||||
t.Errorf("expected %q, got %q", "mixedcase", v)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLCString_String(t *testing.T) {
|
||||
s := LCString("MixedCase")
|
||||
if s.String() != "mixedcase" {
|
||||
t.Errorf("expected %q, got %q", "mixedcase", s.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestLCString_Equal(t *testing.T) {
|
||||
if !LCString("ABC").Equal(LCString("abc")) {
|
||||
t.Error("expected ABC and abc to be equal")
|
||||
}
|
||||
if LCString("ABC").Equal(LCString("xyz")) {
|
||||
t.Error("expected ABC and xyz to not be equal")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLCString_EqualString(t *testing.T) {
|
||||
if !LCString("ABC").EqualString("abc") {
|
||||
t.Error("expected EqualString to match case-insensitively")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUCString_Scan(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input interface{}
|
||||
expected UCString
|
||||
}{
|
||||
{name: "mixed case string", input: "MixedCase", expected: "MIXEDCASE"},
|
||||
{name: "bytes mixed case", input: []byte("FromBytes"), expected: "FROMBYTES"},
|
||||
{name: "nil value", input: nil, expected: ""},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var s UCString
|
||||
if err := s.Scan(tt.input); err != nil {
|
||||
t.Fatalf("Scan failed: %v", err)
|
||||
}
|
||||
if s != tt.expected {
|
||||
t.Errorf("expected %q, got %q", tt.expected, s)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUCString_Scan_InvalidType(t *testing.T) {
|
||||
var s UCString
|
||||
if err := s.Scan(123); err == nil {
|
||||
t.Fatal("expected error scanning int into UCString, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUCString_Value(t *testing.T) {
|
||||
s := UCString("MixedCase")
|
||||
v, err := s.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if v != "MIXEDCASE" {
|
||||
t.Errorf("expected %q, got %q", "MIXEDCASE", v)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUCString_String(t *testing.T) {
|
||||
s := UCString("MixedCase")
|
||||
if s.String() != "MIXEDCASE" {
|
||||
t.Errorf("expected %q, got %q", "MIXEDCASE", s.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestUCString_Equal(t *testing.T) {
|
||||
if !UCString("ABC").Equal(UCString("abc")) {
|
||||
t.Error("expected ABC and abc to be equal")
|
||||
}
|
||||
if UCString("ABC").Equal(UCString("xyz")) {
|
||||
t.Error("expected ABC and xyz to not be equal")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUCString_EqualString(t *testing.T) {
|
||||
if !UCString("ABC").EqualString("abc") {
|
||||
t.Error("expected EqualString to match case-insensitively")
|
||||
}
|
||||
}
|
||||
|
||||
// TestLCString_MarshalJSON_NotFromDB verifies a value constructed directly
|
||||
// in Go (never passed through Scan) still normalizes on JSON marshal.
|
||||
func TestLCString_MarshalJSON_NotFromDB(t *testing.T) {
|
||||
s := LCString("MixedCase")
|
||||
b, err := json.Marshal(s)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
if string(b) != `"mixedcase"` {
|
||||
t.Errorf("expected %s, got %s", `"mixedcase"`, b)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLCString_UnmarshalJSON(t *testing.T) {
|
||||
var s LCString
|
||||
if err := json.Unmarshal([]byte(`"MixedCase"`), &s); err != nil {
|
||||
t.Fatalf("Unmarshal failed: %v", err)
|
||||
}
|
||||
if s != "mixedcase" {
|
||||
t.Errorf("expected %q, got %q", "mixedcase", s)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLCString_JSON_StructField(t *testing.T) {
|
||||
type wrapper struct {
|
||||
Code LCString `json:"code"`
|
||||
}
|
||||
in := wrapper{Code: "MixedCase"}
|
||||
b, err := json.Marshal(in)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
if string(b) != `{"code":"mixedcase"}` {
|
||||
t.Errorf("expected %s, got %s", `{"code":"mixedcase"}`, b)
|
||||
}
|
||||
|
||||
var out wrapper
|
||||
if err := json.Unmarshal([]byte(`{"code":"AnotherMixedCase"}`), &out); err != nil {
|
||||
t.Fatalf("Unmarshal failed: %v", err)
|
||||
}
|
||||
if out.Code != "anothermixedcase" {
|
||||
t.Errorf("expected %q, got %q", "anothermixedcase", out.Code)
|
||||
}
|
||||
}
|
||||
|
||||
// TestUCString_MarshalJSON_NotFromDB verifies a value constructed directly
|
||||
// in Go (never passed through Scan) still normalizes on JSON marshal.
|
||||
func TestUCString_MarshalJSON_NotFromDB(t *testing.T) {
|
||||
s := UCString("MixedCase")
|
||||
b, err := json.Marshal(s)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
if string(b) != `"MIXEDCASE"` {
|
||||
t.Errorf("expected %s, got %s", `"MIXEDCASE"`, b)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUCString_UnmarshalJSON(t *testing.T) {
|
||||
var s UCString
|
||||
if err := json.Unmarshal([]byte(`"MixedCase"`), &s); err != nil {
|
||||
t.Fatalf("Unmarshal failed: %v", err)
|
||||
}
|
||||
if s != "MIXEDCASE" {
|
||||
t.Errorf("expected %q, got %q", "MIXEDCASE", s)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUCString_JSON_StructField(t *testing.T) {
|
||||
type wrapper struct {
|
||||
Code UCString `json:"code"`
|
||||
}
|
||||
in := wrapper{Code: "MixedCase"}
|
||||
b, err := json.Marshal(in)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
if string(b) != `{"code":"MIXEDCASE"}` {
|
||||
t.Errorf("expected %s, got %s", `{"code":"MIXEDCASE"}`, b)
|
||||
}
|
||||
|
||||
var out wrapper
|
||||
if err := json.Unmarshal([]byte(`{"code":"AnotherMixedCase"}`), &out); err != nil {
|
||||
t.Fatalf("Unmarshal failed: %v", err)
|
||||
}
|
||||
if out.Code != "ANOTHERMIXEDCASE" {
|
||||
t.Errorf("expected %q, got %q", "ANOTHERMIXEDCASE", out.Code)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCIString_JSON_PreservesCase confirms CIString needs no custom JSON
|
||||
// methods: it should never normalize case, only its DB Value/Scan and the
|
||||
// Equal/EqualString comparisons apply case-insensitivity.
|
||||
func TestCIString_JSON_PreservesCase(t *testing.T) {
|
||||
s := CIString("MixedCase")
|
||||
b, err := json.Marshal(s)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
if string(b) != `"MixedCase"` {
|
||||
t.Errorf("expected %s, got %s", `"MixedCase"`, b)
|
||||
}
|
||||
|
||||
var out CIString
|
||||
if err := json.Unmarshal([]byte(`"AnotherMixedCase"`), &out); err != nil {
|
||||
t.Fatalf("Unmarshal failed: %v", err)
|
||||
}
|
||||
if out != "AnotherMixedCase" {
|
||||
t.Errorf("expected case to be preserved, got %q", out)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,239 @@
|
||||
package spectypes
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ── PostGIS geometry / geography ─────────────────────────────────────────────
|
||||
|
||||
// SqlGeometry is a nullable PostGIS `geometry` column.
|
||||
//
|
||||
// On Scan it accepts the PostGIS default hex-EWKB text output, a raw GeoJSON
|
||||
// object, or a WKT/EWKT string (e.g. when the column is selected via
|
||||
// ST_AsGeoJSON / ST_AsText). Internally it holds a canonical GeoJSON geometry
|
||||
// object plus the SRID.
|
||||
//
|
||||
// On Value it emits `SRID=<n>;<WKT>` text. PostGIS registers an implicit
|
||||
// text -> geometry cast, so parameterised inserts/updates work without wrapping
|
||||
// the placeholder in a constructor function.
|
||||
//
|
||||
// MarshalJSON emits the GeoJSON geometry object (or null).
|
||||
type SqlGeometry struct {
|
||||
GeoJSON json.RawMessage
|
||||
SRID int
|
||||
Valid bool
|
||||
}
|
||||
|
||||
// SqlGeography is identical to SqlGeometry but maps to a PostGIS `geography`
|
||||
// column. Coordinates are always lon/lat and the default SRID is 4326.
|
||||
type SqlGeography struct {
|
||||
SqlGeometry
|
||||
}
|
||||
|
||||
func (g *SqlGeometry) Scan(value any) error {
|
||||
if value == nil {
|
||||
g.Valid = false
|
||||
g.GeoJSON = nil
|
||||
g.SRID = 0
|
||||
return nil
|
||||
}
|
||||
|
||||
var s string
|
||||
switch v := value.(type) {
|
||||
case string:
|
||||
s = v
|
||||
case []byte:
|
||||
s = string(v)
|
||||
default:
|
||||
return fmt.Errorf("SqlGeometry: cannot scan type %T", value)
|
||||
}
|
||||
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
g.Valid = false
|
||||
g.GeoJSON = nil
|
||||
return nil
|
||||
}
|
||||
|
||||
switch {
|
||||
case strings.HasPrefix(s, "{"):
|
||||
// GeoJSON object.
|
||||
if _, err := geoJSONToGeom([]byte(s)); err != nil {
|
||||
return fmt.Errorf("SqlGeometry: invalid GeoJSON: %w", err)
|
||||
}
|
||||
g.GeoJSON = json.RawMessage(s)
|
||||
g.Valid = true
|
||||
return nil
|
||||
case isHex(s):
|
||||
gj, srid, err := DecodeEWKBHex(s)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SqlGeometry: %w", err)
|
||||
}
|
||||
g.GeoJSON = gj
|
||||
g.SRID = srid
|
||||
g.Valid = true
|
||||
return nil
|
||||
default:
|
||||
// WKT / EWKT text.
|
||||
srid, wkt := splitEWKT(s)
|
||||
gj, err := wktToGeoJSON(wkt)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SqlGeometry: %w", err)
|
||||
}
|
||||
g.GeoJSON = gj
|
||||
g.SRID = srid
|
||||
g.Valid = true
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (g SqlGeometry) Value() (driver.Value, error) {
|
||||
if !g.Valid || len(g.GeoJSON) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
wkt, err := GeoJSONToWKT(g.GeoJSON)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
srid := g.SRID
|
||||
if srid == 0 {
|
||||
srid = 4326
|
||||
}
|
||||
return fmt.Sprintf("SRID=%d;%s", srid, wkt), nil
|
||||
}
|
||||
|
||||
func (g SqlGeometry) MarshalJSON() ([]byte, error) {
|
||||
if !g.Valid || len(g.GeoJSON) == 0 {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
return g.GeoJSON, nil
|
||||
}
|
||||
|
||||
func (g *SqlGeometry) UnmarshalJSON(b []byte) error {
|
||||
s := strings.TrimSpace(string(b))
|
||||
if s == "" || s == "null" {
|
||||
g.Valid = false
|
||||
g.GeoJSON = nil
|
||||
g.SRID = 0
|
||||
return nil
|
||||
}
|
||||
|
||||
if strings.HasPrefix(s, "{") {
|
||||
if _, err := geoJSONToGeom(b); err != nil {
|
||||
return fmt.Errorf("SqlGeometry: invalid GeoJSON: %w", err)
|
||||
}
|
||||
g.GeoJSON = append(json.RawMessage(nil), b...)
|
||||
g.Valid = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// String value: EWKT / WKT / hex-EWKB.
|
||||
var str string
|
||||
if err := json.Unmarshal(b, &str); err != nil {
|
||||
return fmt.Errorf("SqlGeometry: cannot unmarshal %s", b)
|
||||
}
|
||||
str = strings.TrimSpace(str)
|
||||
if str == "" {
|
||||
g.Valid = false
|
||||
g.GeoJSON = nil
|
||||
return nil
|
||||
}
|
||||
if isHex(str) {
|
||||
gj, srid, err := DecodeEWKBHex(str)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SqlGeometry: %w", err)
|
||||
}
|
||||
g.GeoJSON = gj
|
||||
g.SRID = srid
|
||||
g.Valid = true
|
||||
return nil
|
||||
}
|
||||
srid, wkt := splitEWKT(str)
|
||||
gj, err := wktToGeoJSON(wkt)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SqlGeometry: %w", err)
|
||||
}
|
||||
g.GeoJSON = gj
|
||||
g.SRID = srid
|
||||
g.Valid = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// WKT returns the geometry as a plain WKT string (no SRID prefix).
|
||||
func (g SqlGeometry) WKT() string {
|
||||
if !g.Valid || len(g.GeoJSON) == 0 {
|
||||
return ""
|
||||
}
|
||||
wkt, err := GeoJSONToWKT(g.GeoJSON)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return wkt
|
||||
}
|
||||
|
||||
// EWKT returns the geometry as `SRID=<n>;<WKT>`.
|
||||
func (g SqlGeometry) EWKT() string {
|
||||
wkt := g.WKT()
|
||||
if wkt == "" {
|
||||
return ""
|
||||
}
|
||||
srid := g.SRID
|
||||
if srid == 0 {
|
||||
srid = 4326
|
||||
}
|
||||
return fmt.Sprintf("SRID=%d;%s", srid, wkt)
|
||||
}
|
||||
|
||||
// NewSqlGeometryFromGeoJSON builds a SqlGeometry from a GeoJSON geometry object.
|
||||
func NewSqlGeometryFromGeoJSON(geojson []byte, srid int) (SqlGeometry, error) {
|
||||
if _, err := geoJSONToGeom(geojson); err != nil {
|
||||
return SqlGeometry{}, err
|
||||
}
|
||||
return SqlGeometry{GeoJSON: append(json.RawMessage(nil), geojson...), SRID: srid, Valid: true}, nil
|
||||
}
|
||||
|
||||
// NewSqlGeometryFromEWKT builds a SqlGeometry from an EWKT or WKT string.
|
||||
func NewSqlGeometryFromEWKT(ewkt string) (SqlGeometry, error) {
|
||||
srid, wkt := splitEWKT(strings.TrimSpace(ewkt))
|
||||
gj, err := wktToGeoJSON(wkt)
|
||||
if err != nil {
|
||||
return SqlGeometry{}, err
|
||||
}
|
||||
return SqlGeometry{GeoJSON: gj, SRID: srid, Valid: true}, nil
|
||||
}
|
||||
|
||||
// ── helpers ─────────────────────────────────────────────────────────────────
|
||||
|
||||
func isHex(s string) bool {
|
||||
if len(s) < 10 || len(s)%2 != 0 {
|
||||
return false
|
||||
}
|
||||
_, err := hex.DecodeString(s)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// splitEWKT separates an optional `SRID=<n>;` prefix from a WKT body.
|
||||
func splitEWKT(s string) (srid int, wkt string) {
|
||||
if strings.HasPrefix(strings.ToUpper(s), "SRID=") {
|
||||
if idx := strings.Index(s, ";"); idx > 0 {
|
||||
if n, err := strconv.Atoi(strings.TrimSpace(s[5:idx])); err == nil {
|
||||
return n, strings.TrimSpace(s[idx+1:])
|
||||
}
|
||||
}
|
||||
}
|
||||
return 0, s
|
||||
}
|
||||
|
||||
// wktToGeoJSON parses a (subset of) WKT into a GeoJSON geometry object.
|
||||
func wktToGeoJSON(wkt string) ([]byte, error) {
|
||||
g, err := parseWKT(wkt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return geomToGeoJSON(g)
|
||||
}
|
||||
@@ -0,0 +1,126 @@
|
||||
package spectypes
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// SRID=4326;POINT (1 2)
|
||||
const pointHexEWKB = "0101000020E6100000000000000000F03F0000000000000040"
|
||||
|
||||
func TestSqlGeometry_ScanHexEWKB(t *testing.T) {
|
||||
var g SqlGeometry
|
||||
if err := g.Scan(pointHexEWKB); err != nil {
|
||||
t.Fatalf("Scan: %v", err)
|
||||
}
|
||||
if !g.Valid || g.SRID != 4326 {
|
||||
t.Fatalf("got Valid=%v SRID=%d", g.Valid, g.SRID)
|
||||
}
|
||||
if !jsonEqual(t, g.GeoJSON, `{"type":"Point","coordinates":[1,2]}`) {
|
||||
t.Errorf("GeoJSON = %s", g.GeoJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlGeometry_ScanGeoJSON(t *testing.T) {
|
||||
var g SqlGeometry
|
||||
if err := g.Scan(`{"type":"Point","coordinates":[3,4]}`); err != nil {
|
||||
t.Fatalf("Scan: %v", err)
|
||||
}
|
||||
if !g.Valid {
|
||||
t.Fatal("expected valid")
|
||||
}
|
||||
if !jsonEqual(t, g.GeoJSON, `{"type":"Point","coordinates":[3,4]}`) {
|
||||
t.Errorf("GeoJSON = %s", g.GeoJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlGeometry_ScanEWKT(t *testing.T) {
|
||||
var g SqlGeometry
|
||||
if err := g.Scan("SRID=3857;POINT (5 6)"); err != nil {
|
||||
t.Fatalf("Scan: %v", err)
|
||||
}
|
||||
if g.SRID != 3857 {
|
||||
t.Errorf("SRID = %d, want 3857", g.SRID)
|
||||
}
|
||||
if !jsonEqual(t, g.GeoJSON, `{"type":"Point","coordinates":[5,6]}`) {
|
||||
t.Errorf("GeoJSON = %s", g.GeoJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlGeometry_Value(t *testing.T) {
|
||||
g, err := NewSqlGeometryFromGeoJSON([]byte(`{"type":"Point","coordinates":[1,2]}`), 4326)
|
||||
if err != nil {
|
||||
t.Fatalf("New: %v", err)
|
||||
}
|
||||
v, err := g.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value: %v", err)
|
||||
}
|
||||
if v != "SRID=4326;POINT (1 2)" {
|
||||
t.Errorf("Value = %v, want SRID=4326;POINT (1 2)", v)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlGeometry_ValueDefaultsSRID(t *testing.T) {
|
||||
g, _ := NewSqlGeometryFromGeoJSON([]byte(`{"type":"Point","coordinates":[1,2]}`), 0)
|
||||
v, _ := g.Value()
|
||||
if v != "SRID=4326;POINT (1 2)" {
|
||||
t.Errorf("Value = %v", v)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlGeometry_JSON(t *testing.T) {
|
||||
g, _ := NewSqlGeometryFromEWKT("SRID=4326;POINT (1 2)")
|
||||
b, err := json.Marshal(g)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal: %v", err)
|
||||
}
|
||||
if !jsonEqual(t, b, `{"type":"Point","coordinates":[1,2]}`) {
|
||||
t.Errorf("json = %s", b)
|
||||
}
|
||||
|
||||
var back SqlGeometry
|
||||
if err := json.Unmarshal([]byte(`{"type":"Point","coordinates":[7,8]}`), &back); err != nil {
|
||||
t.Fatalf("Unmarshal: %v", err)
|
||||
}
|
||||
if !back.Valid || !jsonEqual(t, back.GeoJSON, `{"type":"Point","coordinates":[7,8]}`) {
|
||||
t.Errorf("unmarshal = %+v", back)
|
||||
}
|
||||
|
||||
// Unmarshal also accepts an EWKT string.
|
||||
var fromStr SqlGeometry
|
||||
if err := json.Unmarshal([]byte(`"SRID=4326;POINT(9 10)"`), &fromStr); err != nil {
|
||||
t.Fatalf("Unmarshal string: %v", err)
|
||||
}
|
||||
if !jsonEqual(t, fromStr.GeoJSON, `{"type":"Point","coordinates":[9,10]}`) {
|
||||
t.Errorf("fromStr = %s", fromStr.GeoJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlGeometry_Null(t *testing.T) {
|
||||
var g SqlGeometry
|
||||
if err := g.Scan(nil); err != nil {
|
||||
t.Fatalf("Scan(nil): %v", err)
|
||||
}
|
||||
if g.Valid {
|
||||
t.Error("expected invalid")
|
||||
}
|
||||
v, err := g.Value()
|
||||
if err != nil || v != nil {
|
||||
t.Errorf("Value = %v, %v", v, err)
|
||||
}
|
||||
b, _ := json.Marshal(g)
|
||||
if string(b) != "null" {
|
||||
t.Errorf("json = %s", b)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlGeography_Embeds(t *testing.T) {
|
||||
var g SqlGeography
|
||||
if err := g.Scan("SRID=4326;POINT (1 2)"); err != nil {
|
||||
t.Fatalf("Scan: %v", err)
|
||||
}
|
||||
if !g.Valid || !jsonEqual(t, g.GeoJSON, `{"type":"Point","coordinates":[1,2]}`) {
|
||||
t.Errorf("geography scan = %+v", g)
|
||||
}
|
||||
}
|
||||
@@ -336,21 +336,21 @@ type (
|
||||
SqlUUID = SqlNull[uuid.UUID]
|
||||
)
|
||||
|
||||
// SqlTimeStamp - Timestamp with custom formatting (YYYY-MM-DDTHH:MM:SS).
|
||||
// SqlTimeStamp - Timestamp serialized as RFC3339 with timezone offset.
|
||||
type SqlTimeStamp struct{ SqlNull[time.Time] }
|
||||
|
||||
func (t SqlTimeStamp) MarshalJSON() ([]byte, error) {
|
||||
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
return []byte(fmt.Sprintf(`"%s"`, t.Val.Format("2006-01-02T15:04:05"))), nil
|
||||
return []byte(fmt.Sprintf(`"%s"`, t.Val.Format(time.RFC3339))), nil
|
||||
}
|
||||
|
||||
func (t *SqlTimeStamp) UnmarshalJSON(b []byte) error {
|
||||
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
|
||||
return err
|
||||
}
|
||||
if t.Valid && (t.Val.IsZero() || t.Val.Format("2006-01-02T15:04:05") == "0001-01-01T00:00:00") {
|
||||
if t.Valid && (t.Val.IsZero() || t.Val.Year() <= 1) {
|
||||
t.Valid = false
|
||||
}
|
||||
return nil
|
||||
@@ -360,7 +360,7 @@ func (t SqlTimeStamp) Value() (driver.Value, error) {
|
||||
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
|
||||
return nil, nil
|
||||
}
|
||||
return t.Val.Format("2006-01-02T15:04:05"), nil
|
||||
return t.Val.Format(time.RFC3339), nil
|
||||
}
|
||||
|
||||
func SqlTimeStampNow() SqlTimeStamp {
|
||||
@@ -425,9 +425,7 @@ func (t SqlTime) MarshalJSON() ([]byte, error) {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
s := t.Val.Format("15:04:05")
|
||||
if s == "00:00:00" {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
|
||||
return []byte(fmt.Sprintf(`"%s"`, s)), nil
|
||||
}
|
||||
|
||||
@@ -435,9 +433,7 @@ func (t *SqlTime) UnmarshalJSON(b []byte) error {
|
||||
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
|
||||
return err
|
||||
}
|
||||
if t.Valid && t.Val.Format("15:04:05") == "00:00:00" {
|
||||
t.Valid = false
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -178,7 +178,7 @@ func TestSqlTimeStamp_JSON(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
expected := `"2024-01-15T10:30:45"`
|
||||
expected := `"2024-01-15T10:30:45Z"`
|
||||
if string(data) != expected {
|
||||
t.Errorf("expected %s, got %s", expected, string(data))
|
||||
}
|
||||
@@ -920,6 +920,411 @@ func TestSqlString_RoundTrip(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlBool_Scan tests SqlBool Scan from various input types.
|
||||
func TestSqlBool_Scan(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input interface{}
|
||||
expected bool
|
||||
valid bool
|
||||
}{
|
||||
{"bool true", true, true, true},
|
||||
{"bool false", false, false, true},
|
||||
{"string true", "true", true, true},
|
||||
{"string 1", "1", true, true},
|
||||
{"int64 1 fallback", int64(1), true, true},
|
||||
{"int64 0 fallback", int64(0), false, true},
|
||||
{"nil", nil, false, false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var b SqlBool
|
||||
if err := b.Scan(tt.input); err != nil {
|
||||
t.Fatalf("Scan failed: %v", err)
|
||||
}
|
||||
if b.Valid != tt.valid {
|
||||
t.Errorf("expected valid=%v, got valid=%v", tt.valid, b.Valid)
|
||||
}
|
||||
if tt.valid && b.Val != tt.expected {
|
||||
t.Errorf("expected %v, got %v", tt.expected, b.Val)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlBool_Value(t *testing.T) {
|
||||
b := NewSqlBool(true)
|
||||
val, err := b.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != true {
|
||||
t.Errorf("expected true, got %v", val)
|
||||
}
|
||||
|
||||
b2 := SqlBool{Valid: false}
|
||||
val2, err := b2.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val2 != nil {
|
||||
t.Errorf("expected nil, got %v", val2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlBool_JSON(t *testing.T) {
|
||||
b := NewSqlBool(true)
|
||||
data, err := json.Marshal(b)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
if string(data) != "true" {
|
||||
t.Errorf("expected true, got %s", string(data))
|
||||
}
|
||||
|
||||
var b2 SqlBool
|
||||
if err := json.Unmarshal([]byte("false"), &b2); err != nil {
|
||||
t.Fatalf("Unmarshal failed: %v", err)
|
||||
}
|
||||
if !b2.Valid || b2.Val != false {
|
||||
t.Errorf("expected valid=true val=false, got valid=%v val=%v", b2.Valid, b2.Val)
|
||||
}
|
||||
|
||||
var b3 SqlBool
|
||||
if err := json.Unmarshal([]byte("null"), &b3); err != nil {
|
||||
t.Fatalf("Unmarshal null failed: %v", err)
|
||||
}
|
||||
if b3.Valid {
|
||||
t.Error("expected invalid after unmarshaling null")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlNull_FromString_EdgeCases tests FromString edge cases shared by all SqlNull instantiations.
|
||||
func TestSqlNull_FromString_EdgeCases(t *testing.T) {
|
||||
t.Run("empty string is null", func(t *testing.T) {
|
||||
var n SqlInt64
|
||||
if err := n.FromString(""); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if n.Valid {
|
||||
t.Error("expected invalid for empty string")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("NULL case-insensitive", func(t *testing.T) {
|
||||
var n SqlString
|
||||
if err := n.FromString("NuLL"); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if n.Valid {
|
||||
t.Error("expected invalid for 'NuLL'")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("whitespace trimmed", func(t *testing.T) {
|
||||
var n SqlInt64
|
||||
if err := n.FromString(" 42 "); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if !n.Valid || n.Val != 42 {
|
||||
t.Errorf("expected valid=true val=42, got valid=%v val=%v", n.Valid, n.Val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid int string stays invalid", func(t *testing.T) {
|
||||
var n SqlInt64
|
||||
if err := n.FromString("not-a-number"); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if n.Valid {
|
||||
t.Error("expected invalid for non-numeric string")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("float string truncated into int type", func(t *testing.T) {
|
||||
var n SqlInt64
|
||||
if err := n.FromString("3.7"); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if !n.Valid || n.Val != 3 {
|
||||
t.Errorf("expected valid=true val=3, got valid=%v val=%v", n.Valid, n.Val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid bool string stays invalid", func(t *testing.T) {
|
||||
var n SqlBool
|
||||
if err := n.FromString("maybe"); err != nil {
|
||||
t.Fatalf("FromString failed: %v", err)
|
||||
}
|
||||
if n.Valid {
|
||||
t.Error("expected invalid for non-bool string")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSqlNull_String tests the String() stringer fallback.
|
||||
func TestSqlNull_String(t *testing.T) {
|
||||
t.Run("invalid returns empty", func(t *testing.T) {
|
||||
n := SqlInt64{Valid: false}
|
||||
if n.String() != "" {
|
||||
t.Errorf("expected empty string, got %q", n.String())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("stringer type delegates", func(t *testing.T) {
|
||||
u := uuid.New()
|
||||
n := NewSqlUUID(u)
|
||||
if n.String() != u.String() {
|
||||
t.Errorf("expected %s, got %s", u.String(), n.String())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-stringer falls back to fmt", func(t *testing.T) {
|
||||
n := NewSqlInt64(42)
|
||||
if n.String() != "42" {
|
||||
t.Errorf("expected 42, got %s", n.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestNewSql_Generic tests the generic NewSql constructor.
|
||||
func TestNewSql_Generic(t *testing.T) {
|
||||
t.Run("exact type match", func(t *testing.T) {
|
||||
n := NewSql[int64](int64(5))
|
||||
if !n.Valid || n.Val != 5 {
|
||||
t.Errorf("expected valid=true val=5, got valid=%v val=%v", n.Valid, n.Val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("nil value", func(t *testing.T) {
|
||||
n := NewSql[int64](nil)
|
||||
if n.Valid {
|
||||
t.Error("expected invalid for nil")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("from another SqlNull", func(t *testing.T) {
|
||||
src := SqlNull[int64]{Val: 9, Valid: true}
|
||||
n := NewSql[int64](src)
|
||||
if !n.Valid || n.Val != 9 {
|
||||
t.Errorf("expected valid=true val=9, got valid=%v val=%v", n.Valid, n.Val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("string conversion fallback", func(t *testing.T) {
|
||||
n := NewSql[string](42)
|
||||
if !n.Valid || n.Val != "42" {
|
||||
t.Errorf("expected valid=true val=42, got valid=%v val=%q", n.Valid, n.Val)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSqlNull_Int64_Conversions tests Int64() across differently-typed SqlNull values.
|
||||
func TestSqlNull_Int64_Conversions(t *testing.T) {
|
||||
if v := (SqlNull[string]{Val: "42", Valid: true}).Int64(); v != 42 {
|
||||
t.Errorf("expected 42, got %d", v)
|
||||
}
|
||||
if v := (SqlNull[bool]{Val: true, Valid: true}).Int64(); v != 1 {
|
||||
t.Errorf("expected 1, got %d", v)
|
||||
}
|
||||
if v := (SqlNull[bool]{Val: false, Valid: true}).Int64(); v != 0 {
|
||||
t.Errorf("expected 0, got %d", v)
|
||||
}
|
||||
if v := (SqlNull[float64]{Val: 3.9, Valid: true}).Int64(); v != 3 {
|
||||
t.Errorf("expected 3, got %d", v)
|
||||
}
|
||||
if v := (SqlNull[int64]{Valid: false}).Int64(); v != 0 {
|
||||
t.Errorf("expected 0 for invalid, got %d", v)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlNull_Float64_Conversions tests Float64() across differently-typed SqlNull values.
|
||||
func TestSqlNull_Float64_Conversions(t *testing.T) {
|
||||
if v := (SqlNull[string]{Val: "3.14", Valid: true}).Float64(); v != 3.14 {
|
||||
t.Errorf("expected 3.14, got %v", v)
|
||||
}
|
||||
if v := (SqlNull[int64]{Val: 10, Valid: true}).Float64(); v != 10.0 {
|
||||
t.Errorf("expected 10.0, got %v", v)
|
||||
}
|
||||
if v := (SqlNull[float64]{Valid: false}).Float64(); v != 0.0 {
|
||||
t.Errorf("expected 0.0 for invalid, got %v", v)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlNull_Bool_Conversions tests Bool() across differently-typed SqlNull values.
|
||||
func TestSqlNull_Bool_Conversions(t *testing.T) {
|
||||
if v := (SqlNull[string]{Val: "YES", Valid: true}).Bool(); v != true {
|
||||
t.Error("expected true for 'YES'")
|
||||
}
|
||||
if v := (SqlNull[string]{Val: "no", Valid: true}).Bool(); v != false {
|
||||
t.Error("expected false for 'no'")
|
||||
}
|
||||
if v := (SqlNull[int]{Val: 1, Valid: true}).Bool(); v != true {
|
||||
t.Error("expected true for int 1")
|
||||
}
|
||||
if v := (SqlNull[int]{Val: 0, Valid: true}).Bool(); v != false {
|
||||
t.Error("expected false for int 0")
|
||||
}
|
||||
if v := (SqlNull[bool]{Valid: false}).Bool(); v != false {
|
||||
t.Error("expected false for invalid")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlNull_Time_NonTimeType verifies Time() returns zero value when T is not time.Time.
|
||||
func TestSqlNull_Time_NonTimeType(t *testing.T) {
|
||||
n := SqlNull[string]{Val: "2024-01-15", Valid: true}
|
||||
if !n.Time().IsZero() {
|
||||
t.Error("expected zero time for non-time.Time SqlNull")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlNull_UUID_NonUUIDType verifies UUID() returns uuid.Nil when T is not uuid.UUID.
|
||||
func TestSqlNull_UUID_NonUUIDType(t *testing.T) {
|
||||
n := SqlNull[string]{Val: "not-a-uuid", Valid: true}
|
||||
if n.UUID() != uuid.Nil {
|
||||
t.Error("expected uuid.Nil for non-uuid.UUID SqlNull")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlTime_Midnight verifies midnight times are serialized as "00:00:00", not null.
|
||||
func TestSqlTime_Midnight(t *testing.T) {
|
||||
midnight := time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC)
|
||||
tm := NewSqlTime(midnight)
|
||||
|
||||
data, err := json.Marshal(tm)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal failed: %v", err)
|
||||
}
|
||||
if string(data) != `"00:00:00"` {
|
||||
t.Errorf("expected \"00:00:00\", got %s", string(data))
|
||||
}
|
||||
|
||||
val, err := tm.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != "00:00:00" {
|
||||
t.Errorf("expected 00:00:00, got %v", val)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlTime_Value_Invalid(t *testing.T) {
|
||||
tm := SqlTime{}
|
||||
val, err := tm.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != nil {
|
||||
t.Errorf("expected nil, got %v", val)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlDate_ZeroValueString verifies String() blanks out sentinel zero dates.
|
||||
func TestSqlDate_ZeroValueString(t *testing.T) {
|
||||
d := SqlDate{SqlNull: SqlNull[time.Time]{Val: time.Time{}, Valid: true}}
|
||||
if d.String() != "" {
|
||||
t.Errorf("expected empty string for zero date, got %q", d.String())
|
||||
}
|
||||
|
||||
sentinel := time.Date(1800, 12, 31, 0, 0, 0, 0, time.UTC)
|
||||
d2 := SqlDate{SqlNull: SqlNull[time.Time]{Val: sentinel, Valid: true}}
|
||||
if d2.String() != "" {
|
||||
t.Errorf("expected empty string for 1800-12-31 sentinel, got %q", d2.String())
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlTimeStamp_Value tests driver.Valuer for SqlTimeStamp, including the pre-year-2 cutoff.
|
||||
func TestSqlTimeStamp_Value(t *testing.T) {
|
||||
t.Run("valid recent timestamp", func(t *testing.T) {
|
||||
ts := NewSqlTimeStamp(time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC))
|
||||
val, err := ts.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != "2024-01-15T10:30:00Z" {
|
||||
t.Errorf("expected 2024-01-15T10:30:00Z, got %v", val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("year 1 is treated as null", func(t *testing.T) {
|
||||
ts := NewSqlTimeStamp(time.Date(1, 1, 1, 0, 0, 0, 0, time.UTC))
|
||||
val, err := ts.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != nil {
|
||||
t.Errorf("expected nil, got %v", val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid is null", func(t *testing.T) {
|
||||
ts := SqlTimeStamp{}
|
||||
val, err := ts.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value failed: %v", err)
|
||||
}
|
||||
if val != nil {
|
||||
t.Errorf("expected nil, got %v", val)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSqlTimeStamp_UnmarshalJSON_YearOneInvalid(t *testing.T) {
|
||||
var ts SqlTimeStamp
|
||||
if err := json.Unmarshal([]byte(`"0001-01-01T00:00:00Z"`), &ts); err != nil {
|
||||
t.Fatalf("Unmarshal failed: %v", err)
|
||||
}
|
||||
if ts.Valid {
|
||||
t.Error("expected invalid for year 0001 timestamp")
|
||||
}
|
||||
}
|
||||
|
||||
// TestTryParseDT tests the internal multi-format date/time parser.
|
||||
func TestTryParseDT(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
}{
|
||||
{"RFC3339", "2024-01-15T10:30:00Z"},
|
||||
{"date only", "2024-01-15"},
|
||||
{"datetime no tz", "2024-01-15T10:30:00"},
|
||||
{"space separated", "2024-01-15 10:30:00"},
|
||||
{"UK date slash", "15/01/2024"},
|
||||
{"UK date dash", "15-01-2024"},
|
||||
{"time only", "10:30:00"},
|
||||
{"short time", "10:30"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
tm, err := tryParseDT(tt.input)
|
||||
if err != nil {
|
||||
t.Fatalf("tryParseDT failed for %q: %v", tt.input, err)
|
||||
}
|
||||
if tm.IsZero() {
|
||||
t.Errorf("expected non-zero time for %q", tt.input)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("invalid format", func(t *testing.T) {
|
||||
_, err := tryParseDT("not a date at all")
|
||||
if err == nil {
|
||||
t.Error("expected error for unparseable string")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestToJSONDT tests RFC3339 formatting helper.
|
||||
func TestToJSONDT(t *testing.T) {
|
||||
dt := time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC)
|
||||
expected := dt.Format(time.RFC3339)
|
||||
if got := ToJSONDT(dt); got != expected {
|
||||
t.Errorf("expected %s, got %s", expected, got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSqlByteArray_Base64_RoundTrip tests complete round-trip: Go -> JSON -> Go -> SQL -> Go
|
||||
func TestSqlByteArray_Base64_RoundTrip(t *testing.T) {
|
||||
original := []byte{0x48, 0x65, 0x6C, 0x6C, 0x6F, 0x20, 0xFF, 0xFE} // "Hello " + binary data
|
||||
@@ -955,4 +1360,3 @@ func TestSqlByteArray_Base64_RoundTrip(t *testing.T) {
|
||||
t.Errorf("Round-trip failed: expected %v, got %v", original, b3.Val)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,308 @@
|
||||
package spectypes
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// pgvector column types beyond the plain `vector` (SqlVector, in
|
||||
// sql_array_types.go): `halfvec`, `sparsevec` and `bit`.
|
||||
|
||||
// parseVectorLiteral parses a pgvector dense literal `[1,2,3]` into []float32.
|
||||
func parseVectorLiteral(s string) ([]float32, error) {
|
||||
s = strings.TrimSpace(s)
|
||||
if !strings.HasPrefix(s, "[") || !strings.HasSuffix(s, "]") {
|
||||
return nil, fmt.Errorf("not a valid vector literal: %q", s)
|
||||
}
|
||||
inner := strings.TrimSpace(s[1 : len(s)-1])
|
||||
if inner == "" {
|
||||
return []float32{}, nil
|
||||
}
|
||||
parts := strings.Split(inner, ",")
|
||||
out := make([]float32, len(parts))
|
||||
for i, p := range parts {
|
||||
f, err := strconv.ParseFloat(strings.TrimSpace(p), 32)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("vector element %d %q: %w", i, p, err)
|
||||
}
|
||||
out[i] = float32(f)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func formatVectorLiteral(vals []float32) string {
|
||||
parts := make([]string, len(vals))
|
||||
for i, v := range vals {
|
||||
parts[i] = strconv.FormatFloat(float64(v), 'f', -1, 32)
|
||||
}
|
||||
return "[" + strings.Join(parts, ",") + "]"
|
||||
}
|
||||
|
||||
// ── SqlHalfVector ────────────────────────────────────────────────────────────
|
||||
|
||||
// SqlHalfVector is a nullable pgvector `halfvec` (half-precision) column, backed
|
||||
// by []float32. Wire format matches `vector`: `[1,2,3]`.
|
||||
type SqlHalfVector struct {
|
||||
Val []float32
|
||||
Valid bool
|
||||
}
|
||||
|
||||
func (v *SqlHalfVector) Scan(value any) error {
|
||||
if value == nil {
|
||||
v.Valid = false
|
||||
v.Val = nil
|
||||
return nil
|
||||
}
|
||||
var s string
|
||||
switch val := value.(type) {
|
||||
case string:
|
||||
s = val
|
||||
case []byte:
|
||||
s = string(val)
|
||||
default:
|
||||
return fmt.Errorf("SqlHalfVector: cannot scan type %T", value)
|
||||
}
|
||||
parsed, err := parseVectorLiteral(s)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SqlHalfVector: %w", err)
|
||||
}
|
||||
v.Val = parsed
|
||||
v.Valid = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (v SqlHalfVector) Value() (driver.Value, error) {
|
||||
if !v.Valid {
|
||||
return nil, nil
|
||||
}
|
||||
return formatVectorLiteral(v.Val), nil
|
||||
}
|
||||
|
||||
func (v SqlHalfVector) MarshalJSON() ([]byte, error) {
|
||||
if !v.Valid {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
return json.Marshal(v.Val)
|
||||
}
|
||||
|
||||
func (v *SqlHalfVector) UnmarshalJSON(b []byte) error {
|
||||
if strings.TrimSpace(string(b)) == "null" {
|
||||
v.Valid = false
|
||||
v.Val = nil
|
||||
return nil
|
||||
}
|
||||
var vals []float32
|
||||
if err := json.Unmarshal(b, &vals); err != nil {
|
||||
return err
|
||||
}
|
||||
v.Val = vals
|
||||
v.Valid = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func NewSqlHalfVector(val []float32) SqlHalfVector {
|
||||
return SqlHalfVector{Val: val, Valid: true}
|
||||
}
|
||||
|
||||
// ── SqlSparseVector ──────────────────────────────────────────────────────────
|
||||
|
||||
// SqlSparseVector is a nullable pgvector `sparsevec` column. Wire format:
|
||||
// `{1:0.5,4:0.2}/8` (1-based indices). JSON:
|
||||
// `{"dim":8,"indices":[1,4],"values":[0.5,0.2]}`.
|
||||
type SqlSparseVector struct {
|
||||
Dim int
|
||||
Indices []int32
|
||||
Values []float32
|
||||
Valid bool
|
||||
}
|
||||
|
||||
func (v *SqlSparseVector) Scan(value any) error {
|
||||
if value == nil {
|
||||
v.Valid = false
|
||||
v.Dim, v.Indices, v.Values = 0, nil, nil
|
||||
return nil
|
||||
}
|
||||
var s string
|
||||
switch val := value.(type) {
|
||||
case string:
|
||||
s = val
|
||||
case []byte:
|
||||
s = string(val)
|
||||
default:
|
||||
return fmt.Errorf("SqlSparseVector: cannot scan type %T", value)
|
||||
}
|
||||
s = strings.TrimSpace(s)
|
||||
slash := strings.LastIndex(s, "/")
|
||||
if !strings.HasPrefix(s, "{") || slash < 0 || !strings.Contains(s[:slash], "}") {
|
||||
return fmt.Errorf("SqlSparseVector: invalid literal %q", s)
|
||||
}
|
||||
dim, err := strconv.Atoi(strings.TrimSpace(s[slash+1:]))
|
||||
if err != nil {
|
||||
return fmt.Errorf("SqlSparseVector: bad dimension: %w", err)
|
||||
}
|
||||
body := strings.TrimSpace(s[1:strings.LastIndex(s, "}")])
|
||||
var idx []int32
|
||||
var vals []float32
|
||||
if body != "" {
|
||||
for _, pair := range strings.Split(body, ",") {
|
||||
kv := strings.SplitN(pair, ":", 2)
|
||||
if len(kv) != 2 {
|
||||
return fmt.Errorf("SqlSparseVector: bad pair %q", pair)
|
||||
}
|
||||
k, err := strconv.Atoi(strings.TrimSpace(kv[0]))
|
||||
if err != nil {
|
||||
return fmt.Errorf("SqlSparseVector: bad index %q: %w", kv[0], err)
|
||||
}
|
||||
f, err := strconv.ParseFloat(strings.TrimSpace(kv[1]), 32)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SqlSparseVector: bad value %q: %w", kv[1], err)
|
||||
}
|
||||
idx = append(idx, int32(k))
|
||||
vals = append(vals, float32(f))
|
||||
}
|
||||
}
|
||||
v.Dim, v.Indices, v.Values, v.Valid = dim, idx, vals, true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (v SqlSparseVector) Value() (driver.Value, error) {
|
||||
if !v.Valid {
|
||||
return nil, nil
|
||||
}
|
||||
pairs := make([]string, len(v.Indices))
|
||||
for i, k := range v.Indices {
|
||||
val := float32(0)
|
||||
if i < len(v.Values) {
|
||||
val = v.Values[i]
|
||||
}
|
||||
pairs[i] = strconv.Itoa(int(k)) + ":" + strconv.FormatFloat(float64(val), 'f', -1, 32)
|
||||
}
|
||||
return "{" + strings.Join(pairs, ",") + "}/" + strconv.Itoa(v.Dim), nil
|
||||
}
|
||||
|
||||
type sparseVectorJSON struct {
|
||||
Dim int `json:"dim"`
|
||||
Indices []int32 `json:"indices"`
|
||||
Values []float32 `json:"values"`
|
||||
}
|
||||
|
||||
func (v SqlSparseVector) MarshalJSON() ([]byte, error) {
|
||||
if !v.Valid {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
return json.Marshal(sparseVectorJSON{Dim: v.Dim, Indices: v.Indices, Values: v.Values})
|
||||
}
|
||||
|
||||
func (v *SqlSparseVector) UnmarshalJSON(b []byte) error {
|
||||
if strings.TrimSpace(string(b)) == "null" {
|
||||
v.Valid = false
|
||||
v.Dim, v.Indices, v.Values = 0, nil, nil
|
||||
return nil
|
||||
}
|
||||
var j sparseVectorJSON
|
||||
if err := json.Unmarshal(b, &j); err != nil {
|
||||
return err
|
||||
}
|
||||
v.Dim, v.Indices, v.Values, v.Valid = j.Dim, j.Indices, j.Values, true
|
||||
return nil
|
||||
}
|
||||
|
||||
func NewSqlSparseVector(dim int, indices []int32, values []float32) SqlSparseVector {
|
||||
return SqlSparseVector{Dim: dim, Indices: indices, Values: values, Valid: true}
|
||||
}
|
||||
|
||||
// ── SqlBitVector ─────────────────────────────────────────────────────────────
|
||||
|
||||
// SqlBitVector is a nullable Postgres `bit(n)` / `varbit` column (used by
|
||||
// pgvector for Hamming/Jaccard distance), backed by []bool. Wire format: a
|
||||
// string of '0'/'1' characters. JSON: a bool array.
|
||||
type SqlBitVector struct {
|
||||
Val []bool
|
||||
Valid bool
|
||||
}
|
||||
|
||||
func (v *SqlBitVector) Scan(value any) error {
|
||||
if value == nil {
|
||||
v.Valid = false
|
||||
v.Val = nil
|
||||
return nil
|
||||
}
|
||||
var s string
|
||||
switch val := value.(type) {
|
||||
case string:
|
||||
s = val
|
||||
case []byte:
|
||||
s = string(val)
|
||||
default:
|
||||
return fmt.Errorf("SqlBitVector: cannot scan type %T", value)
|
||||
}
|
||||
s = strings.TrimSpace(s)
|
||||
out := make([]bool, len(s))
|
||||
for i, c := range s {
|
||||
switch c {
|
||||
case '1':
|
||||
out[i] = true
|
||||
case '0':
|
||||
out[i] = false
|
||||
default:
|
||||
return fmt.Errorf("SqlBitVector: invalid bit %q", string(c))
|
||||
}
|
||||
}
|
||||
v.Val = out
|
||||
v.Valid = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (v SqlBitVector) Value() (driver.Value, error) {
|
||||
if !v.Valid {
|
||||
return nil, nil
|
||||
}
|
||||
var b strings.Builder
|
||||
b.Grow(len(v.Val))
|
||||
for _, bit := range v.Val {
|
||||
if bit {
|
||||
b.WriteByte('1')
|
||||
} else {
|
||||
b.WriteByte('0')
|
||||
}
|
||||
}
|
||||
return b.String(), nil
|
||||
}
|
||||
|
||||
func (v SqlBitVector) MarshalJSON() ([]byte, error) {
|
||||
if !v.Valid {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
return json.Marshal(v.Val)
|
||||
}
|
||||
|
||||
func (v *SqlBitVector) UnmarshalJSON(b []byte) error {
|
||||
s := strings.TrimSpace(string(b))
|
||||
if s == "null" {
|
||||
v.Valid = false
|
||||
v.Val = nil
|
||||
return nil
|
||||
}
|
||||
// Accept both a bool array and a "0101" string.
|
||||
if strings.HasPrefix(s, "\"") {
|
||||
var str string
|
||||
if err := json.Unmarshal(b, &str); err != nil {
|
||||
return err
|
||||
}
|
||||
return v.Scan(str)
|
||||
}
|
||||
var vals []bool
|
||||
if err := json.Unmarshal(b, &vals); err != nil {
|
||||
return err
|
||||
}
|
||||
v.Val = vals
|
||||
v.Valid = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func NewSqlBitVector(val []bool) SqlBitVector {
|
||||
return SqlBitVector{Val: val, Valid: true}
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
package spectypes
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSqlHalfVector_RoundTrip(t *testing.T) {
|
||||
v := NewSqlHalfVector([]float32{1, 2.5, -3})
|
||||
|
||||
dv, err := v.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value: %v", err)
|
||||
}
|
||||
if dv != "[1,2.5,-3]" {
|
||||
t.Errorf("Value = %v, want [1,2.5,-3]", dv)
|
||||
}
|
||||
|
||||
var back SqlHalfVector
|
||||
if err := back.Scan(dv.(string)); err != nil {
|
||||
t.Fatalf("Scan: %v", err)
|
||||
}
|
||||
if !back.Valid || !reflect.DeepEqual(back.Val, v.Val) {
|
||||
t.Errorf("Scan = %+v, want %+v", back, v)
|
||||
}
|
||||
|
||||
b, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal: %v", err)
|
||||
}
|
||||
if string(b) != "[1,2.5,-3]" {
|
||||
t.Errorf("json = %s, want [1,2.5,-3]", b)
|
||||
}
|
||||
|
||||
var fromJSON SqlHalfVector
|
||||
if err := json.Unmarshal(b, &fromJSON); err != nil {
|
||||
t.Fatalf("Unmarshal: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(fromJSON.Val, v.Val) {
|
||||
t.Errorf("json round-trip = %+v", fromJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlHalfVector_Null(t *testing.T) {
|
||||
var v SqlHalfVector
|
||||
if err := v.Scan(nil); err != nil {
|
||||
t.Fatalf("Scan(nil): %v", err)
|
||||
}
|
||||
if v.Valid {
|
||||
t.Error("expected invalid after Scan(nil)")
|
||||
}
|
||||
dv, err := v.Value()
|
||||
if err != nil || dv != nil {
|
||||
t.Errorf("Value = %v, %v; want nil, nil", dv, err)
|
||||
}
|
||||
b, _ := json.Marshal(v)
|
||||
if string(b) != "null" {
|
||||
t.Errorf("json = %s, want null", b)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlSparseVector_RoundTrip(t *testing.T) {
|
||||
v := NewSqlSparseVector(8, []int32{1, 4}, []float32{0.5, 0.2})
|
||||
|
||||
dv, err := v.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value: %v", err)
|
||||
}
|
||||
if dv != "{1:0.5,4:0.2}/8" {
|
||||
t.Errorf("Value = %v, want {1:0.5,4:0.2}/8", dv)
|
||||
}
|
||||
|
||||
var back SqlSparseVector
|
||||
if err := back.Scan(dv.(string)); err != nil {
|
||||
t.Fatalf("Scan: %v", err)
|
||||
}
|
||||
if back.Dim != 8 || !reflect.DeepEqual(back.Indices, []int32{1, 4}) ||
|
||||
!reflect.DeepEqual(back.Values, []float32{0.5, 0.2}) {
|
||||
t.Errorf("Scan = %+v", back)
|
||||
}
|
||||
|
||||
b, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal: %v", err)
|
||||
}
|
||||
want := `{"dim":8,"indices":[1,4],"values":[0.5,0.2]}`
|
||||
if string(b) != want {
|
||||
t.Errorf("json = %s, want %s", b, want)
|
||||
}
|
||||
|
||||
var fromJSON SqlSparseVector
|
||||
if err := json.Unmarshal([]byte(want), &fromJSON); err != nil {
|
||||
t.Fatalf("Unmarshal: %v", err)
|
||||
}
|
||||
if fromJSON.Dim != 8 || !fromJSON.Valid {
|
||||
t.Errorf("json round-trip = %+v", fromJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlSparseVector_ScanInvalid(t *testing.T) {
|
||||
var v SqlSparseVector
|
||||
for _, s := range []string{"[1,2,3]", "{1:0.5}", "{1:0.5}/x", "bad"} {
|
||||
if err := v.Scan(s); err == nil {
|
||||
t.Errorf("Scan(%q) expected error", s)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlBitVector_RoundTrip(t *testing.T) {
|
||||
v := NewSqlBitVector([]bool{true, false, true, true})
|
||||
|
||||
dv, err := v.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value: %v", err)
|
||||
}
|
||||
if dv != "1011" {
|
||||
t.Errorf("Value = %v, want 1011", dv)
|
||||
}
|
||||
|
||||
var back SqlBitVector
|
||||
if err := back.Scan("1011"); err != nil {
|
||||
t.Fatalf("Scan: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(back.Val, v.Val) {
|
||||
t.Errorf("Scan = %+v", back)
|
||||
}
|
||||
|
||||
b, _ := json.Marshal(v)
|
||||
if string(b) != "[true,false,true,true]" {
|
||||
t.Errorf("json = %s", b)
|
||||
}
|
||||
|
||||
// JSON also accepts a "0101" string.
|
||||
var fromStr SqlBitVector
|
||||
if err := json.Unmarshal([]byte(`"1011"`), &fromStr); err != nil {
|
||||
t.Fatalf("Unmarshal string: %v", err)
|
||||
}
|
||||
if !reflect.DeepEqual(fromStr.Val, v.Val) {
|
||||
t.Errorf("string json = %+v", fromStr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlBitVector_ScanInvalid(t *testing.T) {
|
||||
var v SqlBitVector
|
||||
if err := v.Scan("1021"); err == nil {
|
||||
t.Error("expected error for invalid bit")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,123 @@
|
||||
package spectypes
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// pkgPath is the import path of this package, used to recognise spectypes
|
||||
// wrappers by reflection.
|
||||
const pkgPath = "github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||
|
||||
// canonicalSQLNames maps a spectypes wrapper type name to the PostgreSQL type
|
||||
// name it represents. Dimensioned types (vector(1536), geometry(Point,4326))
|
||||
// still need a gorm/bun `type:` tag for the full declaration — this is the
|
||||
// fallback used for metadata and OpenAPI when no tag is present.
|
||||
var canonicalSQLNames = map[string]string{
|
||||
"SqlVector": "vector",
|
||||
"SqlHalfVector": "halfvec",
|
||||
"SqlSparseVector": "sparsevec",
|
||||
"SqlBitVector": "bit",
|
||||
"SqlGeometry": "geometry",
|
||||
"SqlGeography": "geography",
|
||||
"SqlJSONB": "jsonb",
|
||||
"SqlStringArray": "text[]",
|
||||
"SqlInt16Array": "smallint[]",
|
||||
"SqlInt32Array": "integer[]",
|
||||
"SqlInt64Array": "bigint[]",
|
||||
"SqlFloat32Array": "real[]",
|
||||
"SqlFloat64Array": "double precision[]",
|
||||
"SqlBoolArray": "boolean[]",
|
||||
"SqlUUIDArray": "uuid[]",
|
||||
"SqlDate": "date",
|
||||
"SqlTime": "time",
|
||||
"SqlTimeStamp": "timestamp",
|
||||
}
|
||||
|
||||
// sqlNullElemNames maps the element type of a SqlNull[T] alias to a PG type name.
|
||||
var sqlNullElemNames = map[string]string{
|
||||
"int16": "smallint",
|
||||
"int32": "integer",
|
||||
"int64": "bigint",
|
||||
"float64": "double precision",
|
||||
"bool": "boolean",
|
||||
"string": "text",
|
||||
"[]uint8": "bytea",
|
||||
"uuid.UUID": "uuid",
|
||||
"Time": "timestamp",
|
||||
}
|
||||
|
||||
// SQLTypeName returns the canonical PostgreSQL type name for a spectypes wrapper
|
||||
// type, or ("", false) if t is not a recognised spectypes type.
|
||||
func SQLTypeName(t reflect.Type) (string, bool) {
|
||||
for t != nil && t.Kind() == reflect.Pointer {
|
||||
t = t.Elem()
|
||||
}
|
||||
if t == nil || t.PkgPath() != pkgPath {
|
||||
return "", false
|
||||
}
|
||||
|
||||
name := t.Name()
|
||||
if n, ok := canonicalSQLNames[name]; ok {
|
||||
return n, true
|
||||
}
|
||||
|
||||
// SqlNull[T] aliases, e.g. "SqlNull[int16]", "SqlNull[uuid.UUID]".
|
||||
if strings.HasPrefix(name, "SqlNull[") && strings.HasSuffix(name, "]") {
|
||||
elem := name[len("SqlNull[") : len(name)-1]
|
||||
if idx := strings.LastIndex(elem, "."); idx >= 0 {
|
||||
// keep last path segment, e.g. "github.com/google/uuid.UUID" -> "uuid.UUID"
|
||||
if slash := strings.LastIndex(elem[:idx], "/"); slash >= 0 {
|
||||
elem = elem[slash+1:]
|
||||
}
|
||||
}
|
||||
if n, ok := sqlNullElemNames[elem]; ok {
|
||||
return n, true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
return "", false
|
||||
}
|
||||
|
||||
// IsSpatialType reports whether t is a PostGIS geometry/geography wrapper.
|
||||
func IsSpatialType(t reflect.Type) bool {
|
||||
n, ok := SQLTypeName(t)
|
||||
return ok && (n == "geometry" || n == "geography")
|
||||
}
|
||||
|
||||
// IsVectorType reports whether t is a pgvector wrapper (vector/halfvec/sparsevec).
|
||||
func IsVectorType(t reflect.Type) bool {
|
||||
n, ok := SQLTypeName(t)
|
||||
return ok && (n == "vector" || n == "halfvec" || n == "sparsevec")
|
||||
}
|
||||
|
||||
// IsJSONType reports whether t is a spectypes JSON/JSONB wrapper.
|
||||
func IsJSONType(t reflect.Type) bool {
|
||||
n, ok := SQLTypeName(t)
|
||||
return ok && (n == "jsonb" || n == "json")
|
||||
}
|
||||
|
||||
// UnwrapKind returns the reflect.Kind to use when reasoning about a column's
|
||||
// comparability (numeric vs. string vs. other) for filter building. Plain Go
|
||||
// types return their own Kind unchanged. spectypes.SqlNull[T] wrappers (and
|
||||
// types that embed one, such as SqlTimeStamp/SqlDate/SqlTime) always report
|
||||
// reflect.Struct for their own Kind even when T is an int64 or string, which
|
||||
// would otherwise make numeric/text columns look "complex" and force an
|
||||
// unnecessary CAST(... AS TEXT) that defeats native column indexes. For those
|
||||
// wrappers, UnwrapKind returns the Kind of the wrapped value T instead.
|
||||
func UnwrapKind(t reflect.Type) reflect.Kind {
|
||||
for t != nil && t.Kind() == reflect.Pointer {
|
||||
t = t.Elem()
|
||||
}
|
||||
if t == nil {
|
||||
return reflect.Invalid
|
||||
}
|
||||
if t.Kind() != reflect.Struct || t.PkgPath() != pkgPath {
|
||||
return t.Kind()
|
||||
}
|
||||
if f, ok := t.FieldByName("Val"); ok {
|
||||
return f.Type.Kind()
|
||||
}
|
||||
return t.Kind()
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package spectypes
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSQLTypeName(t *testing.T) {
|
||||
cases := []struct {
|
||||
val any
|
||||
want string
|
||||
}{
|
||||
{SqlVector{}, "vector"},
|
||||
{SqlHalfVector{}, "halfvec"},
|
||||
{SqlSparseVector{}, "sparsevec"},
|
||||
{SqlBitVector{}, "bit"},
|
||||
{SqlGeometry{}, "geometry"},
|
||||
{SqlGeography{}, "geography"},
|
||||
{SqlJSONB{}, "jsonb"},
|
||||
{SqlStringArray{}, "text[]"},
|
||||
{SqlString{}, "text"},
|
||||
{SqlInt64{}, "bigint"},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, ok := SQLTypeName(reflect.TypeOf(c.val))
|
||||
if !ok || got != c.want {
|
||||
t.Errorf("SQLTypeName(%T) = %q, %v; want %q", c.val, got, ok, c.want)
|
||||
}
|
||||
}
|
||||
|
||||
// Pointer is unwrapped.
|
||||
if got, ok := SQLTypeName(reflect.TypeOf(&SqlGeometry{})); !ok || got != "geometry" {
|
||||
t.Errorf("pointer: got %q, %v", got, ok)
|
||||
}
|
||||
|
||||
// Non-spectypes type.
|
||||
if _, ok := SQLTypeName(reflect.TypeOf("")); ok {
|
||||
t.Error("expected false for string")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsSpatialType(t *testing.T) {
|
||||
if !IsSpatialType(reflect.TypeOf(SqlGeometry{})) {
|
||||
t.Error("SqlGeometry should be spatial")
|
||||
}
|
||||
if !IsSpatialType(reflect.TypeOf(SqlGeography{})) {
|
||||
t.Error("SqlGeography should be spatial")
|
||||
}
|
||||
if IsSpatialType(reflect.TypeOf(SqlVector{})) {
|
||||
t.Error("SqlVector should not be spatial")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsJSONType(t *testing.T) {
|
||||
if !IsJSONType(reflect.TypeOf(SqlJSONB{})) {
|
||||
t.Error("SqlJSONB should be a JSON type")
|
||||
}
|
||||
if !IsJSONType(reflect.TypeOf(&SqlJSONB{})) {
|
||||
t.Error("*SqlJSONB should be a JSON type (pointer unwrapped)")
|
||||
}
|
||||
for _, v := range []any{SqlGeometry{}, SqlVector{}, SqlString{}, SqlStringArray{}, ""} {
|
||||
if IsJSONType(reflect.TypeOf(v)) {
|
||||
t.Errorf("%T should not be a JSON type", v)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsVectorType(t *testing.T) {
|
||||
for _, v := range []any{SqlVector{}, SqlHalfVector{}, SqlSparseVector{}} {
|
||||
if !IsVectorType(reflect.TypeOf(v)) {
|
||||
t.Errorf("%T should be vector", v)
|
||||
}
|
||||
}
|
||||
if IsVectorType(reflect.TypeOf(SqlBitVector{})) {
|
||||
t.Error("SqlBitVector is not a vector type")
|
||||
}
|
||||
if IsVectorType(reflect.TypeOf(SqlGeometry{})) {
|
||||
t.Error("SqlGeometry is not a vector type")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,697 @@
|
||||
package spectypes
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Minimal self-contained EWKB (PostGIS extended WKB) <-> GeoJSON / WKT codec.
|
||||
// Supports 2D and 3D (Z) geometries of type Point, LineString, Polygon,
|
||||
// MultiPoint, MultiLineString, MultiPolygon and GeometryCollection. The M
|
||||
// dimension is parsed but dropped (GeoJSON has no M). SRID is tracked separately
|
||||
// from the GeoJSON payload (GeoJSON assumes CRS84 / EPSG:4326).
|
||||
|
||||
// EWKB type flag bits (PostGIS).
|
||||
const (
|
||||
ewkbZ = 0x80000000
|
||||
ewkbM = 0x40000000
|
||||
ewkbSRID = 0x20000000
|
||||
)
|
||||
|
||||
// geom is the intermediate geometry representation used by the codec.
|
||||
//
|
||||
// Point -> coord ([]float64, len 2 or 3)
|
||||
// LineString/MultiPt -> line ([][]float64)
|
||||
// Polygon/MultiLine -> poly ([][][]float64)
|
||||
// MultiPolygon -> multi ([][][][]float64)
|
||||
// GeometryCollection -> geoms ([]geom)
|
||||
type geom struct {
|
||||
typ string
|
||||
coord []float64
|
||||
line [][]float64
|
||||
poly [][][]float64
|
||||
multi [][][][]float64
|
||||
geoms []geom
|
||||
}
|
||||
|
||||
// wkbReader consumes an EWKB byte stream.
|
||||
type wkbReader struct {
|
||||
buf []byte
|
||||
pos int
|
||||
}
|
||||
|
||||
func (r *wkbReader) readByte() (byte, error) {
|
||||
if r.pos >= len(r.buf) {
|
||||
return 0, fmt.Errorf("wkb: unexpected end of input")
|
||||
}
|
||||
b := r.buf[r.pos]
|
||||
r.pos++
|
||||
return b, nil
|
||||
}
|
||||
|
||||
func (r *wkbReader) readUint32(bo binary.ByteOrder) (uint32, error) {
|
||||
if r.pos+4 > len(r.buf) {
|
||||
return 0, fmt.Errorf("wkb: unexpected end of input")
|
||||
}
|
||||
v := bo.Uint32(r.buf[r.pos:])
|
||||
r.pos += 4
|
||||
return v, nil
|
||||
}
|
||||
|
||||
func (r *wkbReader) readFloat64(bo binary.ByteOrder) (float64, error) {
|
||||
if r.pos+8 > len(r.buf) {
|
||||
return 0, fmt.Errorf("wkb: unexpected end of input")
|
||||
}
|
||||
v := math.Float64frombits(bo.Uint64(r.buf[r.pos:]))
|
||||
r.pos += 8
|
||||
return v, nil
|
||||
}
|
||||
|
||||
// DecodeEWKBHex decodes a PostGIS hex-EWKB string (the default text
|
||||
// representation of a geometry column) into a GeoJSON geometry object and its
|
||||
// SRID. An SRID of 0 means "unspecified".
|
||||
func DecodeEWKBHex(s string) (geojson []byte, srid int, err error) {
|
||||
s = strings.TrimSpace(s)
|
||||
raw, err := hex.DecodeString(s)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("wkb: invalid hex: %w", err)
|
||||
}
|
||||
return DecodeEWKB(raw)
|
||||
}
|
||||
|
||||
// DecodeEWKB decodes raw PostGIS EWKB bytes into a GeoJSON geometry object and
|
||||
// its SRID.
|
||||
func DecodeEWKB(raw []byte) (geojson []byte, srid int, err error) {
|
||||
r := &wkbReader{buf: raw}
|
||||
g, sr, err := readGeom(r)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
out, err := geomToGeoJSON(g)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return out, sr, nil
|
||||
}
|
||||
|
||||
func readGeom(r *wkbReader) (geom, int, error) {
|
||||
order, err := r.readByte()
|
||||
if err != nil {
|
||||
return geom{}, 0, err
|
||||
}
|
||||
var bo binary.ByteOrder
|
||||
switch order {
|
||||
case 0:
|
||||
bo = binary.BigEndian
|
||||
case 1:
|
||||
bo = binary.LittleEndian
|
||||
default:
|
||||
return geom{}, 0, fmt.Errorf("wkb: invalid byte order %d", order)
|
||||
}
|
||||
|
||||
rawType, err := r.readUint32(bo)
|
||||
if err != nil {
|
||||
return geom{}, 0, err
|
||||
}
|
||||
hasZ := rawType&ewkbZ != 0
|
||||
hasM := rawType&ewkbM != 0
|
||||
hasSRID := rawType&ewkbSRID != 0
|
||||
baseType := rawType & 0xff
|
||||
|
||||
srid := 0
|
||||
if hasSRID {
|
||||
s, err := r.readUint32(bo)
|
||||
if err != nil {
|
||||
return geom{}, 0, err
|
||||
}
|
||||
srid = int(s)
|
||||
}
|
||||
|
||||
dims := 2
|
||||
if hasZ {
|
||||
dims = 3
|
||||
}
|
||||
// M is consumed but not retained.
|
||||
stride := dims
|
||||
if hasM {
|
||||
stride++
|
||||
}
|
||||
|
||||
readCoord := func() ([]float64, error) {
|
||||
c := make([]float64, 0, dims)
|
||||
for i := 0; i < stride; i++ {
|
||||
v, err := r.readFloat64(bo)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if i < dims {
|
||||
c = append(c, v)
|
||||
}
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
readLine := func() ([][]float64, error) {
|
||||
n, err := r.readUint32(bo)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pts := make([][]float64, n)
|
||||
for i := range pts {
|
||||
pts[i], err = readCoord()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return pts, nil
|
||||
}
|
||||
readPoly := func() ([][][]float64, error) {
|
||||
n, err := r.readUint32(bo)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rings := make([][][]float64, n)
|
||||
for i := range rings {
|
||||
rings[i], err = readLine()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return rings, nil
|
||||
}
|
||||
|
||||
switch baseType {
|
||||
case 1: // Point
|
||||
c, err := readCoord()
|
||||
if err != nil {
|
||||
return geom{}, 0, err
|
||||
}
|
||||
return geom{typ: "Point", coord: c}, srid, nil
|
||||
case 2: // LineString
|
||||
l, err := readLine()
|
||||
if err != nil {
|
||||
return geom{}, 0, err
|
||||
}
|
||||
return geom{typ: "LineString", line: l}, srid, nil
|
||||
case 3: // Polygon
|
||||
p, err := readPoly()
|
||||
if err != nil {
|
||||
return geom{}, 0, err
|
||||
}
|
||||
return geom{typ: "Polygon", poly: p}, srid, nil
|
||||
case 4, 5, 6: // Multi*
|
||||
n, err := r.readUint32(bo)
|
||||
if err != nil {
|
||||
return geom{}, 0, err
|
||||
}
|
||||
parts := make([]geom, n)
|
||||
for i := range parts {
|
||||
sub, _, err := readGeom(r)
|
||||
if err != nil {
|
||||
return geom{}, 0, err
|
||||
}
|
||||
parts[i] = sub
|
||||
}
|
||||
switch baseType {
|
||||
case 4:
|
||||
pts := make([][]float64, len(parts))
|
||||
for i := range parts {
|
||||
pts[i] = parts[i].coord
|
||||
}
|
||||
return geom{typ: "MultiPoint", line: pts}, srid, nil
|
||||
case 5:
|
||||
lines := make([][][]float64, len(parts))
|
||||
for i := range parts {
|
||||
lines[i] = parts[i].line
|
||||
}
|
||||
return geom{typ: "MultiLineString", poly: lines}, srid, nil
|
||||
default:
|
||||
polys := make([][][][]float64, len(parts))
|
||||
for i := range parts {
|
||||
polys[i] = parts[i].poly
|
||||
}
|
||||
return geom{typ: "MultiPolygon", multi: polys}, srid, nil
|
||||
}
|
||||
case 7: // GeometryCollection
|
||||
n, err := r.readUint32(bo)
|
||||
if err != nil {
|
||||
return geom{}, 0, err
|
||||
}
|
||||
parts := make([]geom, n)
|
||||
for i := range parts {
|
||||
sub, _, err := readGeom(r)
|
||||
if err != nil {
|
||||
return geom{}, 0, err
|
||||
}
|
||||
parts[i] = sub
|
||||
}
|
||||
return geom{typ: "GeometryCollection", geoms: parts}, srid, nil
|
||||
default:
|
||||
return geom{}, 0, fmt.Errorf("wkb: unsupported geometry type %d", baseType)
|
||||
}
|
||||
}
|
||||
|
||||
// ── GeoJSON ──────────────────────────────────────────────────────────────────
|
||||
|
||||
type geoJSON struct {
|
||||
Type string `json:"type"`
|
||||
Coordinates json.RawMessage `json:"coordinates,omitempty"`
|
||||
Geometries []geoJSON `json:"geometries,omitempty"`
|
||||
}
|
||||
|
||||
func geomToGeoJSON(g geom) ([]byte, error) {
|
||||
gj, err := geomToGeoJSONStruct(g)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return json.Marshal(gj)
|
||||
}
|
||||
|
||||
func geomToGeoJSONStruct(g geom) (geoJSON, error) {
|
||||
var coords any
|
||||
switch g.typ {
|
||||
case "Point":
|
||||
coords = g.coord
|
||||
case "LineString", "MultiPoint":
|
||||
coords = g.line
|
||||
case "Polygon", "MultiLineString":
|
||||
coords = g.poly
|
||||
case "MultiPolygon":
|
||||
coords = g.multi
|
||||
case "GeometryCollection":
|
||||
subs := make([]geoJSON, len(g.geoms))
|
||||
for i := range g.geoms {
|
||||
s, err := geomToGeoJSONStruct(g.geoms[i])
|
||||
if err != nil {
|
||||
return geoJSON{}, err
|
||||
}
|
||||
subs[i] = s
|
||||
}
|
||||
return geoJSON{Type: "GeometryCollection", Geometries: subs}, nil
|
||||
default:
|
||||
return geoJSON{}, fmt.Errorf("wkb: cannot encode geometry type %q", g.typ)
|
||||
}
|
||||
rc, err := json.Marshal(coords)
|
||||
if err != nil {
|
||||
return geoJSON{}, err
|
||||
}
|
||||
return geoJSON{Type: g.typ, Coordinates: rc}, nil
|
||||
}
|
||||
|
||||
func geoJSONToGeom(data []byte) (geom, error) {
|
||||
var gj geoJSON
|
||||
if err := json.Unmarshal(data, &gj); err != nil {
|
||||
return geom{}, fmt.Errorf("geojson: %w", err)
|
||||
}
|
||||
return geoJSONStructToGeom(gj)
|
||||
}
|
||||
|
||||
func geoJSONStructToGeom(gj geoJSON) (geom, error) {
|
||||
switch gj.Type {
|
||||
case "Point":
|
||||
var c []float64
|
||||
if err := json.Unmarshal(gj.Coordinates, &c); err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
return geom{typ: "Point", coord: c}, nil
|
||||
case "LineString", "MultiPoint":
|
||||
var l [][]float64
|
||||
if err := json.Unmarshal(gj.Coordinates, &l); err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
return geom{typ: gj.Type, line: l}, nil
|
||||
case "Polygon", "MultiLineString":
|
||||
var p [][][]float64
|
||||
if err := json.Unmarshal(gj.Coordinates, &p); err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
return geom{typ: gj.Type, poly: p}, nil
|
||||
case "MultiPolygon":
|
||||
var m [][][][]float64
|
||||
if err := json.Unmarshal(gj.Coordinates, &m); err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
return geom{typ: gj.Type, multi: m}, nil
|
||||
case "GeometryCollection":
|
||||
subs := make([]geom, len(gj.Geometries))
|
||||
for i, s := range gj.Geometries {
|
||||
g, err := geoJSONStructToGeom(s)
|
||||
if err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
subs[i] = g
|
||||
}
|
||||
return geom{typ: "GeometryCollection", geoms: subs}, nil
|
||||
default:
|
||||
return geom{}, fmt.Errorf("geojson: unsupported type %q", gj.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// ── WKT ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
// GeoJSONToWKT converts a GeoJSON geometry object to its WKT representation.
|
||||
func GeoJSONToWKT(geojson []byte) (string, error) {
|
||||
g, err := geoJSONToGeom(geojson)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return geomToWKT(g)
|
||||
}
|
||||
|
||||
func fmtNum(f float64) string {
|
||||
return strconv.FormatFloat(f, 'f', -1, 64)
|
||||
}
|
||||
|
||||
func coordWKT(c []float64) string {
|
||||
parts := make([]string, len(c))
|
||||
for i, v := range c {
|
||||
parts[i] = fmtNum(v)
|
||||
}
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
func lineWKT(pts [][]float64) string {
|
||||
parts := make([]string, len(pts))
|
||||
for i, p := range pts {
|
||||
parts[i] = coordWKT(p)
|
||||
}
|
||||
return "(" + strings.Join(parts, ", ") + ")"
|
||||
}
|
||||
|
||||
func polyWKT(rings [][][]float64) string {
|
||||
parts := make([]string, len(rings))
|
||||
for i, r := range rings {
|
||||
parts[i] = lineWKT(r)
|
||||
}
|
||||
return "(" + strings.Join(parts, ", ") + ")"
|
||||
}
|
||||
|
||||
// ── WKT parsing ─────────────────────────────────────────────────────────────
|
||||
|
||||
// parseWKT parses a subset of WKT (2D/3D, no M) into the intermediate geom.
|
||||
func parseWKT(s string) (geom, error) {
|
||||
p := &wktParser{s: s}
|
||||
p.skipSpace()
|
||||
g, err := p.parseGeom()
|
||||
if err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
p.skipSpace()
|
||||
if p.pos != len(p.s) {
|
||||
return geom{}, fmt.Errorf("wkt: trailing input %q", p.s[p.pos:])
|
||||
}
|
||||
return g, nil
|
||||
}
|
||||
|
||||
type wktParser struct {
|
||||
s string
|
||||
pos int
|
||||
}
|
||||
|
||||
func (p *wktParser) skipSpace() {
|
||||
for p.pos < len(p.s) && (p.s[p.pos] == ' ' || p.s[p.pos] == '\t' || p.s[p.pos] == '\n' || p.s[p.pos] == '\r') {
|
||||
p.pos++
|
||||
}
|
||||
}
|
||||
|
||||
func (p *wktParser) parseGeom() (geom, error) {
|
||||
p.skipSpace()
|
||||
start := p.pos
|
||||
for p.pos < len(p.s) && (p.s[p.pos] >= 'A' && p.s[p.pos] <= 'Z' || p.s[p.pos] >= 'a' && p.s[p.pos] <= 'z') {
|
||||
p.pos++
|
||||
}
|
||||
kw := strings.ToUpper(p.s[start:p.pos])
|
||||
p.skipSpace()
|
||||
// Optional Z / M / ZM dimension tag — coordinates carry their own arity.
|
||||
if p.pos < len(p.s) && (p.s[p.pos] == 'Z' || p.s[p.pos] == 'M' || p.s[p.pos] == 'z' || p.s[p.pos] == 'm') {
|
||||
for p.pos < len(p.s) && p.s[p.pos] != '(' {
|
||||
p.pos++
|
||||
}
|
||||
}
|
||||
p.skipSpace()
|
||||
|
||||
switch kw {
|
||||
case "POINT":
|
||||
pts, err := p.parseCoordList()
|
||||
if err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
if len(pts) != 1 {
|
||||
return geom{}, fmt.Errorf("wkt: POINT needs exactly one coordinate")
|
||||
}
|
||||
return geom{typ: "Point", coord: pts[0]}, nil
|
||||
case "LINESTRING":
|
||||
pts, err := p.parseCoordList()
|
||||
if err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
return geom{typ: "LineString", line: pts}, nil
|
||||
case "MULTIPOINT":
|
||||
pts, err := p.parseMultiPoint()
|
||||
if err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
return geom{typ: "MultiPoint", line: pts}, nil
|
||||
case "POLYGON":
|
||||
rings, err := p.parseRingList()
|
||||
if err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
return geom{typ: "Polygon", poly: rings}, nil
|
||||
case "MULTILINESTRING":
|
||||
lines, err := p.parseRingList()
|
||||
if err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
return geom{typ: "MultiLineString", poly: lines}, nil
|
||||
case "MULTIPOLYGON":
|
||||
polys, err := p.parsePolyList()
|
||||
if err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
return geom{typ: "MultiPolygon", multi: polys}, nil
|
||||
case "GEOMETRYCOLLECTION":
|
||||
return p.parseCollection()
|
||||
default:
|
||||
return geom{}, fmt.Errorf("wkt: unsupported geometry %q", kw)
|
||||
}
|
||||
}
|
||||
|
||||
func (p *wktParser) expect(c byte) error {
|
||||
p.skipSpace()
|
||||
if p.pos >= len(p.s) || p.s[p.pos] != c {
|
||||
return fmt.Errorf("wkt: expected %q at offset %d", string(c), p.pos)
|
||||
}
|
||||
p.pos++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *wktParser) peek() byte {
|
||||
p.skipSpace()
|
||||
if p.pos >= len(p.s) {
|
||||
return 0
|
||||
}
|
||||
return p.s[p.pos]
|
||||
}
|
||||
|
||||
// parseCoordList parses `(x y[, x y]...)`.
|
||||
func (p *wktParser) parseCoordList() ([][]float64, error) {
|
||||
if err := p.expect('('); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var out [][]float64
|
||||
for {
|
||||
c, err := p.parseCoord()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, c)
|
||||
if p.peek() == ',' {
|
||||
p.pos++
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
if err := p.expect(')'); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (p *wktParser) parseCoord() ([]float64, error) {
|
||||
p.skipSpace()
|
||||
// Some MULTIPOINT forms wrap each coord in parentheses.
|
||||
wrapped := false
|
||||
if p.peek() == '(' {
|
||||
p.pos++
|
||||
wrapped = true
|
||||
}
|
||||
var nums []float64
|
||||
for {
|
||||
p.skipSpace()
|
||||
start := p.pos
|
||||
for p.pos < len(p.s) {
|
||||
ch := p.s[p.pos]
|
||||
if ch == '-' || ch == '+' || ch == '.' || ch == 'e' || ch == 'E' || (ch >= '0' && ch <= '9') {
|
||||
p.pos++
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
if p.pos == start {
|
||||
break
|
||||
}
|
||||
f, err := strconv.ParseFloat(p.s[start:p.pos], 64)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("wkt: bad number %q", p.s[start:p.pos])
|
||||
}
|
||||
nums = append(nums, f)
|
||||
p.skipSpace()
|
||||
if p.pos < len(p.s) && p.s[p.pos] == ' ' {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if wrapped {
|
||||
if err := p.expect(')'); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if len(nums) < 2 {
|
||||
return nil, fmt.Errorf("wkt: coordinate needs at least 2 numbers")
|
||||
}
|
||||
if len(nums) > 3 {
|
||||
nums = nums[:3]
|
||||
}
|
||||
return nums, nil
|
||||
}
|
||||
|
||||
func (p *wktParser) parseMultiPoint() ([][]float64, error) {
|
||||
if err := p.expect('('); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var out [][]float64
|
||||
for {
|
||||
c, err := p.parseCoord()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, c)
|
||||
if p.peek() == ',' {
|
||||
p.pos++
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
if err := p.expect(')'); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// parseRingList parses `((x y, ...), (...))`.
|
||||
func (p *wktParser) parseRingList() ([][][]float64, error) {
|
||||
if err := p.expect('('); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var out [][][]float64
|
||||
for {
|
||||
ring, err := p.parseCoordList()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, ring)
|
||||
if p.peek() == ',' {
|
||||
p.pos++
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
if err := p.expect(')'); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// parsePolyList parses `(((...)), ((...)))`.
|
||||
func (p *wktParser) parsePolyList() ([][][][]float64, error) {
|
||||
if err := p.expect('('); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var out [][][][]float64
|
||||
for {
|
||||
poly, err := p.parseRingList()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, poly)
|
||||
if p.peek() == ',' {
|
||||
p.pos++
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
if err := p.expect(')'); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (p *wktParser) parseCollection() (geom, error) {
|
||||
if err := p.expect('('); err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
var subs []geom
|
||||
for {
|
||||
g, err := p.parseGeom()
|
||||
if err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
subs = append(subs, g)
|
||||
if p.peek() == ',' {
|
||||
p.pos++
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
if err := p.expect(')'); err != nil {
|
||||
return geom{}, err
|
||||
}
|
||||
return geom{typ: "GeometryCollection", geoms: subs}, nil
|
||||
}
|
||||
|
||||
func geomToWKT(g geom) (string, error) {
|
||||
switch g.typ {
|
||||
case "Point":
|
||||
return "POINT (" + coordWKT(g.coord) + ")", nil
|
||||
case "LineString":
|
||||
return "LINESTRING " + lineWKT(g.line), nil
|
||||
case "MultiPoint":
|
||||
return "MULTIPOINT " + lineWKT(g.line), nil
|
||||
case "Polygon":
|
||||
return "POLYGON " + polyWKT(g.poly), nil
|
||||
case "MultiLineString":
|
||||
return "MULTILINESTRING " + polyWKT(g.poly), nil
|
||||
case "MultiPolygon":
|
||||
parts := make([]string, len(g.multi))
|
||||
for i, p := range g.multi {
|
||||
parts[i] = polyWKT(p)
|
||||
}
|
||||
return "MULTIPOLYGON (" + strings.Join(parts, ", ") + ")", nil
|
||||
case "GeometryCollection":
|
||||
parts := make([]string, len(g.geoms))
|
||||
for i := range g.geoms {
|
||||
w, err := geomToWKT(g.geoms[i])
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
parts[i] = w
|
||||
}
|
||||
return "GEOMETRYCOLLECTION (" + strings.Join(parts, ", ") + ")", nil
|
||||
default:
|
||||
return "", fmt.Errorf("wkt: cannot encode geometry type %q", g.typ)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package spectypes
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDecodeEWKBHex(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
hex string
|
||||
wantSRID int
|
||||
wantJSON string
|
||||
}{
|
||||
{
|
||||
// SRID=4326;POINT(1 2)
|
||||
name: "point with srid",
|
||||
hex: "0101000020E6100000000000000000F03F0000000000000040",
|
||||
wantSRID: 4326,
|
||||
wantJSON: `{"type":"Point","coordinates":[1,2]}`,
|
||||
},
|
||||
{
|
||||
// POINT(1 2) no SRID, little endian
|
||||
name: "point no srid",
|
||||
hex: "0101000000000000000000F03F0000000000000040",
|
||||
wantSRID: 0,
|
||||
wantJSON: `{"type":"Point","coordinates":[1,2]}`,
|
||||
},
|
||||
{
|
||||
// SRID=4326;LINESTRING(0 0, 1 1, 2 2)
|
||||
name: "linestring",
|
||||
hex: "0102000020E610000003000000000000000000000000000000000000000000000000" +
|
||||
"00F03F000000000000F03F00000000000000400000000000000040",
|
||||
wantSRID: 4326,
|
||||
wantJSON: `{"type":"LineString","coordinates":[[0,0],[1,1],[2,2]]}`,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
gj, srid, err := DecodeEWKBHex(tt.hex)
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeEWKBHex: %v", err)
|
||||
}
|
||||
if srid != tt.wantSRID {
|
||||
t.Errorf("srid = %d, want %d", srid, tt.wantSRID)
|
||||
}
|
||||
if !jsonEqual(t, gj, tt.wantJSON) {
|
||||
t.Errorf("geojson = %s, want %s", gj, tt.wantJSON)
|
||||
}
|
||||
// Round-trip through WKT parser.
|
||||
wkt, err := GeoJSONToWKT(gj)
|
||||
if err != nil {
|
||||
t.Fatalf("GeoJSONToWKT: %v", err)
|
||||
}
|
||||
gj2, err := wktToGeoJSON(wkt)
|
||||
if err != nil {
|
||||
t.Fatalf("wktToGeoJSON(%q): %v", wkt, err)
|
||||
}
|
||||
if !jsonEqual(t, gj2, tt.wantJSON) {
|
||||
t.Errorf("round-trip geojson = %s, want %s", gj2, tt.wantJSON)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWKTPolygon(t *testing.T) {
|
||||
src := `POLYGON ((0 0, 4 0, 4 4, 0 4, 0 0), (1 1, 2 1, 2 2, 1 2, 1 1))`
|
||||
gj, err := wktToGeoJSON(src)
|
||||
if err != nil {
|
||||
t.Fatalf("wktToGeoJSON: %v", err)
|
||||
}
|
||||
want := `{"type":"Polygon","coordinates":[[[0,0],[4,0],[4,4],[0,4],[0,0]],[[1,1],[2,1],[2,2],[1,2],[1,1]]]}`
|
||||
if !jsonEqual(t, gj, want) {
|
||||
t.Errorf("geojson = %s, want %s", gj, want)
|
||||
}
|
||||
}
|
||||
|
||||
func jsonEqual(t *testing.T, got []byte, want string) bool {
|
||||
t.Helper()
|
||||
var a, b any
|
||||
if err := json.Unmarshal(got, &a); err != nil {
|
||||
t.Fatalf("unmarshal got %s: %v", got, err)
|
||||
}
|
||||
if err := json.Unmarshal([]byte(want), &b); err != nil {
|
||||
t.Fatalf("unmarshal want %s: %v", want, err)
|
||||
}
|
||||
ab, _ := json.Marshal(a)
|
||||
bb, _ := json.Marshal(b)
|
||||
return string(ab) == string(bb)
|
||||
}
|
||||
@@ -221,7 +221,7 @@ func (h *Handler) handleRequest(conn *Connection, msg *Message) {
|
||||
// handleRead processes a read operation
|
||||
func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContext) {
|
||||
// Execute before hook
|
||||
if err := h.hooks.Execute(BeforeRead, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil {
|
||||
logger.Error("[WebSocketSpec] BeforeRead hook failed: %v", err)
|
||||
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
||||
_ = conn.SendJSON(errResp)
|
||||
@@ -273,7 +273,7 @@ func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContex
|
||||
// handleCreate processes a create operation
|
||||
func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookContext) {
|
||||
// Execute before hook
|
||||
if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
|
||||
logger.Error("[WebSocketSpec] BeforeCreate hook failed: %v", err)
|
||||
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
||||
_ = conn.SendJSON(errResp)
|
||||
@@ -311,7 +311,7 @@ func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookCont
|
||||
// handleUpdate processes an update operation
|
||||
func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookContext) {
|
||||
// Execute before hook
|
||||
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil {
|
||||
logger.Error("[WebSocketSpec] BeforeUpdate hook failed: %v", err)
|
||||
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
||||
_ = conn.SendJSON(errResp)
|
||||
@@ -349,7 +349,7 @@ func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookCont
|
||||
// handleDelete processes a delete operation
|
||||
func (h *Handler) handleDelete(conn *Connection, msg *Message, hookCtx *HookContext) {
|
||||
// Execute before hook
|
||||
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil {
|
||||
logger.Error("[WebSocketSpec] BeforeDelete hook failed: %v", err)
|
||||
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
||||
_ = conn.SendJSON(errResp)
|
||||
@@ -564,7 +564,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) {
|
||||
|
||||
// Apply columns
|
||||
if hookCtx.Options != nil && len(hookCtx.Options.Columns) > 0 {
|
||||
query = query.Column(hookCtx.Options.Columns...)
|
||||
query = common.ApplySelectColumns(query, hookCtx.Model, "", hookCtx.Options.Columns)
|
||||
}
|
||||
|
||||
// Apply preloads (simplified for now)
|
||||
@@ -574,6 +574,18 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// Execute BeforeScan hooks - pass query so hooks (e.g. row-level security) can modify it
|
||||
if hookCtx.Metadata == nil {
|
||||
hookCtx.Metadata = make(map[string]interface{})
|
||||
}
|
||||
hookCtx.Metadata["query"] = query
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
|
||||
return nil, fmt.Errorf("BeforeScan hook failed: %w", err)
|
||||
}
|
||||
if modifiedQuery, ok := hookCtx.Metadata["query"].(common.SelectQuery); ok {
|
||||
query = modifiedQuery
|
||||
}
|
||||
|
||||
// Execute query
|
||||
if err := query.ScanModel(hookCtx.Context); err != nil {
|
||||
return nil, fmt.Errorf("failed to read record: %w", err)
|
||||
@@ -594,7 +606,7 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
// Apply options (simplified implementation)
|
||||
if hookCtx.Options != nil {
|
||||
// Apply filters with OR grouping support
|
||||
query = h.applyFilters(query, hookCtx.Options.Filters)
|
||||
query = h.applyFilters(query, hookCtx.Options.Filters, hookCtx.Model)
|
||||
|
||||
// Apply sorting
|
||||
for _, sort := range hookCtx.Options.Sort {
|
||||
@@ -602,6 +614,10 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
if sort.Direction == "desc" {
|
||||
direction = "DESC"
|
||||
}
|
||||
if expr, jargs, _, ok := common.ResolveJSONColumnExpr(hookCtx.Model, "", sort.Column); ok {
|
||||
query = query.OrderExpr(fmt.Sprintf("%s %s", expr, direction), jargs...)
|
||||
continue
|
||||
}
|
||||
query = query.Order(fmt.Sprintf("%s %s", sort.Column, direction))
|
||||
}
|
||||
|
||||
@@ -620,10 +636,22 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
|
||||
// Apply columns
|
||||
if len(hookCtx.Options.Columns) > 0 {
|
||||
query = query.Column(hookCtx.Options.Columns...)
|
||||
query = common.ApplySelectColumns(query, hookCtx.Model, "", hookCtx.Options.Columns)
|
||||
}
|
||||
}
|
||||
|
||||
// Execute BeforeScan hooks - pass query so hooks (e.g. row-level security) can modify it
|
||||
if hookCtx.Metadata == nil {
|
||||
hookCtx.Metadata = make(map[string]interface{})
|
||||
}
|
||||
hookCtx.Metadata["query"] = query
|
||||
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
|
||||
return nil, nil, fmt.Errorf("BeforeScan hook failed: %w", err)
|
||||
}
|
||||
if modifiedQuery, ok := hookCtx.Metadata["query"].(common.SelectQuery); ok {
|
||||
query = modifiedQuery
|
||||
}
|
||||
|
||||
// Execute query
|
||||
if err := query.ScanModel(hookCtx.Context); err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to read records: %w", err)
|
||||
@@ -641,7 +669,7 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
countQuery := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
||||
if hookCtx.Options != nil {
|
||||
for _, filter := range hookCtx.Options.Filters {
|
||||
cond, args := h.buildFilterCondition(filter)
|
||||
cond, args := h.buildFilterCondition(filter, hookCtx.Model)
|
||||
if cond != "" {
|
||||
countQuery = countQuery.Where(cond, args...)
|
||||
}
|
||||
@@ -752,7 +780,7 @@ func (h *Handler) getMetadata(schema, entity string, model interface{}) map[stri
|
||||
// getOperatorSQL converts filter operator to SQL operator
|
||||
// applyFilters applies all filters with proper grouping for OR logic
|
||||
// Groups consecutive OR filters together to ensure proper query precedence
|
||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
|
||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
||||
if len(filters) == 0 {
|
||||
return query
|
||||
}
|
||||
@@ -772,11 +800,11 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
||||
}
|
||||
|
||||
// Apply the OR group as a single grouped WHERE clause
|
||||
query = h.applyFilterGroup(query, orGroup)
|
||||
query = h.applyFilterGroup(query, orGroup, model)
|
||||
i = j
|
||||
} else {
|
||||
// Single filter with AND logic (or first filter)
|
||||
condition, args := h.buildFilterCondition(filters[i])
|
||||
condition, args := h.buildFilterCondition(filters[i], model)
|
||||
if condition != "" {
|
||||
query = query.Where(condition, args...)
|
||||
}
|
||||
@@ -789,7 +817,7 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
||||
|
||||
// applyFilterGroup applies a group of filters that should be OR'd together
|
||||
// Always wraps them in parentheses and applies as a single WHERE clause
|
||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
|
||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
||||
if len(filters) == 0 {
|
||||
return query
|
||||
}
|
||||
@@ -799,7 +827,7 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
||||
var args []interface{}
|
||||
|
||||
for _, filter := range filters {
|
||||
condition, filterArgs := h.buildFilterCondition(filter)
|
||||
condition, filterArgs := h.buildFilterCondition(filter, model)
|
||||
if condition != "" {
|
||||
conditions = append(conditions, condition)
|
||||
args = append(args, filterArgs...)
|
||||
@@ -820,8 +848,14 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
||||
return query.Where(groupedCondition, args...)
|
||||
}
|
||||
|
||||
// buildFilterCondition builds a filter condition and returns it with args
|
||||
func (h *Handler) buildFilterCondition(filter common.FilterOption) (conditionString string, conditionArgs []interface{}) {
|
||||
// buildFilterCondition builds a filter condition and returns it with args.
|
||||
// model, when non-nil, lets JSON sub-field references (data->>'x', data#>>'{a,b}',
|
||||
// or the dotted data.x shorthand for a JSON column) resolve to a safe,
|
||||
// parameterised expression before the ordinary operator handling below.
|
||||
func (h *Handler) buildFilterCondition(filter common.FilterOption, model interface{}) (conditionString string, conditionArgs []interface{}) {
|
||||
if cond, jargs, ok := common.BuildJSONFilterCondition(model, "", filter.Column, filter.Operator, filter.Value); ok {
|
||||
return cond, jargs
|
||||
}
|
||||
if strings.EqualFold(filter.Operator, "in") {
|
||||
cond, args := common.BuildInCondition(filter.Column, filter.Value)
|
||||
return cond, args
|
||||
@@ -829,6 +863,11 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption) (conditionStr
|
||||
op := strings.ToLower(filter.Operator)
|
||||
if op == "like" || op == "ilike" {
|
||||
operatorSQL := h.getOperatorSQL(filter.Operator)
|
||||
// citext columns are already case-insensitive; casting to TEXT would
|
||||
// switch to case-sensitive matching and defeat a citext index.
|
||||
if reflection.IsCitextColumn(model, filter.Column) {
|
||||
return fmt.Sprintf("%s %s ?", filter.Column, operatorSQL), []interface{}{filter.Value}
|
||||
}
|
||||
return fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, operatorSQL), []interface{}{filter.Value}
|
||||
}
|
||||
operatorSQL := h.getOperatorSQL(filter.Operator)
|
||||
|
||||
@@ -35,6 +35,11 @@ const (
|
||||
// AfterDelete is called after a delete operation
|
||||
AfterDelete HookType = "after_delete"
|
||||
|
||||
// BeforeScan is called right before a read query is executed against the database,
|
||||
// after all filters/sort/pagination have been applied. Use this for row-level
|
||||
// security that needs to modify the query (stored in HookContext.Metadata["query"]).
|
||||
BeforeScan HookType = "before_scan"
|
||||
|
||||
// BeforeSubscribe is called before creating a subscription
|
||||
BeforeSubscribe HookType = "before_subscribe"
|
||||
// AfterSubscribe is called after creating a subscription
|
||||
@@ -54,6 +59,11 @@ const (
|
||||
BeforeDisconnect HookType = "before_disconnect"
|
||||
// AfterDisconnect is called after a connection is closed
|
||||
AfterDisconnect HookType = "after_disconnect"
|
||||
|
||||
// BeforeOp fires immediately before every SQL operation (read, create, update, delete).
|
||||
// Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each
|
||||
// individual SQL-operation hook point, so it runs once per statement executed.
|
||||
BeforeOp HookType = "before_op"
|
||||
)
|
||||
|
||||
// HookContext contains context information for hook execution
|
||||
@@ -197,6 +207,16 @@ func (hr *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
|
||||
// (BeforeRead, BeforeCreate, BeforeUpdate, or BeforeDelete). BeforeOp always runs first
|
||||
// so it can observe/veto every SQL operation regardless of type.
|
||||
func (hr *HookRegistry) ExecuteBeforeOp(hookType HookType, ctx *HookContext) error {
|
||||
if err := hr.Execute(BeforeOp, ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
return hr.Execute(hookType, ctx)
|
||||
}
|
||||
|
||||
// HasHooks checks if any hooks are registered for a hook type
|
||||
func (hr *HookRegistry) HasHooks(hookType HookType) bool {
|
||||
hooks, exists := hr.hooks[hookType]
|
||||
|
||||
@@ -27,25 +27,31 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
||||
return security.LoadSecurityRules(secCtx, securityList)
|
||||
})
|
||||
|
||||
// Hook 2: AfterRead - Apply column-level security (masking)
|
||||
// Hook 2: BeforeScan - Apply row-level security filters
|
||||
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
return security.ApplyRowSecurity(secCtx, securityList)
|
||||
})
|
||||
|
||||
// Hook 3: AfterRead - Apply column-level security (masking)
|
||||
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
return security.ApplyColumnSecurity(secCtx, securityList)
|
||||
})
|
||||
|
||||
// Hook 3 (Optional): Audit logging
|
||||
// Hook 4 (Optional): Audit logging
|
||||
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
return security.LogDataAccess(secCtx)
|
||||
})
|
||||
|
||||
// Hook 4: BeforeUpdate - enforce CanUpdate rule from context/registry
|
||||
// Hook 5: BeforeUpdate - enforce CanUpdate rule from context/registry
|
||||
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
return security.CheckModelUpdateAllowed(secCtx)
|
||||
})
|
||||
|
||||
// Hook 5: BeforeDelete - enforce CanDelete rule from context/registry
|
||||
// Hook 6: BeforeDelete - enforce CanDelete rule from context/registry
|
||||
handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error {
|
||||
secCtx := newSecurityContext(hookCtx)
|
||||
return security.CheckModelDeleteAllowed(secCtx)
|
||||
@@ -71,6 +77,17 @@ func (s *securityContext) GetUserID() (int, bool) {
|
||||
return security.GetUserID(s.ctx.Context)
|
||||
}
|
||||
|
||||
// GetUserRef returns an opaque user identifier for row security lookups.
|
||||
// It prefers the full *security.UserContext (so providers can read JWT claims,
|
||||
// e.g. a UUID subject) and falls back to the int user ID.
|
||||
func (s *securityContext) GetUserRef() (any, bool) {
|
||||
if userCtx, ok := security.GetUserContext(s.ctx.Context); ok {
|
||||
return userCtx, true
|
||||
}
|
||||
userID, ok := security.GetUserID(s.ctx.Context)
|
||||
return userID, ok
|
||||
}
|
||||
|
||||
func (s *securityContext) GetSchema() string {
|
||||
return s.ctx.Schema
|
||||
}
|
||||
|
||||
@@ -1,5 +1,12 @@
|
||||
# @warkypublic/resolvespec-js
|
||||
|
||||
## 1.0.2
|
||||
|
||||
### Patch Changes
|
||||
|
||||
- b587cbd: Forward custom ClientConfig headers on every ResolveSpec and HeaderSpec request. Merge headers case-insensitively and isolate cached clients by URL and effective headers, including authentication and tenant headers.
|
||||
- 7f8982f: fix: added headers and few fixes
|
||||
|
||||
## 1.0.1
|
||||
|
||||
### Patch Changes
|
||||
|
||||
@@ -28,7 +28,7 @@ import { ResolveSpecClient, getResolveSpecClient } from '@warkypublic/resolvespe
|
||||
// Class instantiation
|
||||
const client = new ResolveSpecClient({ baseUrl: 'http://localhost:3000', token: 'your-token' });
|
||||
|
||||
// Or singleton factory (returns cached instance per baseUrl)
|
||||
// Or singleton factory (returns cached instance per baseUrl and effective headers)
|
||||
const client = getResolveSpecClient({ baseUrl: 'http://localhost:3000', token: 'your-token' });
|
||||
|
||||
// Read with filters, sort, pagination
|
||||
@@ -211,3 +211,25 @@ pnpm run lint # eslint
|
||||
## License
|
||||
|
||||
MIT
|
||||
|
||||
### Custom HTTP headers
|
||||
|
||||
Both `ResolveSpecClient` and `HeaderSpecClient` (including their factory functions)
|
||||
accept `headers` in `ClientConfig` and send them on every HTTP request:
|
||||
|
||||
```typescript
|
||||
const client = new ResolveSpecClient({
|
||||
baseUrl: 'http://localhost:3000',
|
||||
token: 'your-token',
|
||||
headers: { 'X-Tenant': 'acme' },
|
||||
});
|
||||
```
|
||||
|
||||
Header names are merged case-insensitively. Custom headers override the default
|
||||
`Content-Type`; a supplied `token` overrides custom `Authorization`, and HeaderSpec
|
||||
query options override matching custom query headers. Without a token, custom
|
||||
`Authorization` is preserved. Configuration is copied at construction; create or
|
||||
retrieve a client with new configuration to change headers. Factory clients are
|
||||
cached by URL and effective headers, keeping different tenants and tokens separate.
|
||||
|
||||
Grid adapters must forward `dataSourceOptions.headers` to this `headers` option.
|
||||
|
||||
Vendored
+1
-1
File diff suppressed because one or more lines are too long
Vendored
+5
-366
@@ -1,366 +1,5 @@
|
||||
export declare interface APIError {
|
||||
code: string;
|
||||
message: string;
|
||||
details?: any;
|
||||
detail?: string;
|
||||
}
|
||||
|
||||
export declare interface APIResponse<T = any> {
|
||||
success: boolean;
|
||||
data: T;
|
||||
metadata?: Metadata;
|
||||
error?: APIError;
|
||||
}
|
||||
|
||||
/**
|
||||
* Build HTTP headers from Options, matching Go's restheadspec handler conventions.
|
||||
*
|
||||
* Header mapping:
|
||||
* - X-Select-Fields: comma-separated columns
|
||||
* - X-Not-Select-Fields: comma-separated omit_columns
|
||||
* - X-FieldFilter-{col}: exact match (eq)
|
||||
* - X-SearchOp-{operator}-{col}: AND filter
|
||||
* - X-SearchOr-{operator}-{col}: OR filter
|
||||
* - X-Sort: +col (asc), -col (desc)
|
||||
* - X-Limit, X-Offset: pagination
|
||||
* - X-Cursor-Forward, X-Cursor-Backward: cursor pagination
|
||||
* - X-Preload: RelationName:field1,field2 pipe-separated
|
||||
* - X-Fetch-RowNumber: row number fetch
|
||||
* - X-CQL-SEL-{col}: computed columns
|
||||
* - X-Custom-SQL-W: custom operators (AND)
|
||||
*/
|
||||
export declare function buildHeaders(options: Options): Record<string, string>;
|
||||
|
||||
export declare interface ClientConfig {
|
||||
baseUrl: string;
|
||||
token?: string;
|
||||
}
|
||||
|
||||
export declare interface Column {
|
||||
name: string;
|
||||
type: string;
|
||||
is_nullable: boolean;
|
||||
is_primary: boolean;
|
||||
is_unique: boolean;
|
||||
has_index: boolean;
|
||||
}
|
||||
|
||||
export declare interface ComputedColumn {
|
||||
name: string;
|
||||
expression: string;
|
||||
}
|
||||
|
||||
export declare type ConnectionState = 'connecting' | 'connected' | 'disconnecting' | 'disconnected' | 'reconnecting';
|
||||
|
||||
export declare interface CustomOperator {
|
||||
name: string;
|
||||
sql: string;
|
||||
}
|
||||
|
||||
/**
|
||||
* Decode a header value that may be base64 encoded with ZIP_ or __ prefix.
|
||||
*/
|
||||
export declare function decodeHeaderValue(value: string): string;
|
||||
|
||||
/**
|
||||
* Encode a value with base64 and ZIP_ prefix for complex header values.
|
||||
*/
|
||||
export declare function encodeHeaderValue(value: string): string;
|
||||
|
||||
export declare interface FilterOption {
|
||||
column: string;
|
||||
operator: Operator | string;
|
||||
value: any;
|
||||
logic_operator?: 'AND' | 'OR';
|
||||
}
|
||||
|
||||
export declare function getHeaderSpecClient(config: ClientConfig): HeaderSpecClient;
|
||||
|
||||
export declare function getResolveSpecClient(config: ClientConfig): ResolveSpecClient;
|
||||
|
||||
export declare function getWebSocketClient(config: WebSocketClientConfig): WebSocketClient;
|
||||
|
||||
/**
|
||||
* HeaderSpec REST client.
|
||||
* Sends query options via HTTP headers instead of request body, matching the Go restheadspec handler.
|
||||
*
|
||||
* HTTP methods: GET=read, POST=create, PUT=update, DELETE=delete
|
||||
*/
|
||||
export declare class HeaderSpecClient {
|
||||
private config;
|
||||
constructor(config: ClientConfig);
|
||||
private buildUrl;
|
||||
private baseHeaders;
|
||||
private fetchWithError;
|
||||
read<T = any>(schema: string, entity: string, id?: string, options?: Options): Promise<APIResponse<T>>;
|
||||
create<T = any>(schema: string, entity: string, data: any, options?: Options): Promise<APIResponse<T>>;
|
||||
update<T = any>(schema: string, entity: string, id: string, data: any, options?: Options): Promise<APIResponse<T>>;
|
||||
delete(schema: string, entity: string, id: string): Promise<APIResponse<void>>;
|
||||
}
|
||||
|
||||
export declare type MessageType = 'request' | 'response' | 'notification' | 'subscription' | 'error' | 'ping' | 'pong';
|
||||
|
||||
export declare interface Metadata {
|
||||
total: number;
|
||||
count: number;
|
||||
filtered: number;
|
||||
limit: number;
|
||||
offset: number;
|
||||
row_number?: number;
|
||||
}
|
||||
|
||||
export declare type Operation = 'read' | 'create' | 'update' | 'delete';
|
||||
|
||||
export declare type Operator = 'eq' | 'neq' | 'gt' | 'gte' | 'lt' | 'lte' | 'like' | 'ilike' | 'in' | 'contains' | 'startswith' | 'endswith' | 'between' | 'between_inclusive' | 'is_null' | 'is_not_null';
|
||||
|
||||
export declare interface Options {
|
||||
preload?: PreloadOption[];
|
||||
columns?: string[];
|
||||
omit_columns?: string[];
|
||||
filters?: FilterOption[];
|
||||
sort?: SortOption[];
|
||||
limit?: number;
|
||||
offset?: number;
|
||||
customOperators?: CustomOperator[];
|
||||
computedColumns?: ComputedColumn[];
|
||||
parameters?: Parameter[];
|
||||
cursor_forward?: string;
|
||||
cursor_backward?: string;
|
||||
fetch_row_number?: string;
|
||||
}
|
||||
|
||||
export declare interface Parameter {
|
||||
name: string;
|
||||
value: string;
|
||||
sequence?: number;
|
||||
}
|
||||
|
||||
export declare interface PreloadOption {
|
||||
relation: string;
|
||||
table_name?: string;
|
||||
columns?: string[];
|
||||
omit_columns?: string[];
|
||||
sort?: SortOption[];
|
||||
filters?: FilterOption[];
|
||||
where?: string;
|
||||
limit?: number;
|
||||
offset?: number;
|
||||
updatable?: boolean;
|
||||
computed_ql?: Record<string, string>;
|
||||
recursive?: boolean;
|
||||
primary_key?: string;
|
||||
related_key?: string;
|
||||
foreign_key?: string;
|
||||
recursive_child_key?: string;
|
||||
sql_joins?: string[];
|
||||
join_aliases?: string[];
|
||||
}
|
||||
|
||||
export declare interface RequestBody {
|
||||
operation: Operation;
|
||||
id?: number | string | string[];
|
||||
data?: any | any[];
|
||||
options?: Options;
|
||||
}
|
||||
|
||||
export declare class ResolveSpecClient {
|
||||
private config;
|
||||
constructor(config: ClientConfig);
|
||||
private buildUrl;
|
||||
private baseHeaders;
|
||||
private fetchWithError;
|
||||
getMetadata(schema: string, entity: string): Promise<APIResponse<TableMetadata>>;
|
||||
read<T = any>(schema: string, entity: string, id?: number | string | string[], options?: Options): Promise<APIResponse<T>>;
|
||||
create<T = any>(schema: string, entity: string, data: any | any[], options?: Options): Promise<APIResponse<T>>;
|
||||
update<T = any>(schema: string, entity: string, data: any | any[], id?: number | string | string[], options?: Options): Promise<APIResponse<T>>;
|
||||
delete(schema: string, entity: string, id: number | string): Promise<APIResponse<void>>;
|
||||
}
|
||||
|
||||
export declare type SortDirection = 'asc' | 'desc' | 'ASC' | 'DESC';
|
||||
|
||||
export declare interface SortOption {
|
||||
column: string;
|
||||
direction: SortDirection;
|
||||
}
|
||||
|
||||
export declare interface Subscription {
|
||||
id: string;
|
||||
entity: string;
|
||||
schema?: string;
|
||||
options?: WSOptions;
|
||||
callback?: (notification: WSNotificationMessage) => void;
|
||||
}
|
||||
|
||||
export declare interface SubscriptionOptions {
|
||||
filters?: FilterOption[];
|
||||
onNotification?: (notification: WSNotificationMessage) => void;
|
||||
}
|
||||
|
||||
export declare interface TableMetadata {
|
||||
schema: string;
|
||||
table: string;
|
||||
columns: Column[];
|
||||
relations: string[];
|
||||
}
|
||||
|
||||
export declare class WebSocketClient {
|
||||
private ws;
|
||||
private config;
|
||||
private messageHandlers;
|
||||
private subscriptions;
|
||||
private eventListeners;
|
||||
private state;
|
||||
private reconnectAttempts;
|
||||
private reconnectTimer;
|
||||
private heartbeatTimer;
|
||||
private isManualClose;
|
||||
constructor(config: WebSocketClientConfig);
|
||||
connect(): Promise<void>;
|
||||
disconnect(): void;
|
||||
request<T = any>(operation: WSOperation, entity: string, options?: {
|
||||
schema?: string;
|
||||
record_id?: string;
|
||||
data?: any;
|
||||
options?: WSOptions;
|
||||
}): Promise<T>;
|
||||
read<T = any>(entity: string, options?: {
|
||||
schema?: string;
|
||||
record_id?: string;
|
||||
filters?: FilterOption[];
|
||||
columns?: string[];
|
||||
sort?: SortOption[];
|
||||
preload?: PreloadOption[];
|
||||
limit?: number;
|
||||
offset?: number;
|
||||
}): Promise<T>;
|
||||
create<T = any>(entity: string, data: any, options?: {
|
||||
schema?: string;
|
||||
}): Promise<T>;
|
||||
update<T = any>(entity: string, id: string, data: any, options?: {
|
||||
schema?: string;
|
||||
}): Promise<T>;
|
||||
delete(entity: string, id: string, options?: {
|
||||
schema?: string;
|
||||
}): Promise<void>;
|
||||
meta<T = any>(entity: string, options?: {
|
||||
schema?: string;
|
||||
}): Promise<T>;
|
||||
subscribe(entity: string, callback: (notification: WSNotificationMessage) => void, options?: {
|
||||
schema?: string;
|
||||
filters?: FilterOption[];
|
||||
}): Promise<string>;
|
||||
unsubscribe(subscriptionId: string): Promise<void>;
|
||||
getSubscriptions(): Subscription[];
|
||||
getState(): ConnectionState;
|
||||
isConnected(): boolean;
|
||||
on<K extends keyof WebSocketClientEvents>(event: K, callback: WebSocketClientEvents[K]): void;
|
||||
off<K extends keyof WebSocketClientEvents>(event: K): void;
|
||||
private handleMessage;
|
||||
private handleResponse;
|
||||
private handleNotification;
|
||||
private send;
|
||||
private startHeartbeat;
|
||||
private stopHeartbeat;
|
||||
private setState;
|
||||
private ensureConnected;
|
||||
private emit;
|
||||
private log;
|
||||
}
|
||||
|
||||
export declare interface WebSocketClientConfig {
|
||||
url: string;
|
||||
reconnect?: boolean;
|
||||
reconnectInterval?: number;
|
||||
maxReconnectAttempts?: number;
|
||||
heartbeatInterval?: number;
|
||||
debug?: boolean;
|
||||
}
|
||||
|
||||
export declare interface WebSocketClientEvents {
|
||||
connect: () => void;
|
||||
disconnect: (event: CloseEvent) => void;
|
||||
error: (error: Error) => void;
|
||||
message: (message: WSMessage) => void;
|
||||
stateChange: (state: ConnectionState) => void;
|
||||
}
|
||||
|
||||
export declare interface WSErrorInfo {
|
||||
code: string;
|
||||
message: string;
|
||||
details?: Record<string, any>;
|
||||
}
|
||||
|
||||
export declare interface WSMessage {
|
||||
id?: string;
|
||||
type: MessageType;
|
||||
operation?: WSOperation;
|
||||
schema?: string;
|
||||
entity?: string;
|
||||
record_id?: string;
|
||||
data?: any;
|
||||
options?: WSOptions;
|
||||
subscription_id?: string;
|
||||
success?: boolean;
|
||||
error?: WSErrorInfo;
|
||||
metadata?: Record<string, any>;
|
||||
timestamp?: string;
|
||||
}
|
||||
|
||||
export declare interface WSNotificationMessage {
|
||||
type: 'notification';
|
||||
operation: WSOperation;
|
||||
subscription_id: string;
|
||||
schema?: string;
|
||||
entity: string;
|
||||
data: any;
|
||||
timestamp: string;
|
||||
}
|
||||
|
||||
export declare type WSOperation = 'read' | 'create' | 'update' | 'delete' | 'subscribe' | 'unsubscribe' | 'meta';
|
||||
|
||||
export declare interface WSOptions {
|
||||
filters?: FilterOption[];
|
||||
columns?: string[];
|
||||
omit_columns?: string[];
|
||||
preload?: PreloadOption[];
|
||||
sort?: SortOption[];
|
||||
limit?: number;
|
||||
offset?: number;
|
||||
parameters?: Parameter[];
|
||||
cursor_forward?: string;
|
||||
cursor_backward?: string;
|
||||
fetch_row_number?: string;
|
||||
}
|
||||
|
||||
export declare interface WSRequestMessage {
|
||||
id: string;
|
||||
type: 'request';
|
||||
operation: WSOperation;
|
||||
schema?: string;
|
||||
entity: string;
|
||||
record_id?: string;
|
||||
data?: any;
|
||||
options?: WSOptions;
|
||||
}
|
||||
|
||||
export declare interface WSResponseMessage {
|
||||
id: string;
|
||||
type: 'response';
|
||||
success: boolean;
|
||||
data?: any;
|
||||
error?: WSErrorInfo;
|
||||
metadata?: Record<string, any>;
|
||||
timestamp: string;
|
||||
}
|
||||
|
||||
export declare interface WSSubscriptionMessage {
|
||||
id: string;
|
||||
type: 'subscription';
|
||||
operation: 'subscribe' | 'unsubscribe';
|
||||
schema?: string;
|
||||
entity: string;
|
||||
options?: WSOptions;
|
||||
subscription_id?: string;
|
||||
}
|
||||
|
||||
export { }
|
||||
export * from './common';
|
||||
export * from './resolvespec';
|
||||
export * from './websocketspec';
|
||||
export * from './headerspec';
|
||||
//# sourceMappingURL=index.d.ts.map
|
||||
Vendored
+426
-463
@@ -1,469 +1,432 @@
|
||||
import { v4 as l } from "uuid";
|
||||
const d = /* @__PURE__ */ new Map();
|
||||
function E(n) {
|
||||
const e = n.baseUrl;
|
||||
let t = d.get(e);
|
||||
return t || (t = new g(n), d.set(e, t)), t;
|
||||
import { v4 as e } from "uuid";
|
||||
import { b64DecodeUnicode as t, b64EncodeUnicode as n } from "@warkypublic/artemis-kit/base64";
|
||||
//#region src/common/http.ts
|
||||
function r(...e) {
|
||||
let t = {};
|
||||
for (let n of e) for (let [e, r] of Object.entries(n)) {
|
||||
for (let n of Object.keys(t)) n.toLowerCase() === e.toLowerCase() && delete t[n];
|
||||
Object.defineProperty(t, e, {
|
||||
value: r,
|
||||
enumerable: !0,
|
||||
configurable: !0,
|
||||
writable: !0
|
||||
});
|
||||
}
|
||||
return t;
|
||||
}
|
||||
class g {
|
||||
constructor(e) {
|
||||
this.config = e;
|
||||
}
|
||||
buildUrl(e, t, s) {
|
||||
let r = `${this.config.baseUrl}/${e}/${t}`;
|
||||
return s && (r += `/${s}`), r;
|
||||
}
|
||||
baseHeaders() {
|
||||
const e = {
|
||||
"Content-Type": "application/json"
|
||||
};
|
||||
return this.config.token && (e.Authorization = `Bearer ${this.config.token}`), e;
|
||||
}
|
||||
async fetchWithError(e, t) {
|
||||
const s = await fetch(e, t), r = await s.json();
|
||||
if (!s.ok)
|
||||
throw new Error(r.error?.message || "An error occurred");
|
||||
return r;
|
||||
}
|
||||
async getMetadata(e, t) {
|
||||
const s = this.buildUrl(e, t);
|
||||
return this.fetchWithError(s, {
|
||||
method: "GET",
|
||||
headers: this.baseHeaders()
|
||||
});
|
||||
}
|
||||
async read(e, t, s, r) {
|
||||
const i = typeof s == "number" || typeof s == "string" ? String(s) : void 0, a = this.buildUrl(e, t, i), c = {
|
||||
operation: "read",
|
||||
id: Array.isArray(s) ? s : void 0,
|
||||
options: r
|
||||
};
|
||||
return this.fetchWithError(a, {
|
||||
method: "POST",
|
||||
headers: this.baseHeaders(),
|
||||
body: JSON.stringify(c)
|
||||
});
|
||||
}
|
||||
async create(e, t, s, r) {
|
||||
const i = this.buildUrl(e, t), a = {
|
||||
operation: "create",
|
||||
data: s,
|
||||
options: r
|
||||
};
|
||||
return this.fetchWithError(i, {
|
||||
method: "POST",
|
||||
headers: this.baseHeaders(),
|
||||
body: JSON.stringify(a)
|
||||
});
|
||||
}
|
||||
async update(e, t, s, r, i) {
|
||||
const a = typeof r == "number" || typeof r == "string" ? String(r) : void 0, c = this.buildUrl(e, t, a), o = {
|
||||
operation: "update",
|
||||
id: Array.isArray(r) ? r : void 0,
|
||||
data: s,
|
||||
options: i
|
||||
};
|
||||
return this.fetchWithError(c, {
|
||||
method: "POST",
|
||||
headers: this.baseHeaders(),
|
||||
body: JSON.stringify(o)
|
||||
});
|
||||
}
|
||||
async delete(e, t, s) {
|
||||
const r = this.buildUrl(e, t, String(s)), i = {
|
||||
operation: "delete"
|
||||
};
|
||||
return this.fetchWithError(r, {
|
||||
method: "POST",
|
||||
headers: this.baseHeaders(),
|
||||
body: JSON.stringify(i)
|
||||
});
|
||||
}
|
||||
function i(e) {
|
||||
return r({ "Content-Type": "application/json" }, e.headers ?? {}, e.token ? { Authorization: `Bearer ${e.token}` } : {});
|
||||
}
|
||||
const f = /* @__PURE__ */ new Map();
|
||||
function _(n) {
|
||||
const e = n.url;
|
||||
let t = f.get(e);
|
||||
return t || (t = new p(n), f.set(e, t)), t;
|
||||
function a(e) {
|
||||
let t = Object.entries(i(e)).map(([e, t]) => [e.toLowerCase(), t]).sort(([e], [t]) => e.localeCompare(t));
|
||||
return JSON.stringify([e.baseUrl, t]);
|
||||
}
|
||||
class p {
|
||||
constructor(e) {
|
||||
this.ws = null, this.messageHandlers = /* @__PURE__ */ new Map(), this.subscriptions = /* @__PURE__ */ new Map(), this.eventListeners = {}, this.state = "disconnected", this.reconnectAttempts = 0, this.reconnectTimer = null, this.heartbeatTimer = null, this.isManualClose = !1, this.config = {
|
||||
url: e.url,
|
||||
reconnect: e.reconnect ?? !0,
|
||||
reconnectInterval: e.reconnectInterval ?? 3e3,
|
||||
maxReconnectAttempts: e.maxReconnectAttempts ?? 10,
|
||||
heartbeatInterval: e.heartbeatInterval ?? 3e4,
|
||||
debug: e.debug ?? !1
|
||||
};
|
||||
}
|
||||
async connect() {
|
||||
if (this.ws?.readyState === WebSocket.OPEN) {
|
||||
this.log("Already connected");
|
||||
return;
|
||||
}
|
||||
return this.isManualClose = !1, this.setState("connecting"), new Promise((e, t) => {
|
||||
try {
|
||||
this.ws = new WebSocket(this.config.url), this.ws.onopen = () => {
|
||||
this.log("Connected to WebSocket server"), this.setState("connected"), this.reconnectAttempts = 0, this.startHeartbeat(), this.emit("connect"), e();
|
||||
}, this.ws.onmessage = (s) => {
|
||||
this.handleMessage(s.data);
|
||||
}, this.ws.onerror = (s) => {
|
||||
this.log("WebSocket error:", s);
|
||||
const r = new Error("WebSocket connection error");
|
||||
this.emit("error", r), t(r);
|
||||
}, this.ws.onclose = (s) => {
|
||||
this.log("WebSocket closed:", s.code, s.reason), this.stopHeartbeat(), this.setState("disconnected"), this.emit("disconnect", s), this.config.reconnect && !this.isManualClose && this.reconnectAttempts < this.config.maxReconnectAttempts && (this.reconnectAttempts++, this.log(`Reconnection attempt ${this.reconnectAttempts}/${this.config.maxReconnectAttempts}`), this.setState("reconnecting"), this.reconnectTimer = setTimeout(() => {
|
||||
this.connect().catch((r) => {
|
||||
this.log("Reconnection failed:", r);
|
||||
});
|
||||
}, this.config.reconnectInterval));
|
||||
};
|
||||
} catch (s) {
|
||||
t(s);
|
||||
}
|
||||
});
|
||||
}
|
||||
disconnect() {
|
||||
this.isManualClose = !0, this.reconnectTimer && (clearTimeout(this.reconnectTimer), this.reconnectTimer = null), this.stopHeartbeat(), this.ws && (this.setState("disconnecting"), this.ws.close(), this.ws = null), this.setState("disconnected"), this.messageHandlers.clear();
|
||||
}
|
||||
async request(e, t, s) {
|
||||
this.ensureConnected();
|
||||
const r = l(), i = {
|
||||
id: r,
|
||||
type: "request",
|
||||
operation: e,
|
||||
entity: t,
|
||||
schema: s?.schema,
|
||||
record_id: s?.record_id,
|
||||
data: s?.data,
|
||||
options: s?.options
|
||||
};
|
||||
return new Promise((a, c) => {
|
||||
this.messageHandlers.set(r, (o) => {
|
||||
o.success ? a(o.data) : c(new Error(o.error?.message || "Request failed"));
|
||||
}), this.send(i), setTimeout(() => {
|
||||
this.messageHandlers.has(r) && (this.messageHandlers.delete(r), c(new Error("Request timeout")));
|
||||
}, 3e4);
|
||||
});
|
||||
}
|
||||
async read(e, t) {
|
||||
return this.request("read", e, {
|
||||
schema: t?.schema,
|
||||
record_id: t?.record_id,
|
||||
options: {
|
||||
filters: t?.filters,
|
||||
columns: t?.columns,
|
||||
sort: t?.sort,
|
||||
preload: t?.preload,
|
||||
limit: t?.limit,
|
||||
offset: t?.offset
|
||||
}
|
||||
});
|
||||
}
|
||||
async create(e, t, s) {
|
||||
return this.request("create", e, {
|
||||
schema: s?.schema,
|
||||
data: t
|
||||
});
|
||||
}
|
||||
async update(e, t, s, r) {
|
||||
return this.request("update", e, {
|
||||
schema: r?.schema,
|
||||
record_id: t,
|
||||
data: s
|
||||
});
|
||||
}
|
||||
async delete(e, t, s) {
|
||||
await this.request("delete", e, {
|
||||
schema: s?.schema,
|
||||
record_id: t
|
||||
});
|
||||
}
|
||||
async meta(e, t) {
|
||||
return this.request("meta", e, {
|
||||
schema: t?.schema
|
||||
});
|
||||
}
|
||||
async subscribe(e, t, s) {
|
||||
this.ensureConnected();
|
||||
const r = l(), i = {
|
||||
id: r,
|
||||
type: "subscription",
|
||||
operation: "subscribe",
|
||||
entity: e,
|
||||
schema: s?.schema,
|
||||
options: {
|
||||
filters: s?.filters
|
||||
}
|
||||
};
|
||||
return new Promise((a, c) => {
|
||||
this.messageHandlers.set(r, (o) => {
|
||||
if (o.success && o.data?.subscription_id) {
|
||||
const h = o.data.subscription_id;
|
||||
this.subscriptions.set(h, {
|
||||
id: h,
|
||||
entity: e,
|
||||
schema: s?.schema,
|
||||
options: { filters: s?.filters },
|
||||
callback: t
|
||||
}), this.log(`Subscribed to ${e} with ID: ${h}`), a(h);
|
||||
} else
|
||||
c(new Error(o.error?.message || "Subscription failed"));
|
||||
}), this.send(i), setTimeout(() => {
|
||||
this.messageHandlers.has(r) && (this.messageHandlers.delete(r), c(new Error("Subscription timeout")));
|
||||
}, 1e4);
|
||||
});
|
||||
}
|
||||
async unsubscribe(e) {
|
||||
this.ensureConnected();
|
||||
const t = l(), s = {
|
||||
id: t,
|
||||
type: "subscription",
|
||||
operation: "unsubscribe",
|
||||
subscription_id: e
|
||||
};
|
||||
return new Promise((r, i) => {
|
||||
this.messageHandlers.set(t, (a) => {
|
||||
a.success ? (this.subscriptions.delete(e), this.log(`Unsubscribed from ${e}`), r()) : i(new Error(a.error?.message || "Unsubscribe failed"));
|
||||
}), this.send(s), setTimeout(() => {
|
||||
this.messageHandlers.has(t) && (this.messageHandlers.delete(t), i(new Error("Unsubscribe timeout")));
|
||||
}, 1e4);
|
||||
});
|
||||
}
|
||||
getSubscriptions() {
|
||||
return Array.from(this.subscriptions.values());
|
||||
}
|
||||
getState() {
|
||||
return this.state;
|
||||
}
|
||||
isConnected() {
|
||||
return this.ws?.readyState === WebSocket.OPEN;
|
||||
}
|
||||
on(e, t) {
|
||||
this.eventListeners[e] = t;
|
||||
}
|
||||
off(e) {
|
||||
delete this.eventListeners[e];
|
||||
}
|
||||
// Private methods
|
||||
handleMessage(e) {
|
||||
try {
|
||||
const t = JSON.parse(e);
|
||||
switch (this.log("Received message:", t), this.emit("message", t), t.type) {
|
||||
case "response":
|
||||
this.handleResponse(t);
|
||||
break;
|
||||
case "notification":
|
||||
this.handleNotification(t);
|
||||
break;
|
||||
case "pong":
|
||||
break;
|
||||
default:
|
||||
this.log("Unknown message type:", t.type);
|
||||
}
|
||||
} catch (t) {
|
||||
this.log("Error parsing message:", t);
|
||||
}
|
||||
}
|
||||
handleResponse(e) {
|
||||
const t = this.messageHandlers.get(e.id);
|
||||
t && (t(e), this.messageHandlers.delete(e.id));
|
||||
}
|
||||
handleNotification(e) {
|
||||
const t = this.subscriptions.get(e.subscription_id);
|
||||
t?.callback && t.callback(e);
|
||||
}
|
||||
send(e) {
|
||||
if (!this.ws || this.ws.readyState !== WebSocket.OPEN)
|
||||
throw new Error("WebSocket is not connected");
|
||||
const t = JSON.stringify(e);
|
||||
this.log("Sending message:", e), this.ws.send(t);
|
||||
}
|
||||
startHeartbeat() {
|
||||
this.heartbeatTimer || (this.heartbeatTimer = setInterval(() => {
|
||||
if (this.isConnected()) {
|
||||
const e = {
|
||||
id: l(),
|
||||
type: "ping"
|
||||
};
|
||||
this.send(e);
|
||||
}
|
||||
}, this.config.heartbeatInterval));
|
||||
}
|
||||
stopHeartbeat() {
|
||||
this.heartbeatTimer && (clearInterval(this.heartbeatTimer), this.heartbeatTimer = null);
|
||||
}
|
||||
setState(e) {
|
||||
this.state !== e && (this.state = e, this.emit("stateChange", e));
|
||||
}
|
||||
ensureConnected() {
|
||||
if (!this.isConnected())
|
||||
throw new Error("WebSocket is not connected. Call connect() first.");
|
||||
}
|
||||
emit(e, ...t) {
|
||||
const s = this.eventListeners[e];
|
||||
s && s(...t);
|
||||
}
|
||||
log(...e) {
|
||||
this.config.debug && console.log("[WebSocketClient]", ...e);
|
||||
}
|
||||
//#endregion
|
||||
//#region src/resolvespec/client.ts
|
||||
var o = /* @__PURE__ */ new Map();
|
||||
function s(e) {
|
||||
let t = a(e), n = o.get(t);
|
||||
return n || (n = new c(e), o.set(t, n)), n;
|
||||
}
|
||||
function v(n) {
|
||||
return typeof btoa == "function" ? "ZIP_" + btoa(n) : "ZIP_" + Buffer.from(n, "utf-8").toString("base64");
|
||||
var c = class {
|
||||
constructor(e) {
|
||||
this.config = {
|
||||
...e,
|
||||
headers: { ...e.headers }
|
||||
};
|
||||
}
|
||||
buildUrl(e, t, n) {
|
||||
let r = `${this.config.baseUrl}/${e}/${t}`;
|
||||
return n && (r += `/${n}`), r;
|
||||
}
|
||||
baseHeaders() {
|
||||
return i(this.config);
|
||||
}
|
||||
async fetchWithError(e, t) {
|
||||
let n = await fetch(e, t), r = await n.json();
|
||||
if (!n.ok) throw Error(r.error?.message || "An error occurred");
|
||||
return r;
|
||||
}
|
||||
async getMetadata(e, t) {
|
||||
let n = this.buildUrl(e, t);
|
||||
return this.fetchWithError(n, {
|
||||
method: "GET",
|
||||
headers: this.baseHeaders()
|
||||
});
|
||||
}
|
||||
async read(e, t, n, r) {
|
||||
let i = typeof n == "number" || typeof n == "string" ? String(n) : void 0, a = this.buildUrl(e, t, i), o = {
|
||||
operation: "read",
|
||||
id: Array.isArray(n) ? n : void 0,
|
||||
options: r
|
||||
};
|
||||
return this.fetchWithError(a, {
|
||||
method: "POST",
|
||||
headers: this.baseHeaders(),
|
||||
body: JSON.stringify(o)
|
||||
});
|
||||
}
|
||||
async create(e, t, n, r) {
|
||||
let i = this.buildUrl(e, t), a = {
|
||||
operation: "create",
|
||||
data: n,
|
||||
options: r
|
||||
};
|
||||
return this.fetchWithError(i, {
|
||||
method: "POST",
|
||||
headers: this.baseHeaders(),
|
||||
body: JSON.stringify(a)
|
||||
});
|
||||
}
|
||||
async update(e, t, n, r, i) {
|
||||
let a = typeof r == "number" || typeof r == "string" ? String(r) : void 0, o = this.buildUrl(e, t, a), s = {
|
||||
operation: "update",
|
||||
id: Array.isArray(r) ? r : void 0,
|
||||
data: n,
|
||||
options: i
|
||||
};
|
||||
return this.fetchWithError(o, {
|
||||
method: "POST",
|
||||
headers: this.baseHeaders(),
|
||||
body: JSON.stringify(s)
|
||||
});
|
||||
}
|
||||
async delete(e, t, n) {
|
||||
let r = this.buildUrl(e, t, String(n));
|
||||
return this.fetchWithError(r, {
|
||||
method: "POST",
|
||||
headers: this.baseHeaders(),
|
||||
body: JSON.stringify({ operation: "delete" })
|
||||
});
|
||||
}
|
||||
}, l = /* @__PURE__ */ new Map();
|
||||
function u(e) {
|
||||
let t = e.url, n = l.get(t);
|
||||
return n || (n = new d(e), l.set(t, n)), n;
|
||||
}
|
||||
function w(n) {
|
||||
let e = n;
|
||||
return e.startsWith("ZIP_") ? (e = e.slice(4).replace(/[\n\r ]/g, ""), e = m(e)) : e.startsWith("__") && (e = e.slice(2).replace(/[\n\r ]/g, ""), e = m(e)), (e.startsWith("ZIP_") || e.startsWith("__")) && (e = w(e)), e;
|
||||
}
|
||||
function m(n) {
|
||||
return typeof atob == "function" ? atob(n) : Buffer.from(n, "base64").toString("utf-8");
|
||||
}
|
||||
function u(n) {
|
||||
const e = {};
|
||||
if (n.columns?.length && (e["X-Select-Fields"] = n.columns.join(",")), n.omit_columns?.length && (e["X-Not-Select-Fields"] = n.omit_columns.join(",")), n.filters?.length)
|
||||
for (const t of n.filters) {
|
||||
const s = t.logic_operator ?? "AND", r = y(t.operator), i = S(t);
|
||||
t.operator === "eq" && s === "AND" ? e[`X-FieldFilter-${t.column}`] = i : s === "OR" ? e[`X-SearchOr-${r}-${t.column}`] = i : e[`X-SearchOp-${r}-${t.column}`] = i;
|
||||
}
|
||||
if (n.sort?.length) {
|
||||
const t = n.sort.map((s) => s.direction.toUpperCase() === "DESC" ? `-${s.column}` : `+${s.column}`);
|
||||
e["X-Sort"] = t.join(",");
|
||||
}
|
||||
if (n.limit !== void 0 && (e["X-Limit"] = String(n.limit)), n.offset !== void 0 && (e["X-Offset"] = String(n.offset)), n.cursor_forward && (e["X-Cursor-Forward"] = n.cursor_forward), n.cursor_backward && (e["X-Cursor-Backward"] = n.cursor_backward), n.preload?.length) {
|
||||
const t = n.preload.map((s) => s.columns?.length ? `${s.relation}:${s.columns.join(",")}` : s.relation);
|
||||
e["X-Preload"] = t.join("|");
|
||||
}
|
||||
if (n.fetch_row_number && (e["X-Fetch-RowNumber"] = n.fetch_row_number), n.computedColumns?.length)
|
||||
for (const t of n.computedColumns)
|
||||
e[`X-CQL-SEL-${t.name}`] = t.expression;
|
||||
if (n.customOperators?.length) {
|
||||
const t = n.customOperators.map(
|
||||
(s) => s.sql
|
||||
);
|
||||
e["X-Custom-SQL-W"] = t.join(" AND ");
|
||||
}
|
||||
return e;
|
||||
}
|
||||
function y(n) {
|
||||
switch (n) {
|
||||
case "eq":
|
||||
return "equals";
|
||||
case "neq":
|
||||
return "notequals";
|
||||
case "gt":
|
||||
return "greaterthan";
|
||||
case "gte":
|
||||
return "greaterthanorequal";
|
||||
case "lt":
|
||||
return "lessthan";
|
||||
case "lte":
|
||||
return "lessthanorequal";
|
||||
case "like":
|
||||
case "ilike":
|
||||
case "contains":
|
||||
return "contains";
|
||||
case "startswith":
|
||||
return "beginswith";
|
||||
case "endswith":
|
||||
return "endswith";
|
||||
case "in":
|
||||
return "in";
|
||||
case "between":
|
||||
return "between";
|
||||
case "between_inclusive":
|
||||
return "betweeninclusive";
|
||||
case "is_null":
|
||||
return "empty";
|
||||
case "is_not_null":
|
||||
return "notempty";
|
||||
default:
|
||||
return n;
|
||||
}
|
||||
}
|
||||
function S(n) {
|
||||
return n.value === null || n.value === void 0 ? "" : Array.isArray(n.value) ? n.value.join(",") : String(n.value);
|
||||
}
|
||||
const b = /* @__PURE__ */ new Map();
|
||||
function C(n) {
|
||||
const e = n.baseUrl;
|
||||
let t = b.get(e);
|
||||
return t || (t = new H(n), b.set(e, t)), t;
|
||||
}
|
||||
class H {
|
||||
constructor(e) {
|
||||
this.config = e;
|
||||
}
|
||||
buildUrl(e, t, s) {
|
||||
let r = `${this.config.baseUrl}/${e}/${t}`;
|
||||
return s && (r += `/${s}`), r;
|
||||
}
|
||||
baseHeaders() {
|
||||
const e = {
|
||||
"Content-Type": "application/json"
|
||||
};
|
||||
return this.config.token && (e.Authorization = `Bearer ${this.config.token}`), e;
|
||||
}
|
||||
async fetchWithError(e, t) {
|
||||
const s = await fetch(e, t), r = await s.json();
|
||||
if (!s.ok)
|
||||
throw new Error(
|
||||
r.error?.message || `${s.statusText} (${s.status})`
|
||||
);
|
||||
return {
|
||||
data: r,
|
||||
success: !0,
|
||||
error: r.error ? r.error : void 0,
|
||||
metadata: {
|
||||
count: s.headers.get("content-range") ? Number(s.headers.get("content-range")?.split("/")[1]) : 0,
|
||||
total: s.headers.get("content-range") ? Number(s.headers.get("content-range")?.split("/")[1]) : 0,
|
||||
filtered: s.headers.get("content-range") ? Number(s.headers.get("content-range")?.split("/")[1]) : 0,
|
||||
offset: s.headers.get("content-range") ? Number(
|
||||
s.headers.get("content-range")?.split("/")[0].split("-")[0]
|
||||
) : 0,
|
||||
limit: s.headers.get("x-limit") ? Number(s.headers.get("x-limit")) : 0
|
||||
}
|
||||
};
|
||||
}
|
||||
async read(e, t, s, r) {
|
||||
const i = this.buildUrl(e, t, s), a = r ? u(r) : {};
|
||||
return this.fetchWithError(i, {
|
||||
method: "GET",
|
||||
headers: { ...this.baseHeaders(), ...a }
|
||||
});
|
||||
}
|
||||
async create(e, t, s, r) {
|
||||
const i = this.buildUrl(e, t), a = r ? u(r) : {};
|
||||
return this.fetchWithError(i, {
|
||||
method: "POST",
|
||||
headers: { ...this.baseHeaders(), ...a },
|
||||
body: JSON.stringify(s)
|
||||
});
|
||||
}
|
||||
async update(e, t, s, r, i) {
|
||||
const a = this.buildUrl(e, t, s), c = i ? u(i) : {};
|
||||
return this.fetchWithError(a, {
|
||||
method: "PUT",
|
||||
headers: { ...this.baseHeaders(), ...c },
|
||||
body: JSON.stringify(r)
|
||||
});
|
||||
}
|
||||
async delete(e, t, s) {
|
||||
const r = this.buildUrl(e, t, s);
|
||||
return this.fetchWithError(r, {
|
||||
method: "DELETE",
|
||||
headers: this.baseHeaders()
|
||||
});
|
||||
}
|
||||
}
|
||||
export {
|
||||
H as HeaderSpecClient,
|
||||
g as ResolveSpecClient,
|
||||
p as WebSocketClient,
|
||||
u as buildHeaders,
|
||||
w as decodeHeaderValue,
|
||||
v as encodeHeaderValue,
|
||||
C as getHeaderSpecClient,
|
||||
E as getResolveSpecClient,
|
||||
_ as getWebSocketClient
|
||||
var d = class {
|
||||
constructor(e) {
|
||||
this.ws = null, this.messageHandlers = /* @__PURE__ */ new Map(), this.subscriptions = /* @__PURE__ */ new Map(), this.eventListeners = {}, this.state = "disconnected", this.reconnectAttempts = 0, this.reconnectTimer = null, this.heartbeatTimer = null, this.isManualClose = !1, this.config = {
|
||||
url: e.url,
|
||||
reconnect: e.reconnect ?? !0,
|
||||
reconnectInterval: e.reconnectInterval ?? 3e3,
|
||||
maxReconnectAttempts: e.maxReconnectAttempts ?? 10,
|
||||
heartbeatInterval: e.heartbeatInterval ?? 3e4,
|
||||
debug: e.debug ?? !1
|
||||
};
|
||||
}
|
||||
async connect() {
|
||||
if (this.ws?.readyState === WebSocket.OPEN) {
|
||||
this.log("Already connected");
|
||||
return;
|
||||
}
|
||||
return this.isManualClose = !1, this.setState("connecting"), new Promise((e, t) => {
|
||||
try {
|
||||
this.ws = new WebSocket(this.config.url), this.ws.onopen = () => {
|
||||
this.log("Connected to WebSocket server"), this.setState("connected"), this.reconnectAttempts = 0, this.startHeartbeat(), this.emit("connect"), e();
|
||||
}, this.ws.onmessage = (e) => {
|
||||
this.handleMessage(e.data);
|
||||
}, this.ws.onerror = (e) => {
|
||||
this.log("WebSocket error:", e);
|
||||
let n = /* @__PURE__ */ Error("WebSocket connection error");
|
||||
this.emit("error", n), t(n);
|
||||
}, this.ws.onclose = (e) => {
|
||||
this.log("WebSocket closed:", e.code, e.reason), this.stopHeartbeat(), this.setState("disconnected"), this.emit("disconnect", e), this.config.reconnect && !this.isManualClose && this.reconnectAttempts < this.config.maxReconnectAttempts && (this.reconnectAttempts++, this.log(`Reconnection attempt ${this.reconnectAttempts}/${this.config.maxReconnectAttempts}`), this.setState("reconnecting"), this.reconnectTimer = setTimeout(() => {
|
||||
this.connect().catch((e) => {
|
||||
this.log("Reconnection failed:", e);
|
||||
});
|
||||
}, this.config.reconnectInterval));
|
||||
};
|
||||
} catch (e) {
|
||||
t(e);
|
||||
}
|
||||
});
|
||||
}
|
||||
disconnect() {
|
||||
this.isManualClose = !0, this.reconnectTimer &&= (clearTimeout(this.reconnectTimer), null), this.stopHeartbeat(), this.ws &&= (this.setState("disconnecting"), this.ws.close(), null), this.setState("disconnected"), this.messageHandlers.clear();
|
||||
}
|
||||
async request(t, n, r) {
|
||||
this.ensureConnected();
|
||||
let i = e(), a = {
|
||||
id: i,
|
||||
type: "request",
|
||||
operation: t,
|
||||
entity: n,
|
||||
schema: r?.schema,
|
||||
record_id: r?.record_id,
|
||||
data: r?.data,
|
||||
options: r?.options
|
||||
};
|
||||
return new Promise((e, t) => {
|
||||
this.messageHandlers.set(i, (n) => {
|
||||
n.success ? e(n.data) : t(Error(n.error?.message || "Request failed"));
|
||||
}), this.send(a), setTimeout(() => {
|
||||
this.messageHandlers.has(i) && (this.messageHandlers.delete(i), t(/* @__PURE__ */ Error("Request timeout")));
|
||||
}, 3e4);
|
||||
});
|
||||
}
|
||||
async read(e, t) {
|
||||
return this.request("read", e, {
|
||||
schema: t?.schema,
|
||||
record_id: t?.record_id,
|
||||
options: {
|
||||
filters: t?.filters,
|
||||
columns: t?.columns,
|
||||
sort: t?.sort,
|
||||
preload: t?.preload,
|
||||
limit: t?.limit,
|
||||
offset: t?.offset
|
||||
}
|
||||
});
|
||||
}
|
||||
async create(e, t, n) {
|
||||
return this.request("create", e, {
|
||||
schema: n?.schema,
|
||||
data: t
|
||||
});
|
||||
}
|
||||
async update(e, t, n, r) {
|
||||
return this.request("update", e, {
|
||||
schema: r?.schema,
|
||||
record_id: t,
|
||||
data: n
|
||||
});
|
||||
}
|
||||
async delete(e, t, n) {
|
||||
await this.request("delete", e, {
|
||||
schema: n?.schema,
|
||||
record_id: t
|
||||
});
|
||||
}
|
||||
async meta(e, t) {
|
||||
return this.request("meta", e, { schema: t?.schema });
|
||||
}
|
||||
async subscribe(t, n, r) {
|
||||
this.ensureConnected();
|
||||
let i = e(), a = {
|
||||
id: i,
|
||||
type: "subscription",
|
||||
operation: "subscribe",
|
||||
entity: t,
|
||||
schema: r?.schema,
|
||||
options: { filters: r?.filters }
|
||||
};
|
||||
return new Promise((e, o) => {
|
||||
this.messageHandlers.set(i, (i) => {
|
||||
if (i.success && i.data?.subscription_id) {
|
||||
let a = i.data.subscription_id;
|
||||
this.subscriptions.set(a, {
|
||||
id: a,
|
||||
entity: t,
|
||||
schema: r?.schema,
|
||||
options: { filters: r?.filters },
|
||||
callback: n
|
||||
}), this.log(`Subscribed to ${t} with ID: ${a}`), e(a);
|
||||
} else o(Error(i.error?.message || "Subscription failed"));
|
||||
}), this.send(a), setTimeout(() => {
|
||||
this.messageHandlers.has(i) && (this.messageHandlers.delete(i), o(/* @__PURE__ */ Error("Subscription timeout")));
|
||||
}, 1e4);
|
||||
});
|
||||
}
|
||||
async unsubscribe(t) {
|
||||
this.ensureConnected();
|
||||
let n = e(), r = {
|
||||
id: n,
|
||||
type: "subscription",
|
||||
operation: "unsubscribe",
|
||||
subscription_id: t
|
||||
};
|
||||
return new Promise((e, i) => {
|
||||
this.messageHandlers.set(n, (n) => {
|
||||
n.success ? (this.subscriptions.delete(t), this.log(`Unsubscribed from ${t}`), e()) : i(Error(n.error?.message || "Unsubscribe failed"));
|
||||
}), this.send(r), setTimeout(() => {
|
||||
this.messageHandlers.has(n) && (this.messageHandlers.delete(n), i(/* @__PURE__ */ Error("Unsubscribe timeout")));
|
||||
}, 1e4);
|
||||
});
|
||||
}
|
||||
getSubscriptions() {
|
||||
return Array.from(this.subscriptions.values());
|
||||
}
|
||||
getState() {
|
||||
return this.state;
|
||||
}
|
||||
isConnected() {
|
||||
return this.ws?.readyState === WebSocket.OPEN;
|
||||
}
|
||||
on(e, t) {
|
||||
this.eventListeners[e] = t;
|
||||
}
|
||||
off(e) {
|
||||
delete this.eventListeners[e];
|
||||
}
|
||||
handleMessage(e) {
|
||||
try {
|
||||
let t = JSON.parse(e);
|
||||
switch (this.log("Received message:", t), this.emit("message", t), t.type) {
|
||||
case "response":
|
||||
this.handleResponse(t);
|
||||
break;
|
||||
case "notification":
|
||||
this.handleNotification(t);
|
||||
break;
|
||||
case "pong": break;
|
||||
default: this.log("Unknown message type:", t.type);
|
||||
}
|
||||
} catch (e) {
|
||||
this.log("Error parsing message:", e);
|
||||
}
|
||||
}
|
||||
handleResponse(e) {
|
||||
let t = this.messageHandlers.get(e.id);
|
||||
t && (t(e), this.messageHandlers.delete(e.id));
|
||||
}
|
||||
handleNotification(e) {
|
||||
let t = this.subscriptions.get(e.subscription_id);
|
||||
t?.callback && t.callback(e);
|
||||
}
|
||||
send(e) {
|
||||
if (!this.ws || this.ws.readyState !== WebSocket.OPEN) throw Error("WebSocket is not connected");
|
||||
let t = JSON.stringify(e);
|
||||
this.log("Sending message:", e), this.ws.send(t);
|
||||
}
|
||||
startHeartbeat() {
|
||||
this.heartbeatTimer ||= setInterval(() => {
|
||||
if (this.isConnected()) {
|
||||
let t = {
|
||||
id: e(),
|
||||
type: "ping"
|
||||
};
|
||||
this.send(t);
|
||||
}
|
||||
}, this.config.heartbeatInterval);
|
||||
}
|
||||
stopHeartbeat() {
|
||||
this.heartbeatTimer &&= (clearInterval(this.heartbeatTimer), null);
|
||||
}
|
||||
setState(e) {
|
||||
this.state !== e && (this.state = e, this.emit("stateChange", e));
|
||||
}
|
||||
ensureConnected() {
|
||||
if (!this.isConnected()) throw Error("WebSocket is not connected. Call connect() first.");
|
||||
}
|
||||
emit(e, ...t) {
|
||||
let n = this.eventListeners[e];
|
||||
n && n(...t);
|
||||
}
|
||||
log(...e) {
|
||||
this.config.debug && console.log("[WebSocketClient]", ...e);
|
||||
}
|
||||
};
|
||||
//#endregion
|
||||
//#region src/headerspec/client.ts
|
||||
function f(e) {
|
||||
return "ZIP_" + n(e);
|
||||
}
|
||||
function p(e) {
|
||||
let t = e;
|
||||
return t.startsWith("ZIP_") ? (t = t.slice(4).replace(/[\n\r ]/g, ""), t = m(t)) : t.startsWith("__") && (t = t.slice(2).replace(/[\n\r ]/g, ""), t = m(t)), (t.startsWith("ZIP_") || t.startsWith("__")) && (t = p(t)), t;
|
||||
}
|
||||
function m(e) {
|
||||
return t(e);
|
||||
}
|
||||
function h(e) {
|
||||
let t = {};
|
||||
if (e.columns?.length && (t["X-Select-Fields"] = e.columns.join(",")), e.omit_columns?.length && (t["X-Not-Select-Fields"] = e.omit_columns.join(",")), e.filters?.length) for (let n of e.filters) {
|
||||
let e = n.logic_operator ?? "AND", r = g(n.operator), i = _(n);
|
||||
n.operator === "eq" && e === "AND" ? t[`X-FieldFilter-${n.column}`] = i : e === "OR" ? t[`X-SearchOr-${r}-${n.column}`] = i : t[`X-SearchOp-${r}-${n.column}`] = i;
|
||||
}
|
||||
if (e.sort?.length && (t["X-Sort"] = e.sort.map((e) => e.direction.toUpperCase() === "DESC" ? `-${e.column}` : `+${e.column}`).join(",")), e.limit !== void 0 && (t["X-Limit"] = String(e.limit)), e.offset !== void 0 && (t["X-Offset"] = String(e.offset)), e.cursor_forward && (t["X-Cursor-Forward"] = e.cursor_forward), e.cursor_backward && (t["X-Cursor-Backward"] = e.cursor_backward), e.preload?.length && (t["X-Preload"] = e.preload.map((e) => e.columns?.length ? `${e.relation}:${e.columns.join(",")}` : e.relation).join("|")), e.fetch_row_number && (t["X-Fetch-RowNumber"] = e.fetch_row_number), e.computedColumns?.length) for (let n of e.computedColumns) t[`X-CQL-SEL-${n.name}`] = n.expression;
|
||||
return e.customOperators?.length && (t["X-Custom-SQL-W"] = e.customOperators.map((e) => e.sql).join(" AND ")), t;
|
||||
}
|
||||
function g(e) {
|
||||
switch (e) {
|
||||
case "eq": return "equals";
|
||||
case "neq": return "notequals";
|
||||
case "gt": return "greaterthan";
|
||||
case "gte": return "greaterthanorequal";
|
||||
case "lt": return "lessthan";
|
||||
case "lte": return "lessthanorequal";
|
||||
case "like":
|
||||
case "ilike":
|
||||
case "contains": return "contains";
|
||||
case "startswith": return "beginswith";
|
||||
case "endswith": return "endswith";
|
||||
case "in": return "in";
|
||||
case "between": return "between";
|
||||
case "between_inclusive": return "betweeninclusive";
|
||||
case "is_null": return "empty";
|
||||
case "is_not_null": return "notempty";
|
||||
default: return e;
|
||||
}
|
||||
}
|
||||
function _(e) {
|
||||
return e.value === null || e.value === void 0 ? "" : Array.isArray(e.value) ? e.value.join(",") : String(e.value);
|
||||
}
|
||||
var v = /* @__PURE__ */ new Map();
|
||||
function y(e) {
|
||||
let t = a(e), n = v.get(t);
|
||||
return n || (n = new b(e), v.set(t, n)), n;
|
||||
}
|
||||
var b = class {
|
||||
constructor(e) {
|
||||
this.config = {
|
||||
...e,
|
||||
headers: { ...e.headers }
|
||||
};
|
||||
}
|
||||
buildUrl(e, t, n) {
|
||||
let r = `${this.config.baseUrl}/${e}/${t}`;
|
||||
return n && (r += `/${n}`), r;
|
||||
}
|
||||
baseHeaders() {
|
||||
return i(this.config);
|
||||
}
|
||||
async fetchWithError(e, t) {
|
||||
let n = await fetch(e, t), r = await n.json();
|
||||
if (!n.ok) throw Error(r.error?.message || `${n.statusText} (${n.status})`);
|
||||
return {
|
||||
data: r,
|
||||
success: !0,
|
||||
error: r.error ? r.error : void 0,
|
||||
metadata: {
|
||||
count: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[1]) : 0,
|
||||
total: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[1]) : 0,
|
||||
filtered: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[1]) : 0,
|
||||
offset: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[0].split("-")[0]) : 0,
|
||||
limit: n.headers.get("x-limit") ? Number(n.headers.get("x-limit")) : 0
|
||||
}
|
||||
};
|
||||
}
|
||||
async read(e, t, n, i) {
|
||||
let a = this.buildUrl(e, t, n), o = i ? h(i) : {};
|
||||
return this.fetchWithError(a, {
|
||||
method: "GET",
|
||||
headers: r(this.baseHeaders(), o)
|
||||
});
|
||||
}
|
||||
async create(e, t, n, i) {
|
||||
let a = this.buildUrl(e, t), o = i ? h(i) : {};
|
||||
return this.fetchWithError(a, {
|
||||
method: "POST",
|
||||
headers: r(this.baseHeaders(), o),
|
||||
body: JSON.stringify(n)
|
||||
});
|
||||
}
|
||||
async update(e, t, n, i, a) {
|
||||
let o = this.buildUrl(e, t, n), s = a ? h(a) : {};
|
||||
return this.fetchWithError(o, {
|
||||
method: "PUT",
|
||||
headers: r(this.baseHeaders(), s),
|
||||
body: JSON.stringify(i)
|
||||
});
|
||||
}
|
||||
async delete(e, t, n) {
|
||||
let r = this.buildUrl(e, t, n);
|
||||
return this.fetchWithError(r, {
|
||||
method: "DELETE",
|
||||
headers: this.baseHeaders()
|
||||
});
|
||||
}
|
||||
};
|
||||
//#endregion
|
||||
export { b as HeaderSpecClient, c as ResolveSpecClient, d as WebSocketClient, h as buildHeaders, p as decodeHeaderValue, f as encodeHeaderValue, y as getHeaderSpecClient, s as getResolveSpecClient, u as getWebSocketClient };
|
||||
|
||||
+14
-12
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@warkypublic/resolvespec-js",
|
||||
"version": "1.0.1",
|
||||
"version": "1.0.2",
|
||||
"description": "TypeScript client library for ResolveSpec REST, HeaderSpec, and WebSocket APIs",
|
||||
"type": "module",
|
||||
"main": "./dist/index.cjs",
|
||||
@@ -38,20 +38,22 @@
|
||||
"author": "Hein (Warkanum) Puth",
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"uuid": "^13.0.0"
|
||||
"@warkypublic/artemis-kit": "^1.0.10",
|
||||
"uuid": "^14.0.2"
|
||||
},
|
||||
"devDependencies": {
|
||||
"@changesets/cli": "^2.29.8",
|
||||
"@changesets/cli": "^3.0.3",
|
||||
"@eslint/js": "^10.0.1",
|
||||
"@types/jsdom": "^27.0.0",
|
||||
"eslint": "^10.0.0",
|
||||
"globals": "^17.3.0",
|
||||
"jsdom": "^28.1.0",
|
||||
"typescript": "^5.9.3",
|
||||
"typescript-eslint": "^8.55.0",
|
||||
"vite": "^7.3.1",
|
||||
"vite-plugin-dts": "^4.5.4",
|
||||
"vitest": "^4.0.18"
|
||||
"@types/jsdom": "^30.0.0",
|
||||
"@types/node": "^26.6.2",
|
||||
"eslint": "^10.11.0",
|
||||
"globals": "^17.12.0",
|
||||
"jsdom": "^30.1.1",
|
||||
"typescript": "^6.0.3",
|
||||
"typescript-eslint": "^8.70.1",
|
||||
"vite": "^8.3.0",
|
||||
"vite-plugin-dts": "^5.1.1",
|
||||
"vitest": "^5.0.1"
|
||||
},
|
||||
"engines": {
|
||||
"node": ">=18"
|
||||
|
||||
Generated
+1283
-1293
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user