mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-09-29 11:32:01 +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
|
.PHONY: test test-unit test-integration docker-up docker-down clean
|
||||||
|
|
||||||
|
GOLANGCI_LINT := $(shell go env GOPATH)/bin/golangci-lint
|
||||||
|
|
||||||
# Run all unit tests
|
# Run all unit tests
|
||||||
test-unit:
|
test-unit:
|
||||||
@echo "Running unit tests..."
|
@echo "Running unit tests..."
|
||||||
@@ -49,7 +51,9 @@ release-version: ## Create and push a release with specific version (use: make r
|
|||||||
|
|
||||||
lint: ## Run linter
|
lint: ## Run linter
|
||||||
@echo "Running linter..."
|
@echo "Running linter..."
|
||||||
@if command -v golangci-lint > /dev/null; then \
|
@if [ -x "$(GOLANGCI_LINT)" ]; then \
|
||||||
|
"$(GOLANGCI_LINT)" run --config=.golangci.json; \
|
||||||
|
elif command -v golangci-lint > /dev/null; then \
|
||||||
golangci-lint run --config=.golangci.json; \
|
golangci-lint run --config=.golangci.json; \
|
||||||
else \
|
else \
|
||||||
echo "golangci-lint not installed. Install with: go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest"; \
|
echo "golangci-lint not installed. Install with: go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest"; \
|
||||||
@@ -58,7 +62,9 @@ lint: ## Run linter
|
|||||||
|
|
||||||
lintfix: ## Run linter
|
lintfix: ## Run linter
|
||||||
@echo "Running linter..."
|
@echo "Running linter..."
|
||||||
@if command -v golangci-lint > /dev/null; then \
|
@if [ -x "$(GOLANGCI_LINT)" ]; then \
|
||||||
|
"$(GOLANGCI_LINT)" run --config=.golangci.json --fix; \
|
||||||
|
elif command -v golangci-lint > /dev/null; then \
|
||||||
golangci-lint run --config=.golangci.json --fix; \
|
golangci-lint run --config=.golangci.json --fix; \
|
||||||
else \
|
else \
|
||||||
echo "golangci-lint not installed. Install with: go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest"; \
|
echo "golangci-lint not installed. Install with: go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest"; \
|
||||||
|
|||||||
@@ -135,6 +135,73 @@ For complete documentation including setup, headers, lifecycle hooks, cursor pag
|
|||||||
|
|
||||||
For detailed examples of reading data, cursor pagination, recursive CRUD operations, filtering, sorting, and more, see [pkg/resolvespec/README.md](pkg/resolvespec/README.md).
|
For detailed examples of reading data, cursor pagination, recursive CRUD operations, filtering, sorting, and more, see [pkg/resolvespec/README.md](pkg/resolvespec/README.md).
|
||||||
|
|
||||||
|
## PostGIS & Vector (PostgreSQL only)
|
||||||
|
|
||||||
|
First-class support for PostGIS geometry/geography and pgvector columns in `resolvespec` + `restheadspec`. No extra dependencies. On non-Postgres databases the spatial/vector operators simply don't match.
|
||||||
|
|
||||||
|
### Column types (`pkg/spectypes`)
|
||||||
|
|
||||||
|
| Go type | SQL type | Wire / JSON |
|
||||||
|
|--------------------|--------------|--------------------------------------------------------|
|
||||||
|
| `SqlGeometry` | `geometry` | JSON in/out = **GeoJSON**; also accepts EWKT / hex-EWKB |
|
||||||
|
| `SqlGeography` | `geography` | same as `SqlGeometry` |
|
||||||
|
| `SqlVector` | `vector` | `[]float32` ⇄ `[1,2,3]` |
|
||||||
|
| `SqlHalfVector` | `halfvec` | `[]float32` ⇄ `[1,2,3]` |
|
||||||
|
| `SqlSparseVector` | `sparsevec` | `{"dim":8,"indices":[1,4],"values":[0.5,0.2]}` |
|
||||||
|
| `SqlBitVector` | `bit`/`varbit` | bool array or `"1011"` string |
|
||||||
|
|
||||||
|
- Geometry `Value()` emits `SRID=<n>;<WKT>` (PostGIS implicit text→geometry cast; no wrapper function needed).
|
||||||
|
- Declare dimensioned types with a tag: `gorm:"type:vector(1536)"` — the tag wins over the canonical name in metadata/OpenAPI.
|
||||||
|
- Metadata endpoint and OpenAPI schema report `geometry`/`vector`/`halfvec`/`sparsevec`/`bit`.
|
||||||
|
|
||||||
|
### Spatial filter operators
|
||||||
|
|
||||||
|
`value` is a geometry (GeoJSON object, EWKT string, or hex-EWKB) unless noted.
|
||||||
|
|
||||||
|
| Operator | Value shape |
|
||||||
|
|----------|-------------|
|
||||||
|
| `st_intersects`, `st_contains`, `st_within`, `st_covers`, `st_coveredby`, `st_overlaps`, `st_touches`, `st_crosses`, `st_equals`, `st_disjoint` | geometry |
|
||||||
|
| `st_dwithin` | `{"geom": <geometry>, "distance": <meters>}` |
|
||||||
|
| `bbox` (alias `&&`) | geometry, or `{"bbox":[minx,miny,maxx,maxy],"srid":4326}` |
|
||||||
|
|
||||||
|
### Vector similarity filter operators
|
||||||
|
|
||||||
|
| Operator | pgvector op | Value shape |
|
||||||
|
|----------|-------------|-------------|
|
||||||
|
| `l2_within` / `euclidean_within` | `<->` | `{"vector":[...], "distance": <n>}` |
|
||||||
|
| `cosine_within` | `<=>` | same (also `"lt"`/`"lte"`/`"gt"`/`"gte"` instead of `"distance"`) |
|
||||||
|
| `ip_within` / `inner_within` | `<#>` | same |
|
||||||
|
|
||||||
|
### KNN search (ordering + distance column)
|
||||||
|
|
||||||
|
**resolvespec** — `options.vector_search`:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{ "options": { "vector_search": {
|
||||||
|
"column": "embedding",
|
||||||
|
"vector": [0.1, 0.2, 0.3],
|
||||||
|
"metric": "cosine",
|
||||||
|
"as": "_distance",
|
||||||
|
"direction": "asc"
|
||||||
|
}}}
|
||||||
|
```
|
||||||
|
|
||||||
|
Orders rows by distance; when `as` is set, returns the distance as an extra column (all model columns are auto-selected).
|
||||||
|
`metric`: `l2` (default) | `cosine` | `ip`.
|
||||||
|
|
||||||
|
**restheadspec** — headers:
|
||||||
|
|
||||||
|
```HTTP
|
||||||
|
X-Vector-Search-embedding: cosine
|
||||||
|
X-Vector-Search-Vector: [0.1,0.2,0.3]
|
||||||
|
X-Vector-Search-As: _distance
|
||||||
|
X-Vector-Search-Dir: asc
|
||||||
|
```
|
||||||
|
|
||||||
|
Spatial/vector filters via headers: `X-SpatialFilter-<col>` / `X-VectorFilter-<col>` with a JSON operator object, e.g.
|
||||||
|
`X-SpatialFilter-geom: {"op":"st_dwithin","geom":"SRID=4326;POINT(0 0)","distance":1000}`
|
||||||
|
(optional `"logic":"or"`).
|
||||||
|
|
||||||
## Installation
|
## Installation
|
||||||
|
|
||||||
```Shell
|
```Shell
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ require (
|
|||||||
github.com/eclipse/paho.mqtt.golang v1.5.1
|
github.com/eclipse/paho.mqtt.golang v1.5.1
|
||||||
github.com/getsentry/sentry-go v0.46.2
|
github.com/getsentry/sentry-go v0.46.2
|
||||||
github.com/glebarez/sqlite v1.11.0
|
github.com/glebarez/sqlite v1.11.0
|
||||||
|
github.com/golang-jwt/jwt/v5 v5.3.1
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/gorilla/mux v1.8.1
|
github.com/gorilla/mux v1.8.1
|
||||||
github.com/gorilla/websocket v1.5.3
|
github.com/gorilla/websocket v1.5.3
|
||||||
@@ -32,13 +33,13 @@ require (
|
|||||||
github.com/uptrace/bun/driver/sqliteshim v1.2.16
|
github.com/uptrace/bun/driver/sqliteshim v1.2.16
|
||||||
github.com/uptrace/bunrouter v1.0.23
|
github.com/uptrace/bunrouter v1.0.23
|
||||||
go.mongodb.org/mongo-driver v1.17.9
|
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 v1.43.0
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc 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/sdk v1.44.0
|
||||||
go.opentelemetry.io/otel/trace v1.43.0
|
go.opentelemetry.io/otel/trace v1.44.0
|
||||||
go.uber.org/zap v1.28.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/oauth2 v0.36.0
|
||||||
golang.org/x/time v0.15.0
|
golang.org/x/time v0.15.0
|
||||||
gorm.io/driver/postgres v1.6.0
|
gorm.io/driver/postgres v1.6.0
|
||||||
@@ -61,7 +62,6 @@ require (
|
|||||||
github.com/containerd/platforms v0.2.1 // indirect
|
github.com/containerd/platforms v0.2.1 // indirect
|
||||||
github.com/cpuguy83/dockercfg v0.3.2 // indirect
|
github.com/cpuguy83/dockercfg v0.3.2 // indirect
|
||||||
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
|
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
|
||||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
|
|
||||||
github.com/distribution/reference v0.6.0 // indirect
|
github.com/distribution/reference v0.6.0 // indirect
|
||||||
github.com/docker/docker v28.5.1+incompatible // indirect
|
github.com/docker/docker v28.5.1+incompatible // indirect
|
||||||
github.com/docker/go-connections v0.6.0 // indirect
|
github.com/docker/go-connections v0.6.0 // indirect
|
||||||
@@ -137,22 +137,21 @@ require (
|
|||||||
github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 // indirect
|
github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 // indirect
|
||||||
github.com/yusufpapurcu/wmi v1.2.4 // indirect
|
github.com/yusufpapurcu/wmi v1.2.4 // indirect
|
||||||
go.opentelemetry.io/auto/sdk v1.2.1 // indirect
|
go.opentelemetry.io/auto/sdk v1.2.1 // indirect
|
||||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 // indirect
|
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0 // indirect
|
||||||
go.opentelemetry.io/otel/metric v1.43.0 // indirect
|
go.opentelemetry.io/otel/metric v1.44.0 // indirect
|
||||||
go.opentelemetry.io/proto/otlp v1.10.0 // indirect
|
go.opentelemetry.io/proto/otlp v1.10.0 // indirect
|
||||||
go.uber.org/atomic v1.11.0 // indirect
|
go.uber.org/atomic v1.11.0 // indirect
|
||||||
go.uber.org/multierr v1.11.0 // indirect
|
go.uber.org/multierr v1.11.0 // indirect
|
||||||
go.yaml.in/yaml/v2 v2.4.4 // indirect
|
go.yaml.in/yaml/v2 v2.4.4 // indirect
|
||||||
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
||||||
golang.org/x/exp v0.0.0-20260508232706-74f9aab9d74a // indirect
|
golang.org/x/mod v0.38.0 // indirect
|
||||||
golang.org/x/mod v0.36.0 // indirect
|
golang.org/x/net v0.58.0 // indirect
|
||||||
golang.org/x/net v0.54.0 // indirect
|
golang.org/x/sync v0.22.0 // indirect
|
||||||
golang.org/x/sync v0.20.0 // indirect
|
golang.org/x/sys v0.47.0 // indirect
|
||||||
golang.org/x/sys v0.44.0 // indirect
|
golang.org/x/text v0.41.0 // indirect
|
||||||
golang.org/x/text v0.37.0 // indirect
|
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||||
google.golang.org/genproto/googleapis/api v0.0.0-20260519071638-aa98bba5eb94 // indirect
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260519071638-aa98bba5eb94 // indirect
|
google.golang.org/grpc v1.83.2 // indirect
|
||||||
google.golang.org/grpc v1.81.1 // indirect
|
|
||||||
google.golang.org/protobuf v1.36.11 // indirect
|
google.golang.org/protobuf v1.36.11 // indirect
|
||||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||||
modernc.org/libc v1.72.3 // 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.0/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.1/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.1/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.11.1/go.mod h1:a6xsAQUZg+VsS3TJ05SRp524Hs4pZ/AeFSr5ENf0Yjo=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.11.1/go.mod h1:a6xsAQUZg+VsS3TJ05SRp524Hs4pZ/AeFSr5ENf0Yjo=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.18.0 h1:Gt0j3wceWMwPmiazCa8MzMA0MfhmPIz0Qp0FJ6qcM0U=
|
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.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 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.3.1/go.mod h1:uE9zaUfEQT/nbQjVi2IblCG9iaLtZsuYZ8ne+PuQ02M=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.6.0/go.mod h1:9kIvujWAA58nmPmWB1m23fyWic1kYZMxD9CxaWn4Qpg=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.6.0/go.mod h1:9kIvujWAA58nmPmWB1m23fyWic1kYZMxD9CxaWn4Qpg=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.10.1 h1:B+blDbyVIG3WaikNxPnhPiJ1MThR03b3vKGtER95TP4=
|
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.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 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.3.0/go.mod h1:okt5dMMTOFjX/aovMlrjvvXoPMBVSPzk9185BT0+eZM=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.5.2/go.mod h1:yInRyqWXAuaPrgI7p70+lDDgh3mlBohis29jGMISnmc=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.5.2/go.mod h1:yInRyqWXAuaPrgI7p70+lDDgh3mlBohis29jGMISnmc=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.8.0/go.mod h1:4OG6tQ9EOP/MT0NMjDlRzWoVFxfu9rN9B2X+tlSVktg=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.8.0/go.mod h1:4OG6tQ9EOP/MT0NMjDlRzWoVFxfu9rN9B2X+tlSVktg=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.11.1 h1:FPKJS1T+clwv+OLGt13a8UjqeRuh0O4SJ3lUriThc+4=
|
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.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 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.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 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.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 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 h1:UQHMgLO+TxOElx5B5HZ4hJQsoJ/PvUvKRhJHDQXO8P8=
|
||||||
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E=
|
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.1.1/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
github.com/AzureAD/microsoft-authentication-library-for-go v1.1.1/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.2.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
github.com/AzureAD/microsoft-authentication-library-for-go v1.2.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.4.2 h1:oygO0locgZJe7PpYPXT5A29ZkwJaPqcva7BVeemZOZs=
|
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.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 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 h1:OcvFkGmslmlZibjAjaHm3L//6LiuBgolP7OputlJIzU=
|
||||||
github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU=
|
github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU=
|
||||||
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
|
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
|
||||||
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
|
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
|
||||||
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
||||||
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
||||||
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf h1:TqhNAT4zKbTdLa62d2HDBFdvgSbIGB3eJE8HqhgiL9I=
|
|
||||||
github.com/bradfitz/gomemcache v0.0.0-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 h1:6Gpm9YYUEQx2T9zMsYolQhr6sjwwGtFitSA0pQsa7a8=
|
||||||
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c/go.mod h1:r5xuitiExdLAJ09PR7vBVENGvp4ZuTBeWTGtxuX3K+c=
|
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c/go.mod h1:r5xuitiExdLAJ09PR7vBVENGvp4ZuTBeWTGtxuX3K+c=
|
||||||
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
|
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
|
||||||
@@ -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 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY=
|
||||||
github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4=
|
github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4=
|
||||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
|
||||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM=
|
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/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 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk=
|
||||||
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
|
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
|
||||||
|
github.com/dlclark/regexp2 v1.11.0 h1:G/nrcoOa7ZXlpoa/91N3X7mM3r8eIlMBBJZvsz/mxKI=
|
||||||
|
github.com/dlclark/regexp2 v1.11.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8=
|
||||||
github.com/dnaeon/go-vcr v1.1.0/go.mod h1:M7tiix8f0r6mKKJ3Yq/kqU1OYf3MnfmBWVbPx/yU9ko=
|
github.com/dnaeon/go-vcr v1.1.0/go.mod h1:M7tiix8f0r6mKKJ3Yq/kqU1OYf3MnfmBWVbPx/yU9ko=
|
||||||
github.com/dnaeon/go-vcr v1.2.0/go.mod h1:R4UdLID7HZT3taECzJs4YgbbH6PIGXB6W/sc5OLb6RQ=
|
github.com/dnaeon/go-vcr v1.2.0/go.mod h1:R4UdLID7HZT3taECzJs4YgbbH6PIGXB6W/sc5OLb6RQ=
|
||||||
github.com/docker/docker v28.5.1+incompatible h1:Bm8DchhSD2J6PsFzxC35TZo4TLGR2PdW/E69rU45NhM=
|
github.com/docker/docker v28.5.1+incompatible h1:Bm8DchhSD2J6PsFzxC35TZo4TLGR2PdW/E69rU45NhM=
|
||||||
@@ -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/felixge/httpsnoop v1.0.4/go.mod h1:m8KPJKqk1gH5J9DgRY2ASl2lWCfGKXixSwevea8zH2U=
|
||||||
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
|
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
|
||||||
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
||||||
github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k=
|
|
||||||
github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0=
|
|
||||||
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
|
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
|
||||||
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
|
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 h1:1jhYwrKGa3sIpo/y5iDNXS5wDoT7I1KNzMHrnK6ojns=
|
||||||
github.com/getsentry/sentry-go v0.46.2/go.mod h1:evVbw2qotNUdYG8KxXbAdjOQWWvWIwKxpjdZZIvcIPw=
|
github.com/getsentry/sentry-go v0.46.2/go.mod h1:evVbw2qotNUdYG8KxXbAdjOQWWvWIwKxpjdZZIvcIPw=
|
||||||
github.com/glebarez/go-sqlite v1.22.0 h1:uAcMJhaA6r3LHMTFgP0SifzgXg46yJkgxqyuyec+ruQ=
|
github.com/glebarez/go-sqlite v1.22.0 h1:uAcMJhaA6r3LHMTFgP0SifzgXg46yJkgxqyuyec+ruQ=
|
||||||
@@ -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-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
|
||||||
github.com/go-ole/go-ole v1.2.6 h1:/Fpf6oFPoeFik9ty7siob0G6Ke8QvQEuVcuChpwXzpY=
|
github.com/go-ole/go-ole v1.2.6 h1:/Fpf6oFPoeFik9ty7siob0G6Ke8QvQEuVcuChpwXzpY=
|
||||||
github.com/go-ole/go-ole v1.2.6/go.mod h1:pprOEPIfldk/42T2oK7lQ4v4JSDwmV0As9GaiUsvbm0=
|
github.com/go-ole/go-ole v1.2.6/go.mod h1:pprOEPIfldk/42T2oK7lQ4v4JSDwmV0As9GaiUsvbm0=
|
||||||
github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs=
|
|
||||||
github.com/go-viper/mapstructure/v2 v2.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 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro=
|
||||||
github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
|
github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
|
||||||
github.com/golang-jwt/jwt/v5 v5.0.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
github.com/golang-jwt/jwt/v5 v5.0.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||||
github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||||
github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo=
|
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.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 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 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
|
||||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0=
|
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0=
|
||||||
github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A=
|
github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A=
|
||||||
@@ -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.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||||
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||||
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
||||||
github.com/google/jsonschema-go v0.4.2 h1:tmrUohrwoLZZS/P3x7ex0WAVknEkBZM46iALbcqoRA8=
|
|
||||||
github.com/google/jsonschema-go v0.4.2/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE=
|
|
||||||
github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0=
|
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/jsonschema-go v0.4.3/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE=
|
||||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
|
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
|
||||||
@@ -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/sessions v1.2.1/go.mod h1:dk2InVEVJ0sfLlnXv9EAgkf6ecYs/i80K/zI+bUmuGM=
|
||||||
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
||||||
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
||||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.2 h1:8Tjv8EJ+pM1xP8mK6egEbD1OgnVTyacbefKhmbLhIhU=
|
|
||||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.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 h1:5VipnvEpbqr2gA2VbM+nYVbkIF28c5ZQfqCBQ5g2xfk=
|
||||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0/go.mod h1:Hyl3n6Twe1hvtd9XUXDec4pTvgMSEixRuQKPTMH2bNs=
|
github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0/go.mod h1:Hyl3n6Twe1hvtd9XUXDec4pTvgMSEixRuQKPTMH2bNs=
|
||||||
github.com/hashicorp/go-uuid v1.0.2/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro=
|
github.com/hashicorp/go-uuid v1.0.2/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro=
|
||||||
@@ -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/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
||||||
github.com/jackc/pgx/v5 v5.8.0 h1:TYPDoleBBme0xGSAX3/+NujXXtpZn9HBONkQC7IEZSo=
|
|
||||||
github.com/jackc/pgx/v5 v5.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 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw=
|
||||||
github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
|
github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
|
||||||
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
||||||
@@ -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 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
|
||||||
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
|
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
|
||||||
github.com/kisielk/sqlstruct v0.0.0-20201105191214-5f3e10d3ab46/go.mod h1:yyMNCyc/Ib3bDTKd379tNMpB/7/H5TjM2Y9QJ5THLbE=
|
github.com/kisielk/sqlstruct v0.0.0-20201105191214-5f3e10d3ab46/go.mod h1:yyMNCyc/Ib3bDTKd379tNMpB/7/H5TjM2Y9QJ5THLbE=
|
||||||
github.com/klauspost/compress v1.18.2 h1:iiPHWW0YrcFgpBYhsA6D1+fqHssJscY/Tm/y2Uqnapk=
|
|
||||||
github.com/klauspost/compress v1.18.2/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4=
|
|
||||||
github.com/klauspost/compress v1.18.6 h1:2jupLlAwFm95+YDR+NwD2MEfFO9d4z4Prjl1XXDjuao=
|
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/compress v1.18.6/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||||
|
github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE=
|
||||||
|
github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
|
||||||
github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
|
github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
|
||||||
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
||||||
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
||||||
@@ -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/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0/go.mod h1:zJYVVT2jmtg6P3p1VtQj7WsuWi/y4VnjVBn7F8KPB3I=
|
||||||
github.com/magiconair/properties v1.8.10 h1:s31yESBquKXCV9a/ScB3ESkOjUYYv+X0rg8SYxI99mE=
|
github.com/magiconair/properties v1.8.10 h1:s31yESBquKXCV9a/ScB3ESkOjUYYv+X0rg8SYxI99mE=
|
||||||
github.com/magiconair/properties v1.8.10/go.mod h1:Dhd985XPs7jluiymwWYZ0G4Z61jb3vdS329zhj2hYo0=
|
github.com/magiconair/properties v1.8.10/go.mod h1:Dhd985XPs7jluiymwWYZ0G4Z61jb3vdS329zhj2hYo0=
|
||||||
github.com/mark3labs/mcp-go v0.46.0 h1:8KRibF4wcKejbLsHxCA/QBVUr5fQ9nwz/n8lGqmaALo=
|
|
||||||
github.com/mark3labs/mcp-go v0.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 h1:PZhQvd+5xrT43cUoiaKn/hDcvLUhcLc1twSEKYPTcTA=
|
||||||
github.com/mark3labs/mcp-go v0.54.0/go.mod h1:+8WclSK1ZUweCP3hvktSji8n8ABG/95QaEkeVE/Uwas=
|
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 h1:j8l17JJ9i6VGPUFUYoTUKPSgKe/83EYU2zBC7YNKMw4=
|
||||||
github.com/mattn/go-isatty v0.0.22/go.mod h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4=
|
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 h1:3VSe+xafpbzsLbdr2AWlAZk9yRHiBhTBakioXaCKTF8=
|
||||||
github.com/mattn/go-sqlite3 v1.14.44/go.mod h1:pjEuOr8IwzLJP2MfGeTb0A35jauH+C2kbHKBr7yXKVQ=
|
github.com/mattn/go-sqlite3 v1.14.44/go.mod h1:pjEuOr8IwzLJP2MfGeTb0A35jauH+C2kbHKBr7yXKVQ=
|
||||||
github.com/microsoft/go-mssqldb v1.8.2/go.mod h1:vp38dT33FGfVotRiTmDo3bFyaHq+p3LektQrjTULowo=
|
github.com/microsoft/go-mssqldb v1.8.2/go.mod h1:vp38dT33FGfVotRiTmDo3bFyaHq+p3LektQrjTULowo=
|
||||||
github.com/microsoft/go-mssqldb v1.9.5 h1:orwya0X/5bsL1o+KasupTkk2eNTNFkTQG0BEe/HxCn0=
|
|
||||||
github.com/microsoft/go-mssqldb v1.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 h1:pHEt+Qz6YFPWqREq10mqSE524QQo+/QremwTCQht7TY=
|
||||||
github.com/microsoft/go-mssqldb v1.10.0/go.mod h1:mnG7lGa9iYJbzJqGCXyuQCegStKMr3kogDLD6+bmggg=
|
github.com/microsoft/go-mssqldb v1.10.0/go.mod h1:mnG7lGa9iYJbzJqGCXyuQCegStKMr3kogDLD6+bmggg=
|
||||||
github.com/moby/docker-image-spec v1.3.1 h1:jMKff3w6PgbfSa69GfNg+zN/XLhfXJGnEx3Nl2EsFP0=
|
github.com/moby/docker-image-spec v1.3.1 h1:jMKff3w6PgbfSa69GfNg+zN/XLhfXJGnEx3Nl2EsFP0=
|
||||||
@@ -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/mochi-mqtt/server/v2 v2.7.9/go.mod h1:lZD3j35AVNqJL5cezlnSkuG05c0FCHSsfAKSPBOSbqc=
|
||||||
github.com/modocache/gover v0.0.0-20171022184752-b58185e213c5/go.mod h1:caMODM3PzxT8aQXRPkAt8xlV/e7d7w8GM5g0fa5F0D8=
|
github.com/modocache/gover v0.0.0-20171022184752-b58185e213c5/go.mod h1:caMODM3PzxT8aQXRPkAt8xlV/e7d7w8GM5g0fa5F0D8=
|
||||||
github.com/montanaflynn/stats v0.7.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
github.com/montanaflynn/stats v0.7.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
||||||
github.com/montanaflynn/stats v0.7.1 h1:etflOAAHORrCC44V+aR6Ftzort912ZU+YLiSTuV8eaE=
|
|
||||||
github.com/montanaflynn/stats v0.7.1/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
|
||||||
github.com/montanaflynn/stats v0.9.0 h1:tsBJ0RXwph9BmAuFoCmqGv6e8xa0MENQ8m0ptKq29mQ=
|
github.com/montanaflynn/stats v0.9.0 h1:tsBJ0RXwph9BmAuFoCmqGv6e8xa0MENQ8m0ptKq29mQ=
|
||||||
github.com/montanaflynn/stats v0.9.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
github.com/montanaflynn/stats v0.9.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
||||||
github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A=
|
github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A=
|
||||||
github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc=
|
github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc=
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
||||||
github.com/nats-io/nats.go v1.48.0 h1:pSFyXApG+yWU/TgbKCjmm5K4wrHu86231/w84qRVR+U=
|
|
||||||
github.com/nats-io/nats.go v1.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 h1:n3avV4VBsCgsdwh71TppsTwtv+QdPs7ntSKM8qJLGsc=
|
||||||
github.com/nats-io/nats.go v1.52.0/go.mod h1:26HypzazeOkyO3/mqd1zZd53STJN0EjCYF9Uy2ZOBno=
|
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 h1:JACV5jRVO9V856KOapQ7x+EY8Jo3qw1vJt/9Jpwzkk4=
|
||||||
github.com/nats-io/nkeys v0.4.15/go.mod h1:CpMchTXC9fxA5zrMo4KpySxNjiDVvr8ANOSZdiNfUrs=
|
github.com/nats-io/nkeys v0.4.15/go.mod h1:CpMchTXC9fxA5zrMo4KpySxNjiDVvr8ANOSZdiNfUrs=
|
||||||
github.com/nats-io/nuid v1.0.1 h1:5iA8DT8V7q8WK2EScv2padNa/rTESc1KdnPw4TC2paw=
|
github.com/nats-io/nuid v1.0.1 h1:5iA8DT8V7q8WK2EScv2padNa/rTESc1KdnPw4TC2paw=
|
||||||
@@ -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/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
|
||||||
github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040=
|
github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040=
|
||||||
github.com/opencontainers/image-spec v1.1.1/go.mod h1:qpqAh3Dmcf36wStyyWU+kCeDgrGnAve2nCC8+7h8Q0M=
|
github.com/opencontainers/image-spec v1.1.1/go.mod h1:qpqAh3Dmcf36wStyyWU+kCeDgrGnAve2nCC8+7h8Q0M=
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4=
|
|
||||||
github.com/pelletier/go-toml/v2 v2.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 h1:MYEvvGnQjeNkRF1qUuGolNtNExTDwct51yp7olPtrEc=
|
||||||
github.com/pelletier/go-toml/v2 v2.3.1/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
|
github.com/pelletier/go-toml/v2 v2.3.1/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
|
||||||
github.com/pingcap/errors v0.11.4 h1:lFuQV/oaUMGcD2tqt+01ROSmJs75VG1ToEOkZIZ4nE4=
|
github.com/pingcap/errors v0.11.4 h1:lFuQV/oaUMGcD2tqt+01ROSmJs75VG1ToEOkZIZ4nE4=
|
||||||
@@ -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_golang v1.23.2/go.mod h1:Tb1a6LWHB3/SPIzCoaDXI4I8UHKeFTEQ1YCr+0Gyqmg=
|
||||||
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
||||||
github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
|
github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
|
||||||
github.com/prometheus/common v0.67.4 h1:yR3NqWO1/UyO1w2PhUvXlGQs/PtFmoveVO0KZ4+Lvsc=
|
|
||||||
github.com/prometheus/common v0.67.4/go.mod h1:gP0fq6YjjNCLssJCQp0yk4M8W6ikLURwkdd/YKtTbyI=
|
|
||||||
github.com/prometheus/common v0.67.5 h1:pIgK94WWlQt1WLwAC5j2ynLaBRDiinoAb86HZHTUGI4=
|
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/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 h1:XwbrGOIplXW/AU3YhIhLODXMJYyC1isLFfYCsTEycfc=
|
||||||
github.com/prometheus/procfs v0.20.1/go.mod h1:o9EMBZGRyvDrSPH1RqdxhojkuXstoe4UlK79eF5TGGo=
|
github.com/prometheus/procfs v0.20.1/go.mod h1:o9EMBZGRyvDrSPH1RqdxhojkuXstoe4UlK79eF5TGGo=
|
||||||
github.com/puzpuzpuz/xsync/v3 v3.5.1 h1:GJYJZwO6IdxN/IKbneznS6yPkVC+c3zyY/j19c++5Fg=
|
github.com/puzpuzpuz/xsync/v3 v3.5.1 h1:GJYJZwO6IdxN/IKbneznS6yPkVC+c3zyY/j19c++5Fg=
|
||||||
github.com/puzpuzpuz/xsync/v3 v3.5.1/go.mod h1:VjzYrABPabuM4KyBh1Ftq6u8nhwY5tBPKP9jpmh0nnA=
|
github.com/puzpuzpuz/xsync/v3 v3.5.1/go.mod h1:VjzYrABPabuM4KyBh1Ftq6u8nhwY5tBPKP9jpmh0nnA=
|
||||||
github.com/redis/go-redis/v9 v9.17.2 h1:P2EGsA4qVIM3Pp+aPocCJ7DguDHhqrXNhVcEp4ViluI=
|
|
||||||
github.com/redis/go-redis/v9 v9.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 h1:XPVaaPSnG6RhYf7p+rmSa9zZfeVAnWsH5h3lxthOm/k=
|
||||||
github.com/redis/go-redis/v9 v9.19.0/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA=
|
github.com/redis/go-redis/v9 v9.19.0/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA=
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
||||||
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
||||||
github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4=
|
github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4=
|
||||||
github.com/rogpeppe/go-internal v1.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII=
|
|
||||||
github.com/rogpeppe/go-internal v1.13.1/go.mod h1:uMEvuHeurkdAXX61udpOXGD/AzZDWNMNyH2VO9fmH0o=
|
|
||||||
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
|
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
|
||||||
github.com/rs/xid v1.4.0 h1:qd7wPTDkN6KQx2VmMBLrpHkiyQwgFXRnkOLacUiaSNY=
|
github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc=
|
||||||
github.com/rs/xid v1.4.0/go.mod h1:trrq9SKmegXys3aeAKXMUTdJsYXVwGY3RLcfgqegfbg=
|
|
||||||
github.com/rs/xid v1.6.0 h1:fV591PaemRlL6JfRxGDEPl69wICngIQ3shQtzfy2gxU=
|
github.com/rs/xid v1.6.0 h1:fV591PaemRlL6JfRxGDEPl69wICngIQ3shQtzfy2gxU=
|
||||||
github.com/rs/xid v1.6.0/go.mod h1:7XoLgs4eV+QndskICGsho+ADou8ySMSjJKDIan90Nz0=
|
github.com/rs/xid v1.6.0/go.mod h1:7XoLgs4eV+QndskICGsho+ADou8ySMSjJKDIan90Nz0=
|
||||||
github.com/sagikazarmark/locafero v0.12.0 h1:/NQhBAkUb4+fH1jivKHWusDYFjMOOKU88eegjfxfHb4=
|
github.com/sagikazarmark/locafero v0.12.0 h1:/NQhBAkUb4+fH1jivKHWusDYFjMOOKU88eegjfxfHb4=
|
||||||
@@ -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 h1:pSdJYLOVgLE8YdUY2FHQ1Fxu+aMnb6JfVz1mxk7OeMU=
|
||||||
github.com/testcontainers/testcontainers-go v0.40.0/go.mod h1:FSXV5KQtX2HAMlm7U3APNyLkkap35zNLxukw9oBi/MY=
|
github.com/testcontainers/testcontainers-go v0.40.0/go.mod h1:FSXV5KQtX2HAMlm7U3APNyLkkap35zNLxukw9oBi/MY=
|
||||||
github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
||||||
github.com/tidwall/gjson v1.18.0 h1:FIDeeyB800efLX89e5a8Y0BNH+LOngJyGrIWxG2FKQY=
|
|
||||||
github.com/tidwall/gjson v1.18.0/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
|
||||||
github.com/tidwall/gjson v1.19.0 h1:xwxm7n691Uf3u5OFjzngavjGTh55KX5q/9w9xHW88JU=
|
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/gjson v1.19.0/go.mod h1:V37/opeE/JbLUOfH0QTXiNez2l0RUjYUhpT4szFQAfc=
|
||||||
github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
||||||
@@ -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/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 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.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 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.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 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.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 h1:M6Dh5kkDWFbUWBrOsIE1g1zdZ5JbSytTD4piFRBOUAI=
|
||||||
github.com/uptrace/bun/driver/sqliteshim v1.2.16/go.mod h1:iKdJ06P3XS+pwKcONjSIK07bbhksH3lWsw3mpfr0+bY=
|
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=
|
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/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
|
||||||
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
|
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
|
||||||
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
||||||
go.mongodb.org/mongo-driver v1.17.6 h1:87JUG1wZfWsr6rIz3ZmpH90rL5tea7O3IHuSwHUpsss=
|
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
|
||||||
go.mongodb.org/mongo-driver v1.17.6/go.mod h1:Hy04i7O2kC4RS06ZrhPRqj/u4DTYkFDAAccj+rVKqgQ=
|
github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s=
|
||||||
go.mongodb.org/mongo-driver v1.17.9 h1:IexDdCuuNJ3BHrELgBlyaH9p60JXAvdzWR128q+U5tU=
|
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.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 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
|
||||||
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
|
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.61.0 h1:F7Jx+6hwnZ41NSFTO5q4LYDtJRXBf2PD0rNBkeB/lus=
|
||||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0/go.mod h1:p8pYQP+m5XfbZm9fxtSKAbM6oIllS7s2AfxrChvc7iw=
|
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.61.0/go.mod h1:UHB22Z8QsdRDrnAtX4PntOl36ajSxcdUMt1sF7Y6E7Q=
|
||||||
go.opentelemetry.io/otel v1.38.0 h1:RkfdswUDRimDg0m2Az18RKOsnI8UDzppJAtj01/Ymk8=
|
go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU=
|
||||||
go.opentelemetry.io/otel v1.38.0/go.mod h1:zcmtmQ1+YmQM9wrNsTGV/q/uyusom3P8RxwExxkZhjM=
|
go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc=
|
||||||
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/otel/exporters/otlp/otlptrace v1.43.0 h1:88Y4s2C8oTui1LGM6bTWkw0ICGcOLCAI5l6zsD1j20k=
|
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 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 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/otlptracegrpc v1.43.0/go.mod h1:AGmbycVGEsRx9mXMZ75CsOyhSP6MFIcj/6dnG+vhVjk=
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.19.0 h1:IeMeyr1aBvBiPVYihXIaeIZba6b8E1bYp7lbdxK8CQg=
|
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.19.0 h1:IeMeyr1aBvBiPVYihXIaeIZba6b8E1bYp7lbdxK8CQg=
|
||||||
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.19.0/go.mod h1:oVdCUtjq9MK9BlS7TtucsQwUcXcymNiEDjgDD2jMtZU=
|
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.19.0/go.mod h1:oVdCUtjq9MK9BlS7TtucsQwUcXcymNiEDjgDD2jMtZU=
|
||||||
go.opentelemetry.io/otel/metric v1.38.0 h1:Kl6lzIYGAh5M159u9NgiRkmoMKjvbsKtYRwgfrA6WpA=
|
go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc=
|
||||||
go.opentelemetry.io/otel/metric v1.38.0/go.mod h1:kB5n/QoRM8YwmUahxvI3bO34eVtQf2i4utNVLr9gEmI=
|
go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo=
|
||||||
go.opentelemetry.io/otel/metric v1.43.0 h1:d7638QeInOnuwOONPp4JAOGfbCEpYb+K6DVWvdxGzgM=
|
go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58=
|
||||||
go.opentelemetry.io/otel/metric v1.43.0/go.mod h1:RDnPtIxvqlgO8GRW18W6Z/4P462ldprJtfxHxyKd2PY=
|
go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0=
|
||||||
go.opentelemetry.io/otel/sdk v1.38.0 h1:l48sr5YbNf2hpCUj/FoGhW9yDkl+Ma+LrVl8qaM5b+E=
|
go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI=
|
||||||
go.opentelemetry.io/otel/sdk v1.38.0/go.mod h1:ghmNdGlVemJI3+ZB5iDEuk4bWA3GkTpW+DOoZMYBVVg=
|
go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA=
|
||||||
go.opentelemetry.io/otel/sdk v1.43.0 h1:pi5mE86i5rTeLXqoF/hhiBtUNcrAGHLKQdhg4h4V9Dg=
|
go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk=
|
||||||
go.opentelemetry.io/otel/sdk v1.43.0/go.mod h1:P+IkVU3iWukmiit/Yf9AWvpyRDlUeBaRg6Y+C58QHzg=
|
go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE=
|
||||||
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/proto/otlp v1.10.0 h1:IQRWgT5srOCYfiWnpqUYz9CVmbO8bFmKcwYxpuCSL2g=
|
go.opentelemetry.io/proto/otlp v1.10.0 h1:IQRWgT5srOCYfiWnpqUYz9CVmbO8bFmKcwYxpuCSL2g=
|
||||||
go.opentelemetry.io/proto/otlp v1.10.0/go.mod h1:/CV4QoCR/S9yaPj8utp3lvQPoqMtxXdzn7ozvvozVqk=
|
go.opentelemetry.io/proto/otlp v1.10.0/go.mod h1:/CV4QoCR/S9yaPj8utp3lvQPoqMtxXdzn7ozvvozVqk=
|
||||||
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
|
go.uber.org/atomic v1.11.0 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/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||||
go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0=
|
go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0=
|
||||||
go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
||||||
go.uber.org/zap v1.27.1 h1:08RqriUEv8+ArZRYSTXy1LeBScaMpVSTBhCeaZYfMYc=
|
|
||||||
go.uber.org/zap v1.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 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo=
|
||||||
go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q=
|
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 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
||||||
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
||||||
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
|
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
|
||||||
@@ -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.22.0/go.mod h1:vr6Su+7cTlO45qkww3VDJlzDn0ctJvRgYbC2NvXHt+M=
|
||||||
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
||||||
golang.org/x/crypto v0.24.0/go.mod h1:Z1PMYSOR5nyMcyAVAIQSKCDwalqy85Aqn1x3Ws4L5DM=
|
golang.org/x/crypto v0.24.0/go.mod h1:Z1PMYSOR5nyMcyAVAIQSKCDwalqy85Aqn1x3Ws4L5DM=
|
||||||
golang.org/x/crypto v0.46.0 h1:cKRW/pmt1pKAfetfu+RCEvjvZkA9RimPbh7bhFjGVBU=
|
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
||||||
golang.org/x/crypto v0.46.0/go.mod h1:Evb/oLKmMraqjZ2iQTwDwvCtJkczlDuTmdJXoZVzqU0=
|
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
||||||
golang.org/x/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/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
|
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
|
||||||
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||||
golang.org/x/mod v0.9.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
golang.org/x/mod v0.9.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||||
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||||
golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||||
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||||
golang.org/x/mod v0.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
|
golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk=
|
||||||
golang.org/x/mod v0.31.0/go.mod h1:43JraMp9cGx1Rx3AqioxrbrhNsLl2l/iNAvuBkrezpg=
|
golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40=
|
||||||
golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4=
|
|
||||||
golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ=
|
|
||||||
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||||
golang.org/x/net v0.0.0-20200114155413-6afb5195e5aa/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
golang.org/x/net v0.0.0-20200114155413-6afb5195e5aa/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||||
@@ -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.24.0/go.mod h1:2Q7sJY5mzlzWjKtYUEXSlBWCdyaioyXzRB2RtU8KVE8=
|
||||||
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
||||||
golang.org/x/net v0.26.0/go.mod h1:5YKkiSynbBIh3p6iOc/vibscux0x38BZDkn8sCUPxHE=
|
golang.org/x/net v0.26.0/go.mod h1:5YKkiSynbBIh3p6iOc/vibscux0x38BZDkn8sCUPxHE=
|
||||||
golang.org/x/net v0.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU=
|
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
|
||||||
golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY=
|
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
|
||||||
golang.org/x/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/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
|
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/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q=
|
||||||
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
@@ -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.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||||
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||||
golang.org/x/sync v0.9.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
golang.org/x/sync v0.9.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||||
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
|
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||||
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||||
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
|
||||||
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
|
||||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
@@ -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.19.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||||
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||||
golang.org/x/sys v0.21.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
golang.org/x/sys v0.21.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||||
golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk=
|
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||||
golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||||
golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ=
|
|
||||||
golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
|
||||||
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
|
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
|
||||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||||
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
|
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
|
||||||
@@ -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.19.0/go.mod h1:2CuTdWZ7KHSQwUzKva0cbMg6q2DMI3Mmxp+gKJbskEk=
|
||||||
golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
|
golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
|
||||||
golang.org/x/term v0.21.0/go.mod h1:ooXLefLobQVslOqselCNF4SxFAaoS6KujMbsGzSDmX0=
|
golang.org/x/term v0.21.0/go.mod h1:ooXLefLobQVslOqselCNF4SxFAaoS6KujMbsGzSDmX0=
|
||||||
golang.org/x/term v0.38.0 h1:PQ5pkm/rLO6HnxFR7N2lJHOZX6Kez5Y1gDSJla6jo7Q=
|
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
||||||
golang.org/x/term v0.38.0/go.mod h1:bSEAKrOT1W+VSu9TSCMtoGEOUcKxOKgl3LE5QEF/xVg=
|
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
|
||||||
golang.org/x/term v0.43.0 h1:S4RLU2sB31O/NCl+zFN9Aru9A/Cq2aqKpTZJ6B+DwT4=
|
|
||||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||||
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
||||||
@@ -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.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||||
golang.org/x/text v0.16.0/go.mod h1:GhwF1Be+LQoKShO3cGOHzqOgRrGaYc9AvblQOmPVHnI=
|
golang.org/x/text v0.16.0/go.mod h1:GhwF1Be+LQoKShO3cGOHzqOgRrGaYc9AvblQOmPVHnI=
|
||||||
golang.org/x/text v0.20.0/go.mod h1:D4IsuqiFMhST5bX19pQ9ikHC2GsaKyk/oF+pn3ducp4=
|
golang.org/x/text v0.20.0/go.mod h1:D4IsuqiFMhST5bX19pQ9ikHC2GsaKyk/oF+pn3ducp4=
|
||||||
golang.org/x/text v0.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU=
|
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
|
||||||
golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY=
|
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
|
||||||
golang.org/x/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/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
|
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
|
||||||
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
|
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
|
||||||
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
||||||
@@ -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.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
|
||||||
golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58=
|
golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58=
|
||||||
golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
|
golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
|
||||||
golang.org/x/tools v0.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA=
|
golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE=
|
||||||
golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc=
|
golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk=
|
||||||
golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8=
|
|
||||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk=
|
|
||||||
gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E=
|
|
||||||
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
|
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
|
||||||
google.golang.org/genproto/googleapis/api v0.0.0-20250825161204-c5933d9347a5 h1:BIRfGDEjiHRrk0QKZe3Xv2ieMhtgRGeLcZQ0mIVn4EY=
|
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
|
||||||
google.golang.org/genproto/googleapis/api v0.0.0-20250825161204-c5933d9347a5/go.mod h1:j3QtIyytwqGr1JUDtYXwtMXWPKsEa5LtzIFN1Wn5WvE=
|
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa h1:Kjn0N0tCrDgiAFW+lGO4JZ3ck44CehvJQMAwj9QF0G8=
|
||||||
google.golang.org/genproto/googleapis/api v0.0.0-20260519071638-aa98bba5eb94 h1:DddG61lE5LkX6144z22i0gma9BMBs5aZ9B8lZLobxyw=
|
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:q4lMZS6kskjT5HvCPrnnypcDPVJqT/f4nfxmkE7gryY=
|
||||||
google.golang.org/genproto/googleapis/api v0.0.0-20260519071638-aa98bba5eb94/go.mod h1:1dCETSCY2YKZNXQE3h4fun3TYwF5p8jejRKZgfWAgAY=
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk=
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20250825161204-c5933d9347a5 h1:eaY8u2EuxbRv7c3NiGK0/NedzVsCcV6hDuU5qPX5EGE=
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20250825161204-c5933d9347a5/go.mod h1:M4/wBTSeyLxupu3W3tJtOgB14jILAS/XWPSSa3TAlJc=
|
google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU=
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260519071638-aa98bba5eb94 h1:eZCjr/aAF8c5ccm5pb6T4EXgIei5MlAAPWPJk+5ArfY=
|
google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8=
|
||||||
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=
|
|
||||||
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
||||||
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||||
@@ -644,37 +533,28 @@ gorm.io/gorm v1.31.1 h1:7CA8FTFz/gRfgqgpeKIBcervUn3xSyPUmr6B2WXJ7kg=
|
|||||||
gorm.io/gorm v1.31.1/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs=
|
gorm.io/gorm v1.31.1/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs=
|
||||||
gotest.tools/v3 v3.5.2 h1:7koQfIKdy+I8UTetycgUqXWSDwpgv193Ka+qRsmBY8Q=
|
gotest.tools/v3 v3.5.2 h1:7koQfIKdy+I8UTetycgUqXWSDwpgv193Ka+qRsmBY8Q=
|
||||||
gotest.tools/v3 v3.5.2/go.mod h1:LtdLGcnqToBH83WByAAi/wiwSFCArdFIUV/xxN4pcjA=
|
gotest.tools/v3 v3.5.2/go.mod h1:LtdLGcnqToBH83WByAAi/wiwSFCArdFIUV/xxN4pcjA=
|
||||||
modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis=
|
|
||||||
modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
|
|
||||||
modernc.org/cc/v4 v4.28.2 h1:3tQ0lf2ADtoby2EtSP+J7IE2SHwEJdP8ioR59wx7XpY=
|
modernc.org/cc/v4 v4.28.2 h1:3tQ0lf2ADtoby2EtSP+J7IE2SHwEJdP8ioR59wx7XpY=
|
||||||
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
|
modernc.org/cc/v4 v4.28.2/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI=
|
||||||
modernc.org/ccgo/v4 v4.30.1/go.mod h1:bIOeI1JL54Utlxn+LwrFyjCx2n2RDiYEaJVSrgdrRfM=
|
|
||||||
modernc.org/ccgo/v4 v4.34.0 h1:yRLPFZieg532OT4rp4JFNIVcquwalMX26G95WQDqwCQ=
|
modernc.org/ccgo/v4 v4.34.0 h1:yRLPFZieg532OT4rp4JFNIVcquwalMX26G95WQDqwCQ=
|
||||||
modernc.org/fileutil v1.3.40 h1:ZGMswMNc9JOCrcrakF1HrvmergNLAmxOPjizirpfqBA=
|
modernc.org/ccgo/v4 v4.34.0/go.mod h1:AS5WYMyBakQ+fhsHhtP8mWB82KTGPkNNJDGfGQCe0/A=
|
||||||
modernc.org/fileutil v1.3.40/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc=
|
|
||||||
modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM=
|
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 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI=
|
||||||
modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
|
modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
|
||||||
modernc.org/gc/v3 v3.1.1 h1:k8T3gkXWY9sEiytKhcgyiZ2L0DTyCQ/nvX+LoCljoRE=
|
|
||||||
modernc.org/gc/v3 v3.1.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 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 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks=
|
||||||
modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI=
|
modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI=
|
||||||
modernc.org/libc v1.67.4 h1:zZGmCMUVPORtKv95c2ReQN5VDjvkoRm9GWPTEPuvlWg=
|
|
||||||
modernc.org/libc v1.67.4/go.mod h1:QvvnnJ5P7aitu0ReNpVIEyesuhmDLQ8kaEoyMjIFZJA=
|
|
||||||
modernc.org/libc v1.72.3 h1:ZnDF4tXn4NBXFutMMQC4vtbTFSXhhKzR73fv0beZEAU=
|
modernc.org/libc v1.72.3 h1:ZnDF4tXn4NBXFutMMQC4vtbTFSXhhKzR73fv0beZEAU=
|
||||||
modernc.org/libc v1.72.3/go.mod h1:dn0dZNnnn1clLyvRxLxYExxiKRZIRENOfqQ8XEeg4Qs=
|
modernc.org/libc v1.72.3/go.mod h1:dn0dZNnnn1clLyvRxLxYExxiKRZIRENOfqQ8XEeg4Qs=
|
||||||
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
|
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
|
||||||
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
|
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
|
||||||
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
|
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
|
||||||
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
|
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
|
||||||
modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8=
|
|
||||||
modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
|
|
||||||
modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg=
|
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 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=
|
||||||
modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
|
modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
|
||||||
modernc.org/sqlite v1.42.2 h1:7hkZUNJvJFN2PgfUdjni9Kbvd4ef4mNLOu0B9FGxM74=
|
|
||||||
modernc.org/sqlite v1.42.2/go.mod h1:+VkC6v3pLOAE0A0uVucQEcbVW0I5nHCeDaBf+DpsQT8=
|
|
||||||
modernc.org/sqlite v1.50.1 h1:l+cQvn0sd0zJJtfygGHuQJ5AjlrwXmWPw4KP3ZMwr9w=
|
modernc.org/sqlite v1.50.1 h1:l+cQvn0sd0zJJtfygGHuQJ5AjlrwXmWPw4KP3ZMwr9w=
|
||||||
modernc.org/sqlite v1.50.1/go.mod h1:tcNzv5p84E0skkmJn038y+hWJbLQXQqEnQfeh5r2JLM=
|
modernc.org/sqlite v1.50.1/go.mod h1:tcNzv5p84E0skkmJn038y+hWJbLQXQqEnQfeh5r2JLM=
|
||||||
modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
|
modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
|
||||||
|
|||||||
@@ -339,7 +339,7 @@ func (b *BunSelectQuery) Column(columns ...string) common.SelectQuery {
|
|||||||
|
|
||||||
func (b *BunSelectQuery) ColumnExpr(query string, args ...interface{}) common.SelectQuery {
|
func (b *BunSelectQuery) ColumnExpr(query string, args ...interface{}) common.SelectQuery {
|
||||||
if len(args) > 0 {
|
if len(args) > 0 {
|
||||||
b.query = b.query.ColumnExpr(query, args)
|
b.query = b.query.ColumnExpr(query, args...)
|
||||||
} else {
|
} else {
|
||||||
b.query = b.query.ColumnExpr(query)
|
b.query = b.query.ColumnExpr(query)
|
||||||
}
|
}
|
||||||
@@ -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 {
|
b.query = b.query.Relation(relation, func(sq *bun.SelectQuery) *bun.SelectQuery {
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
err := logger.HandlePanic("BunSelectQuery.PreloadRelation", r)
|
_ = logger.HandlePanic("BunSelectQuery.PreloadRelation", r)
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
if len(apply) == 0 {
|
if len(apply) == 0 {
|
||||||
|
|||||||
@@ -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"`
|
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) {
|
func TestPgSQLAdapterRecordsSchemaEntityTableMetrics(t *testing.T) {
|
||||||
db, mock, err := sqlmock.New()
|
db, mock, err := sqlmock.New()
|
||||||
require.NoError(t, err)
|
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, "query_metrics_bun_user", calls[0].entity)
|
||||||
assert.Equal(t, "metrics_bun_users", calls[0].table)
|
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
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -979,3 +979,61 @@ t.Errorf("AddTablePrefixToColumns(%q, %q) = %q; want %q", tt.where, tt.tableName
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
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"`
|
CursorBackward string `json:"cursor_backward"`
|
||||||
FetchRowNumber *string `json:"fetch_row_number"`
|
FetchRowNumber *string `json:"fetch_row_number"`
|
||||||
|
|
||||||
|
// VectorSearch performs a pgvector nearest-neighbour ordering (KNN) and
|
||||||
|
// optionally returns the computed distance as an extra column.
|
||||||
|
VectorSearch *VectorSearchOption `json:"vector_search"`
|
||||||
|
|
||||||
// Join table aliases (used for validation of prefixed columns in filters/sorts)
|
// Join table aliases (used for validation of prefixed columns in filters/sorts)
|
||||||
// Not serialized to JSON as it's internal validation state
|
// Not serialized to JSON as it's internal validation state
|
||||||
JoinAliases []string `json:"-"`
|
JoinAliases []string `json:"-"`
|
||||||
@@ -90,6 +94,41 @@ type SortOption struct {
|
|||||||
Direction string `json:"direction"`
|
Direction string `json:"direction"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// PrimaryKeySortColumn is a sentinel SortOption.Column value that gets
|
||||||
|
// resolved to a model's actual primary key column name at query time.
|
||||||
|
// This lets a single global default sort (e.g. set once via
|
||||||
|
// Handler.SetDefaultSort) work across models with different primary keys,
|
||||||
|
// e.g. common.SortOption{Column: common.PrimaryKeySortColumn, Direction: "asc"}.
|
||||||
|
const PrimaryKeySortColumn = "$pk"
|
||||||
|
|
||||||
|
// ResolveSortColumns returns a copy of sort with any PrimaryKeySortColumn
|
||||||
|
// entries replaced by pkName. If pkName is empty, matching entries are
|
||||||
|
// dropped since there is no column to sort by.
|
||||||
|
func ResolveSortColumns(sort []SortOption, pkName string) []SortOption {
|
||||||
|
resolved := make([]SortOption, 0, len(sort))
|
||||||
|
for _, s := range sort {
|
||||||
|
if s.Column == PrimaryKeySortColumn {
|
||||||
|
if pkName == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
s.Column = pkName
|
||||||
|
}
|
||||||
|
resolved = append(resolved, s)
|
||||||
|
}
|
||||||
|
return resolved
|
||||||
|
}
|
||||||
|
|
||||||
|
// VectorSearchOption describes a pgvector KNN search: order rows by the distance
|
||||||
|
// between Column and Vector using Metric, and (when As is set) select that
|
||||||
|
// distance as an additional result column.
|
||||||
|
type VectorSearchOption struct {
|
||||||
|
Column string `json:"column"`
|
||||||
|
Vector []float32 `json:"vector"`
|
||||||
|
Metric string `json:"metric"` // "l2" (default) | "cosine" | "ip"
|
||||||
|
As string `json:"as"` // distance column alias; default "_distance"
|
||||||
|
Direction string `json:"direction"` // "asc" (default) | "desc"
|
||||||
|
}
|
||||||
|
|
||||||
type CustomOperator struct {
|
type CustomOperator struct {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
SQL string `json:"sql"`
|
SQL string `json:"sql"`
|
||||||
|
|||||||
@@ -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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// JSON-traversing references (data->>'x', data#>>'{a,b}', or the dotted
|
||||||
|
// data.x shorthand): validate the base column, and for the ambiguous
|
||||||
|
// dotted form require that the base is actually a JSON column.
|
||||||
|
if ref, isJSON := ParseColumnRef(column); isJSON {
|
||||||
|
if ref.Ambiguous && !reflection.IsJSONColumn(v.model, ref.Base) {
|
||||||
|
return fmt.Errorf("invalid column '%s': '%s' is not a JSON column", column, ref.Base)
|
||||||
|
}
|
||||||
|
if _, exists := v.validColumns[strings.ToLower(ref.Base)]; !exists {
|
||||||
|
return fmt.Errorf("invalid column '%s': column does not exist in model", column)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Extract source column name (remove JSON operators like ->> or ->)
|
// Extract source column name (remove JSON operators like ->> or ->)
|
||||||
sourceColumn := reflection.ExtractSourceColumn(column)
|
sourceColumn := reflection.ExtractSourceColumn(column)
|
||||||
|
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestExtractSourceColumn(t *testing.T) {
|
func TestExtractSourceColumn(t *testing.T) {
|
||||||
@@ -124,3 +125,35 @@ func TestValidateColumnWithJSONOperators(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestValidateColumn_JSONPathsAndDottedShorthand(t *testing.T) {
|
||||||
|
type Model struct {
|
||||||
|
ID int64 `json:"id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Data spectypes.SqlJSONB `json:"data"`
|
||||||
|
}
|
||||||
|
v := NewColumnValidator(Model{})
|
||||||
|
|
||||||
|
valid := []string{
|
||||||
|
"data->>'city'",
|
||||||
|
"data->'addr'->>'city'",
|
||||||
|
"data#>>'{addr,city}'",
|
||||||
|
"data.addr.city", // dotted shorthand, base is JSON -> allowed
|
||||||
|
"(data->>'age')::int", // cast + paren
|
||||||
|
}
|
||||||
|
for _, c := range valid {
|
||||||
|
if err := v.ValidateColumn(c); err != nil {
|
||||||
|
t.Errorf("ValidateColumn(%q) = %v, want nil", c, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
invalid := []string{
|
||||||
|
"nope->>'city'", // base column does not exist
|
||||||
|
"name.first", // dotted shorthand but 'name' is not a JSON column
|
||||||
|
}
|
||||||
|
for _, c := range invalid {
|
||||||
|
if err := v.ValidateColumn(c); err == nil {
|
||||||
|
t.Errorf("ValidateColumn(%q) = nil, want error", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -197,7 +197,7 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
|||||||
hookCtx.Tx = tx
|
hookCtx.Tx = tx
|
||||||
|
|
||||||
// Execute BeforeQueryList hook (inside transaction)
|
// Execute BeforeQueryList hook (inside transaction)
|
||||||
if err := h.hooks.Execute(BeforeQueryList, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeQueryList, hookCtx); err != nil {
|
||||||
logger.Error("BeforeQueryList hook failed: %v", err)
|
logger.Error("BeforeQueryList hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
return err
|
||||||
@@ -261,7 +261,7 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
|||||||
|
|
||||||
// Execute BeforeSQLExec hook
|
// Execute BeforeSQLExec hook
|
||||||
hookCtx.SQLQuery = sqlquery
|
hookCtx.SQLQuery = sqlquery
|
||||||
if err := h.hooks.Execute(BeforeSQLExec, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
|
||||||
logger.Error("BeforeSQLExec hook failed: %v", err)
|
logger.Error("BeforeSQLExec hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
return err
|
||||||
@@ -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))
|
w.Header().Set("Content-Range", fmt.Sprintf("items %d-%d/%d", respOffset, respOffset+len(dbobjlist), total))
|
||||||
logger.Info("Serving: Records %d of %d", len(dbobjlist), total)
|
logger.Info("Serving: Records %d of %d", len(dbobjlist), total)
|
||||||
|
|
||||||
// Execute BeforeResponse hook
|
// Execute BeforeResponse hook. The transaction has already committed by
|
||||||
|
// this point, so hooks must use the pooled connection rather than the
|
||||||
|
// now-dead tx.
|
||||||
|
hookCtx.Tx = h.db
|
||||||
hookCtx.Result = dbobjlist
|
hookCtx.Result = dbobjlist
|
||||||
hookCtx.Total = total
|
hookCtx.Total = total
|
||||||
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
||||||
@@ -560,7 +563,7 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
|||||||
hookCtx.Tx = tx
|
hookCtx.Tx = tx
|
||||||
|
|
||||||
// Execute BeforeQuery hook (inside transaction)
|
// Execute BeforeQuery hook (inside transaction)
|
||||||
if err := h.hooks.Execute(BeforeQuery, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeQuery, hookCtx); err != nil {
|
||||||
logger.Error("BeforeQuery hook failed: %v", err)
|
logger.Error("BeforeQuery hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
return err
|
||||||
@@ -579,7 +582,7 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
|||||||
sqlquery = hookCtx.SQLQuery
|
sqlquery = hookCtx.SQLQuery
|
||||||
|
|
||||||
// Execute BeforeSQLExec hook
|
// Execute BeforeSQLExec hook
|
||||||
if err := h.hooks.Execute(BeforeSQLExec, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
|
||||||
logger.Error("BeforeSQLExec hook failed: %v", err)
|
logger.Error("BeforeSQLExec hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
return err
|
||||||
@@ -631,7 +634,10 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute BeforeResponse hook
|
// Execute BeforeResponse hook. The transaction has already committed by
|
||||||
|
// this point, so hooks must use the pooled connection rather than the
|
||||||
|
// now-dead tx.
|
||||||
|
hookCtx.Tx = h.db
|
||||||
hookCtx.Result = dbobj
|
hookCtx.Result = dbobj
|
||||||
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
||||||
logger.Error("BeforeResponse hook failed: %v", err)
|
logger.Error("BeforeResponse hook failed: %v", err)
|
||||||
|
|||||||
@@ -28,6 +28,10 @@ const (
|
|||||||
|
|
||||||
// Response hooks (before response is sent)
|
// Response hooks (before response is sent)
|
||||||
BeforeResponse HookType = "before_response"
|
BeforeResponse HookType = "before_response"
|
||||||
|
|
||||||
|
// BeforeOp fires immediately before every SQL operation (query, query list, SQL exec).
|
||||||
|
// It fires at each individual SQL-operation hook point, so it runs once per statement executed.
|
||||||
|
BeforeOp HookType = "before_op"
|
||||||
)
|
)
|
||||||
|
|
||||||
// HookContext contains all the data available to a hook
|
// HookContext contains all the data available to a hook
|
||||||
@@ -130,6 +134,16 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
|
||||||
|
// (BeforeQuery, BeforeQueryList, or BeforeSQLExec). BeforeOp always runs first so it can
|
||||||
|
// observe/veto every SQL operation regardless of type.
|
||||||
|
func (r *HookRegistry) ExecuteBeforeOp(hookType HookType, ctx *HookContext) error {
|
||||||
|
if err := r.Execute(BeforeOp, ctx); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return r.Execute(hookType, ctx)
|
||||||
|
}
|
||||||
|
|
||||||
// Clear removes all hooks for the specified type
|
// Clear removes all hooks for the specified type
|
||||||
func (r *HookRegistry) Clear(hookType HookType) {
|
func (r *HookRegistry) Clear(hookType HookType) {
|
||||||
delete(r.hooks, hookType)
|
delete(r.hooks, hookType)
|
||||||
|
|||||||
@@ -51,7 +51,7 @@ func (h *Handler) ParseParameters(r *http.Request) *RequestParameters {
|
|||||||
FieldFilters: make(map[string]string),
|
FieldFilters: make(map[string]string),
|
||||||
SearchFilters: make(map[string]string),
|
SearchFilters: make(map[string]string),
|
||||||
SearchOps: make(map[string]FilterOperator),
|
SearchOps: make(map[string]FilterOperator),
|
||||||
Limit: 20, // Default limit
|
Limit: 100000, // Default limit
|
||||||
Offset: 0, // Default offset
|
Offset: 0, // Default offset
|
||||||
ResponseFormat: "simple", // Default format
|
ResponseFormat: "simple", // Default format
|
||||||
ComplexAPI: false, // Default to simple API
|
ComplexAPI: false, // Default to simple API
|
||||||
|
|||||||
@@ -71,6 +71,16 @@ func (f *funcSpecSecurityContext) GetUserID() (int, bool) {
|
|||||||
return int(f.ctx.UserContext.UserID), true
|
return int(f.ctx.UserContext.UserID), true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetUserRef returns an opaque user identifier for row security lookups.
|
||||||
|
// It returns the full *security.UserContext so providers can read JWT claims
|
||||||
|
// (e.g. a UUID subject) instead of relying on the int user ID.
|
||||||
|
func (f *funcSpecSecurityContext) GetUserRef() (any, bool) {
|
||||||
|
if f.ctx.UserContext == nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
return f.ctx.UserContext, true
|
||||||
|
}
|
||||||
|
|
||||||
func (f *funcSpecSecurityContext) GetSchema() string {
|
func (f *funcSpecSecurityContext) GetSchema() string {
|
||||||
// funcspec doesn't have a schema concept, extract from SQL query or use default
|
// funcspec doesn't have a schema concept, extract from SQL query or use default
|
||||||
return "public"
|
return "public"
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
"sync"
|
"sync"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ModelRules defines the permissions and security settings for a model
|
// 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 {
|
func GetDefaultRegistry() *DefaultModelRegistry {
|
||||||
|
for i := 0; i < lockRetryAttempts; i++ {
|
||||||
|
if registriesMutex.TryRLock() {
|
||||||
|
defer registriesMutex.RUnlock()
|
||||||
|
return defaultRegistry
|
||||||
|
}
|
||||||
|
time.Sleep(lockRetryDelay)
|
||||||
|
}
|
||||||
return defaultRegistry
|
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) {
|
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()
|
defer registriesMutex.Unlock()
|
||||||
|
|
||||||
foundAt := -1
|
foundAt := -1
|
||||||
@@ -90,8 +123,34 @@ func AddRegistry(registry *DefaultModelRegistry) {
|
|||||||
registries = append(registries, registry)
|
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 {
|
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()
|
defer r.mutex.Unlock()
|
||||||
|
|
||||||
if _, exists := r.models[name]; exists {
|
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) {
|
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()
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
model, exists := r.models[name]
|
model, exists := r.models[name]
|
||||||
@@ -149,7 +210,9 @@ func (r *DefaultModelRegistry) GetModel(name string) (interface{}, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
|
func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
|
||||||
r.mutex.RLock()
|
if !r.tryRLock() {
|
||||||
|
return make(map[string]interface{})
|
||||||
|
}
|
||||||
defer r.mutex.RUnlock()
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
result := make(map[string]interface{})
|
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
|
// GetModels returns a list of all models from all registries
|
||||||
// Models are collected in registry order, with duplicates included
|
// Models are collected in registry order, with duplicates included
|
||||||
func GetModels() []interface{} {
|
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()
|
defer registriesMutex.RUnlock()
|
||||||
|
|
||||||
var models []interface{}
|
var models []interface{}
|
||||||
seen := make(map[string]bool)
|
seen := make(map[string]bool)
|
||||||
|
|
||||||
for _, registry := range registries {
|
for _, registry := range registries {
|
||||||
registry.mutex.RLock()
|
if !registry.tryRLock() {
|
||||||
|
continue
|
||||||
|
}
|
||||||
for name, model := range registry.models {
|
for name, model := range registry.models {
|
||||||
// Only add the first occurrence of each model name
|
// Only add the first occurrence of each model name
|
||||||
if !seen[name] {
|
if !seen[name] {
|
||||||
|
|||||||
+54
-6
@@ -313,7 +313,7 @@ func (h *Handler) handleRequest(client *Client, msg *Message) {
|
|||||||
// handleRead processes a read operation
|
// handleRead processes a read operation
|
||||||
func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext) {
|
func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext) {
|
||||||
// Execute before hook
|
// 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)
|
logger.Error("[MQTTSpec] BeforeRead hook failed: %v", err)
|
||||||
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
||||||
return
|
return
|
||||||
@@ -356,7 +356,7 @@ func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext)
|
|||||||
// handleCreate processes a create operation
|
// handleCreate processes a create operation
|
||||||
func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContext) {
|
func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContext) {
|
||||||
// Execute before hook
|
// 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)
|
logger.Error("[MQTTSpec] BeforeCreate hook failed: %v", err)
|
||||||
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
||||||
return
|
return
|
||||||
@@ -390,7 +390,7 @@ func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContex
|
|||||||
// handleUpdate processes an update operation
|
// handleUpdate processes an update operation
|
||||||
func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContext) {
|
func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContext) {
|
||||||
// Execute before hook
|
// 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)
|
logger.Error("[MQTTSpec] BeforeUpdate hook failed: %v", err)
|
||||||
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
||||||
return
|
return
|
||||||
@@ -424,7 +424,7 @@ func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContex
|
|||||||
// handleDelete processes a delete operation
|
// handleDelete processes a delete operation
|
||||||
func (h *Handler) handleDelete(client *Client, msg *Message, hookCtx *HookContext) {
|
func (h *Handler) handleDelete(client *Client, msg *Message, hookCtx *HookContext) {
|
||||||
// Execute before hook
|
// 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)
|
logger.Error("[MQTTSpec] BeforeDelete hook failed: %v", err)
|
||||||
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
|
||||||
return
|
return
|
||||||
@@ -676,7 +676,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) {
|
|||||||
|
|
||||||
// Apply columns
|
// Apply columns
|
||||||
if hookCtx.Options != nil && len(hookCtx.Options.Columns) > 0 {
|
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)
|
// 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
|
// Execute query
|
||||||
if err := query.ScanModel(hookCtx.Context); err != nil {
|
if err := query.ScanModel(hookCtx.Context); err != nil {
|
||||||
return nil, fmt.Errorf("failed to read record: %w", err)
|
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 {
|
if hookCtx.Options != nil {
|
||||||
// Apply filters
|
// Apply filters
|
||||||
for _, filter := range hookCtx.Options.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)
|
op := strings.ToLower(filter.Operator)
|
||||||
if op == "like" || op == "ilike" {
|
if op == "like" || op == "ilike" {
|
||||||
|
// 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)
|
query = query.Where(fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
query = query.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
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" {
|
if sort.Direction == "desc" {
|
||||||
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))
|
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
|
// Apply columns
|
||||||
if len(hookCtx.Options.Columns) > 0 {
|
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
|
// Execute query
|
||||||
if err := query.ScanModel(hookCtx.Context); err != nil {
|
if err := query.ScanModel(hookCtx.Context); err != nil {
|
||||||
return nil, nil, fmt.Errorf("failed to read records: %w", err)
|
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)
|
countQuery := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
||||||
if hookCtx.Options != nil {
|
if hookCtx.Options != nil {
|
||||||
for _, filter := range hookCtx.Options.Filters {
|
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)
|
op := strings.ToLower(filter.Operator)
|
||||||
if op == "like" || op == "ilike" {
|
if op == "like" || op == "ilike" {
|
||||||
|
// 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)
|
countQuery = countQuery.Where(fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
countQuery = countQuery.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
countQuery = countQuery.Where(fmt.Sprintf("%s %s ?", filter.Column, h.getOperatorSQL(filter.Operator)), filter.Value)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -34,6 +34,7 @@ const (
|
|||||||
AfterUpdate = websocketspec.AfterUpdate
|
AfterUpdate = websocketspec.AfterUpdate
|
||||||
BeforeDelete = websocketspec.BeforeDelete
|
BeforeDelete = websocketspec.BeforeDelete
|
||||||
AfterDelete = websocketspec.AfterDelete
|
AfterDelete = websocketspec.AfterDelete
|
||||||
|
BeforeScan = websocketspec.BeforeScan
|
||||||
|
|
||||||
// Subscription hooks
|
// Subscription hooks
|
||||||
BeforeSubscribe = websocketspec.BeforeSubscribe
|
BeforeSubscribe = websocketspec.BeforeSubscribe
|
||||||
@@ -46,6 +47,9 @@ const (
|
|||||||
AfterConnect = websocketspec.AfterConnect
|
AfterConnect = websocketspec.AfterConnect
|
||||||
BeforeDisconnect = websocketspec.BeforeDisconnect
|
BeforeDisconnect = websocketspec.BeforeDisconnect
|
||||||
AfterDisconnect = websocketspec.AfterDisconnect
|
AfterDisconnect = websocketspec.AfterDisconnect
|
||||||
|
|
||||||
|
// BeforeOp fires immediately before every SQL operation (read, create, update, delete)
|
||||||
|
BeforeOp = websocketspec.BeforeOp
|
||||||
)
|
)
|
||||||
|
|
||||||
// NewHookRegistry creates a new hook registry
|
// NewHookRegistry creates a new hook registry
|
||||||
|
|||||||
@@ -27,25 +27,31 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
|||||||
return security.LoadSecurityRules(secCtx, 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 {
|
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
return security.ApplyColumnSecurity(secCtx, securityList)
|
return security.ApplyColumnSecurity(secCtx, securityList)
|
||||||
})
|
})
|
||||||
|
|
||||||
// Hook 3 (Optional): Audit logging
|
// Hook 4 (Optional): Audit logging
|
||||||
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
return security.LogDataAccess(secCtx)
|
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 {
|
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
return security.CheckModelUpdateAllowed(secCtx)
|
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 {
|
handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
return security.CheckModelDeleteAllowed(secCtx)
|
return security.CheckModelDeleteAllowed(secCtx)
|
||||||
@@ -71,6 +77,17 @@ func (s *securityContext) GetUserID() (int, bool) {
|
|||||||
return security.GetUserID(s.ctx.Context)
|
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 {
|
func (s *securityContext) GetSchema() string {
|
||||||
return s.ctx.Schema
|
return s.ctx.Schema
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||||
)
|
)
|
||||||
|
|
||||||
// OpenAPISpec represents the OpenAPI 3.0 specification structure
|
// OpenAPISpec represents the OpenAPI 3.0 specification structure
|
||||||
@@ -440,6 +441,28 @@ func (g *Generator) generatePropertySchema(field reflect.StructField) *Schema {
|
|||||||
schema.Description = desc
|
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() {
|
switch fieldType.Kind() {
|
||||||
case reflect.String:
|
case reflect.String:
|
||||||
schema.Type = "string"
|
schema.Type = "string"
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||||
)
|
)
|
||||||
|
|
||||||
type PrimaryKeyNameProvider interface {
|
type PrimaryKeyNameProvider interface {
|
||||||
@@ -438,6 +439,63 @@ func GetSQLModelColumns(model any) []string {
|
|||||||
return columns
|
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
|
// collectSQLColumnsFromType recursively collects SQL column names from a struct type
|
||||||
// scanOnlyEmbedded indicates if we're inside a scan-only embedded struct
|
// scanOnlyEmbedded indicates if we're inside a scan-only embedded struct
|
||||||
func collectSQLColumnsFromType(typ reflect.Type, columns *[]string, scanOnlyEmbedded bool) {
|
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")
|
// Parse JSON tag (format: "name,omitempty")
|
||||||
parts := strings.Split(jsonTag, ",")
|
parts := strings.Split(jsonTag, ",")
|
||||||
if parts[0] == sourceColName {
|
if parts[0] == sourceColName {
|
||||||
return field.Type.Kind()
|
return spectypes.UnwrapKind(field.Type)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check field name (case-insensitive)
|
// Check field name (case-insensitive)
|
||||||
if strings.EqualFold(field.Name, sourceColName) {
|
if strings.EqualFold(field.Name, sourceColName) {
|
||||||
return field.Type.Kind()
|
return spectypes.UnwrapKind(field.Type)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check snake_case conversion
|
// Check snake_case conversion
|
||||||
snakeCaseName := ToSnakeCase(field.Name)
|
snakeCaseName := ToSnakeCase(field.Name)
|
||||||
if snakeCaseName == sourceColName {
|
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 (
|
import (
|
||||||
"reflect"
|
"reflect"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Test models for GORM
|
// 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) {
|
func TestIsColumnWritableWithEmbedded(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
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 =============
|
// ============= Tests for relation functions =============
|
||||||
|
|
||||||
// Models for relation testing
|
// 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
|
// Filters
|
||||||
query = h.applyFilters(query, options.Filters)
|
query = h.applyFilters(query, options.Filters, model)
|
||||||
|
|
||||||
// Custom operators
|
// Custom operators
|
||||||
for _, customOp := range options.CustomOperators {
|
for _, customOp := range options.CustomOperators {
|
||||||
@@ -340,18 +340,28 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
|
|||||||
|
|
||||||
var data interface{}
|
var data interface{}
|
||||||
if id != "" {
|
if id != "" {
|
||||||
singleResult := reflect.New(modelType).Interface()
|
pkName := reflection.GetPrimaryKeyName(model)
|
||||||
pkName := reflection.GetPrimaryKeyName(singleResult)
|
|
||||||
query = query.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
|
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 {
|
if err == sql.ErrNoRows {
|
||||||
return nil, nil, fmt.Errorf("record not found")
|
return nil, nil, fmt.Errorf("record not found")
|
||||||
}
|
}
|
||||||
return nil, nil, fmt.Errorf("query error: %w", err)
|
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 {
|
} 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)
|
return nil, nil, fmt.Errorf("query error: %w", err)
|
||||||
}
|
}
|
||||||
data = reflect.ValueOf(modelPtr).Elem().Interface()
|
data = reflect.ValueOf(modelPtr).Elem().Interface()
|
||||||
@@ -741,8 +751,10 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
|
|||||||
return recordToDelete, nil
|
return recordToDelete, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// applyFilters applies all filters with OR grouping logic.
|
// applyFilters applies all filters with OR grouping logic. model, when
|
||||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption) common.SelectQuery {
|
// 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 {
|
if len(filters) == 0 {
|
||||||
return query
|
return query
|
||||||
}
|
}
|
||||||
@@ -758,10 +770,10 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
|||||||
orGroup = append(orGroup, filters[j])
|
orGroup = append(orGroup, filters[j])
|
||||||
j++
|
j++
|
||||||
}
|
}
|
||||||
query = h.applyFilterGroup(query, orGroup)
|
query = h.applyFilterGroup(query, orGroup, model)
|
||||||
i = j
|
i = j
|
||||||
} else {
|
} else {
|
||||||
condition, args := h.buildFilterCondition(filters[i])
|
condition, args := h.buildFilterCondition(filters[i], model)
|
||||||
if condition != "" {
|
if condition != "" {
|
||||||
query = query.Where(condition, args...)
|
query = query.Where(condition, args...)
|
||||||
}
|
}
|
||||||
@@ -772,12 +784,12 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
|||||||
return query
|
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 conditions []string
|
||||||
var args []interface{}
|
var args []interface{}
|
||||||
|
|
||||||
for _, filter := range filters {
|
for _, filter := range filters {
|
||||||
condition, filterArgs := h.buildFilterCondition(filter)
|
condition, filterArgs := h.buildFilterCondition(filter, model)
|
||||||
if condition != "" {
|
if condition != "" {
|
||||||
conditions = append(conditions, condition)
|
conditions = append(conditions, condition)
|
||||||
args = append(args, filterArgs...)
|
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...)
|
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 {
|
switch filter.Operator {
|
||||||
case "eq", "=":
|
case "eq", "=":
|
||||||
return fmt.Sprintf("%s = ?", filter.Column), []interface{}{filter.Value}
|
return fmt.Sprintf("%s = ?", filter.Column), []interface{}{filter.Value}
|
||||||
@@ -808,9 +827,9 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption) (condition st
|
|||||||
case "lte", "<=":
|
case "lte", "<=":
|
||||||
return fmt.Sprintf("%s <= ?", filter.Column), []interface{}{filter.Value}
|
return fmt.Sprintf("%s <= ?", filter.Column), []interface{}{filter.Value}
|
||||||
case "like":
|
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":
|
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":
|
case "in":
|
||||||
condition, args := common.BuildInCondition(filter.Column, filter.Value)
|
condition, args := common.BuildInCondition(filter.Column, filter.Value)
|
||||||
return condition, args
|
return condition, args
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ package resolvemcp
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"runtime/debug"
|
||||||
|
|
||||||
"github.com/gorilla/mux"
|
"github.com/gorilla/mux"
|
||||||
"github.com/uptrace/bun"
|
"github.com/uptrace/bun"
|
||||||
@@ -25,6 +26,7 @@ import (
|
|||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -82,11 +84,20 @@ func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler) {
|
|||||||
// - GET {basePath}/sse — SSE connection endpoint
|
// - GET {basePath}/sse — SSE connection endpoint
|
||||||
// - POST {basePath}/message — JSON-RPC message endpoint
|
// - POST {basePath}/message — JSON-RPC message endpoint
|
||||||
func SetupBunRouterRoutes(router *bunrouter.Router, handler *Handler) {
|
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
|
basePath := handler.config.BasePath
|
||||||
h := handler.SSEServer()
|
h := handler.SSEServer()
|
||||||
|
|
||||||
router.GET(basePath+"/sse", bunrouter.HTTPHandler(h))
|
router.GET(basePath+"/sse", bunrouter.HTTPHandler(h))
|
||||||
|
logger.Info("Registered resolvemcp bunrouter route GET %s/sse", basePath)
|
||||||
|
|
||||||
router.POST(basePath+"/message", bunrouter.HTTPHandler(h))
|
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.
|
// 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)
|
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 {
|
func (s *securityContext) GetSchema() string {
|
||||||
return s.ctx.Schema
|
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).
|
// Skip relation fields (slice or user-defined struct that isn't time.Time).
|
||||||
fieldType, found := modelType.FieldByName(d.Name)
|
fieldType, found := modelType.FieldByName(d.Name)
|
||||||
|
var unwrappedType reflect.Type
|
||||||
|
isSQLType := false
|
||||||
if found {
|
if found {
|
||||||
ft := fieldType.Type
|
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()
|
ft = ft.Elem()
|
||||||
}
|
}
|
||||||
isUserStruct := ft.Kind() == reflect.Struct && ft.Name() != "Time" && ft.PkgPath() != ""
|
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)
|
info.relationNames = append(info.relationNames, jsonName)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -104,6 +110,9 @@ func buildModelInfo(schema, entity string, model interface{}) modelInfo {
|
|||||||
|
|
||||||
// Derive Go type name, unwrapping pointer if needed.
|
// Derive Go type name, unwrapping pointer if needed.
|
||||||
goType := d.DataType
|
goType := d.DataType
|
||||||
|
if isSQLType {
|
||||||
|
goType = unwrappedType.Name()
|
||||||
|
}
|
||||||
if goType == "" && found {
|
if goType == "" && found {
|
||||||
ft := fieldType.Type
|
ft := fieldType.Type
|
||||||
for ft.Kind() == reflect.Pointer {
|
for ft.Kind() == reflect.Pointer {
|
||||||
@@ -125,7 +134,7 @@ func buildModelInfo(schema, entity string, model interface{}) modelInfo {
|
|||||||
isPrimary: isPrimary,
|
isPrimary: isPrimary,
|
||||||
isUnique: d.SQLKey == "unique" || d.SQLKey == "uniqueindex",
|
isUnique: d.SQLKey == "unique" || d.SQLKey == "uniqueindex",
|
||||||
isFK: d.SQLKey == "foreign_key",
|
isFK: d.SQLKey == "foreign_key",
|
||||||
nullable: d.Nullable,
|
nullable: isSQLType || d.Nullable,
|
||||||
}
|
}
|
||||||
info.columns = append(info.columns, ci)
|
info.columns = append(info.columns, ci)
|
||||||
}
|
}
|
||||||
@@ -134,6 +143,25 @@ func buildModelInfo(schema, entity string, model interface{}) modelInfo {
|
|||||||
return info
|
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.
|
// fieldJSONName returns the JSON tag name for a struct field, falling back to the field name.
|
||||||
func fieldJSONName(modelType reflect.Type, fieldName string) string {
|
func fieldJSONName(modelType reflect.Type, fieldName string) string {
|
||||||
field, ok := modelType.FieldByName(fieldName)
|
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 ?",
|
expectedCondition: "CAST(email AS TEXT) LIKE ?",
|
||||||
expectedArgsCount: 1,
|
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 {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
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 {
|
if condition != tt.expectedCondition {
|
||||||
t.Errorf("Expected condition '%s', got '%s'", tt.expectedCondition, condition)
|
t.Errorf("Expected condition '%s', got '%s'", tt.expectedCondition, condition)
|
||||||
|
|||||||
+520
-154
File diff suppressed because it is too large
Load Diff
@@ -5,6 +5,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestNewHandler(t *testing.T) {
|
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) {
|
func TestGetDatabase(t *testing.T) {
|
||||||
handler := NewHandler(nil, nil)
|
handler := NewHandler(nil, nil)
|
||||||
db := handler.GetDatabase()
|
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) {
|
func TestIsNullable(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
|
|||||||
@@ -3,6 +3,8 @@ package resolvespec
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
@@ -34,6 +36,11 @@ const (
|
|||||||
|
|
||||||
// Scan/Execute operation hooks (for query building)
|
// Scan/Execute operation hooks (for query building)
|
||||||
BeforeScan HookType = "before_scan"
|
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
|
// HookContext contains all the data available to a hook
|
||||||
@@ -77,6 +84,7 @@ type HookFunc func(*HookContext) error
|
|||||||
// HookRegistry manages all registered hooks
|
// HookRegistry manages all registered hooks
|
||||||
type HookRegistry struct {
|
type HookRegistry struct {
|
||||||
hooks map[HookType][]HookFunc
|
hooks map[HookType][]HookFunc
|
||||||
|
mutex sync.RWMutex
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewHookRegistry creates a new hook registry
|
// 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
|
// Register adds a new hook for the specified hook type
|
||||||
func (r *HookRegistry) Register(hookType HookType, hook HookFunc) {
|
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 {
|
if r.hooks == nil {
|
||||||
r.hooks = make(map[HookType][]HookFunc)
|
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
|
// Execute runs all hooks for the specified type in order
|
||||||
// If any hook returns an error, execution stops and the error is returned
|
// If any hook returns an error, execution stops and the error is returned
|
||||||
func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
||||||
hooks, exists := r.hooks[hookType]
|
if !r.tryRLock() {
|
||||||
if !exists || len(hooks) == 0 {
|
return fmt.Errorf("hook execution failed: registry locked")
|
||||||
|
}
|
||||||
|
hooks := append([]HookFunc(nil), r.hooks[hookType]...)
|
||||||
|
r.mutex.RUnlock()
|
||||||
|
|
||||||
|
if len(hooks) == 0 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -128,20 +179,47 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
|
||||||
|
// (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
|
// Clear removes all hooks for the specified type
|
||||||
func (r *HookRegistry) Clear(hookType HookType) {
|
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)
|
delete(r.hooks, hookType)
|
||||||
logger.Info("Cleared all resolvespec hooks for %s", hookType)
|
logger.Info("Cleared all resolvespec hooks for %s", hookType)
|
||||||
}
|
}
|
||||||
|
|
||||||
// ClearAll removes all registered hooks
|
// ClearAll removes all registered hooks
|
||||||
func (r *HookRegistry) ClearAll() {
|
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)
|
r.hooks = make(map[HookType][]HookFunc)
|
||||||
logger.Info("Cleared all resolvespec hooks")
|
logger.Info("Cleared all resolvespec hooks")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Count returns the number of hooks registered for a specific type
|
// Count returns the number of hooks registered for a specific type
|
||||||
func (r *HookRegistry) Count(hookType HookType) int {
|
func (r *HookRegistry) Count(hookType HookType) int {
|
||||||
|
if !r.tryRLock() {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
if hooks, exists := r.hooks[hookType]; exists {
|
if hooks, exists := r.hooks[hookType]; exists {
|
||||||
return len(hooks)
|
return len(hooks)
|
||||||
}
|
}
|
||||||
@@ -155,6 +233,11 @@ func (r *HookRegistry) HasHooks(hookType HookType) bool {
|
|||||||
|
|
||||||
// GetAllHookTypes returns all hook types that have registered hooks
|
// GetAllHookTypes returns all hook types that have registered hooks
|
||||||
func (r *HookRegistry) GetAllHookTypes() []HookType {
|
func (r *HookRegistry) GetAllHookTypes() []HookType {
|
||||||
|
if !r.tryRLock() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
types := make([]HookType, 0, len(r.hooks))
|
types := make([]HookType, 0, len(r.hooks))
|
||||||
for hookType := range r.hooks {
|
for hookType := range r.hooks {
|
||||||
types = append(types, hookType)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -2,6 +2,7 @@ package resolvespec
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"runtime/debug"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/gorilla/mux"
|
"github.com/gorilla/mux"
|
||||||
@@ -12,6 +13,7 @@ import (
|
|||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/router"
|
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/router"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -244,6 +246,11 @@ func wrapBunRouterHandler(handler bunrouter.HandlerFunc, authMiddleware Middlewa
|
|||||||
// Accepts bunrouter.Router or bunrouter.Group
|
// Accepts bunrouter.Router or bunrouter.Group
|
||||||
// authMiddleware is optional - if provided, routes will be protected with the middleware
|
// authMiddleware is optional - if provided, routes will be protected with the middleware
|
||||||
func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware MiddlewareFunc) {
|
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
|
// CORS config
|
||||||
corsConfig := common.DefaultCORSConfig()
|
corsConfig := common.DefaultCORSConfig()
|
||||||
@@ -269,6 +276,13 @@ func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware M
|
|||||||
|
|
||||||
// Loop through each registered model and create explicit routes
|
// Loop through each registered model and create explicit routes
|
||||||
for fullName := range allModels {
|
for fullName := range allModels {
|
||||||
|
func() {
|
||||||
|
defer func() {
|
||||||
|
if rec := recover(); rec != nil {
|
||||||
|
logger.Error("panic registering resolvespec routes for model %s: %v\n%s", fullName, rec, debug.Stack())
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
// Parse the full name (e.g., "public.users" or just "users")
|
// Parse the full name (e.g., "public.users" or just "users")
|
||||||
schema, entity := parseModelName(fullName)
|
schema, entity := parseModelName(fullName)
|
||||||
|
|
||||||
@@ -375,6 +389,9 @@ func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware M
|
|||||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||||
return nil
|
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
|
// Hook 1: BeforeRead - Load security rules
|
||||||
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
|
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
|
if security.IsModelSecurityDisabled(secCtx) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
return security.LoadSecurityRules(secCtx, securityList)
|
return security.LoadSecurityRules(secCtx, securityList)
|
||||||
})
|
})
|
||||||
|
|
||||||
// Hook 2: BeforeScan - Apply row-level security filters
|
// Hook 2: BeforeScan - Apply row-level security filters
|
||||||
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
|
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
|
if security.ShouldSkipRowSecurity(secCtx, hookCtx.Operation) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
return security.ApplyRowSecurity(secCtx, securityList)
|
return security.ApplyRowSecurity(secCtx, securityList)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -78,6 +84,17 @@ func (s *securityContext) GetUserID() (int, bool) {
|
|||||||
return security.GetUserID(s.ctx.Context)
|
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 {
|
func (s *securityContext) GetSchema() string {
|
||||||
return s.ctx.Schema
|
return s.ctx.Schema
|
||||||
}
|
}
|
||||||
@@ -86,6 +103,10 @@ func (s *securityContext) GetEntity() string {
|
|||||||
return s.ctx.Entity
|
return s.ctx.Entity
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *securityContext) GetOperation() string {
|
||||||
|
return s.ctx.Operation
|
||||||
|
}
|
||||||
|
|
||||||
func (s *securityContext) GetModel() interface{} {
|
func (s *securityContext) GetModel() interface{} {
|
||||||
return s.ctx.Model
|
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).
|
Search with specific operators (AND logic).
|
||||||
|
|
||||||
**Supported Operators:**
|
**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)
|
- `beginswith` / `startswith` - Starts with (case-insensitive)
|
||||||
- `endswith` - Ends with (case-insensitive)
|
- `endswith` - Ends with (case-insensitive)
|
||||||
- `equals` / `eq` - Exact match
|
- `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`
|
**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).
|
For complete header documentation, see [HEADERS.md](HEADERS.md).
|
||||||
|
|
||||||
## Lifecycle Hooks
|
## 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"
|
"testing"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"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.
|
// 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))
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+420
-120
@@ -32,6 +32,7 @@ type Handler struct {
|
|||||||
nestedProcessor *common.NestedCUDProcessor
|
nestedProcessor *common.NestedCUDProcessor
|
||||||
fallbackHandler FallbackHandler
|
fallbackHandler FallbackHandler
|
||||||
openAPIGenerator func() (string, error)
|
openAPIGenerator func() (string, error)
|
||||||
|
defaultSort map[string][]common.SortOption
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewHandler creates a new API handler with database and registry abstractions
|
// NewHandler creates a new API handler with database and registry abstractions
|
||||||
@@ -64,6 +65,33 @@ func (h *Handler) SetFallbackHandler(fallback FallbackHandler) {
|
|||||||
h.fallbackHandler = fallback
|
h.fallbackHandler = fallback
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetDefaultSort sets the sort order applied to list ("read") results when a
|
||||||
|
// request does not specify its own sort (no x-sort header / sort query param).
|
||||||
|
// Pass an empty schema and name to set a global default used by any model
|
||||||
|
// without a more specific default.
|
||||||
|
func (h *Handler) SetDefaultSort(schema, name string, sort ...common.SortOption) {
|
||||||
|
if h.defaultSort == nil {
|
||||||
|
h.defaultSort = make(map[string][]common.SortOption)
|
||||||
|
}
|
||||||
|
h.defaultSort[defaultSortKey(schema, name)] = sort
|
||||||
|
}
|
||||||
|
|
||||||
|
// getDefaultSort returns the configured default sort for a model, falling
|
||||||
|
// back to the global default (registered with an empty schema and name).
|
||||||
|
func (h *Handler) getDefaultSort(schema, name string) []common.SortOption {
|
||||||
|
if h.defaultSort == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if sort := h.defaultSort[defaultSortKey(schema, name)]; len(sort) > 0 {
|
||||||
|
return sort
|
||||||
|
}
|
||||||
|
return h.defaultSort[defaultSortKey("", "")]
|
||||||
|
}
|
||||||
|
|
||||||
|
func defaultSortKey(schema, name string) string {
|
||||||
|
return fmt.Sprintf("%s.%s", schema, name)
|
||||||
|
}
|
||||||
|
|
||||||
// handlePanic is a helper function to handle panics with stack traces
|
// handlePanic is a helper function to handle panics with stack traces
|
||||||
func (h *Handler) handlePanic(w common.ResponseWriter, method string, err interface{}) {
|
func (h *Handler) handlePanic(w common.ResponseWriter, method string, err interface{}) {
|
||||||
stack := debug.Stack()
|
stack := debug.Stack()
|
||||||
@@ -205,8 +233,18 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
validId, _ := strconv.ParseInt(id, 10, 64)
|
validId, _ := strconv.ParseInt(id, 10, 64)
|
||||||
if validId > 0 {
|
updateID := id
|
||||||
h.handleUpdate(ctx, w, id, nil, data, options)
|
isUpdate := validId > 0
|
||||||
|
if !isUpdate {
|
||||||
|
// No valid /:id in the URL - check whether the body itself carries
|
||||||
|
// a valid primary key value and treat this as an update if so.
|
||||||
|
if pkID, ok := h.extractPrimaryKeyFromBody(model, data); ok && pkID != "0" {
|
||||||
|
updateID = pkID
|
||||||
|
isUpdate = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if isUpdate {
|
||||||
|
h.handleUpdate(ctx, w, updateID, nil, data, options)
|
||||||
} else {
|
} else {
|
||||||
h.handleCreate(ctx, w, data, options)
|
h.handleCreate(ctx, w, data, options)
|
||||||
}
|
}
|
||||||
@@ -243,6 +281,49 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// extractPrimaryKeyFromBody looks for a valid primary key value inside a
|
||||||
|
// decoded (single-record) POST body, keyed by the model's primary key column
|
||||||
|
// or its JSON equivalent. It returns the string form of that value and true
|
||||||
|
// if one was found and is non-empty/non-zero; otherwise ("", false).
|
||||||
|
func (h *Handler) extractPrimaryKeyFromBody(model interface{}, data interface{}) (string, bool) {
|
||||||
|
dataMap, ok := data.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
// Batch payloads (slices) aren't eligible for this implicit-update detection.
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
pkCol := reflection.GetPrimaryKeyName(model)
|
||||||
|
if pkCol == "" {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
val, exists := dataMap[pkCol]
|
||||||
|
if !exists {
|
||||||
|
modelType := reflection.GetPointerElement(reflect.TypeOf(model))
|
||||||
|
for jsonKey, col := range reflection.BuildJSONToDBColumnMap(modelType) {
|
||||||
|
if col == pkCol {
|
||||||
|
val, exists = dataMap[jsonKey]
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !exists || val == nil || reflection.IsEmptyValue(val) {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
switch v := val.(type) {
|
||||||
|
case float64:
|
||||||
|
if v <= 0 {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
return strconv.FormatInt(int64(v), 10), true
|
||||||
|
case string:
|
||||||
|
return v, true
|
||||||
|
default:
|
||||||
|
return fmt.Sprintf("%v", v), true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// HandleGet processes GET requests for metadata
|
// HandleGet processes GET requests for metadata
|
||||||
func (h *Handler) HandleGet(w common.ResponseWriter, r common.Request, params map[string]string) {
|
func (h *Handler) HandleGet(w common.ResponseWriter, r common.Request, params map[string]string) {
|
||||||
// Capture panics and return error response
|
// Capture panics and return error response
|
||||||
@@ -327,26 +408,6 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
options.SingleRecordAsObject = false
|
options.SingleRecordAsObject = false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute BeforeRead hooks
|
|
||||||
hookCtx := &HookContext{
|
|
||||||
Context: ctx,
|
|
||||||
Handler: h,
|
|
||||||
Schema: schema,
|
|
||||||
Entity: entity,
|
|
||||||
TableName: tableName,
|
|
||||||
Model: model,
|
|
||||||
Options: options,
|
|
||||||
ID: id,
|
|
||||||
Writer: w,
|
|
||||||
Tx: h.db,
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := h.hooks.Execute(BeforeRead, hookCtx); err != nil {
|
|
||||||
logger.Error("BeforeRead hook failed: %v", err)
|
|
||||||
h.sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Validate and unwrap model type to get base struct
|
// Validate and unwrap model type to get base struct
|
||||||
modelType := reflect.TypeOf(model)
|
modelType := reflect.TypeOf(model)
|
||||||
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
@@ -364,9 +425,44 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
|
|
||||||
logger.Info("Reading records from %s.%s", schema, entity)
|
logger.Info("Reading records from %s.%s", schema, entity)
|
||||||
|
|
||||||
|
hookCtx := &HookContext{
|
||||||
|
Context: ctx,
|
||||||
|
Handler: h,
|
||||||
|
Schema: schema,
|
||||||
|
Entity: entity,
|
||||||
|
TableName: tableName,
|
||||||
|
Model: model,
|
||||||
|
Operation: "read",
|
||||||
|
Options: options,
|
||||||
|
ID: id,
|
||||||
|
Writer: w,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Everything below runs inside a single transaction so that the BeforeRead/BeforeScan
|
||||||
|
// hooks (which may set session-scoped RLS GUCs via SET LOCAL) execute on the same
|
||||||
|
// physical connection as the queries they are meant to protect. Under connection
|
||||||
|
// pooling, firing a hook against h.db and then querying against h.db again may hand
|
||||||
|
// out two different connections, silently bypassing RLS.
|
||||||
|
var (
|
||||||
|
total int
|
||||||
|
fetchedRowNumber *int64
|
||||||
|
statusCode int
|
||||||
|
errCode string
|
||||||
|
errMsg string
|
||||||
|
)
|
||||||
|
|
||||||
|
txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error {
|
||||||
|
hookCtx.Tx = tx
|
||||||
|
|
||||||
|
if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil {
|
||||||
|
statusCode, errCode, errMsg = http.StatusBadRequest, "hook_error", "Hook execution failed"
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
options = hookCtx.Options
|
||||||
|
|
||||||
// Start with Model() using the slice pointer to avoid "Model(nil)" errors in Count()
|
// Start with Model() using the slice pointer to avoid "Model(nil)" errors in Count()
|
||||||
// Bun's Model() accepts both single pointers and slice pointers
|
// Bun's Model() accepts both single pointers and slice pointers
|
||||||
query := h.db.NewSelect().Model(modelPtr)
|
query := tx.NewSelect().Model(modelPtr)
|
||||||
|
|
||||||
// Only set Table() if the model doesn't provide a table name via the underlying type
|
// Only set Table() if the model doesn't provide a table name via the underlying type
|
||||||
// Create a temporary instance to check for TableNameProvider
|
// Create a temporary instance to check for TableNameProvider
|
||||||
@@ -377,7 +473,12 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
|
|
||||||
// If we have computed columns/expressions but options.Columns is empty,
|
// If we have computed columns/expressions but options.Columns is empty,
|
||||||
// populate it with all model columns first since computed columns are additions
|
// populate it with all model columns first since computed columns are additions
|
||||||
if len(options.Columns) == 0 && (len(options.ComputedQL) > 0 || len(options.ComputedColumns) > 0) {
|
vectorSearchActive := options.VectorSearch != nil &&
|
||||||
|
options.VectorSearch.Column != "" && len(options.VectorSearch.Vector) > 0
|
||||||
|
|
||||||
|
if len(options.Columns) == 0 &&
|
||||||
|
(len(options.ComputedQL) > 0 || len(options.ComputedColumns) > 0 ||
|
||||||
|
(vectorSearchActive && options.VectorSearch.As != "")) {
|
||||||
logger.Debug("Populating options.Columns with all model columns since computed columns are additions")
|
logger.Debug("Populating options.Columns with all model columns since computed columns are additions")
|
||||||
options.Columns = reflection.GetSQLModelColumns(model)
|
options.Columns = reflection.GetSQLModelColumns(model)
|
||||||
}
|
}
|
||||||
@@ -424,12 +525,46 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
// Apply column selection
|
// Apply column selection
|
||||||
if len(options.Columns) > 0 {
|
if len(options.Columns) > 0 {
|
||||||
logger.Debug("Selecting columns: %v", options.Columns)
|
logger.Debug("Selecting columns: %v", options.Columns)
|
||||||
|
selectAlias := reflection.ExtractTableNameOnly(tableName)
|
||||||
for _, col := range options.Columns {
|
for _, col := range options.Columns {
|
||||||
|
// JSON sub-field selection (data->>'x', data.x, data#>>'{a,b}'):
|
||||||
|
// emit a parameterised expression aliased to a stable name.
|
||||||
|
if expr, jargs, alias, ok := common.ResolveJSONColumnExpr(model, selectAlias, col); ok {
|
||||||
|
if !reflection.HasColumn(model, alias) {
|
||||||
|
logger.Warn("Skipping JSON select column %q: model has no scan target for alias %q", col, alias)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
query = query.ColumnExpr(expr+" AS "+common.QuoteIdent(alias), jargs...)
|
||||||
|
continue
|
||||||
|
}
|
||||||
query = query.Column(reflection.ExtractSourceColumn(col))
|
query = query.Column(reflection.ExtractSourceColumn(col))
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// pgvector KNN search: order by distance to the query vector and,
|
||||||
|
// optionally, return that distance as an extra column. Postgres only.
|
||||||
|
if vectorSearchActive {
|
||||||
|
vs := options.VectorSearch
|
||||||
|
op := common.VectorOperator(vs.Metric)
|
||||||
|
lit, litErr := common.VectorLiteral(vs.Vector)
|
||||||
|
if litErr != nil {
|
||||||
|
logger.Error("Invalid vector search vector: %v", litErr)
|
||||||
|
statusCode, errCode, errMsg = http.StatusBadRequest, "invalid_vector_search", "Invalid vector search vector"
|
||||||
|
return litErr
|
||||||
|
}
|
||||||
|
col := common.QuoteIdent(vs.Column)
|
||||||
|
dir := "ASC"
|
||||||
|
if strings.EqualFold(vs.Direction, "desc") {
|
||||||
|
dir = "DESC"
|
||||||
|
}
|
||||||
|
if vs.As != "" {
|
||||||
|
query = query.ColumnExpr(fmt.Sprintf("(%s %s ?) AS %s", col, op, common.QuoteIdent(vs.As)), lit)
|
||||||
|
}
|
||||||
|
query = query.OrderExpr(fmt.Sprintf("%s %s ? %s", col, op, dir), lit)
|
||||||
|
logger.Debug("Applying vector search on %s (%s)", vs.Column, op)
|
||||||
|
}
|
||||||
|
|
||||||
// Apply expand (Just expand to Preload for now)
|
// Apply expand (Just expand to Preload for now)
|
||||||
for _, expand := range options.Expand {
|
for _, expand := range options.Expand {
|
||||||
logger.Debug("Applying expand: %s", expand.Relation)
|
logger.Debug("Applying expand: %s", expand.Relation)
|
||||||
@@ -483,9 +618,9 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
fixedWhere, err := common.ValidateAndFixPreloadWhere(preload.Where, preload.Relation)
|
fixedWhere, err := common.ValidateAndFixPreloadWhere(preload.Where, preload.Relation)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("Invalid preload WHERE clause for relation '%s': %v", preload.Relation, err)
|
logger.Error("Invalid preload WHERE clause for relation '%s': %v", preload.Relation, err)
|
||||||
h.sendError(w, http.StatusBadRequest, "invalid_preload_where",
|
statusCode, errCode, errMsg = http.StatusBadRequest, "invalid_preload_where",
|
||||||
fmt.Sprintf("Invalid preload WHERE clause for relation '%s'", preload.Relation), err)
|
fmt.Sprintf("Invalid preload WHERE clause for relation '%s'", preload.Relation)
|
||||||
return
|
return err
|
||||||
}
|
}
|
||||||
preload.Where = fixedWhere
|
preload.Where = fixedWhere
|
||||||
}
|
}
|
||||||
@@ -540,12 +675,12 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
|
|
||||||
// Apply the OR group as a single grouped condition
|
// Apply the OR group as a single grouped condition
|
||||||
logger.Debug("Applying OR filter group with %d conditions", len(orFilters))
|
logger.Debug("Applying OR filter group with %d conditions", len(orFilters))
|
||||||
query = h.applyOrFilterGroup(query, orFilters, orCastInfo, tableName)
|
query = h.applyOrFilterGroup(query, orFilters, orCastInfo, tableName, model)
|
||||||
i = j
|
i = j
|
||||||
} else {
|
} else {
|
||||||
// Single AND filter - apply normally
|
// Single AND filter - apply normally
|
||||||
logger.Debug("Applying filter: %s %s %v (needsCast=%v, logic=%s)", filter.Column, filter.Operator, filter.Value, castInfo.NeedsCast, logicOp)
|
logger.Debug("Applying filter: %s %s %v (needsCast=%v, logic=%s)", filter.Column, filter.Operator, filter.Value, castInfo.NeedsCast, logicOp)
|
||||||
query = h.applyFilter(query, *filter, tableName, castInfo.NeedsCast, logicOp)
|
query = h.applyFilter(query, *filter, tableName, castInfo.NeedsCast, logicOp, model)
|
||||||
i++
|
i++
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -602,7 +737,6 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
|
|
||||||
// Handle FetchRowNumber before applying ID filter
|
// Handle FetchRowNumber before applying ID filter
|
||||||
// This must happen before the query to get the row position, then filter by PK
|
// This must happen before the query to get the row position, then filter by PK
|
||||||
var fetchedRowNumber *int64
|
|
||||||
var fetchRowNumberPKValue string
|
var fetchRowNumberPKValue string
|
||||||
if options.FetchRowNumber != nil && *options.FetchRowNumber != "" {
|
if options.FetchRowNumber != nil && *options.FetchRowNumber != "" {
|
||||||
pkName := reflection.GetPrimaryKeyName(model)
|
pkName := reflection.GetPrimaryKeyName(model)
|
||||||
@@ -610,11 +744,11 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
|
|
||||||
logger.Debug("FetchRowNumber: Fetching row number for PK %s = %s", pkName, fetchRowNumberPKValue)
|
logger.Debug("FetchRowNumber: Fetching row number for PK %s = %s", pkName, fetchRowNumberPKValue)
|
||||||
|
|
||||||
rowNum, err := h.FetchRowNumber(ctx, tableName, pkName, fetchRowNumberPKValue, options, model)
|
rowNum, err := h.FetchRowNumber(ctx, tx, tableName, pkName, fetchRowNumberPKValue, options, model)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("Failed to fetch row number: %v", err)
|
logger.Error("Failed to fetch row number: %v", err)
|
||||||
h.sendError(w, http.StatusBadRequest, "fetch_rownumber_error", "Failed to fetch row number", err)
|
statusCode, errCode, errMsg = http.StatusBadRequest, "fetch_rownumber_error", "Failed to fetch row number"
|
||||||
return
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
fetchedRowNumber = &rowNum
|
fetchedRowNumber = &rowNum
|
||||||
@@ -632,6 +766,11 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
query = query.Where(fmt.Sprintf("%s.%s = ?", common.QuoteIdent(tableAlias), common.QuoteIdent(pkName)), id)
|
query = query.Where(fmt.Sprintf("%s.%s = ?", common.QuoteIdent(tableAlias), common.QuoteIdent(pkName)), id)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Fall back to the configured default sort when the request didn't specify one
|
||||||
|
if len(options.Sort) == 0 {
|
||||||
|
options.Sort = common.ResolveSortColumns(h.getDefaultSort(schema, entity), reflection.GetPrimaryKeyName(model))
|
||||||
|
}
|
||||||
|
|
||||||
// Apply sorting
|
// Apply sorting
|
||||||
tableAlias := reflection.ExtractTableNameOnly(tableName)
|
tableAlias := reflection.ExtractTableNameOnly(tableName)
|
||||||
for _, sort := range options.Sort {
|
for _, sort := range options.Sort {
|
||||||
@@ -641,8 +780,12 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
}
|
}
|
||||||
logger.Debug("Applying sort: %s %s", sort.Column, direction)
|
logger.Debug("Applying sort: %s %s", sort.Column, direction)
|
||||||
|
|
||||||
// Check if it's an expression (enclosed in brackets) - use directly without quoting
|
// JSON sub-field reference (data->>'x', data#>>'{a,b}', or dotted
|
||||||
if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
// shorthand when the base is a JSON column) - resolve to a safe
|
||||||
|
// parameterised expression before the generic branches.
|
||||||
|
if expr, jargs, _, ok := common.ResolveJSONColumnExpr(model, tableAlias, sort.Column); ok {
|
||||||
|
query = query.OrderExpr(fmt.Sprintf("%s %s", expr, direction), jargs...)
|
||||||
|
} else if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||||
// For expressions, pass as raw SQL to prevent auto-quoting
|
// For expressions, pass as raw SQL to prevent auto-quoting
|
||||||
query = query.OrderExpr(fmt.Sprintf("%s %s", sort.Column, direction))
|
query = query.OrderExpr(fmt.Sprintf("%s %s", sort.Column, direction))
|
||||||
} else if strings.Contains(sort.Column, ".") {
|
} else if strings.Contains(sort.Column, ".") {
|
||||||
@@ -655,7 +798,6 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Get total count before pagination (unless skip count is requested)
|
// Get total count before pagination (unless skip count is requested)
|
||||||
var total int
|
|
||||||
if !options.SkipCount {
|
if !options.SkipCount {
|
||||||
// Try to get from cache first (unless SkipCache is true)
|
// Try to get from cache first (unless SkipCache is true)
|
||||||
var cachedTotalData *cachedTotal
|
var cachedTotalData *cachedTotal
|
||||||
@@ -703,8 +845,8 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
count, err := query.Count(ctx)
|
count, err := query.Count(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("Error counting records: %v", err)
|
logger.Error("Error counting records: %v", err)
|
||||||
h.sendError(w, http.StatusInternalServerError, "query_error", "Error counting records", err)
|
statusCode, errCode, errMsg = http.StatusInternalServerError, "query_error", "Error counting records"
|
||||||
return
|
return err
|
||||||
}
|
}
|
||||||
total = count
|
total = count
|
||||||
logger.Debug("Total records (from query): %d", total)
|
logger.Debug("Total records (from query): %d", total)
|
||||||
@@ -764,8 +906,8 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
cursorFilter, err := options.GetCursorFilter(tableName, pkName, modelColumns, expandJoins)
|
cursorFilter, err := options.GetCursorFilter(tableName, pkName, modelColumns, expandJoins)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("Error building cursor filter: %v", err)
|
logger.Error("Error building cursor filter: %v", err)
|
||||||
h.sendError(w, http.StatusBadRequest, "cursor_error", "Invalid cursor pagination", err)
|
statusCode, errCode, errMsg = http.StatusBadRequest, "cursor_error", "Invalid cursor pagination"
|
||||||
return
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Apply cursor filter to query
|
// Apply cursor filter to query
|
||||||
@@ -780,10 +922,10 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
|
|
||||||
// Execute BeforeScan hooks - pass query chain so hooks can modify it
|
// Execute BeforeScan hooks - pass query chain so hooks can modify it
|
||||||
hookCtx.Query = query
|
hookCtx.Query = query
|
||||||
if err := h.hooks.Execute(BeforeScan, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
|
||||||
logger.Error("BeforeScan hook failed: %v", err)
|
logger.Error("BeforeScan hook failed: %v", err)
|
||||||
h.sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
statusCode, errCode, errMsg = http.StatusBadRequest, "hook_error", "Hook execution failed"
|
||||||
return
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Use potentially modified query from hook context
|
// Use potentially modified query from hook context
|
||||||
@@ -794,7 +936,18 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
// Execute query - modelPtr was already created earlier
|
// Execute query - modelPtr was already created earlier
|
||||||
if err := query.ScanModel(ctx); err != nil {
|
if err := query.ScanModel(ctx); err != nil {
|
||||||
logger.Error("Error executing query: %v", err)
|
logger.Error("Error executing query: %v", err)
|
||||||
h.sendError(w, http.StatusInternalServerError, "query_error", "Error executing query", err)
|
statusCode, errCode, errMsg = http.StatusInternalServerError, "query_error", "Error executing query"
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
if txErr != nil {
|
||||||
|
if statusCode == 0 {
|
||||||
|
statusCode, errCode, errMsg = http.StatusInternalServerError, "query_error", "Error executing query"
|
||||||
|
}
|
||||||
|
h.sendError(w, statusCode, errCode, errMsg, txErr)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -839,7 +992,8 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
logger.Debug("FetchRowNumber: Row number %d set in metadata", *fetchedRowNumber)
|
logger.Debug("FetchRowNumber: Row number %d set in metadata", *fetchedRowNumber)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute AfterRead hooks
|
// Execute AfterRead hooks (runs after the transaction commits, against the pooled db)
|
||||||
|
hookCtx.Tx = h.db
|
||||||
hookCtx.Result = modelPtr
|
hookCtx.Result = modelPtr
|
||||||
hookCtx.Error = nil
|
hookCtx.Error = nil
|
||||||
|
|
||||||
@@ -978,7 +1132,7 @@ func (h *Handler) applyPreloadWithRecursion(query common.SelectQuery, preload co
|
|||||||
// Apply filters
|
// Apply filters
|
||||||
if len(preload.Filters) > 0 {
|
if len(preload.Filters) > 0 {
|
||||||
for _, filter := range preload.Filters {
|
for _, filter := range preload.Filters {
|
||||||
sq = h.applyFilter(sq, filter, "", false, "AND")
|
sq = h.applyFilter(sq, filter, "", false, "AND", nil)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1133,7 +1287,6 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
|||||||
|
|
||||||
logger.Info("Creating record in %s.%s", schema, entity)
|
logger.Info("Creating record in %s.%s", schema, entity)
|
||||||
|
|
||||||
// Execute BeforeCreate hooks
|
|
||||||
hookCtx := &HookContext{
|
hookCtx := &HookContext{
|
||||||
Context: ctx,
|
Context: ctx,
|
||||||
Handler: h,
|
Handler: h,
|
||||||
@@ -1141,31 +1294,42 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
|||||||
Entity: entity,
|
Entity: entity,
|
||||||
TableName: tableName,
|
TableName: tableName,
|
||||||
Model: model,
|
Model: model,
|
||||||
|
Operation: "create",
|
||||||
Options: options,
|
Options: options,
|
||||||
Data: data,
|
Data: data,
|
||||||
Writer: w,
|
Writer: w,
|
||||||
Tx: h.db,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil {
|
// Everything below (including the BeforeCreate hook) runs inside a single
|
||||||
logger.Error("BeforeCreate hook failed: %v", err)
|
// transaction so that session-scoped RLS GUCs set by the hook execute on the
|
||||||
h.sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
// same physical connection as the inserts they are meant to protect.
|
||||||
return
|
var (
|
||||||
|
dataSlice []interface{}
|
||||||
|
originalDataMaps []map[string]interface{}
|
||||||
|
statusCode int
|
||||||
|
errCode string
|
||||||
|
errMsg string
|
||||||
|
)
|
||||||
|
|
||||||
|
// Process all items in a transaction
|
||||||
|
results := make([]interface{}, 0)
|
||||||
|
txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error {
|
||||||
|
hookCtx.Tx = tx
|
||||||
|
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
|
||||||
|
statusCode, errCode, errMsg = http.StatusBadRequest, "hook_error", "Hook execution failed"
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Use potentially modified data from hook context
|
// Use potentially modified data from hook context
|
||||||
data = hookCtx.Data
|
data = hookCtx.Data
|
||||||
|
|
||||||
// Normalize data to slice for unified processing
|
// Normalize data to slice for unified processing
|
||||||
dataSlice := h.normalizeToSlice(data)
|
dataSlice = h.normalizeToSlice(data)
|
||||||
logger.Debug("Processing %d item(s) for creation", len(dataSlice))
|
logger.Debug("Processing %d item(s) for creation", len(dataSlice))
|
||||||
|
|
||||||
// Store original data maps for merging later
|
// Store original data maps for merging later
|
||||||
originalDataMaps := make([]map[string]interface{}, 0, len(dataSlice))
|
originalDataMaps = make([]map[string]interface{}, 0, len(dataSlice))
|
||||||
|
|
||||||
// Process all items in a transaction
|
|
||||||
results := make([]interface{}, 0, len(dataSlice))
|
|
||||||
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
|
|
||||||
// Create temporary nested processor with transaction
|
// Create temporary nested processor with transaction
|
||||||
txNestedProcessor := common.NewNestedCUDProcessor(tx, h.registry, h)
|
txNestedProcessor := common.NewNestedCUDProcessor(tx, h.registry, h)
|
||||||
|
|
||||||
@@ -1230,13 +1394,14 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
|||||||
Entity: entity,
|
Entity: entity,
|
||||||
TableName: tableName,
|
TableName: tableName,
|
||||||
Model: model,
|
Model: model,
|
||||||
|
Operation: "create",
|
||||||
Options: options,
|
Options: options,
|
||||||
Data: modelValue,
|
Data: modelValue,
|
||||||
Writer: w,
|
Writer: w,
|
||||||
Query: query,
|
Query: query,
|
||||||
Tx: tx,
|
Tx: tx,
|
||||||
}
|
}
|
||||||
if err := h.hooks.Execute(BeforeScan, itemHookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeScan, itemHookCtx); err != nil {
|
||||||
return fmt.Errorf("BeforeScan hook failed for item %d: %w", i, err)
|
return fmt.Errorf("BeforeScan hook failed for item %d: %w", i, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1266,9 +1431,12 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
|||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
|
|
||||||
if err != nil {
|
if txErr != nil {
|
||||||
logger.Error("Error creating records: %v", err)
|
if statusCode == 0 {
|
||||||
h.sendError(w, http.StatusInternalServerError, "create_error", "Error creating records", err)
|
statusCode, errCode, errMsg = http.StatusInternalServerError, "create_error", "Error creating records"
|
||||||
|
}
|
||||||
|
logger.Error("Error creating records: %v", txErr)
|
||||||
|
h.sendError(w, statusCode, errCode, errMsg, txErr)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1284,7 +1452,10 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute AfterCreate hooks
|
// Execute AfterCreate hooks (runs after the transaction commits, against the
|
||||||
|
// pooled db — hookCtx.Tx was pointed at the now-closed transaction inside the
|
||||||
|
// RunInTransaction closure above and must not be reused here).
|
||||||
|
hookCtx.Tx = h.db
|
||||||
var responseData interface{}
|
var responseData interface{}
|
||||||
if len(mergedResults) == 1 {
|
if len(mergedResults) == 1 {
|
||||||
responseData = mergedResults[0]
|
responseData = mergedResults[0]
|
||||||
@@ -1366,7 +1537,35 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
|||||||
// Create temporary nested processor with transaction
|
// Create temporary nested processor with transaction
|
||||||
txNestedProcessor := common.NewNestedCUDProcessor(tx, h.registry, h)
|
txNestedProcessor := common.NewNestedCUDProcessor(tx, h.registry, h)
|
||||||
|
|
||||||
// First, read the existing record from the database
|
// Execute BeforeUpdate hooks inside transaction, before any queries run.
|
||||||
|
// BeforeUpdate hooks may set session-scoped RLS GUCs (via SET LOCAL);
|
||||||
|
// they must run before the existence-check select so that select is
|
||||||
|
// also subject to RLS on this connection/transaction.
|
||||||
|
hookCtx = &HookContext{
|
||||||
|
Context: ctx,
|
||||||
|
Handler: h,
|
||||||
|
Schema: schema,
|
||||||
|
Entity: entity,
|
||||||
|
TableName: tableName,
|
||||||
|
Tx: tx,
|
||||||
|
Model: model,
|
||||||
|
Operation: "update",
|
||||||
|
Options: options,
|
||||||
|
ID: id,
|
||||||
|
Data: dataMap,
|
||||||
|
Writer: w,
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil {
|
||||||
|
return fmt.Errorf("BeforeUpdate hook failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use potentially modified data from hook context
|
||||||
|
if modifiedData, ok := hookCtx.Data.(map[string]interface{}); ok {
|
||||||
|
dataMap = modifiedData
|
||||||
|
}
|
||||||
|
|
||||||
|
// Now read the existing record from the database
|
||||||
existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
|
existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
|
||||||
selectQuery := tx.NewSelect().Model(existingRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
|
selectQuery := tx.NewSelect().Model(existingRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
|
||||||
if err := selectQuery.ScanModel(ctx); err != nil {
|
if err := selectQuery.ScanModel(ctx); err != nil {
|
||||||
@@ -1398,30 +1597,6 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
|||||||
nestedRelations = relations
|
nestedRelations = relations
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute BeforeUpdate hooks inside transaction
|
|
||||||
hookCtx = &HookContext{
|
|
||||||
Context: ctx,
|
|
||||||
Handler: h,
|
|
||||||
Schema: schema,
|
|
||||||
Entity: entity,
|
|
||||||
TableName: tableName,
|
|
||||||
Tx: tx,
|
|
||||||
Model: model,
|
|
||||||
Options: options,
|
|
||||||
ID: id,
|
|
||||||
Data: dataMap,
|
|
||||||
Writer: w,
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
|
|
||||||
return fmt.Errorf("BeforeUpdate hook failed: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Use potentially modified data from hook context
|
|
||||||
if modifiedData, ok := hookCtx.Data.(map[string]interface{}); ok {
|
|
||||||
dataMap = modifiedData
|
|
||||||
}
|
|
||||||
|
|
||||||
// Merge only non-null and non-empty values from the incoming request into the existing record
|
// Merge only non-null and non-empty values from the incoming request into the existing record
|
||||||
for key, newValue := range dataMap {
|
for key, newValue := range dataMap {
|
||||||
// Skip if the value is nil
|
// Skip if the value is nil
|
||||||
@@ -1458,7 +1633,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
|||||||
// Execute BeforeScan hooks - pass query chain so hooks can modify it
|
// Execute BeforeScan hooks - pass query chain so hooks can modify it
|
||||||
hookCtx.Query = query
|
hookCtx.Query = query
|
||||||
hookCtx.Tx = tx
|
hookCtx.Tx = tx
|
||||||
if err := h.hooks.Execute(BeforeScan, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
|
||||||
return fmt.Errorf("BeforeScan hook failed: %w", err)
|
return fmt.Errorf("BeforeScan hook failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1494,6 +1669,23 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
|||||||
// Fetch the updated record after the transaction commits to capture any trigger changes
|
// Fetch the updated record after the transaction commits to capture any trigger changes
|
||||||
fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface()
|
fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface()
|
||||||
selectQuery := h.db.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
|
selectQuery := h.db.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
|
||||||
|
|
||||||
|
// Execute BeforeScan hooks so row security is re-applied to the post-update
|
||||||
|
// re-fetch, same as it is for the initial read and the update query itself.
|
||||||
|
// Without this, the re-fetch can return a row the caller isn't authorized to see.
|
||||||
|
// 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.Query = selectQuery
|
||||||
|
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
|
||||||
|
logger.Error("BeforeScan hook failed: %v", err)
|
||||||
|
h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if modifiedQuery, ok := hookCtx.Query.(common.SelectQuery); ok {
|
||||||
|
selectQuery = modifiedQuery
|
||||||
|
}
|
||||||
|
|
||||||
if err := selectQuery.ScanModel(ctx); err != nil {
|
if err := selectQuery.ScanModel(ctx); err != nil {
|
||||||
logger.Error("Failed to fetch updated record: %v", err)
|
logger.Error("Failed to fetch updated record: %v", err)
|
||||||
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err)
|
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err)
|
||||||
@@ -1555,12 +1747,13 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
|
|||||||
Entity: entity,
|
Entity: entity,
|
||||||
TableName: tableName,
|
TableName: tableName,
|
||||||
Model: model,
|
Model: model,
|
||||||
|
Operation: "delete",
|
||||||
ID: itemID,
|
ID: itemID,
|
||||||
Writer: w,
|
Writer: w,
|
||||||
Tx: tx,
|
Tx: tx,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil {
|
||||||
logger.Error("BeforeDelete hook failed for ID %s: %v", itemID, err)
|
logger.Error("BeforeDelete hook failed for ID %s: %v", itemID, err)
|
||||||
return fmt.Errorf("delete not allowed for ID %s: %w", itemID, err)
|
return fmt.Errorf("delete not allowed for ID %s: %w", itemID, err)
|
||||||
}
|
}
|
||||||
@@ -1629,12 +1822,13 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
|
|||||||
Entity: entity,
|
Entity: entity,
|
||||||
TableName: tableName,
|
TableName: tableName,
|
||||||
Model: model,
|
Model: model,
|
||||||
|
Operation: "delete",
|
||||||
ID: itemIDStr,
|
ID: itemIDStr,
|
||||||
Writer: w,
|
Writer: w,
|
||||||
Tx: tx,
|
Tx: tx,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil {
|
||||||
logger.Error("BeforeDelete hook failed for ID %v: %v", itemID, err)
|
logger.Error("BeforeDelete hook failed for ID %v: %v", itemID, err)
|
||||||
return fmt.Errorf("delete not allowed for ID %v: %w", itemID, err)
|
return fmt.Errorf("delete not allowed for ID %v: %w", itemID, err)
|
||||||
}
|
}
|
||||||
@@ -1687,12 +1881,13 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
|
|||||||
Entity: entity,
|
Entity: entity,
|
||||||
TableName: tableName,
|
TableName: tableName,
|
||||||
Model: model,
|
Model: model,
|
||||||
|
Operation: "delete",
|
||||||
ID: itemIDStr,
|
ID: itemIDStr,
|
||||||
Writer: w,
|
Writer: w,
|
||||||
Tx: tx,
|
Tx: tx,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil {
|
||||||
logger.Error("BeforeDelete hook failed for ID %v: %v", itemID, err)
|
logger.Error("BeforeDelete hook failed for ID %v: %v", itemID, err)
|
||||||
return fmt.Errorf("delete not allowed for ID %v: %w", itemID, err)
|
return fmt.Errorf("delete not allowed for ID %v: %w", itemID, err)
|
||||||
}
|
}
|
||||||
@@ -1771,13 +1966,14 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
|
|||||||
Entity: entity,
|
Entity: entity,
|
||||||
TableName: tableName,
|
TableName: tableName,
|
||||||
Model: model,
|
Model: model,
|
||||||
|
Operation: "delete",
|
||||||
ID: id,
|
ID: id,
|
||||||
Writer: w,
|
Writer: w,
|
||||||
Tx: h.db,
|
Tx: h.db,
|
||||||
Data: recordToDelete,
|
Data: recordToDelete,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil {
|
||||||
logger.Error("BeforeDelete hook failed: %v", err)
|
logger.Error("BeforeDelete hook failed: %v", err)
|
||||||
h.sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
h.sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return
|
return
|
||||||
@@ -1788,7 +1984,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
|
|||||||
|
|
||||||
// Execute BeforeScan hooks - pass query chain so hooks can modify it
|
// Execute BeforeScan hooks - pass query chain so hooks can modify it
|
||||||
hookCtx.Query = query
|
hookCtx.Query = query
|
||||||
if err := h.hooks.Execute(BeforeScan, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
|
||||||
logger.Error("BeforeScan hook failed: %v", err)
|
logger.Error("BeforeScan hook failed: %v", err)
|
||||||
h.sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
h.sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return
|
return
|
||||||
@@ -2168,16 +2364,11 @@ func (h *Handler) qualifyColumnName(columnName, fullTableName string) string {
|
|||||||
return fmt.Sprintf("%s.%s", tableOnly, columnName)
|
return fmt.Sprintf("%s.%s", tableOnly, columnName)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, tableName string, needsCast bool, logicOp string) common.SelectQuery {
|
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, tableName string, needsCast bool, logicOp string, model interface{}) common.SelectQuery {
|
||||||
// Qualify the column name with table name if not already qualified
|
// Qualify the column name with table name if not already qualified
|
||||||
rawQualifiedColumn := h.qualifyColumnName(filter.Column, tableName)
|
rawQualifiedColumn := h.qualifyColumnName(filter.Column, tableName)
|
||||||
qualifiedColumn := rawQualifiedColumn
|
qualifiedColumn := rawQualifiedColumn
|
||||||
|
|
||||||
// Apply casting to text if needed for non-numeric columns or non-numeric values
|
|
||||||
if needsCast {
|
|
||||||
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Helper function to apply the correct Where method based on logic operator
|
// Helper function to apply the correct Where method based on logic operator
|
||||||
applyWhere := func(condition string, args ...interface{}) common.SelectQuery {
|
applyWhere := func(condition string, args ...interface{}) common.SelectQuery {
|
||||||
if logicOp == "OR" {
|
if logicOp == "OR" {
|
||||||
@@ -2186,6 +2377,26 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
|
|||||||
return query.Where(condition, args...)
|
return query.Where(condition, args...)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// JSON sub-field access (data->>'x', data#>>'{a,b}', or the dotted data.x
|
||||||
|
// shorthand when "data" is a JSON column): resolve to a safe, parameterised
|
||||||
|
// expression before the ordinary column handling below.
|
||||||
|
tableAlias := reflection.ExtractTableNameOnly(tableName)
|
||||||
|
if cond, jargs, ok := common.BuildJSONFilterCondition(model, tableAlias, filter.Column, filter.Operator, filter.Value); ok {
|
||||||
|
return applyWhere(cond, jargs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply casting to text if needed for non-numeric columns or non-numeric values
|
||||||
|
if needsCast {
|
||||||
|
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
|
||||||
|
}
|
||||||
|
|
||||||
|
// citext columns already compare case-insensitively; casting to TEXT for
|
||||||
|
// LIKE/ILIKE would switch to case-sensitive matching and defeat a citext index.
|
||||||
|
likeColumn := rawQualifiedColumn
|
||||||
|
if !reflection.IsCitextColumn(model, filter.Column) {
|
||||||
|
likeColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
|
||||||
|
}
|
||||||
|
|
||||||
switch strings.ToLower(filter.Operator) {
|
switch strings.ToLower(filter.Operator) {
|
||||||
case "eq", "equals":
|
case "eq", "equals":
|
||||||
return applyWhere(fmt.Sprintf("%s = ?", qualifiedColumn), filter.Value)
|
return applyWhere(fmt.Sprintf("%s = ?", qualifiedColumn), filter.Value)
|
||||||
@@ -2200,11 +2411,14 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
|
|||||||
case "lte", "less_than_equals", "le":
|
case "lte", "less_than_equals", "le":
|
||||||
return applyWhere(fmt.Sprintf("%s <= ?", qualifiedColumn), filter.Value)
|
return applyWhere(fmt.Sprintf("%s <= ?", qualifiedColumn), filter.Value)
|
||||||
case "like":
|
case "like":
|
||||||
// Always cast to TEXT for LIKE/ILIKE to support date/time/timestamp columns
|
// Cast to TEXT for LIKE to support date/time/timestamp columns; citext
|
||||||
return applyWhere(fmt.Sprintf("CAST(%s AS TEXT) LIKE ?", rawQualifiedColumn), filter.Value)
|
// columns are compared natively (see likeColumn above).
|
||||||
|
return applyWhere(fmt.Sprintf("%s LIKE ?", likeColumn), filter.Value)
|
||||||
case "ilike":
|
case "ilike":
|
||||||
// Always cast to TEXT for LIKE/ILIKE to support date/time/timestamp columns
|
// Cast to TEXT for ILIKE to support date/time/timestamp columns; citext
|
||||||
return applyWhere(fmt.Sprintf("CAST(%s AS TEXT) ILIKE ?", rawQualifiedColumn), filter.Value)
|
// columns are compared natively (see likeColumn above) since citext is
|
||||||
|
// already case-insensitive.
|
||||||
|
return applyWhere(fmt.Sprintf("%s ILIKE ?", likeColumn), filter.Value)
|
||||||
case "in":
|
case "in":
|
||||||
cond, inArgs := common.BuildInCondition(qualifiedColumn, filter.Value)
|
cond, inArgs := common.BuildInCondition(qualifiedColumn, filter.Value)
|
||||||
if cond == "" {
|
if cond == "" {
|
||||||
@@ -2238,6 +2452,18 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
|
|||||||
colName := h.qualifyColumnName(filter.Column, tableName)
|
colName := h.qualifyColumnName(filter.Column, tableName)
|
||||||
return applyWhere(fmt.Sprintf("(%s IS NOT NULL AND %s != '')", colName, colName))
|
return applyWhere(fmt.Sprintf("(%s IS NOT NULL AND %s != '')", colName, colName))
|
||||||
default:
|
default:
|
||||||
|
if common.IsSpatialOperator(filter.Operator) {
|
||||||
|
if cond, sargs, ok := common.BuildSpatialCondition(rawQualifiedColumn, filter.Operator, filter.Value); ok {
|
||||||
|
return applyWhere(cond, sargs...)
|
||||||
|
}
|
||||||
|
return query
|
||||||
|
}
|
||||||
|
if common.IsVectorOperator(filter.Operator) {
|
||||||
|
if cond, vargs, ok := common.BuildVectorCondition(rawQualifiedColumn, filter.Operator, filter.Value); ok {
|
||||||
|
return applyWhere(cond, vargs...)
|
||||||
|
}
|
||||||
|
return query
|
||||||
|
}
|
||||||
logger.Warn("Unknown filter operator: %s, defaulting to equals", filter.Operator)
|
logger.Warn("Unknown filter operator: %s, defaulting to equals", filter.Operator)
|
||||||
return applyWhere(fmt.Sprintf("%s = ?", qualifiedColumn), filter.Value)
|
return applyWhere(fmt.Sprintf("%s = ?", qualifiedColumn), filter.Value)
|
||||||
}
|
}
|
||||||
@@ -2245,24 +2471,37 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
|
|||||||
|
|
||||||
// applyOrFilterGroup applies a group of OR filters as a single grouped condition
|
// applyOrFilterGroup applies a group of OR filters as a single grouped condition
|
||||||
// This ensures OR conditions are properly grouped with parentheses to prevent OR logic from escaping
|
// This ensures OR conditions are properly grouped with parentheses to prevent OR logic from escaping
|
||||||
func (h *Handler) applyOrFilterGroup(query common.SelectQuery, filters []*common.FilterOption, castInfo []ColumnCastInfo, tableName string) common.SelectQuery {
|
func (h *Handler) applyOrFilterGroup(query common.SelectQuery, filters []*common.FilterOption, castInfo []ColumnCastInfo, tableName string, model interface{}) common.SelectQuery {
|
||||||
if len(filters) == 0 {
|
if len(filters) == 0 {
|
||||||
return query
|
return query
|
||||||
}
|
}
|
||||||
|
|
||||||
|
tableAlias := reflection.ExtractTableNameOnly(tableName)
|
||||||
|
|
||||||
// Build individual filter conditions
|
// Build individual filter conditions
|
||||||
conditions := []string{}
|
conditions := []string{}
|
||||||
args := []interface{}{}
|
args := []interface{}{}
|
||||||
|
|
||||||
for i, filter := range filters {
|
for i, filter := range filters {
|
||||||
|
// JSON sub-field access: resolve to a safe parameterised condition first.
|
||||||
|
if cond, jargs, ok := common.BuildJSONFilterCondition(model, tableAlias, filter.Column, filter.Operator, filter.Value); ok {
|
||||||
|
conditions = append(conditions, cond)
|
||||||
|
args = append(args, jargs...)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
// Qualify the column name with table name if not already qualified
|
// Qualify the column name with table name if not already qualified
|
||||||
rawQualifiedColumn := h.qualifyColumnName(filter.Column, tableName)
|
rawQualifiedColumn := h.qualifyColumnName(filter.Column, tableName)
|
||||||
qualifiedColumn := rawQualifiedColumn
|
qualifiedColumn := rawQualifiedColumn
|
||||||
|
|
||||||
op := strings.ToLower(filter.Operator)
|
op := strings.ToLower(filter.Operator)
|
||||||
if op == "like" || op == "ilike" {
|
if op == "like" || op == "ilike" {
|
||||||
// Always cast to TEXT for LIKE/ILIKE to support date/time/timestamp columns
|
// Cast to TEXT for LIKE/ILIKE to support date/time/timestamp columns.
|
||||||
|
// citext columns are left native: they're already case-insensitive and
|
||||||
|
// casting would defeat a citext index.
|
||||||
|
if !reflection.IsCitextColumn(model, filter.Column) {
|
||||||
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
|
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
|
||||||
|
}
|
||||||
} else if castInfo[i].NeedsCast {
|
} else if castInfo[i].NeedsCast {
|
||||||
// Apply casting to text if needed for non-numeric columns or non-numeric values
|
// Apply casting to text if needed for non-numeric columns or non-numeric values
|
||||||
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
|
qualifiedColumn = fmt.Sprintf("CAST(%s AS TEXT)", rawQualifiedColumn)
|
||||||
@@ -2337,6 +2576,20 @@ func (h *Handler) buildFilterCondition(qualifiedColumn string, filter *common.Fi
|
|||||||
colName := h.qualifyColumnName(filter.Column, tableName)
|
colName := h.qualifyColumnName(filter.Column, tableName)
|
||||||
return fmt.Sprintf("(%s IS NOT NULL AND %s != '')", colName, colName), nil
|
return fmt.Sprintf("(%s IS NOT NULL AND %s != '')", colName, colName), nil
|
||||||
default:
|
default:
|
||||||
|
if common.IsSpatialOperator(filter.Operator) {
|
||||||
|
rawCol := h.qualifyColumnName(filter.Column, tableName)
|
||||||
|
if cond, sargs, ok := common.BuildSpatialCondition(rawCol, filter.Operator, filter.Value); ok {
|
||||||
|
return cond, sargs
|
||||||
|
}
|
||||||
|
return "", nil
|
||||||
|
}
|
||||||
|
if common.IsVectorOperator(filter.Operator) {
|
||||||
|
rawCol := h.qualifyColumnName(filter.Column, tableName)
|
||||||
|
if cond, vargs, ok := common.BuildVectorCondition(rawCol, filter.Operator, filter.Value); ok {
|
||||||
|
return cond, vargs
|
||||||
|
}
|
||||||
|
return "", nil
|
||||||
|
}
|
||||||
logger.Warn("Unknown filter operator: %s, defaulting to equals", filter.Operator)
|
logger.Warn("Unknown filter operator: %s, defaulting to equals", filter.Operator)
|
||||||
return fmt.Sprintf("%s = ?", qualifiedColumn), []interface{}{filter.Value}
|
return fmt.Sprintf("%s = ?", qualifiedColumn), []interface{}{filter.Value}
|
||||||
}
|
}
|
||||||
@@ -2459,10 +2712,18 @@ func (h *Handler) generateMetadata(schema, entity string, model interface{}) *co
|
|||||||
jsonName = field.Name
|
jsonName = field.Name
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if this is a relation field (slice or struct, but not time.Time)
|
columnType := field.Type
|
||||||
if field.Type.Kind() == reflect.Slice ||
|
isSQLType := false
|
||||||
(field.Type.Kind() == reflect.Struct && field.Type.Name() != "Time") ||
|
if unwrappedType, ok := unwrapSQLType(field.Type); ok {
|
||||||
(field.Type.Kind() == reflect.Pointer && field.Type.Elem().Kind() == reflect.Struct && field.Type.Elem().Name() != "Time") {
|
columnType = unwrappedType
|
||||||
|
isSQLType = true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this is a relation field (slice or struct, but not time.Time).
|
||||||
|
// spectypes SQL values are columns even when their Go representation is a
|
||||||
|
// struct or a slice.
|
||||||
|
if !isSQLType && (columnType.Kind() == reflect.Slice ||
|
||||||
|
(columnType.Kind() == reflect.Struct && columnType.Name() != "Time")) {
|
||||||
metadata.Relations = append(metadata.Relations, jsonName)
|
metadata.Relations = append(metadata.Relations, jsonName)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -2483,8 +2744,8 @@ func (h *Handler) generateMetadata(schema, entity string, model interface{}) *co
|
|||||||
|
|
||||||
column := common.Column{
|
column := common.Column{
|
||||||
Name: columnName,
|
Name: columnName,
|
||||||
Type: h.getColumnType(field.Type),
|
Type: h.getColumnType(columnType),
|
||||||
IsNullable: h.isNullable(field),
|
IsNullable: isSQLType || h.isNullable(field),
|
||||||
IsPrimary: strings.Contains(gormTag, "primaryKey") || strings.Contains(gormTag, "primary_key"),
|
IsPrimary: strings.Contains(gormTag, "primaryKey") || strings.Contains(gormTag, "primary_key"),
|
||||||
IsUnique: strings.Contains(gormTag, "unique"),
|
IsUnique: strings.Contains(gormTag, "unique"),
|
||||||
HasIndex: strings.Contains(gormTag, "index"),
|
HasIndex: strings.Contains(gormTag, "index"),
|
||||||
@@ -2515,6 +2776,27 @@ func (h *Handler) getColumnType(t reflect.Type) string {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// unwrapSQLType returns the value type wrapped by a spectypes SQL value. These
|
||||||
|
// types represent 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
|
||||||
|
}
|
||||||
|
|
||||||
func (h *Handler) isNullable(field reflect.StructField) bool {
|
func (h *Handler) isNullable(field reflect.StructField) bool {
|
||||||
return field.Type.Kind() == reflect.Pointer
|
return field.Type.Kind() == reflect.Pointer
|
||||||
}
|
}
|
||||||
@@ -2613,13 +2895,18 @@ func (h *Handler) buildDetailFields(model interface{}) []reflection.ModelFieldDe
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Skip relation fields (slices, structs that aren't time.Time, ptrs to struct)
|
// Skip relation fields (slices and structs that aren't time.Time). spectypes
|
||||||
|
// SQL values are columns, not relations.
|
||||||
ft := field.Type
|
ft := field.Type
|
||||||
if ft.Kind() == reflect.Pointer {
|
isSQLType := false
|
||||||
|
if unwrappedType, ok := unwrapSQLType(ft); ok {
|
||||||
|
ft = unwrappedType
|
||||||
|
isSQLType = true
|
||||||
|
} else if ft.Kind() == reflect.Pointer {
|
||||||
ft = ft.Elem()
|
ft = ft.Elem()
|
||||||
}
|
}
|
||||||
if ft.Kind() == reflect.Slice ||
|
if !isSQLType && (ft.Kind() == reflect.Slice ||
|
||||||
(ft.Kind() == reflect.Struct && ft.Name() != "Time") {
|
(ft.Kind() == reflect.Struct && ft.Name() != "Time")) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2647,7 +2934,7 @@ func (h *Handler) buildDetailFields(model interface{}) []reflection.ModelFieldDe
|
|||||||
sqlKey = "unique"
|
sqlKey = "unique"
|
||||||
}
|
}
|
||||||
|
|
||||||
nullable := field.Type.Kind() == reflect.Pointer
|
nullable := isSQLType || field.Type.Kind() == reflect.Pointer
|
||||||
if strings.Contains(gormLower, "not null") {
|
if strings.Contains(gormLower, "not null") {
|
||||||
nullable = false
|
nullable = false
|
||||||
} else if strings.Contains(gormLower, "nullable") || strings.Contains(gormLower, ",null") {
|
} else if strings.Contains(gormLower, "nullable") || strings.Contains(gormLower, ",null") {
|
||||||
@@ -2656,7 +2943,7 @@ func (h *Handler) buildDetailFields(model interface{}) []reflection.ModelFieldDe
|
|||||||
|
|
||||||
fields = append(fields, reflection.ModelFieldDetail{
|
fields = append(fields, reflection.ModelFieldDetail{
|
||||||
Name: jsonName,
|
Name: jsonName,
|
||||||
DataType: h.getColumnType(field.Type),
|
DataType: h.getColumnType(ft),
|
||||||
SQLName: sqlName,
|
SQLName: sqlName,
|
||||||
SQLDataType: sqlDataType,
|
SQLDataType: sqlDataType,
|
||||||
SQLKey: sqlKey,
|
SQLKey: sqlKey,
|
||||||
@@ -2722,8 +3009,21 @@ func (h *Handler) sendFormattedResponse(w common.ResponseWriter, data interface{
|
|||||||
switch options.ResponseFormat {
|
switch options.ResponseFormat {
|
||||||
case "simple":
|
case "simple":
|
||||||
// Simple format: just return the data array
|
// Simple format: just return the data array
|
||||||
|
jsonData, err := json.Marshal(data)
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("Failed to marshal JSON response: %v", err)
|
||||||
|
w.WriteHeader(http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if string(jsonData) == "null" {
|
||||||
|
if options.SingleRecordAsObject {
|
||||||
|
jsonData = []byte("{}")
|
||||||
|
} else {
|
||||||
|
jsonData = []byte("[]")
|
||||||
|
}
|
||||||
|
}
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
if err := w.WriteJSON(data); err != nil {
|
if _, err := w.Write(jsonData); err != nil {
|
||||||
logger.Error("Failed to write JSON response: %v", err)
|
logger.Error("Failed to write JSON response: %v", err)
|
||||||
}
|
}
|
||||||
case "syncfusion":
|
case "syncfusion":
|
||||||
@@ -2811,7 +3111,7 @@ func (h *Handler) sendError(w common.ResponseWriter, statusCode int, code, messa
|
|||||||
|
|
||||||
// FetchRowNumber calculates the row number of a specific record based on sorting and filtering
|
// FetchRowNumber calculates the row number of a specific record based on sorting and filtering
|
||||||
// Returns the 1-based row number of the record with the given primary key value
|
// Returns the 1-based row number of the record with the given primary key value
|
||||||
func (h *Handler) FetchRowNumber(ctx context.Context, tableName string, pkName string, pkValue string, options ExtendedRequestOptions, model any) (int64, error) {
|
func (h *Handler) FetchRowNumber(ctx context.Context, db common.Database, tableName string, pkName string, pkValue string, options ExtendedRequestOptions, model any) (int64, error) {
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
logger.Error("Panic during FetchRowNumber: %v", r)
|
logger.Error("Panic during FetchRowNumber: %v", r)
|
||||||
@@ -2895,7 +3195,7 @@ func (h *Handler) FetchRowNumber(ctx context.Context, tableName string, pkName s
|
|||||||
RN int64 `bun:"rn"`
|
RN int64 `bun:"rn"`
|
||||||
}
|
}
|
||||||
logger.Debug("[FetchRowNumber] BEFORE Query call - about to execute raw query")
|
logger.Debug("[FetchRowNumber] BEFORE Query call - about to execute raw query")
|
||||||
err := h.db.Query(ctx, &result, queryStr, pkValue)
|
err := db.Query(ctx, &result, queryStr, pkValue)
|
||||||
logger.Debug("[FetchRowNumber] AFTER Query call - query completed with %d results, err: %v", len(result), err)
|
logger.Debug("[FetchRowNumber] AFTER Query call - query completed with %d results, err: %v", len(result), err)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("failed to fetch row number: %w", err)
|
return 0, fmt.Errorf("failed to fetch row number: %w", err)
|
||||||
|
|||||||
@@ -182,6 +182,22 @@ func (h *Handler) parseOptionsFromHeaders(r common.Request, model interface{}) E
|
|||||||
h.parseSearchOp(&options, key, decodedValue, "AND")
|
h.parseSearchOp(&options, key, decodedValue, "AND")
|
||||||
case strings.HasPrefix(key, "x-searchcols"):
|
case strings.HasPrefix(key, "x-searchcols"):
|
||||||
options.SearchColumns = h.parseCommaSeparated(decodedValue)
|
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"):
|
case strings.HasPrefix(key, "x-custom-sql-w"):
|
||||||
if options.CustomSQLWhere != "" {
|
if options.CustomSQLWhere != "" {
|
||||||
options.CustomSQLWhere = fmt.Sprintf("%s AND (%s)", options.CustomSQLWhere, decodedValue)
|
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
|
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
|
// parseSelectFields parses x-select-fields header
|
||||||
func (h *Handler) parseSelectFields(options *ExtendedRequestOptions, value string) {
|
func (h *Handler) parseSelectFields(options *ExtendedRequestOptions, value string) {
|
||||||
if value == "" {
|
if value == "" {
|
||||||
@@ -1365,6 +1458,20 @@ func (h *Handler) ValidateAndAdjustFilterForColumnType(filter *common.FilterOpti
|
|||||||
return ColumnCastInfo{NeedsCast: false, IsNumericType: false}
|
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)
|
colType := reflection.GetColumnTypeFromModel(model, filter.Column)
|
||||||
if colType == reflect.Invalid {
|
if colType == reflect.Invalid {
|
||||||
// Column not found in model, no casting needed
|
// 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}
|
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
|
// Check if the input value is numeric
|
||||||
valueIsNumeric := false
|
valueIsNumeric := false
|
||||||
if strVal, ok := filter.Value.(string); ok {
|
if strVal, ok := filter.Value.(string); ok {
|
||||||
|
|||||||
@@ -3,6 +3,8 @@ package restheadspec
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
@@ -34,6 +36,11 @@ const (
|
|||||||
|
|
||||||
// Scan/Execute operation hooks
|
// Scan/Execute operation hooks
|
||||||
BeforeScan HookType = "before_scan"
|
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
|
// HookContext contains all the data available to a hook
|
||||||
@@ -84,6 +91,7 @@ type HookFunc func(*HookContext) error
|
|||||||
// HookRegistry manages all registered hooks
|
// HookRegistry manages all registered hooks
|
||||||
type HookRegistry struct {
|
type HookRegistry struct {
|
||||||
hooks map[HookType][]HookFunc
|
hooks map[HookType][]HookFunc
|
||||||
|
mutex sync.RWMutex
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewHookRegistry creates a new hook registry
|
// 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
|
// Register adds a new hook for the specified hook type
|
||||||
func (r *HookRegistry) Register(hookType HookType, hook HookFunc) {
|
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 {
|
if r.hooks == nil {
|
||||||
r.hooks = make(map[HookType][]HookFunc)
|
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
|
// Execute runs all hooks for the specified type in order
|
||||||
// If any hook returns an error, execution stops and the error is returned
|
// If any hook returns an error, execution stops and the error is returned
|
||||||
func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
||||||
hooks, exists := r.hooks[hookType]
|
if !r.tryRLock() {
|
||||||
if !exists || len(hooks) == 0 {
|
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)
|
// logger.Debug("No hooks registered for %s", hookType)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -137,20 +188,47 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
|
||||||
|
// (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
|
// Clear removes all hooks for the specified type
|
||||||
func (r *HookRegistry) Clear(hookType HookType) {
|
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)
|
delete(r.hooks, hookType)
|
||||||
logger.Info("Cleared all hooks for %s", hookType)
|
logger.Info("Cleared all hooks for %s", hookType)
|
||||||
}
|
}
|
||||||
|
|
||||||
// ClearAll removes all registered hooks
|
// ClearAll removes all registered hooks
|
||||||
func (r *HookRegistry) ClearAll() {
|
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)
|
r.hooks = make(map[HookType][]HookFunc)
|
||||||
logger.Info("Cleared all hooks")
|
logger.Info("Cleared all hooks")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Count returns the number of hooks registered for a specific type
|
// Count returns the number of hooks registered for a specific type
|
||||||
func (r *HookRegistry) Count(hookType HookType) int {
|
func (r *HookRegistry) Count(hookType HookType) int {
|
||||||
|
if !r.tryRLock() {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
if hooks, exists := r.hooks[hookType]; exists {
|
if hooks, exists := r.hooks[hookType]; exists {
|
||||||
return len(hooks)
|
return len(hooks)
|
||||||
}
|
}
|
||||||
@@ -164,6 +242,11 @@ func (r *HookRegistry) HasHooks(hookType HookType) bool {
|
|||||||
|
|
||||||
// GetAllHookTypes returns all hook types that have registered hooks
|
// GetAllHookTypes returns all hook types that have registered hooks
|
||||||
func (r *HookRegistry) GetAllHookTypes() []HookType {
|
func (r *HookRegistry) GetAllHookTypes() []HookType {
|
||||||
|
if !r.tryRLock() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
types := make([]HookType, 0, len(r.hooks))
|
types := make([]HookType, 0, len(r.hooks))
|
||||||
for hookType := range r.hooks {
|
for hookType := range r.hooks {
|
||||||
types = append(types, hookType)
|
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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -55,6 +55,7 @@ package restheadspec
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"runtime/debug"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/gorilla/mux"
|
"github.com/gorilla/mux"
|
||||||
@@ -308,6 +309,11 @@ func wrapBunRouterHandler(handler bunrouter.HandlerFunc, authMiddleware Middlewa
|
|||||||
// Accepts bunrouter.Router or bunrouter.Group
|
// Accepts bunrouter.Router or bunrouter.Group
|
||||||
// authMiddleware is optional - if provided, routes will be protected with the middleware
|
// authMiddleware is optional - if provided, routes will be protected with the middleware
|
||||||
func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware MiddlewareFunc) {
|
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
|
// CORS config
|
||||||
corsConfig := common.DefaultCORSConfig()
|
corsConfig := common.DefaultCORSConfig()
|
||||||
@@ -333,6 +339,13 @@ func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware M
|
|||||||
|
|
||||||
// Loop through each registered model and create explicit routes
|
// Loop through each registered model and create explicit routes
|
||||||
for fullName := range allModels {
|
for fullName := range allModels {
|
||||||
|
func() {
|
||||||
|
defer func() {
|
||||||
|
if rec := recover(); rec != nil {
|
||||||
|
logger.Error("panic registering restheadspec routes for model %s: %v\n%s", fullName, rec, debug.Stack())
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
// Parse the full name (e.g., "public.users" or just "users")
|
// Parse the full name (e.g., "public.users" or just "users")
|
||||||
schema, entity := parseModelName(fullName)
|
schema, entity := parseModelName(fullName)
|
||||||
|
|
||||||
@@ -498,6 +511,9 @@ func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware M
|
|||||||
handler.HandleGet(respAdapter, reqAdapter, params)
|
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||||
return nil
|
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)
|
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 {
|
func (s *securityContext) GetSchema() string {
|
||||||
return s.ctx.Schema
|
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
|
// GetRowSecurity delegates to the row security provider
|
||||||
func (c *CompositeSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
|
func (c *CompositeSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
|
||||||
return c.rowSec.GetRowSecurity(ctx, userID, schema, table)
|
return c.rowSec.GetRowSecurity(ctx, userRef, schema, table)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Optional interface implementations (if wrapped providers support them)
|
// Optional interface implementations (if wrapped providers support them)
|
||||||
|
|||||||
@@ -79,7 +79,7 @@ type mockRowSec struct {
|
|||||||
supportsCache bool
|
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
|
return m.rowSec, m.err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1597,6 +1597,8 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
|
|||||||
client_name VARCHAR(255),
|
client_name VARCHAR(255),
|
||||||
grant_types TEXT[] DEFAULT ARRAY['authorization_code'],
|
grant_types TEXT[] DEFAULT ARRAY['authorization_code'],
|
||||||
allowed_scopes TEXT[] DEFAULT ARRAY['openid','profile','email'],
|
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,
|
is_active BOOLEAN DEFAULT true,
|
||||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||||
);
|
);
|
||||||
@@ -1634,13 +1636,15 @@ DECLARE
|
|||||||
BEGIN
|
BEGIN
|
||||||
v_client_id := p_data->>'client_id';
|
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 (
|
VALUES (
|
||||||
v_client_id,
|
v_client_id,
|
||||||
ARRAY(SELECT jsonb_array_elements_text(p_data->'redirect_uris')),
|
ARRAY(SELECT jsonb_array_elements_text(p_data->'redirect_uris')),
|
||||||
p_data->>'client_name',
|
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->'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;
|
RETURNING to_jsonb(oauth_clients.*) INTO v_row;
|
||||||
|
|
||||||
|
|||||||
@@ -100,6 +100,8 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
|
|||||||
client_name VARCHAR(255),
|
client_name VARCHAR(255),
|
||||||
grant_types TEXT, -- JSON-encoded []string
|
grant_types TEXT, -- JSON-encoded []string
|
||||||
allowed_scopes 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,
|
is_active BOOLEAN DEFAULT 1,
|
||||||
created_at TIMESTAMP
|
created_at TIMESTAMP
|
||||||
);
|
);
|
||||||
|
|||||||
+76
-20
@@ -14,6 +14,11 @@ import (
|
|||||||
type SecurityContext interface {
|
type SecurityContext interface {
|
||||||
GetContext() context.Context
|
GetContext() context.Context
|
||||||
GetUserID() (int, bool)
|
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
|
GetSchema() string
|
||||||
GetEntity() string
|
GetEntity() string
|
||||||
GetModel() interface{}
|
GetModel() interface{}
|
||||||
@@ -45,8 +50,13 @@ func loadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error
|
|||||||
// return err
|
// return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Load row security rules using the provider
|
// Load row security rules using the provider. Row security uses the opaque
|
||||||
_, err = securityList.LoadRowSecurity(secCtx.GetContext(), userID, schema, tablename, false)
|
// 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 {
|
if err != nil {
|
||||||
logger.Warn("Failed to load row security: %v", err)
|
logger.Warn("Failed to load row security: %v", err)
|
||||||
// Don't fail the request if no security rules exist
|
// 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)
|
// applyRowSecurity applies row-level security filters to the query (generic version)
|
||||||
func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
||||||
userID, ok := secCtx.GetUserID()
|
userRef, ok := secCtx.GetUserRef()
|
||||||
if !ok {
|
if !ok {
|
||||||
|
userID, idOK := secCtx.GetUserID()
|
||||||
|
if !idOK {
|
||||||
return nil // No user context, skip
|
return nil // No user context, skip
|
||||||
}
|
}
|
||||||
|
userRef = userID
|
||||||
|
}
|
||||||
|
|
||||||
schema := secCtx.GetSchema()
|
schema := secCtx.GetSchema()
|
||||||
tablename := secCtx.GetEntity()
|
tablename := secCtx.GetEntity()
|
||||||
|
|
||||||
// Get row security template
|
// Get row security template
|
||||||
rowSec, err := securityList.GetRowSecurityTemplate(userID, schema, tablename)
|
rowSec, err := securityList.GetRowSecurityTemplate(userRef, schema, tablename)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// No row security defined, allow query to proceed
|
// 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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check if user has a blocking rule
|
// Check if user has a blocking rule
|
||||||
if rowSec.HasBlock {
|
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)
|
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
|
// Generate the WHERE clause from template
|
||||||
whereClause := rowSec.GetTemplate(pkName, modelType)
|
whereClause := rowSec.GetTemplate(pkName, modelType)
|
||||||
|
|
||||||
logger.Info("Applying row security filter for user %d on %s.%s: %s",
|
logger.Info("Applying row security filter for user %v on %s.%s: %s",
|
||||||
userID, schema, tablename, whereClause)
|
userRef, schema, tablename, whereClause)
|
||||||
|
|
||||||
// Apply the WHERE clause to the query
|
// Apply the WHERE clause to the query
|
||||||
query := secCtx.GetQuery()
|
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
|
// ApplyRowSecurity is a public wrapper for applyRowSecurity that accepts a SecurityContext
|
||||||
// This allows other packages to apply row-level security using the generic interface
|
// This allows other packages to apply row-level security using the generic interface
|
||||||
func ApplyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
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)
|
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
|
// ApplyColumnSecurity is a public wrapper for applyColumnSecurity that accepts a SecurityContext
|
||||||
// This allows other packages to apply column-level security using the generic interface
|
// This allows other packages to apply column-level security using the generic interface
|
||||||
func ApplyColumnSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
func ApplyColumnSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
||||||
@@ -289,18 +331,8 @@ func checkModelDeleteAllowed(secCtx SecurityContext) error {
|
|||||||
// 7. Guest (UserID == 0) → return "authentication required".
|
// 7. Guest (UserID == 0) → return "authentication required".
|
||||||
// 8. Authenticated user → allow (operation-specific checks remain in BeforeUpdate/BeforeDelete).
|
// 8. Authenticated user → allow (operation-specific checks remain in BeforeUpdate/BeforeDelete).
|
||||||
func CheckModelAuthAllowed(secCtx SecurityContext, operation string) error {
|
func CheckModelAuthAllowed(secCtx SecurityContext, operation string) error {
|
||||||
rules, ok := GetModelRulesFromContext(secCtx.GetContext())
|
rules, ok := resolveModelRules(secCtx)
|
||||||
if !ok {
|
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
|
// Model not registered - fall through to auth check
|
||||||
userID, _ := secCtx.GetUserID()
|
userID, _ := secCtx.GetUserID()
|
||||||
if userID == 0 {
|
if userID == 0 {
|
||||||
@@ -308,7 +340,6 @@ func CheckModelAuthAllowed(secCtx SecurityContext, operation string) error {
|
|||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
if rules.SecurityDisabled {
|
if rules.SecurityDisabled {
|
||||||
return nil
|
return nil
|
||||||
@@ -333,6 +364,31 @@ func CheckModelAuthAllowed(secCtx SecurityContext, operation string) error {
|
|||||||
return nil
|
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.
|
// CheckModelUpdateAllowed is the public wrapper for checkModelUpdateAllowed.
|
||||||
func CheckModelUpdateAllowed(secCtx SecurityContext) error {
|
func CheckModelUpdateAllowed(secCtx SecurityContext) error {
|
||||||
return checkModelUpdateAllowed(secCtx)
|
return checkModelUpdateAllowed(secCtx)
|
||||||
|
|||||||
@@ -26,6 +26,10 @@ func (m *mockSecurityContext) GetUserID() (int, bool) {
|
|||||||
return m.userID, m.hasUser
|
return m.userID, m.hasUser
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (m *mockSecurityContext) GetUserRef() (any, bool) {
|
||||||
|
return m.userID, m.hasUser
|
||||||
|
}
|
||||||
|
|
||||||
func (m *mockSecurityContext) GetSchema() string {
|
func (m *mockSecurityContext) GetSchema() string {
|
||||||
return m.schema
|
return m.schema
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -121,8 +121,12 @@ type ColumnSecurityProvider interface {
|
|||||||
|
|
||||||
// RowSecurityProvider handles row-level security (filtering)
|
// RowSecurityProvider handles row-level security (filtering)
|
||||||
type RowSecurityProvider interface {
|
type RowSecurityProvider interface {
|
||||||
// GetRowSecurity loads row security rules for a user and entity
|
// GetRowSecurity loads row security rules for a user and entity.
|
||||||
GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error)
|
// 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
|
// SecurityProvider is the main interface combining all security concerns
|
||||||
|
|||||||
+399
-10
@@ -3,8 +3,11 @@ package security
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
|
"crypto/rsa"
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
|
"crypto/subtle"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
|
"encoding/hex"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -12,6 +15,9 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/golang-jwt/jwt/v5"
|
||||||
|
"golang.org/x/oauth2"
|
||||||
)
|
)
|
||||||
|
|
||||||
// OAuthServerConfig configures the MCP-standard OAuth2 authorization server.
|
// OAuthServerConfig configures the MCP-standard OAuth2 authorization server.
|
||||||
@@ -44,6 +50,15 @@ type OAuthServerConfig struct {
|
|||||||
|
|
||||||
// AuthCodeTTL is the auth code lifetime. Defaults to 2 minutes.
|
// AuthCodeTTL is the auth code lifetime. Defaults to 2 minutes.
|
||||||
AuthCodeTTL time.Duration
|
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).
|
// oauthClient is a dynamically registered OAuth2 client (RFC 7591).
|
||||||
@@ -53,6 +68,14 @@ type oauthClient struct {
|
|||||||
ClientName string `json:"client_name,omitempty"`
|
ClientName string `json:"client_name,omitempty"`
|
||||||
GrantTypes []string `json:"grant_types"`
|
GrantTypes []string `json:"grant_types"`
|
||||||
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
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.
|
// pendingAuth tracks an in-progress authorization code exchange.
|
||||||
@@ -85,13 +108,25 @@ type externalProvider struct {
|
|||||||
// The server exposes these RFC-compliant endpoints:
|
// The server exposes these RFC-compliant endpoints:
|
||||||
//
|
//
|
||||||
// GET /.well-known/oauth-authorization-server RFC 8414 — server metadata discovery
|
// 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
|
// POST /oauth/register RFC 7591 — dynamic client registration
|
||||||
// GET /oauth/authorize OAuth 2.1 + PKCE — start authorization
|
// GET /oauth/authorize OAuth 2.1 + PKCE — start authorization
|
||||||
// POST /oauth/authorize Direct login form submission
|
// 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/revoke RFC 7009 — token revocation
|
||||||
// POST /oauth/introspect RFC 7662 — token introspection
|
// 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
|
// 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 {
|
type OAuthServer struct {
|
||||||
cfg OAuthServerConfig
|
cfg OAuthServerConfig
|
||||||
auth *DatabaseAuthenticator // nil = only external providers
|
auth *DatabaseAuthenticator // nil = only external providers
|
||||||
@@ -102,6 +137,9 @@ type OAuthServer struct {
|
|||||||
pending map[string]*pendingAuth // provider_state → pending (external flow)
|
pending map[string]*pendingAuth // provider_state → pending (external flow)
|
||||||
codes map[string]*pendingAuth // auth_code → pending (post-auth)
|
codes map[string]*pendingAuth // auth_code → pending (post-auth)
|
||||||
|
|
||||||
|
signingKey *rsa.PrivateKey
|
||||||
|
signingKeyID string
|
||||||
|
|
||||||
done chan struct{} // closed by Close() to stop background goroutines
|
done chan struct{} // closed by Close() to stop background goroutines
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -130,14 +168,33 @@ func NewOAuthServer(cfg OAuthServerConfig, auth *DatabaseAuthenticator) *OAuthSe
|
|||||||
}
|
}
|
||||||
// Normalize issuer: remove trailing slash to ensure consistent endpoint URL construction.
|
// Normalize issuer: remove trailing slash to ensure consistent endpoint URL construction.
|
||||||
cfg.Issuer = strings.TrimSuffix(cfg.Issuer, "/")
|
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{
|
s := &OAuthServer{
|
||||||
cfg: cfg,
|
cfg: cfg,
|
||||||
auth: auth,
|
auth: auth,
|
||||||
clients: make(map[string]*oauthClient),
|
clients: make(map[string]*oauthClient),
|
||||||
pending: make(map[string]*pendingAuth),
|
pending: make(map[string]*pendingAuth),
|
||||||
codes: make(map[string]*pendingAuth),
|
codes: make(map[string]*pendingAuth),
|
||||||
|
signingKey: signingKey,
|
||||||
done: make(chan struct{}),
|
done: make(chan struct{}),
|
||||||
}
|
}
|
||||||
|
if signingKey != nil {
|
||||||
|
s.signingKeyID = rsaKeyID(&signingKey.PublicKey)
|
||||||
|
}
|
||||||
go s.cleanupExpired()
|
go s.cleanupExpired()
|
||||||
return s
|
return s
|
||||||
}
|
}
|
||||||
@@ -178,11 +235,15 @@ func (s *OAuthServer) ProviderCallbackPath() string {
|
|||||||
func (s *OAuthServer) HTTPHandler() http.Handler {
|
func (s *OAuthServer) HTTPHandler() http.Handler {
|
||||||
mux := http.NewServeMux()
|
mux := http.NewServeMux()
|
||||||
mux.HandleFunc("/.well-known/oauth-authorization-server", s.metadataHandler)
|
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/register", s.registerHandler)
|
||||||
mux.HandleFunc("/oauth/authorize", s.authorizeHandler)
|
mux.HandleFunc("/oauth/authorize", s.authorizeHandler)
|
||||||
mux.HandleFunc("/oauth/token", s.tokenHandler)
|
mux.HandleFunc("/oauth/token", s.tokenHandler)
|
||||||
mux.HandleFunc("/oauth/revoke", s.revokeHandler)
|
mux.HandleFunc("/oauth/revoke", s.revokeHandler)
|
||||||
mux.HandleFunc("/oauth/introspect", s.introspectHandler)
|
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)
|
mux.HandleFunc(s.cfg.ProviderCallbackPath, s.providerCallbackHandler)
|
||||||
return mux
|
return mux
|
||||||
}
|
}
|
||||||
@@ -217,25 +278,127 @@ func (s *OAuthServer) cleanupExpired() {
|
|||||||
// RFC 8414 — Server metadata
|
// 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
|
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,
|
"issuer": issuer,
|
||||||
"authorization_endpoint": issuer + "/oauth/authorize",
|
"authorization_endpoint": issuer + "/oauth/authorize",
|
||||||
"token_endpoint": issuer + "/oauth/token",
|
"token_endpoint": issuer + "/oauth/token",
|
||||||
"registration_endpoint": issuer + "/oauth/register",
|
"registration_endpoint": issuer + "/oauth/register",
|
||||||
"revocation_endpoint": issuer + "/oauth/revoke",
|
"revocation_endpoint": issuer + "/oauth/revoke",
|
||||||
"introspection_endpoint": issuer + "/oauth/introspect",
|
"introspection_endpoint": issuer + "/oauth/introspect",
|
||||||
|
"userinfo_endpoint": issuer + "/oauth/userinfo",
|
||||||
|
"jwks_uri": issuer + "/oauth/jwks.json",
|
||||||
"scopes_supported": s.cfg.DefaultScopes,
|
"scopes_supported": s.cfg.DefaultScopes,
|
||||||
"response_types_supported": []string{"code"},
|
"response_types_supported": []string{"code"},
|
||||||
"grant_types_supported": []string{"authorization_code", "refresh_token"},
|
"grant_types_supported": grantTypes,
|
||||||
"code_challenge_methods_supported": []string{"S256"},
|
"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")
|
w.Header().Set("Content-Type", "application/json")
|
||||||
json.NewEncoder(w).Encode(meta) //nolint:errcheck
|
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
|
// RFC 7591 — Dynamic client registration
|
||||||
// --------------------------------------------------------------------------
|
// --------------------------------------------------------------------------
|
||||||
@@ -250,6 +413,7 @@ func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
|
|||||||
ClientName string `json:"client_name"`
|
ClientName string `json:"client_name"`
|
||||||
GrantTypes []string `json:"grant_types"`
|
GrantTypes []string `json:"grant_types"`
|
||||||
AllowedScopes []string `json:"allowed_scopes"`
|
AllowedScopes []string `json:"allowed_scopes"`
|
||||||
|
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"`
|
||||||
}
|
}
|
||||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||||
writeOAuthError(w, "invalid_request", "malformed JSON", http.StatusBadRequest)
|
writeOAuthError(w, "invalid_request", "malformed JSON", http.StatusBadRequest)
|
||||||
@@ -272,12 +436,37 @@ func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
|
|||||||
http.Error(w, "server error", http.StatusInternalServerError)
|
http.Error(w, "server error", http.StatusInternalServerError)
|
||||||
return
|
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{
|
client := &oauthClient{
|
||||||
ClientID: clientID,
|
ClientID: clientID,
|
||||||
RedirectURIs: req.RedirectURIs,
|
RedirectURIs: req.RedirectURIs,
|
||||||
ClientName: req.ClientName,
|
ClientName: req.ClientName,
|
||||||
GrantTypes: grantTypes,
|
GrantTypes: grantTypes,
|
||||||
AllowedScopes: allowedScopes,
|
AllowedScopes: allowedScopes,
|
||||||
|
ClientSecretHash: secretHash,
|
||||||
|
TokenEndpointAuthMethod: authMethod,
|
||||||
}
|
}
|
||||||
|
|
||||||
if s.cfg.PersistClients && s.auth != nil {
|
if s.cfg.PersistClients && s.auth != nil {
|
||||||
@@ -287,6 +476,8 @@ func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
|
|||||||
ClientName: client.ClientName,
|
ClientName: client.ClientName,
|
||||||
GrantTypes: client.GrantTypes,
|
GrantTypes: client.GrantTypes,
|
||||||
AllowedScopes: client.AllowedScopes,
|
AllowedScopes: client.AllowedScopes,
|
||||||
|
ClientSecretHash: client.ClientSecretHash,
|
||||||
|
TokenEndpointAuthMethod: client.TokenEndpointAuthMethod,
|
||||||
}
|
}
|
||||||
if _, err := s.auth.OAuthRegisterClient(r.Context(), dbClient); err != nil {
|
if _, err := s.auth.OAuthRegisterClient(r.Context(), dbClient); err != nil {
|
||||||
http.Error(w, "server error", http.StatusInternalServerError)
|
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.clients[clientID] = client
|
||||||
s.mu.Unlock()
|
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.Header().Set("Content-Type", "application/json")
|
||||||
w.WriteHeader(http.StatusCreated)
|
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)
|
s.handleAuthCodeGrant(w, r)
|
||||||
case "refresh_token":
|
case "refresh_token":
|
||||||
s.handleRefreshGrant(w, r)
|
s.handleRefreshGrant(w, r)
|
||||||
|
case "client_credentials":
|
||||||
|
s.handleClientCredentialsGrant(w, r)
|
||||||
default:
|
default:
|
||||||
writeOAuthError(w, "unsupported_grant_type", "", http.StatusBadRequest)
|
writeOAuthError(w, "unsupported_grant_type", "", http.StatusBadRequest)
|
||||||
}
|
}
|
||||||
@@ -593,6 +802,15 @@ func (s *OAuthServer) handleAuthCodeGrant(w http.ResponseWriter, r *http.Request
|
|||||||
return
|
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 sessionToken string
|
||||||
var refreshToken string
|
var refreshToken string
|
||||||
var scopes []string
|
var scopes []string
|
||||||
@@ -647,12 +865,13 @@ func (s *OAuthServer) handleAuthCodeGrant(w http.ResponseWriter, r *http.Request
|
|||||||
scopes = pending.Scopes
|
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) {
|
func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request) {
|
||||||
refreshToken := r.FormValue("refresh_token")
|
refreshToken := r.FormValue("refresh_token")
|
||||||
providerName := r.FormValue("provider")
|
providerName := r.FormValue("provider")
|
||||||
|
clientID := r.FormValue("client_id")
|
||||||
if refreshToken == "" {
|
if refreshToken == "" {
|
||||||
writeOAuthError(w, "invalid_request", "refresh_token required", http.StatusBadRequest)
|
writeOAuthError(w, "invalid_request", "refresh_token required", http.StatusBadRequest)
|
||||||
return
|
return
|
||||||
@@ -666,7 +885,7 @@ func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request)
|
|||||||
writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest)
|
writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
s.writeOAuthToken(w, loginResp.Token, loginResp.RefreshToken, nil)
|
s.writeOAuthToken(w, r, loginResp.Token, loginResp.RefreshToken, clientID, nil, false)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -676,13 +895,86 @@ func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request)
|
|||||||
writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest)
|
writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
s.writeOAuthToken(w, loginResp.Token, loginResp.RefreshToken, nil)
|
s.writeOAuthToken(w, r, loginResp.Token, loginResp.RefreshToken, clientID, nil, false)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
writeOAuthError(w, "invalid_grant", "no provider available for refresh", http.StatusBadRequest)
|
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
|
// RFC 7009 — Token revocation
|
||||||
// --------------------------------------------------------------------------
|
// --------------------------------------------------------------------------
|
||||||
@@ -879,7 +1171,10 @@ func oauthSliceContains(slice []string, s string) bool {
|
|||||||
return false
|
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())
|
expiresIn := int64(s.cfg.AccessTokenTTL.Seconds())
|
||||||
resp := map[string]interface{}{
|
resp := map[string]interface{}{
|
||||||
"access_token": accessToken,
|
"access_token": accessToken,
|
||||||
@@ -892,12 +1187,106 @@ func (s *OAuthServer) writeOAuthToken(w http.ResponseWriter, accessToken, refres
|
|||||||
if len(scopes) > 0 {
|
if len(scopes) > 0 {
|
||||||
resp["scope"] = strings.Join(scopes, " ")
|
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("Content-Type", "application/json")
|
||||||
w.Header().Set("Cache-Control", "no-store")
|
w.Header().Set("Cache-Control", "no-store")
|
||||||
w.Header().Set("Pragma", "no-cache")
|
w.Header().Set("Pragma", "no-cache")
|
||||||
json.NewEncoder(w).Encode(resp) //nolint:errcheck
|
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) {
|
func writeOAuthError(w http.ResponseWriter, errCode, description string, status int) {
|
||||||
resp := map[string]string{"error": errCode}
|
resp := map[string]string{"error": errCode}
|
||||||
if description != "" {
|
if description != "" {
|
||||||
|
|||||||
@@ -14,6 +14,8 @@ type OAuthServerClient struct {
|
|||||||
ClientName string `json:"client_name,omitempty"`
|
ClientName string `json:"client_name,omitempty"`
|
||||||
GrantTypes []string `json:"grant_types"`
|
GrantTypes []string `json:"grant_types"`
|
||||||
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
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.
|
// OAuthCode is a short-lived authorization code.
|
||||||
|
|||||||
@@ -15,6 +15,15 @@ import (
|
|||||||
// Array columns (redirect_uris, grant_types, allowed_scopes, scopes) are
|
// Array columns (redirect_uris, grant_types, allowed_scopes, scopes) are
|
||||||
// JSON-encoded TEXT instead of native Postgres arrays.
|
// 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) {
|
func (a *DatabaseAuthenticator) oauthRegisterClientDirect(ctx context.Context, client *OAuthServerClient) (*OAuthServerClient, error) {
|
||||||
grantTypes := client.GrantTypes
|
grantTypes := client.GrantTypes
|
||||||
if len(grantTypes) == 0 {
|
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)
|
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 {
|
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||||
query := rewritePlaceholders(db, fmt.Sprintf(
|
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))
|
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
|
return err
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -55,18 +69,20 @@ func (a *DatabaseAuthenticator) oauthRegisterClientDirect(ctx context.Context, c
|
|||||||
ClientName: client.ClientName,
|
ClientName: client.ClientName,
|
||||||
GrantTypes: grantTypes,
|
GrantTypes: grantTypes,
|
||||||
AllowedScopes: allowedScopes,
|
AllowedScopes: allowedScopes,
|
||||||
|
ClientSecretHash: client.ClientSecretHash,
|
||||||
|
TokenEndpointAuthMethod: authMethod,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *DatabaseAuthenticator) oauthGetClientDirect(ctx context.Context, clientID string) (*OAuthServerClient, error) {
|
func (a *DatabaseAuthenticator) oauthGetClientDirect(ctx context.Context, clientID string) (*OAuthServerClient, error) {
|
||||||
var redirectURIsJSON, grantTypesJSON, allowedScopesJSON sql.NullString
|
var redirectURIsJSON, grantTypesJSON, allowedScopesJSON sql.NullString
|
||||||
var clientName sql.NullString
|
var clientName, clientSecretHash, authMethod sql.NullString
|
||||||
|
|
||||||
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
|
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||||
query := rewritePlaceholders(db, fmt.Sprintf(
|
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))
|
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 err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
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)
|
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 {
|
if redirectURIsJSON.Valid {
|
||||||
_ = json.Unmarshal([]byte(redirectURIsJSON.String), &result.RedirectURIs)
|
_ = 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
|
// If username is provided, get user's credentials
|
||||||
var allowCredentials []PasskeyCredentialDescriptor
|
var allowCredentials []PasskeyCredentialDescriptor
|
||||||
if username != "" {
|
if username != "" {
|
||||||
var creds []struct {
|
var creds []passkeyCredential
|
||||||
ID string `json:"credential_id"`
|
|
||||||
Transports []string `json:"transports"`
|
|
||||||
}
|
|
||||||
|
|
||||||
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyGetCredsByUsername) {
|
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyGetCredsByUsername) {
|
||||||
_, directCreds, err := p.getCredsByUsernameDirect(ctx, username)
|
_, directCreds, err := p.getCredsByUsernameDirect(ctx, username)
|
||||||
|
|||||||
@@ -39,7 +39,7 @@ func (p *DatabasePasskeyProvider) storeCredentialDirect(ctx context.Context, par
|
|||||||
var exists int
|
var exists int
|
||||||
checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT 1 FROM %s WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials))
|
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 {
|
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) {
|
} else if !errors.Is(err, sql.ErrNoRows) {
|
||||||
return err
|
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))
|
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 err := db.QueryRowContext(ctx, userCheckQuery, params.UserID).Scan(&userExists); err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
return fmt.Errorf("User not found")
|
return fmt.Errorf("user not found")
|
||||||
}
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -78,7 +78,7 @@ func (p *DatabasePasskeyProvider) getCredentialDirect(ctx context.Context, crede
|
|||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
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)
|
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))
|
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 err := db.QueryRowContext(ctx, query, credentialIDB64).Scan(&oldCounter); err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
return fmt.Errorf("Credential not found")
|
return fmt.Errorf("credential not found")
|
||||||
}
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -188,7 +188,7 @@ func (p *DatabasePasskeyProvider) deleteCredentialDirect(ctx context.Context, us
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if rows == 0 {
|
if rows == 0 {
|
||||||
return fmt.Errorf("Credential not found")
|
return fmt.Errorf("credential not found")
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
@@ -206,24 +206,19 @@ func (p *DatabasePasskeyProvider) updateNameDirect(ctx context.Context, userID i
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if rows == 0 {
|
if rows == 0 {
|
||||||
return fmt.Errorf("Credential not found")
|
return fmt.Errorf("credential not found")
|
||||||
}
|
}
|
||||||
return nil
|
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"`
|
ID string `json:"credential_id"`
|
||||||
Transports []string `json:"transports"`
|
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))
|
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 {
|
if err := db.QueryRowContext(ctx, userQuery, username, true).Scan(&userID); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -236,7 +231,7 @@ func (p *DatabasePasskeyProvider) getCredsByUsernameDirect(ctx context.Context,
|
|||||||
}
|
}
|
||||||
defer rows.Close()
|
defer rows.Close()
|
||||||
|
|
||||||
creds = make([]credT, 0)
|
creds = make([]passkeyCredential, 0)
|
||||||
for rows.Next() {
|
for rows.Next() {
|
||||||
var credID string
|
var credID string
|
||||||
var transportsJSON sql.NullString
|
var transportsJSON sql.NullString
|
||||||
@@ -247,13 +242,13 @@ func (p *DatabasePasskeyProvider) getCredsByUsernameDirect(ctx context.Context,
|
|||||||
if transportsJSON.Valid && transportsJSON.String != "" {
|
if transportsJSON.Valid && transportsJSON.String != "" {
|
||||||
_ = json.Unmarshal([]byte(transportsJSON.String), &transports)
|
_ = 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()
|
return rows.Err()
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
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)
|
return 0, nil, fmt.Errorf("failed to get credentials: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -34,7 +34,10 @@ type RowSecurity struct {
|
|||||||
Tablename string `json:"tablename"`
|
Tablename string `json:"tablename"`
|
||||||
Template string `json:"template"`
|
Template string `json:"template"`
|
||||||
HasBlock bool `json:"has_block"`
|
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 {
|
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, "{PrimaryKeyName}", pPrimaryKeyName)
|
||||||
str = strings.ReplaceAll(str, "{TableName}", m.Tablename)
|
str = strings.ReplaceAll(str, "{TableName}", m.Tablename)
|
||||||
str = strings.ReplaceAll(str, "{SchemaName}", m.Schema)
|
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
|
return str
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -413,7 +416,7 @@ func (m *SecurityList) ClearSecurity(pUserID int, pSchema, pTablename string) er
|
|||||||
return nil
|
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 {
|
if m.provider == nil {
|
||||||
return RowSecurity{}, fmt.Errorf("security provider not set")
|
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 {
|
if m.RowSecurity == nil {
|
||||||
m.RowSecurity = make(map[string]RowSecurity, 0)
|
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
|
// 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 {
|
if err != nil {
|
||||||
return RowSecurity{}, fmt.Errorf("GetRowSecurity failed: %v", err)
|
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
|
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")()
|
defer logger.CatchPanic("GetRowSecurityTemplate")()
|
||||||
|
|
||||||
if m.RowSecurity == nil {
|
if m.RowSecurity == nil {
|
||||||
@@ -446,7 +449,7 @@ func (m *SecurityList) GetRowSecurityTemplate(pUserID int, pSchema, pTablename s
|
|||||||
m.RowSecurityMutex.RLock()
|
m.RowSecurityMutex.RLock()
|
||||||
defer m.RowSecurityMutex.RUnlock()
|
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 {
|
if !ok {
|
||||||
return RowSecurity{}, fmt.Errorf("no row security data")
|
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
|
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
|
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 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
|
Token: userCtx.SessionID, // New session token from stored procedure
|
||||||
User: &userCtx,
|
User: &userCtx,
|
||||||
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
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
|
// JWTAuthenticator provides JWT token-based authentication
|
||||||
@@ -907,17 +923,29 @@ func (p *DatabaseRowSecurityProvider) reconnectDB() error {
|
|||||||
return nil
|
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) {
|
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.RowSecurity) {
|
||||||
return RowSecurity{}, ErrDirectModeUnsupported
|
return RowSecurity{}, ErrDirectModeUnsupported
|
||||||
}
|
}
|
||||||
|
|
||||||
var template string
|
// resolvespec_row_security's p_user_id is a scalar integer. GetUserRef() may
|
||||||
var hasBlock bool
|
// 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 {
|
runQuery := func() error {
|
||||||
query := fmt.Sprintf(`SELECT p_template, p_block FROM %s($1, $2, $3)`, p.sqlNames.RowSecurity)
|
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()
|
err := runQuery()
|
||||||
if isDBClosed(err) {
|
if isDBClosed(err) {
|
||||||
@@ -932,9 +960,9 @@ func (p *DatabaseRowSecurityProvider) GetRowSecurity(ctx context.Context, userID
|
|||||||
return RowSecurity{
|
return RowSecurity{
|
||||||
Schema: schema,
|
Schema: schema,
|
||||||
Tablename: table,
|
Tablename: table,
|
||||||
UserID: userID,
|
UserID: userRef,
|
||||||
Template: template,
|
Template: template.String,
|
||||||
HasBlock: hasBlock,
|
HasBlock: hasBlock.Bool,
|
||||||
}, nil
|
}, 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)
|
key := fmt.Sprintf("%s.%s", schema, table)
|
||||||
|
|
||||||
if p.blocked[key] {
|
if p.blocked[key] {
|
||||||
return RowSecurity{
|
return RowSecurity{
|
||||||
Schema: schema,
|
Schema: schema,
|
||||||
Tablename: table,
|
Tablename: table,
|
||||||
UserID: userID,
|
UserID: userRef,
|
||||||
HasBlock: true,
|
HasBlock: true,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
@@ -985,7 +1013,7 @@ func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userID i
|
|||||||
return RowSecurity{
|
return RowSecurity{
|
||||||
Schema: schema,
|
Schema: schema,
|
||||||
Tablename: table,
|
Tablename: table,
|
||||||
UserID: userID,
|
UserID: userRef,
|
||||||
Template: template,
|
Template: template,
|
||||||
HasBlock: false,
|
HasBlock: false,
|
||||||
}, nil
|
}, nil
|
||||||
|
|||||||
@@ -23,8 +23,8 @@ import (
|
|||||||
// than introducing a mismatch between modes.
|
// than introducing a mismatch between modes.
|
||||||
|
|
||||||
var (
|
var (
|
||||||
errUsernameExists = errors.New("Username already exists")
|
errUsernameExists = errors.New("username already exists")
|
||||||
errEmailExists = errors.New("Email already exists")
|
errEmailExists = errors.New("email already exists")
|
||||||
)
|
)
|
||||||
|
|
||||||
func (a *DatabaseAuthenticator) loginDirect(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
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 err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
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)
|
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) {
|
func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req RegisterRequest) (*LoginResponse, error) {
|
||||||
if req.Username == "" {
|
if req.Username == "" {
|
||||||
return nil, fmt.Errorf("Username is required")
|
return nil, fmt.Errorf("username is required")
|
||||||
}
|
}
|
||||||
if req.Email == "" {
|
if req.Email == "" {
|
||||||
return nil, fmt.Errorf("Email is required")
|
return nil, fmt.Errorf("email is required")
|
||||||
}
|
}
|
||||||
if req.Password == "" {
|
if req.Password == "" {
|
||||||
return nil, fmt.Errorf("Password is required")
|
return nil, fmt.Errorf("password is required")
|
||||||
}
|
}
|
||||||
|
|
||||||
rolesStr := strings.Join(req.Roles, ",")
|
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)
|
return fmt.Errorf("logout query failed: %w", err)
|
||||||
}
|
}
|
||||||
if rows == 0 {
|
if rows == 0 {
|
||||||
return fmt.Errorf("Session not found")
|
return fmt.Errorf("session not found")
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Token != "" {
|
if req.Token != "" {
|
||||||
@@ -222,7 +222,7 @@ func (a *DatabaseAuthenticator) sessionDirect(ctx context.Context, token string)
|
|||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
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)
|
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 err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
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)
|
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)
|
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) {
|
func TestDatabaseAuthenticatorReconnectsClosedDBPaths(t *testing.T) {
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ func (p *DatabaseTwoFactorProvider) enable2FADirect(ctx context.Context, userID
|
|||||||
if rows, err := res.RowsAffected(); err != nil {
|
if rows, err := res.RowsAffected(); err != nil {
|
||||||
return err
|
return err
|
||||||
} else if rows == 0 {
|
} 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))
|
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 {
|
if rows, err := res.RowsAffected(); err != nil {
|
||||||
return err
|
return err
|
||||||
} else if rows == 0 {
|
} 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))
|
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 err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
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)
|
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 err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
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)
|
return "", fmt.Errorf("get 2FA secret query failed: %w", err)
|
||||||
}
|
}
|
||||||
@@ -101,7 +101,7 @@ func (p *DatabaseTwoFactorProvider) regenerateBackupCodesDirect(ctx context.Cont
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if count == 0 {
|
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))
|
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
|
return err
|
||||||
}
|
}
|
||||||
if used {
|
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))
|
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]
|
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] }
|
type SqlTimeStamp struct{ SqlNull[time.Time] }
|
||||||
|
|
||||||
func (t SqlTimeStamp) MarshalJSON() ([]byte, error) {
|
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)) {
|
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("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 {
|
func (t *SqlTimeStamp) UnmarshalJSON(b []byte) error {
|
||||||
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
|
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
|
||||||
return err
|
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
|
t.Valid = false
|
||||||
}
|
}
|
||||||
return nil
|
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)) {
|
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
return t.Val.Format("2006-01-02T15:04:05"), nil
|
return t.Val.Format(time.RFC3339), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func SqlTimeStampNow() SqlTimeStamp {
|
func SqlTimeStampNow() SqlTimeStamp {
|
||||||
@@ -425,9 +425,7 @@ func (t SqlTime) MarshalJSON() ([]byte, error) {
|
|||||||
return []byte("null"), nil
|
return []byte("null"), nil
|
||||||
}
|
}
|
||||||
s := t.Val.Format("15:04:05")
|
s := t.Val.Format("15:04:05")
|
||||||
if s == "00:00:00" {
|
|
||||||
return []byte("null"), nil
|
|
||||||
}
|
|
||||||
return []byte(fmt.Sprintf(`"%s"`, s)), 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 {
|
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if t.Valid && t.Val.Format("15:04:05") == "00:00:00" {
|
|
||||||
t.Valid = false
|
|
||||||
}
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -178,7 +178,7 @@ func TestSqlTimeStamp_JSON(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Marshal failed: %v", err)
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
}
|
}
|
||||||
expected := `"2024-01-15T10:30:45"`
|
expected := `"2024-01-15T10:30:45Z"`
|
||||||
if string(data) != expected {
|
if string(data) != expected {
|
||||||
t.Errorf("expected %s, got %s", expected, string(data))
|
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
|
// TestSqlByteArray_Base64_RoundTrip tests complete round-trip: Go -> JSON -> Go -> SQL -> Go
|
||||||
func TestSqlByteArray_Base64_RoundTrip(t *testing.T) {
|
func TestSqlByteArray_Base64_RoundTrip(t *testing.T) {
|
||||||
original := []byte{0x48, 0x65, 0x6C, 0x6C, 0x6F, 0x20, 0xFF, 0xFE} // "Hello " + binary data
|
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)
|
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
|
// handleRead processes a read operation
|
||||||
func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContext) {
|
func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContext) {
|
||||||
// Execute before hook
|
// 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)
|
logger.Error("[WebSocketSpec] BeforeRead hook failed: %v", err)
|
||||||
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
||||||
_ = conn.SendJSON(errResp)
|
_ = conn.SendJSON(errResp)
|
||||||
@@ -273,7 +273,7 @@ func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContex
|
|||||||
// handleCreate processes a create operation
|
// handleCreate processes a create operation
|
||||||
func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookContext) {
|
func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookContext) {
|
||||||
// Execute before hook
|
// 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)
|
logger.Error("[WebSocketSpec] BeforeCreate hook failed: %v", err)
|
||||||
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
||||||
_ = conn.SendJSON(errResp)
|
_ = conn.SendJSON(errResp)
|
||||||
@@ -311,7 +311,7 @@ func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookCont
|
|||||||
// handleUpdate processes an update operation
|
// handleUpdate processes an update operation
|
||||||
func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookContext) {
|
func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookContext) {
|
||||||
// Execute before hook
|
// 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)
|
logger.Error("[WebSocketSpec] BeforeUpdate hook failed: %v", err)
|
||||||
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
||||||
_ = conn.SendJSON(errResp)
|
_ = conn.SendJSON(errResp)
|
||||||
@@ -349,7 +349,7 @@ func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookCont
|
|||||||
// handleDelete processes a delete operation
|
// handleDelete processes a delete operation
|
||||||
func (h *Handler) handleDelete(conn *Connection, msg *Message, hookCtx *HookContext) {
|
func (h *Handler) handleDelete(conn *Connection, msg *Message, hookCtx *HookContext) {
|
||||||
// Execute before hook
|
// 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)
|
logger.Error("[WebSocketSpec] BeforeDelete hook failed: %v", err)
|
||||||
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
|
||||||
_ = conn.SendJSON(errResp)
|
_ = conn.SendJSON(errResp)
|
||||||
@@ -564,7 +564,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) {
|
|||||||
|
|
||||||
// Apply columns
|
// Apply columns
|
||||||
if hookCtx.Options != nil && len(hookCtx.Options.Columns) > 0 {
|
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)
|
// 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
|
// Execute query
|
||||||
if err := query.ScanModel(hookCtx.Context); err != nil {
|
if err := query.ScanModel(hookCtx.Context); err != nil {
|
||||||
return nil, fmt.Errorf("failed to read record: %w", err)
|
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)
|
// Apply options (simplified implementation)
|
||||||
if hookCtx.Options != nil {
|
if hookCtx.Options != nil {
|
||||||
// Apply filters with OR grouping support
|
// Apply filters with OR grouping support
|
||||||
query = h.applyFilters(query, hookCtx.Options.Filters)
|
query = h.applyFilters(query, hookCtx.Options.Filters, hookCtx.Model)
|
||||||
|
|
||||||
// Apply sorting
|
// Apply sorting
|
||||||
for _, sort := range hookCtx.Options.Sort {
|
for _, sort := range hookCtx.Options.Sort {
|
||||||
@@ -602,6 +614,10 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
|||||||
if sort.Direction == "desc" {
|
if sort.Direction == "desc" {
|
||||||
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))
|
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
|
// Apply columns
|
||||||
if len(hookCtx.Options.Columns) > 0 {
|
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
|
// Execute query
|
||||||
if err := query.ScanModel(hookCtx.Context); err != nil {
|
if err := query.ScanModel(hookCtx.Context); err != nil {
|
||||||
return nil, nil, fmt.Errorf("failed to read records: %w", err)
|
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)
|
countQuery := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
||||||
if hookCtx.Options != nil {
|
if hookCtx.Options != nil {
|
||||||
for _, filter := range hookCtx.Options.Filters {
|
for _, filter := range hookCtx.Options.Filters {
|
||||||
cond, args := h.buildFilterCondition(filter)
|
cond, args := h.buildFilterCondition(filter, hookCtx.Model)
|
||||||
if cond != "" {
|
if cond != "" {
|
||||||
countQuery = countQuery.Where(cond, args...)
|
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
|
// getOperatorSQL converts filter operator to SQL operator
|
||||||
// applyFilters applies all filters with proper grouping for OR logic
|
// applyFilters applies all filters with proper grouping for OR logic
|
||||||
// Groups consecutive OR filters together to ensure proper query precedence
|
// 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 {
|
if len(filters) == 0 {
|
||||||
return query
|
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
|
// Apply the OR group as a single grouped WHERE clause
|
||||||
query = h.applyFilterGroup(query, orGroup)
|
query = h.applyFilterGroup(query, orGroup, model)
|
||||||
i = j
|
i = j
|
||||||
} else {
|
} else {
|
||||||
// Single filter with AND logic (or first filter)
|
// Single filter with AND logic (or first filter)
|
||||||
condition, args := h.buildFilterCondition(filters[i])
|
condition, args := h.buildFilterCondition(filters[i], model)
|
||||||
if condition != "" {
|
if condition != "" {
|
||||||
query = query.Where(condition, args...)
|
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
|
// applyFilterGroup applies a group of filters that should be OR'd together
|
||||||
// Always wraps them in parentheses and applies as a single WHERE clause
|
// 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 {
|
if len(filters) == 0 {
|
||||||
return query
|
return query
|
||||||
}
|
}
|
||||||
@@ -799,7 +827,7 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
|||||||
var args []interface{}
|
var args []interface{}
|
||||||
|
|
||||||
for _, filter := range filters {
|
for _, filter := range filters {
|
||||||
condition, filterArgs := h.buildFilterCondition(filter)
|
condition, filterArgs := h.buildFilterCondition(filter, model)
|
||||||
if condition != "" {
|
if condition != "" {
|
||||||
conditions = append(conditions, condition)
|
conditions = append(conditions, condition)
|
||||||
args = append(args, filterArgs...)
|
args = append(args, filterArgs...)
|
||||||
@@ -820,8 +848,14 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
|||||||
return query.Where(groupedCondition, args...)
|
return query.Where(groupedCondition, args...)
|
||||||
}
|
}
|
||||||
|
|
||||||
// buildFilterCondition builds a filter condition and returns it with args
|
// buildFilterCondition builds a filter condition and returns it with args.
|
||||||
func (h *Handler) buildFilterCondition(filter common.FilterOption) (conditionString string, conditionArgs []interface{}) {
|
// 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") {
|
if strings.EqualFold(filter.Operator, "in") {
|
||||||
cond, args := common.BuildInCondition(filter.Column, filter.Value)
|
cond, args := common.BuildInCondition(filter.Column, filter.Value)
|
||||||
return cond, args
|
return cond, args
|
||||||
@@ -829,6 +863,11 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption) (conditionStr
|
|||||||
op := strings.ToLower(filter.Operator)
|
op := strings.ToLower(filter.Operator)
|
||||||
if op == "like" || op == "ilike" {
|
if op == "like" || op == "ilike" {
|
||||||
operatorSQL := h.getOperatorSQL(filter.Operator)
|
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}
|
return fmt.Sprintf("CAST(%s AS TEXT) %s ?", filter.Column, operatorSQL), []interface{}{filter.Value}
|
||||||
}
|
}
|
||||||
operatorSQL := h.getOperatorSQL(filter.Operator)
|
operatorSQL := h.getOperatorSQL(filter.Operator)
|
||||||
|
|||||||
@@ -35,6 +35,11 @@ const (
|
|||||||
// AfterDelete is called after a delete operation
|
// AfterDelete is called after a delete operation
|
||||||
AfterDelete HookType = "after_delete"
|
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 is called before creating a subscription
|
||||||
BeforeSubscribe HookType = "before_subscribe"
|
BeforeSubscribe HookType = "before_subscribe"
|
||||||
// AfterSubscribe is called after creating a subscription
|
// AfterSubscribe is called after creating a subscription
|
||||||
@@ -54,6 +59,11 @@ const (
|
|||||||
BeforeDisconnect HookType = "before_disconnect"
|
BeforeDisconnect HookType = "before_disconnect"
|
||||||
// AfterDisconnect is called after a connection is closed
|
// AfterDisconnect is called after a connection is closed
|
||||||
AfterDisconnect HookType = "after_disconnect"
|
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
|
// HookContext contains context information for hook execution
|
||||||
@@ -197,6 +207,16 @@ func (hr *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
|||||||
return nil
|
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
|
// HasHooks checks if any hooks are registered for a hook type
|
||||||
func (hr *HookRegistry) HasHooks(hookType HookType) bool {
|
func (hr *HookRegistry) HasHooks(hookType HookType) bool {
|
||||||
hooks, exists := hr.hooks[hookType]
|
hooks, exists := hr.hooks[hookType]
|
||||||
|
|||||||
@@ -27,25 +27,31 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
|||||||
return security.LoadSecurityRules(secCtx, 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 {
|
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
return security.ApplyColumnSecurity(secCtx, securityList)
|
return security.ApplyColumnSecurity(secCtx, securityList)
|
||||||
})
|
})
|
||||||
|
|
||||||
// Hook 3 (Optional): Audit logging
|
// Hook 4 (Optional): Audit logging
|
||||||
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
return security.LogDataAccess(secCtx)
|
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 {
|
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
return security.CheckModelUpdateAllowed(secCtx)
|
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 {
|
handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
return security.CheckModelDeleteAllowed(secCtx)
|
return security.CheckModelDeleteAllowed(secCtx)
|
||||||
@@ -71,6 +77,17 @@ func (s *securityContext) GetUserID() (int, bool) {
|
|||||||
return security.GetUserID(s.ctx.Context)
|
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 {
|
func (s *securityContext) GetSchema() string {
|
||||||
return s.ctx.Schema
|
return s.ctx.Schema
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,5 +1,12 @@
|
|||||||
# @warkypublic/resolvespec-js
|
# @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
|
## 1.0.1
|
||||||
|
|
||||||
### Patch Changes
|
### Patch Changes
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ import { ResolveSpecClient, getResolveSpecClient } from '@warkypublic/resolvespe
|
|||||||
// Class instantiation
|
// Class instantiation
|
||||||
const client = new ResolveSpecClient({ baseUrl: 'http://localhost:3000', token: 'your-token' });
|
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' });
|
const client = getResolveSpecClient({ baseUrl: 'http://localhost:3000', token: 'your-token' });
|
||||||
|
|
||||||
// Read with filters, sort, pagination
|
// Read with filters, sort, pagination
|
||||||
@@ -211,3 +211,25 @@ pnpm run lint # eslint
|
|||||||
## License
|
## License
|
||||||
|
|
||||||
MIT
|
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 {
|
export * from './common';
|
||||||
code: string;
|
export * from './resolvespec';
|
||||||
message: string;
|
export * from './websocketspec';
|
||||||
details?: any;
|
export * from './headerspec';
|
||||||
detail?: string;
|
//# sourceMappingURL=index.d.ts.map
|
||||||
}
|
|
||||||
|
|
||||||
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 { }
|
|
||||||
Vendored
+221
-258
@@ -1,53 +1,75 @@
|
|||||||
import { v4 as l } from "uuid";
|
import { v4 as e } from "uuid";
|
||||||
const d = /* @__PURE__ */ new Map();
|
import { b64DecodeUnicode as t, b64EncodeUnicode as n } from "@warkypublic/artemis-kit/base64";
|
||||||
function E(n) {
|
//#region src/common/http.ts
|
||||||
const e = n.baseUrl;
|
function r(...e) {
|
||||||
let t = d.get(e);
|
let t = {};
|
||||||
return t || (t = new g(n), d.set(e, t)), 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
|
||||||
|
});
|
||||||
}
|
}
|
||||||
class g {
|
return t;
|
||||||
|
}
|
||||||
|
function i(e) {
|
||||||
|
return r({ "Content-Type": "application/json" }, e.headers ?? {}, e.token ? { Authorization: `Bearer ${e.token}` } : {});
|
||||||
|
}
|
||||||
|
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]);
|
||||||
|
}
|
||||||
|
//#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;
|
||||||
|
}
|
||||||
|
var c = class {
|
||||||
constructor(e) {
|
constructor(e) {
|
||||||
this.config = e;
|
this.config = {
|
||||||
|
...e,
|
||||||
|
headers: { ...e.headers }
|
||||||
|
};
|
||||||
}
|
}
|
||||||
buildUrl(e, t, s) {
|
buildUrl(e, t, n) {
|
||||||
let r = `${this.config.baseUrl}/${e}/${t}`;
|
let r = `${this.config.baseUrl}/${e}/${t}`;
|
||||||
return s && (r += `/${s}`), r;
|
return n && (r += `/${n}`), r;
|
||||||
}
|
}
|
||||||
baseHeaders() {
|
baseHeaders() {
|
||||||
const e = {
|
return i(this.config);
|
||||||
"Content-Type": "application/json"
|
|
||||||
};
|
|
||||||
return this.config.token && (e.Authorization = `Bearer ${this.config.token}`), e;
|
|
||||||
}
|
}
|
||||||
async fetchWithError(e, t) {
|
async fetchWithError(e, t) {
|
||||||
const s = await fetch(e, t), r = await s.json();
|
let n = await fetch(e, t), r = await n.json();
|
||||||
if (!s.ok)
|
if (!n.ok) throw Error(r.error?.message || "An error occurred");
|
||||||
throw new Error(r.error?.message || "An error occurred");
|
|
||||||
return r;
|
return r;
|
||||||
}
|
}
|
||||||
async getMetadata(e, t) {
|
async getMetadata(e, t) {
|
||||||
const s = this.buildUrl(e, t);
|
let n = this.buildUrl(e, t);
|
||||||
return this.fetchWithError(s, {
|
return this.fetchWithError(n, {
|
||||||
method: "GET",
|
method: "GET",
|
||||||
headers: this.baseHeaders()
|
headers: this.baseHeaders()
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
async read(e, t, s, r) {
|
async read(e, t, n, r) {
|
||||||
const i = typeof s == "number" || typeof s == "string" ? String(s) : void 0, a = this.buildUrl(e, t, i), c = {
|
let i = typeof n == "number" || typeof n == "string" ? String(n) : void 0, a = this.buildUrl(e, t, i), o = {
|
||||||
operation: "read",
|
operation: "read",
|
||||||
id: Array.isArray(s) ? s : void 0,
|
id: Array.isArray(n) ? n : void 0,
|
||||||
options: r
|
options: r
|
||||||
};
|
};
|
||||||
return this.fetchWithError(a, {
|
return this.fetchWithError(a, {
|
||||||
method: "POST",
|
method: "POST",
|
||||||
headers: this.baseHeaders(),
|
headers: this.baseHeaders(),
|
||||||
body: JSON.stringify(c)
|
body: JSON.stringify(o)
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
async create(e, t, s, r) {
|
async create(e, t, n, r) {
|
||||||
const i = this.buildUrl(e, t), a = {
|
let i = this.buildUrl(e, t), a = {
|
||||||
operation: "create",
|
operation: "create",
|
||||||
data: s,
|
data: n,
|
||||||
options: r
|
options: r
|
||||||
};
|
};
|
||||||
return this.fetchWithError(i, {
|
return this.fetchWithError(i, {
|
||||||
@@ -56,37 +78,33 @@ class g {
|
|||||||
body: JSON.stringify(a)
|
body: JSON.stringify(a)
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
async update(e, t, s, r, i) {
|
async update(e, t, n, r, i) {
|
||||||
const a = typeof r == "number" || typeof r == "string" ? String(r) : void 0, c = this.buildUrl(e, t, a), o = {
|
let a = typeof r == "number" || typeof r == "string" ? String(r) : void 0, o = this.buildUrl(e, t, a), s = {
|
||||||
operation: "update",
|
operation: "update",
|
||||||
id: Array.isArray(r) ? r : void 0,
|
id: Array.isArray(r) ? r : void 0,
|
||||||
data: s,
|
data: n,
|
||||||
options: i
|
options: i
|
||||||
};
|
};
|
||||||
return this.fetchWithError(c, {
|
return this.fetchWithError(o, {
|
||||||
method: "POST",
|
method: "POST",
|
||||||
headers: this.baseHeaders(),
|
headers: this.baseHeaders(),
|
||||||
body: JSON.stringify(o)
|
body: JSON.stringify(s)
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
async delete(e, t, s) {
|
async delete(e, t, n) {
|
||||||
const r = this.buildUrl(e, t, String(s)), i = {
|
let r = this.buildUrl(e, t, String(n));
|
||||||
operation: "delete"
|
|
||||||
};
|
|
||||||
return this.fetchWithError(r, {
|
return this.fetchWithError(r, {
|
||||||
method: "POST",
|
method: "POST",
|
||||||
headers: this.baseHeaders(),
|
headers: this.baseHeaders(),
|
||||||
body: JSON.stringify(i)
|
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;
|
||||||
}
|
}
|
||||||
const f = /* @__PURE__ */ new Map();
|
var d = class {
|
||||||
function _(n) {
|
|
||||||
const e = n.url;
|
|
||||||
let t = f.get(e);
|
|
||||||
return t || (t = new p(n), f.set(e, t)), t;
|
|
||||||
}
|
|
||||||
class p {
|
|
||||||
constructor(e) {
|
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 = {
|
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,
|
url: e.url,
|
||||||
@@ -106,44 +124,44 @@ class p {
|
|||||||
try {
|
try {
|
||||||
this.ws = new WebSocket(this.config.url), this.ws.onopen = () => {
|
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.log("Connected to WebSocket server"), this.setState("connected"), this.reconnectAttempts = 0, this.startHeartbeat(), this.emit("connect"), e();
|
||||||
}, this.ws.onmessage = (s) => {
|
}, this.ws.onmessage = (e) => {
|
||||||
this.handleMessage(s.data);
|
this.handleMessage(e.data);
|
||||||
}, this.ws.onerror = (s) => {
|
}, this.ws.onerror = (e) => {
|
||||||
this.log("WebSocket error:", s);
|
this.log("WebSocket error:", e);
|
||||||
const r = new Error("WebSocket connection error");
|
let n = /* @__PURE__ */ Error("WebSocket connection error");
|
||||||
this.emit("error", r), t(r);
|
this.emit("error", n), t(n);
|
||||||
}, this.ws.onclose = (s) => {
|
}, this.ws.onclose = (e) => {
|
||||||
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.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((r) => {
|
this.connect().catch((e) => {
|
||||||
this.log("Reconnection failed:", r);
|
this.log("Reconnection failed:", e);
|
||||||
});
|
});
|
||||||
}, this.config.reconnectInterval));
|
}, this.config.reconnectInterval));
|
||||||
};
|
};
|
||||||
} catch (s) {
|
} catch (e) {
|
||||||
t(s);
|
t(e);
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
disconnect() {
|
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();
|
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(e, t, s) {
|
async request(t, n, r) {
|
||||||
this.ensureConnected();
|
this.ensureConnected();
|
||||||
const r = l(), i = {
|
let i = e(), a = {
|
||||||
id: r,
|
id: i,
|
||||||
type: "request",
|
type: "request",
|
||||||
operation: e,
|
operation: t,
|
||||||
entity: t,
|
entity: n,
|
||||||
schema: s?.schema,
|
schema: r?.schema,
|
||||||
record_id: s?.record_id,
|
record_id: r?.record_id,
|
||||||
data: s?.data,
|
data: r?.data,
|
||||||
options: s?.options
|
options: r?.options
|
||||||
};
|
};
|
||||||
return new Promise((a, c) => {
|
return new Promise((e, t) => {
|
||||||
this.messageHandlers.set(r, (o) => {
|
this.messageHandlers.set(i, (n) => {
|
||||||
o.success ? a(o.data) : c(new Error(o.error?.message || "Request failed"));
|
n.success ? e(n.data) : t(Error(n.error?.message || "Request failed"));
|
||||||
}), this.send(i), setTimeout(() => {
|
}), this.send(a), setTimeout(() => {
|
||||||
this.messageHandlers.has(r) && (this.messageHandlers.delete(r), c(new Error("Request timeout")));
|
this.messageHandlers.has(i) && (this.messageHandlers.delete(i), t(/* @__PURE__ */ Error("Request timeout")));
|
||||||
}, 3e4);
|
}, 3e4);
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
@@ -161,73 +179,68 @@ class p {
|
|||||||
}
|
}
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
async create(e, t, s) {
|
async create(e, t, n) {
|
||||||
return this.request("create", e, {
|
return this.request("create", e, {
|
||||||
schema: s?.schema,
|
schema: n?.schema,
|
||||||
data: t
|
data: t
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
async update(e, t, s, r) {
|
async update(e, t, n, r) {
|
||||||
return this.request("update", e, {
|
return this.request("update", e, {
|
||||||
schema: r?.schema,
|
schema: r?.schema,
|
||||||
record_id: t,
|
record_id: t,
|
||||||
data: s
|
data: n
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
async delete(e, t, s) {
|
async delete(e, t, n) {
|
||||||
await this.request("delete", e, {
|
await this.request("delete", e, {
|
||||||
schema: s?.schema,
|
schema: n?.schema,
|
||||||
record_id: t
|
record_id: t
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
async meta(e, t) {
|
async meta(e, t) {
|
||||||
return this.request("meta", e, {
|
return this.request("meta", e, { schema: t?.schema });
|
||||||
schema: t?.schema
|
|
||||||
});
|
|
||||||
}
|
}
|
||||||
async subscribe(e, t, s) {
|
async subscribe(t, n, r) {
|
||||||
this.ensureConnected();
|
this.ensureConnected();
|
||||||
const r = l(), i = {
|
let i = e(), a = {
|
||||||
id: r,
|
id: i,
|
||||||
type: "subscription",
|
type: "subscription",
|
||||||
operation: "subscribe",
|
operation: "subscribe",
|
||||||
entity: e,
|
entity: t,
|
||||||
schema: s?.schema,
|
schema: r?.schema,
|
||||||
options: {
|
options: { filters: r?.filters }
|
||||||
filters: s?.filters
|
|
||||||
}
|
|
||||||
};
|
};
|
||||||
return new Promise((a, c) => {
|
return new Promise((e, o) => {
|
||||||
this.messageHandlers.set(r, (o) => {
|
this.messageHandlers.set(i, (i) => {
|
||||||
if (o.success && o.data?.subscription_id) {
|
if (i.success && i.data?.subscription_id) {
|
||||||
const h = o.data.subscription_id;
|
let a = i.data.subscription_id;
|
||||||
this.subscriptions.set(h, {
|
this.subscriptions.set(a, {
|
||||||
id: h,
|
id: a,
|
||||||
entity: e,
|
entity: t,
|
||||||
schema: s?.schema,
|
schema: r?.schema,
|
||||||
options: { filters: s?.filters },
|
options: { filters: r?.filters },
|
||||||
callback: t
|
callback: n
|
||||||
}), this.log(`Subscribed to ${e} with ID: ${h}`), a(h);
|
}), this.log(`Subscribed to ${t} with ID: ${a}`), e(a);
|
||||||
} else
|
} else o(Error(i.error?.message || "Subscription failed"));
|
||||||
c(new Error(o.error?.message || "Subscription failed"));
|
}), this.send(a), setTimeout(() => {
|
||||||
}), this.send(i), setTimeout(() => {
|
this.messageHandlers.has(i) && (this.messageHandlers.delete(i), o(/* @__PURE__ */ Error("Subscription timeout")));
|
||||||
this.messageHandlers.has(r) && (this.messageHandlers.delete(r), c(new Error("Subscription timeout")));
|
|
||||||
}, 1e4);
|
}, 1e4);
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
async unsubscribe(e) {
|
async unsubscribe(t) {
|
||||||
this.ensureConnected();
|
this.ensureConnected();
|
||||||
const t = l(), s = {
|
let n = e(), r = {
|
||||||
id: t,
|
id: n,
|
||||||
type: "subscription",
|
type: "subscription",
|
||||||
operation: "unsubscribe",
|
operation: "unsubscribe",
|
||||||
subscription_id: e
|
subscription_id: t
|
||||||
};
|
};
|
||||||
return new Promise((r, i) => {
|
return new Promise((e, i) => {
|
||||||
this.messageHandlers.set(t, (a) => {
|
this.messageHandlers.set(n, (n) => {
|
||||||
a.success ? (this.subscriptions.delete(e), this.log(`Unsubscribed from ${e}`), r()) : i(new Error(a.error?.message || "Unsubscribe failed"));
|
n.success ? (this.subscriptions.delete(t), this.log(`Unsubscribed from ${t}`), e()) : i(Error(n.error?.message || "Unsubscribe failed"));
|
||||||
}), this.send(s), setTimeout(() => {
|
}), this.send(r), setTimeout(() => {
|
||||||
this.messageHandlers.has(t) && (this.messageHandlers.delete(t), i(new Error("Unsubscribe timeout")));
|
this.messageHandlers.has(n) && (this.messageHandlers.delete(n), i(/* @__PURE__ */ Error("Unsubscribe timeout")));
|
||||||
}, 1e4);
|
}, 1e4);
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
@@ -246,10 +259,9 @@ class p {
|
|||||||
off(e) {
|
off(e) {
|
||||||
delete this.eventListeners[e];
|
delete this.eventListeners[e];
|
||||||
}
|
}
|
||||||
// Private methods
|
|
||||||
handleMessage(e) {
|
handleMessage(e) {
|
||||||
try {
|
try {
|
||||||
const t = JSON.parse(e);
|
let t = JSON.parse(e);
|
||||||
switch (this.log("Received message:", t), this.emit("message", t), t.type) {
|
switch (this.log("Received message:", t), this.emit("message", t), t.type) {
|
||||||
case "response":
|
case "response":
|
||||||
this.handleResponse(t);
|
this.handleResponse(t);
|
||||||
@@ -257,213 +269,164 @@ class p {
|
|||||||
case "notification":
|
case "notification":
|
||||||
this.handleNotification(t);
|
this.handleNotification(t);
|
||||||
break;
|
break;
|
||||||
case "pong":
|
case "pong": break;
|
||||||
break;
|
default: this.log("Unknown message type:", t.type);
|
||||||
default:
|
|
||||||
this.log("Unknown message type:", t.type);
|
|
||||||
}
|
}
|
||||||
} catch (t) {
|
} catch (e) {
|
||||||
this.log("Error parsing message:", t);
|
this.log("Error parsing message:", e);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
handleResponse(e) {
|
handleResponse(e) {
|
||||||
const t = this.messageHandlers.get(e.id);
|
let t = this.messageHandlers.get(e.id);
|
||||||
t && (t(e), this.messageHandlers.delete(e.id));
|
t && (t(e), this.messageHandlers.delete(e.id));
|
||||||
}
|
}
|
||||||
handleNotification(e) {
|
handleNotification(e) {
|
||||||
const t = this.subscriptions.get(e.subscription_id);
|
let t = this.subscriptions.get(e.subscription_id);
|
||||||
t?.callback && t.callback(e);
|
t?.callback && t.callback(e);
|
||||||
}
|
}
|
||||||
send(e) {
|
send(e) {
|
||||||
if (!this.ws || this.ws.readyState !== WebSocket.OPEN)
|
if (!this.ws || this.ws.readyState !== WebSocket.OPEN) throw Error("WebSocket is not connected");
|
||||||
throw new Error("WebSocket is not connected");
|
let t = JSON.stringify(e);
|
||||||
const t = JSON.stringify(e);
|
|
||||||
this.log("Sending message:", e), this.ws.send(t);
|
this.log("Sending message:", e), this.ws.send(t);
|
||||||
}
|
}
|
||||||
startHeartbeat() {
|
startHeartbeat() {
|
||||||
this.heartbeatTimer || (this.heartbeatTimer = setInterval(() => {
|
this.heartbeatTimer ||= setInterval(() => {
|
||||||
if (this.isConnected()) {
|
if (this.isConnected()) {
|
||||||
const e = {
|
let t = {
|
||||||
id: l(),
|
id: e(),
|
||||||
type: "ping"
|
type: "ping"
|
||||||
};
|
};
|
||||||
this.send(e);
|
this.send(t);
|
||||||
}
|
}
|
||||||
}, this.config.heartbeatInterval));
|
}, this.config.heartbeatInterval);
|
||||||
}
|
}
|
||||||
stopHeartbeat() {
|
stopHeartbeat() {
|
||||||
this.heartbeatTimer && (clearInterval(this.heartbeatTimer), this.heartbeatTimer = null);
|
this.heartbeatTimer &&= (clearInterval(this.heartbeatTimer), null);
|
||||||
}
|
}
|
||||||
setState(e) {
|
setState(e) {
|
||||||
this.state !== e && (this.state = e, this.emit("stateChange", e));
|
this.state !== e && (this.state = e, this.emit("stateChange", e));
|
||||||
}
|
}
|
||||||
ensureConnected() {
|
ensureConnected() {
|
||||||
if (!this.isConnected())
|
if (!this.isConnected()) throw Error("WebSocket is not connected. Call connect() first.");
|
||||||
throw new Error("WebSocket is not connected. Call connect() first.");
|
|
||||||
}
|
}
|
||||||
emit(e, ...t) {
|
emit(e, ...t) {
|
||||||
const s = this.eventListeners[e];
|
let n = this.eventListeners[e];
|
||||||
s && s(...t);
|
n && n(...t);
|
||||||
}
|
}
|
||||||
log(...e) {
|
log(...e) {
|
||||||
this.config.debug && console.log("[WebSocketClient]", ...e);
|
this.config.debug && console.log("[WebSocketClient]", ...e);
|
||||||
}
|
}
|
||||||
|
};
|
||||||
|
//#endregion
|
||||||
|
//#region src/headerspec/client.ts
|
||||||
|
function f(e) {
|
||||||
|
return "ZIP_" + n(e);
|
||||||
}
|
}
|
||||||
function v(n) {
|
function p(e) {
|
||||||
return typeof btoa == "function" ? "ZIP_" + btoa(n) : "ZIP_" + Buffer.from(n, "utf-8").toString("base64");
|
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 w(n) {
|
function m(e) {
|
||||||
let e = n;
|
return t(e);
|
||||||
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) {
|
function h(e) {
|
||||||
return typeof atob == "function" ? atob(n) : Buffer.from(n, "base64").toString("utf-8");
|
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;
|
||||||
}
|
}
|
||||||
function u(n) {
|
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;
|
||||||
const e = {};
|
return e.customOperators?.length && (t["X-Custom-SQL-W"] = e.customOperators.map((e) => e.sql).join(" AND ")), t;
|
||||||
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) {
|
function g(e) {
|
||||||
const t = n.sort.map((s) => s.direction.toUpperCase() === "DESC" ? `-${s.column}` : `+${s.column}`);
|
switch (e) {
|
||||||
e["X-Sort"] = t.join(",");
|
case "eq": return "equals";
|
||||||
}
|
case "neq": return "notequals";
|
||||||
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) {
|
case "gt": return "greaterthan";
|
||||||
const t = n.preload.map((s) => s.columns?.length ? `${s.relation}:${s.columns.join(",")}` : s.relation);
|
case "gte": return "greaterthanorequal";
|
||||||
e["X-Preload"] = t.join("|");
|
case "lt": return "lessthan";
|
||||||
}
|
case "lte": return "lessthanorequal";
|
||||||
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 "like":
|
||||||
case "ilike":
|
case "ilike":
|
||||||
case "contains":
|
case "contains": return "contains";
|
||||||
return "contains";
|
case "startswith": return "beginswith";
|
||||||
case "startswith":
|
case "endswith": return "endswith";
|
||||||
return "beginswith";
|
case "in": return "in";
|
||||||
case "endswith":
|
case "between": return "between";
|
||||||
return "endswith";
|
case "between_inclusive": return "betweeninclusive";
|
||||||
case "in":
|
case "is_null": return "empty";
|
||||||
return "in";
|
case "is_not_null": return "notempty";
|
||||||
case "between":
|
default: return e;
|
||||||
return "between";
|
|
||||||
case "between_inclusive":
|
|
||||||
return "betweeninclusive";
|
|
||||||
case "is_null":
|
|
||||||
return "empty";
|
|
||||||
case "is_not_null":
|
|
||||||
return "notempty";
|
|
||||||
default:
|
|
||||||
return n;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
function S(n) {
|
function _(e) {
|
||||||
return n.value === null || n.value === void 0 ? "" : Array.isArray(n.value) ? n.value.join(",") : String(n.value);
|
return e.value === null || e.value === void 0 ? "" : Array.isArray(e.value) ? e.value.join(",") : String(e.value);
|
||||||
}
|
}
|
||||||
const b = /* @__PURE__ */ new Map();
|
var v = /* @__PURE__ */ new Map();
|
||||||
function C(n) {
|
function y(e) {
|
||||||
const e = n.baseUrl;
|
let t = a(e), n = v.get(t);
|
||||||
let t = b.get(e);
|
return n || (n = new b(e), v.set(t, n)), n;
|
||||||
return t || (t = new H(n), b.set(e, t)), t;
|
|
||||||
}
|
}
|
||||||
class H {
|
var b = class {
|
||||||
constructor(e) {
|
constructor(e) {
|
||||||
this.config = e;
|
this.config = {
|
||||||
|
...e,
|
||||||
|
headers: { ...e.headers }
|
||||||
|
};
|
||||||
}
|
}
|
||||||
buildUrl(e, t, s) {
|
buildUrl(e, t, n) {
|
||||||
let r = `${this.config.baseUrl}/${e}/${t}`;
|
let r = `${this.config.baseUrl}/${e}/${t}`;
|
||||||
return s && (r += `/${s}`), r;
|
return n && (r += `/${n}`), r;
|
||||||
}
|
}
|
||||||
baseHeaders() {
|
baseHeaders() {
|
||||||
const e = {
|
return i(this.config);
|
||||||
"Content-Type": "application/json"
|
|
||||||
};
|
|
||||||
return this.config.token && (e.Authorization = `Bearer ${this.config.token}`), e;
|
|
||||||
}
|
}
|
||||||
async fetchWithError(e, t) {
|
async fetchWithError(e, t) {
|
||||||
const s = await fetch(e, t), r = await s.json();
|
let n = await fetch(e, t), r = await n.json();
|
||||||
if (!s.ok)
|
if (!n.ok) throw Error(r.error?.message || `${n.statusText} (${n.status})`);
|
||||||
throw new Error(
|
|
||||||
r.error?.message || `${s.statusText} (${s.status})`
|
|
||||||
);
|
|
||||||
return {
|
return {
|
||||||
data: r,
|
data: r,
|
||||||
success: !0,
|
success: !0,
|
||||||
error: r.error ? r.error : void 0,
|
error: r.error ? r.error : void 0,
|
||||||
metadata: {
|
metadata: {
|
||||||
count: s.headers.get("content-range") ? Number(s.headers.get("content-range")?.split("/")[1]) : 0,
|
count: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[1]) : 0,
|
||||||
total: s.headers.get("content-range") ? Number(s.headers.get("content-range")?.split("/")[1]) : 0,
|
total: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[1]) : 0,
|
||||||
filtered: s.headers.get("content-range") ? Number(s.headers.get("content-range")?.split("/")[1]) : 0,
|
filtered: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[1]) : 0,
|
||||||
offset: s.headers.get("content-range") ? Number(
|
offset: n.headers.get("content-range") ? Number(n.headers.get("content-range")?.split("/")[0].split("-")[0]) : 0,
|
||||||
s.headers.get("content-range")?.split("/")[0].split("-")[0]
|
limit: n.headers.get("x-limit") ? Number(n.headers.get("x-limit")) : 0
|
||||||
) : 0,
|
|
||||||
limit: s.headers.get("x-limit") ? Number(s.headers.get("x-limit")) : 0
|
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
async read(e, t, s, r) {
|
async read(e, t, n, i) {
|
||||||
const i = this.buildUrl(e, t, s), a = r ? u(r) : {};
|
let a = this.buildUrl(e, t, n), o = i ? h(i) : {};
|
||||||
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, {
|
return this.fetchWithError(a, {
|
||||||
method: "PUT",
|
method: "GET",
|
||||||
headers: { ...this.baseHeaders(), ...c },
|
headers: r(this.baseHeaders(), o)
|
||||||
body: JSON.stringify(r)
|
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
async delete(e, t, s) {
|
async create(e, t, n, i) {
|
||||||
const r = this.buildUrl(e, t, s);
|
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, {
|
return this.fetchWithError(r, {
|
||||||
method: "DELETE",
|
method: "DELETE",
|
||||||
headers: this.baseHeaders()
|
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
|
|
||||||
};
|
};
|
||||||
|
//#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",
|
"name": "@warkypublic/resolvespec-js",
|
||||||
"version": "1.0.1",
|
"version": "1.0.2",
|
||||||
"description": "TypeScript client library for ResolveSpec REST, HeaderSpec, and WebSocket APIs",
|
"description": "TypeScript client library for ResolveSpec REST, HeaderSpec, and WebSocket APIs",
|
||||||
"type": "module",
|
"type": "module",
|
||||||
"main": "./dist/index.cjs",
|
"main": "./dist/index.cjs",
|
||||||
@@ -38,20 +38,22 @@
|
|||||||
"author": "Hein (Warkanum) Puth",
|
"author": "Hein (Warkanum) Puth",
|
||||||
"license": "MIT",
|
"license": "MIT",
|
||||||
"dependencies": {
|
"dependencies": {
|
||||||
"uuid": "^13.0.0"
|
"@warkypublic/artemis-kit": "^1.0.10",
|
||||||
|
"uuid": "^14.0.2"
|
||||||
},
|
},
|
||||||
"devDependencies": {
|
"devDependencies": {
|
||||||
"@changesets/cli": "^2.29.8",
|
"@changesets/cli": "^3.0.3",
|
||||||
"@eslint/js": "^10.0.1",
|
"@eslint/js": "^10.0.1",
|
||||||
"@types/jsdom": "^27.0.0",
|
"@types/jsdom": "^30.0.0",
|
||||||
"eslint": "^10.0.0",
|
"@types/node": "^26.6.2",
|
||||||
"globals": "^17.3.0",
|
"eslint": "^10.11.0",
|
||||||
"jsdom": "^28.1.0",
|
"globals": "^17.12.0",
|
||||||
"typescript": "^5.9.3",
|
"jsdom": "^30.1.1",
|
||||||
"typescript-eslint": "^8.55.0",
|
"typescript": "^6.0.3",
|
||||||
"vite": "^7.3.1",
|
"typescript-eslint": "^8.70.1",
|
||||||
"vite-plugin-dts": "^4.5.4",
|
"vite": "^8.3.0",
|
||||||
"vitest": "^4.0.18"
|
"vite-plugin-dts": "^5.1.1",
|
||||||
|
"vitest": "^5.0.1"
|
||||||
},
|
},
|
||||||
"engines": {
|
"engines": {
|
||||||
"node": ">=18"
|
"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