Compare commits

...
8 Commits
Author SHA1 Message Date
Hein 234aac9770 feat(metrics): bound HTTP path labels, add reset, custom push endpoint and JSON pull
- normalize the HTTP path label (ServeMux pattern, custom normalizer, ID
  collapsing) and cap distinct values via HTTPMaxPaths (default 1024)
- add Reset, PushAndReset, ResetHandler and reset-on-push options
- add POST push to a custom endpoint (text or json) with optional reset
- add JSONHandler for JSON pull
- honour Config.Enabled; log Pushgateway push failures
2026-10-07 12:09:11 +02:00
Hein 8cff3bde85 test(wrap_bunrouter): add tests for route param preservation
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 1m28s
Tests / Unit Tests (push) Successful in 1m30s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m53s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m9s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m10s
Tests / Race Detector (push) Successful in 3m39s
2026-10-05 16:44:44 +02:00
Hein 3e6224698c fix(handler): enforce single record return for ID queries 2026-10-05 16:04:57 +02:00
Hein aec87a81e7 fix(bun): ignore scanonly columns in ExcludeColumn
Tests / Integration Tests (push) Skipped
Tests / Unit Tests (push) Successful in 1m37s
Tests / Race Detector (push) Successful in 3m52s
Build , Vet Test, and Lint / Build (push) Successful in 1m33s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m16s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m19s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m19s
Bun's ExcludeColumn errors with "can't find column" for scanonly fields
because they are not in the table's writable fields. Filter the exclude
list to writable bun fields so models with scanonly buffers can insert
and update again. Add tests for the adapter and reflection.
2026-10-05 14:10:52 +02:00
warkanum 9235292586 fix(crud): skip generated and read-only columns on insert and update
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 1m34s
Tests / Unit Tests (push) Successful in 1m42s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m11s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m25s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m28s
Tests / Race Detector (push) Successful in 4m6s
Read-merge-update wrote every model column back, so GENERATED ALWAYS
columns failed with SQLSTATE 428C9. Add a bun 'generated' tag option,
reflection.NonWritableColumns/RemoveNonWritableColumns, and apply them
in resolvespec, restheadspec, websocketspec, mqttspec, resolvemcp and
the nested CUD processor. Add ExcludeColumn to InsertQuery/UpdateQuery
for model-based writes.
2026-10-02 22:43:45 +02:00
warkanum 23f10387c5 ci(release): fix rust toolchain setup and dart publish validation warnings
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 3m15s
Tests / Unit Tests (push) Successful in 3m26s
Build , Vet Test, and Lint / Lint Code (push) Successful in 4m45s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 4m55s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 4m55s
Tests / Race Detector (push) Successful in 7m6s
2026-10-01 21:20:04 +02:00
warkanum 0d3ad9e4fd ci(tests): disable integration tests job and install psql client 2026-10-01 21:16:34 +02:00
warkanum f5d232d971 ci: move workflows to Gitea and add client release workflow
Tests / Integration Tests (push) Failing after 1m39s
Build , Vet Test, and Lint / Build (push) Successful in 2m3s
Tests / Unit Tests (push) Successful in 2m8s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m51s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m57s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m59s
Tests / Race Detector (push) Successful in 5m9s
Move .github/workflows to .gitea/workflows with Gitea-compatible action
versions, fix make_tag outputs/major bump, and add release_clients.yml to
build and publish all clients to the Gitea package registries. Rename the Go
client module to git.warky.dev/wdevs and add LICENSE/CHANGELOG for Dart.
2026-10-01 21:11:41 +02:00
36 changed files with 1917 additions and 79 deletions
@@ -1,6 +1,3 @@
# This workflow will build a golang project
# For more information see: https://docs.github.com/en/actions/automating-builds-and-tests/building-and-testing-go
name: Create Go Release (Tag Versioning)
on:
@@ -26,7 +23,9 @@ jobs:
steps:
- name: Checkout repository
uses: actions/checkout@v2
uses: actions/checkout@v4
with:
fetch-depth: 0
- name: Set up Git
run: |
@@ -38,7 +37,7 @@ jobs:
run: |
git fetch --tags
latest_tag=$(git describe --tags `git rev-list --tags --max-count=1`)
echo "::set-output name=tag::$latest_tag"
echo "tag=${latest_tag}" >> "${GITHUB_OUTPUT}"
- name: Determine new tag version
id: new_tag
@@ -57,7 +56,7 @@ jobs:
((minor++))
patch=0
;;
"release")
"major")
((major++))
minor=0
patch=0
@@ -68,15 +67,11 @@ jobs:
;;
esac
new_tag="v$major.$minor.$patch"
echo "::set-output name=tag::$new_tag"
echo "tag=${new_tag}" >> "${GITHUB_OUTPUT}"
- name: Create tag
run: |
git tag -a ${{ steps.new_tag.outputs.tag }} -m "Tagging ${{ steps.new_tag.outputs.tag }} for release"
- name: Push changes
uses: ad-m/github-push-action@master
with:
github_token: ${{ secrets.BITECH_GITHUB_TOKEN }}
force: true
tags: true
- name: Push tag
run: git push origin ${{ steps.new_tag.outputs.tag }}
+268
View File
@@ -0,0 +1,268 @@
name: Release Clients
on:
workflow_dispatch:
inputs:
version:
description: "Client version (e.g. 1.4.0)"
required: true
type: string
publish:
description: "Publish packages to Gitea (untick for a build/test dry run)"
required: true
default: true
type: boolean
env:
VERSION_INPUT: ${{ github.event.inputs.version }}
PUBLISH: ${{ github.event.inputs.publish }}
SERVER_URL: ${{ github.server_url }}
OWNER: ${{ github.repository_owner }}
REGISTRY_USER: ${{ secrets.PACKAGE_REGISTRY_USERNAME || vars.PACKAGE_REGISTRY_USERNAME }}
TOKEN: ${{ secrets.PACKAGE_REGISTRY_TOKEN || vars.PACKAGE_REGISTRY_TOKEN }}
jobs:
validate:
name: Validate version
runs-on: ubuntu-latest
outputs:
version: ${{ steps.v.outputs.version }}
steps:
- id: v
run: |
version="${VERSION_INPUT#v}"
if ! [[ "$version" =~ ^[0-9]+\.[0-9]+\.[0-9]+(-[0-9A-Za-z.-]+)?$ ]]; then
echo "Invalid version: $VERSION_INPUT" >&2
exit 1
fi
echo "version=${version}" >> "${GITHUB_OUTPUT}"
js:
name: JS (npm)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-js
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- uses: actions/setup-node@v4
with:
node-version: "22"
- name: Enable pnpm
run: corepack enable
- name: Install
run: pnpm install --frozen-lockfile
- name: Test
run: pnpm test
- name: Set version
run: npm version "$VERSION" --no-git-tag-version --allow-same-version
- name: Build
run: pnpm build
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
run: |
host="${SERVER_URL#*://}"
registry="${SERVER_URL}/api/packages/${OWNER}/npm/"
npm config set "@warkypublic:registry" "$registry"
npm config set "//${host}/api/packages/${OWNER}/npm/:_authToken" "$TOKEN"
npm publish --registry "$registry"
python:
name: Python (PyPI)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-python
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Install
run: pip install -e ".[dev]" build twine
- name: Test
run: pytest
- name: Set version
run: sed -i -E "s/^version = \".*\"/version = \"${VERSION}\"/" pyproject.toml
- name: Build
run: python -m build
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
run: |
twine upload \
--repository-url "${SERVER_URL}/api/packages/${OWNER}/pypi" \
-u "$REGISTRY_USER" -p "$TOKEN" \
dist/*
rust:
name: Rust (Cargo)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-rs
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- name: Set up Rust
uses: dtolnay/rust-toolchain@stable
- name: Test
run: cargo test
- name: Set version
run: sed -i -E '0,/^version = ".*"/s//version = "'"${VERSION}"'"/' Cargo.toml
- name: Package
run: cargo package --allow-dirty
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
env:
CARGO_REGISTRIES_GITEA_INDEX: sparse+${{ github.server_url }}/api/packages/${{ github.repository_owner }}/cargo/
run: |
export CARGO_REGISTRIES_GITEA_TOKEN="Bearer ${TOKEN}"
cargo publish --registry gitea --allow-dirty
dotnet:
name: C# (NuGet)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-cs
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- uses: actions/setup-dotnet@v4
with:
dotnet-version: "8.0.x"
- name: Test
run: dotnet test tests/ResolveSpec.Tests.csproj
- name: Pack
run: dotnet pack src/ResolveSpec.csproj -c Release -p:Version="$VERSION" -o out
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
run: |
dotnet nuget push out/*.nupkg \
--source "${SERVER_URL}/api/packages/${OWNER}/nuget/index.json" \
--api-key "$TOKEN" \
--skip-duplicate
go:
name: Go (Go registry)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-go
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- uses: actions/setup-go@v5
with:
go-version-file: clients/resolvespec-go/go.mod
- name: Test
run: go test ./...
- name: Build module zip
run: |
python3 - <<'PY'
import os, re, zipfile
version = "v" + os.environ["VERSION"]
module = re.search(r"^module\s+(\S+)", open("go.mod").read(), re.M).group(1)
prefix = f"{module}@{version}/"
with zipfile.ZipFile("../resolvespec-go.zip", "w", zipfile.ZIP_DEFLATED) as z:
for root, dirs, files in os.walk("."):
dirs[:] = [d for d in dirs if d != ".git"]
for f in files:
path = os.path.join(root, f)
z.write(path, prefix + os.path.relpath(path, "."))
PY
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
run: |
curl -f -X PUT \
--user "${REGISTRY_USER}:${TOKEN}" \
--upload-file ../resolvespec-go.zip \
"${SERVER_URL}/api/packages/${OWNER}/go/upload"
dart:
name: Dart (Pub)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-dart
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- uses: dart-lang/setup-dart@v1
- name: Install
run: dart pub get
- name: Analyze
run: dart analyze
- name: Test
run: dart test
- name: Set version and registry
run: |
sed -i -E "s/^version: .*/version: ${VERSION}/" pubspec.yaml
sed -i -E "s#^publish_to: .*#publish_to: ${SERVER_URL}/api/packages/${OWNER}/pub#" pubspec.yaml
if ! grep -q "^## ${VERSION}\$" CHANGELOG.md; then
{ head -n 1 CHANGELOG.md; printf '\n## %s\n\n- Release %s.\n' "$VERSION" "$VERSION"; tail -n +2 CHANGELOG.md; } > CHANGELOG.tmp
mv CHANGELOG.tmp CHANGELOG.md
fi
# pub warns about a dirty git tree; commit the stamped files locally (never pushed)
git -c user.name=ci -c user.email=ci@localhost commit -q -am "ci: stamp dart version ${VERSION}"
- name: Dry run
if: ${{ env.PUBLISH != 'true' }}
run: dart pub publish --dry-run
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
run: |
dart pub token add "${SERVER_URL}/api/packages/${OWNER}/pub" --env-var TOKEN
dart pub publish --force
@@ -9,9 +9,9 @@ jobs:
name: Unit Tests
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- name: Set up Go
uses: actions/setup-go@v6
uses: actions/setup-go@v5
with:
go-version: "1.24"
- name: Run unit tests
@@ -22,7 +22,7 @@ jobs:
go test ./pkg/resolvespec ./pkg/restheadspec -coverprofile=coverage.out
go tool cover -html=coverage.out -o coverage.html
- name: Upload coverage
uses: actions/upload-artifact@v5
uses: actions/upload-artifact@v3
continue-on-error: true
with:
name: coverage-report
@@ -31,15 +31,16 @@ jobs:
name: Race Detector
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- name: Set up Go
uses: actions/setup-go@v6
uses: actions/setup-go@v5
with:
go-version: "1.24"
- name: Run unit tests with the race detector
run: go test -race -count=1 ./pkg/...
integration-tests:
name: Integration Tests
if: false # disabled for now
runs-on: ubuntu-latest
services:
postgres:
@@ -56,46 +57,51 @@ jobs:
ports:
- 5432:5432
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- name: Set up Go
uses: actions/setup-go@v6
uses: actions/setup-go@v5
with:
go-version: "1.24"
- name: Install PostgreSQL client
run: |
SUDO=""; [ "$(id -u)" -ne 0 ] && SUDO="sudo"
$SUDO apt-get update -qq
$SUDO apt-get install -y -qq postgresql-client
- name: Create test databases
env:
PGPASSWORD: postgres
run: |
psql -h localhost -U postgres -c "CREATE DATABASE resolvespec_test;"
psql -h localhost -U postgres -c "CREATE DATABASE restheadspec_test;"
psql -h postgres -U postgres -c "CREATE DATABASE resolvespec_test;"
psql -h postgres -U postgres -c "CREATE DATABASE restheadspec_test;"
- name: Run resolvespec integration tests
continue-on-error: true
env:
TEST_DATABASE_URL: "host=localhost user=postgres password=postgres dbname=resolvespec_test port=5432 sslmode=disable"
TEST_DATABASE_URL: "host=postgres user=postgres password=postgres dbname=resolvespec_test port=5432 sslmode=disable"
run: go test -tags=integration ./pkg/resolvespec -v -coverprofile=coverage-resolvespec-integration.out
- name: Run restheadspec integration tests
continue-on-error: true
env:
TEST_DATABASE_URL: "host=localhost user=postgres password=postgres dbname=restheadspec_test port=5432 sslmode=disable"
TEST_DATABASE_URL: "host=postgres user=postgres password=postgres dbname=restheadspec_test port=5432 sslmode=disable"
run: go test -tags=integration ./pkg/restheadspec -v -coverprofile=coverage-restheadspec-integration.out
- name: Generate integration coverage
continue-on-error: true
env:
TEST_DATABASE_URL: "host=localhost user=postgres password=postgres dbname=resolvespec_test port=5432 sslmode=disable"
TEST_DATABASE_URL: "host=postgres user=postgres password=postgres dbname=resolvespec_test port=5432 sslmode=disable"
run: |
go tool cover -html=coverage-resolvespec-integration.out -o coverage-resolvespec-integration.html
go tool cover -html=coverage-restheadspec-integration.out -o coverage-restheadspec-integration.html
- name: Upload resolvespec integration coverage
uses: actions/upload-artifact@v5
uses: actions/upload-artifact@v3
continue-on-error: true
with:
name: resolvespec-integration-coverage-report
path: coverage-resolvespec-integration.html
- name: Upload restheadspec integration coverage
uses: actions/upload-artifact@v5
uses: actions/upload-artifact@v3
continue-on-error: true
with:
name: integration-coverage-restheadspec-report
path: coverage-restheadspec-integration
path: coverage-restheadspec-integration.html
+5
View File
@@ -0,0 +1,5 @@
# Changelog
## 0.1.0
- Initial release: ResolveSpec (JSON body) and FunctionSpec client.
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2026 Hein
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+1
View File
@@ -1,6 +1,7 @@
name: resolvespec
description: Client for ResolveSpec (JSON body) and FunctionSpec endpoints.
version: 0.1.0
repository: https://git.warky.dev/wdevs/ResolveSpec
publish_to: none
environment:
+1 -1
View File
@@ -1,6 +1,6 @@
# resolvespec-go
Go client for ResolveSpec (JSON body) and FunctionSpec. Module: `github.com/bitechdev/ResolveSpec/clients/resolvespec-go`. Stdlib only.
Go client for ResolveSpec (JSON body) and FunctionSpec. Module: `git.warky.dev/wdevs/ResolveSpec/clients/resolvespec-go`. Stdlib only.
## Clients
+1 -1
View File
@@ -1,3 +1,3 @@
module github.com/bitechdev/ResolveSpec/clients/resolvespec-go
module git.warky.dev/wdevs/ResolveSpec/clients/resolvespec-go
go 1.22
+2 -2
View File
@@ -109,8 +109,8 @@ require (
github.com/pkg/errors v0.9.1 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect
github.com/power-devops/perfstat v0.0.0-20210106213030-5aafc221ea8c // indirect
github.com/prometheus/client_model v0.6.2 // indirect
github.com/prometheus/common v0.67.5 // indirect
github.com/prometheus/client_model v0.6.2
github.com/prometheus/common v0.67.5
github.com/prometheus/procfs v0.20.1 // indirect
github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
+37
View File
@@ -10,6 +10,7 @@ import (
"time"
"github.com/uptrace/bun"
"github.com/uptrace/bun/schema"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
@@ -1507,6 +1508,35 @@ func (b *BunInsertQuery) OnConflict(action string) common.InsertQuery {
return b
}
// bunWritableExcludes drops columns bun already leaves out of INSERT/UPDATE
// (scanonly fields) or does not know, since bun's ExcludeColumn errors with
// "can't find column" for anything that is not in the table's writable fields.
func bunWritableExcludes(model bun.Model, columns []string) []string {
tm, ok := model.(interface{ Table() *schema.Table })
if !ok || tm.Table() == nil {
return columns
}
table := tm.Table()
writable := make(map[string]struct{}, len(table.Fields))
for _, f := range table.Fields {
writable[f.Name] = struct{}{}
}
out := make([]string, 0, len(columns))
for _, c := range columns {
if _, ok := writable[c]; ok || c == "*" {
out = append(out, c)
}
}
return out
}
func (b *BunInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery {
if columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 {
b.query = b.query.ExcludeColumn(columns...)
}
return b
}
func (b *BunInsertQuery) Returning(columns ...string) common.InsertQuery {
if len(columns) > 0 {
b.query = b.query.Returning(strings.Join(columns, ", "))
@@ -1619,6 +1649,13 @@ func (b *BunUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuer
return b
}
func (b *BunUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
if columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 {
b.query = b.query.ExcludeColumn(columns...)
}
return b
}
func (b *BunUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery {
b.query = b.query.Where(query, args...)
return b
@@ -0,0 +1,100 @@
package database
import (
"database/sql"
"strings"
"testing"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/pgdialect"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
)
// adhocBuffer mirrors the real-world DBAdhocBuffer: scanonly fields with both
// bun and gorm read-only tags.
type adhocBuffer struct {
CQL1 string `json:"cql1,omitempty" gorm:"->" bun:",scanonly"`
CQL2 string `json:"cql2,omitempty" gorm:"->" bun:",scanonly"`
RowNumber int64 `json:"_rownumber,omitempty" gorm:"-" bun:",scanonly"`
RecordError string `json:"_error,omitempty" gorm:"-" bun:",scanonly"`
}
type excludeModel struct {
bun.BaseModel `bun:"table:public.crmnote,alias:crmnote"`
ID int `json:"id" bun:"id,pk"`
Note string `json:"note" bun:"note,type:citext,"`
Norm string `json:"norm" bun:"norm,generated"`
adhocBuffer `json:",omitempty" bun:",scanonly"`
}
func newExcludeDB() *bun.DB {
return bun.NewDB(&sql.DB{}, pgdialect.New())
}
// TestBunExcludeColumnWithNonWritableColumns feeds the reflection output
// straight into the adapter, as the handlers do, for insert and update.
func TestBunExcludeColumnWithNonWritableColumns(t *testing.T) {
db := newExcludeDB()
m := &excludeModel{}
cols := reflection.NonWritableColumns(m)
if len(cols) == 0 {
t.Fatal("expected non-writable columns")
}
ins := &BunInsertQuery{query: db.NewInsert().Model(m)}
ins.ExcludeColumn(cols...)
insSQL, err := ins.query.AppendQuery(db.QueryGen(), nil)
if err != nil {
t.Fatalf("insert: %v", err)
}
upd := &BunUpdateQuery{query: db.NewUpdate().Model(m).Where("id = 1")}
upd.ExcludeColumn(cols...)
updSQL, err := upd.query.AppendQuery(db.QueryGen(), nil)
if err != nil {
t.Fatalf("update: %v", err)
}
for name, q := range map[string]string{"insert": string(insSQL), "update": string(updSQL)} {
for _, bad := range []string{"cql1", "cql2", "_rownumber", "_error", "norm"} {
if strings.Contains(q, `"`+bad+`"`) {
t.Errorf("%s writes non-writable column %s: %s", name, bad, q)
}
}
if !strings.Contains(q, `"note"`) {
t.Errorf("%s dropped writable column note: %s", name, q)
}
}
}
func TestBunExcludeColumnIgnoresUnknownAndKeepsWritable(t *testing.T) {
db := newExcludeDB()
m := &excludeModel{}
ins := &BunInsertQuery{query: db.NewInsert().Model(m)}
ins.ExcludeColumn("does_not_exist", "note")
q, err := ins.query.AppendQuery(db.QueryGen(), nil)
if err != nil {
t.Fatal(err)
}
if strings.Contains(string(q), `"note"`) {
t.Errorf("writable column note should have been excluded: %s", q)
}
}
func TestBunExcludeColumnOnlyNonWritable(t *testing.T) {
db := newExcludeDB()
ins := &BunInsertQuery{query: db.NewInsert().Model(&excludeModel{})}
ins.ExcludeColumn("cql1") // everything filtered out: must not error or panic
if _, err := ins.query.AppendQuery(db.QueryGen(), nil); err != nil {
t.Fatal(err)
}
}
func TestBunExcludeColumnWithoutModel(t *testing.T) {
db := newExcludeDB()
ins := &BunInsertQuery{query: db.NewInsert()}
ins.ExcludeColumn("cql1") // no model yet: must not panic
}
+14
View File
@@ -751,6 +751,13 @@ func (g *GormInsertQuery) OnConflict(action string) common.InsertQuery {
return g
}
func (g *GormInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery {
if len(columns) > 0 {
g.db = g.db.Omit(columns...)
}
return g
}
func (g *GormInsertQuery) Returning(columns ...string) common.InsertQuery {
g.returningColumns = columns
return g
@@ -930,6 +937,13 @@ func (g *GormUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQue
return g
}
func (g *GormUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
if len(columns) > 0 {
g.db = g.db.Omit(columns...)
}
return g
}
func (g *GormUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery {
g.db = g.db.Where(query, args...)
return g
+14
View File
@@ -691,6 +691,13 @@ func (p *PgSQLInsertQuery) OnConflict(action string) common.InsertQuery {
return p
}
func (p *PgSQLInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery {
for _, col := range columns {
delete(p.values, col)
}
return p
}
func (p *PgSQLInsertQuery) Returning(columns ...string) common.InsertQuery {
p.returning = columns
return p
@@ -850,6 +857,13 @@ func (p *PgSQLUpdateQuery) Set(column string, value interface{}) common.UpdateQu
return p
}
func (p *PgSQLUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
for _, col := range columns {
delete(p.sets, col)
}
return p
}
func (p *PgSQLUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuery {
pkName := ""
if p.model != nil {
+4
View File
@@ -81,6 +81,8 @@ type InsertQuery interface {
Table(table string) InsertQuery
Value(column string, value interface{}) InsertQuery
OnConflict(action string) InsertQuery
// ExcludeColumn omits columns from a Model()-based INSERT (e.g. generated columns).
ExcludeColumn(columns ...string) InsertQuery
Returning(columns ...string) InsertQuery
// Execution
@@ -94,6 +96,8 @@ type UpdateQuery interface {
Table(table string) UpdateQuery
Set(column string, value interface{}) UpdateQuery
SetMap(values map[string]interface{}) UpdateQuery
// ExcludeColumn omits columns from a Model()-based UPDATE (e.g. generated columns).
ExcludeColumn(columns ...string) UpdateQuery
Where(query string, args ...interface{}) UpdateQuery
Returning(columns ...string) UpdateQuery
+6 -2
View File
@@ -116,7 +116,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
case "insert", "create", "add":
// Only perform insert if we have data to insert
if hasData {
id, err := p.processInsert(ctx, regularData, tableName)
id, err := p.processInsert(ctx, regularData, model, tableName)
if err != nil {
logger.Error("Insert failed for table=%s, data=%+v, error=%v", tableName, regularData, err)
return nil, fmt.Errorf("insert failed: %w", err)
@@ -148,7 +148,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
return result, nil
}
if hasData {
rows, err := p.processUpdate(ctx, regularData, tableName, data[pkName])
rows, err := p.processUpdate(ctx, regularData, model, tableName, data[pkName])
if err != nil {
logger.Error("Update failed for table=%s, id=%v, data=%+v, error=%v", tableName, data[pkName], regularData, err)
return nil, fmt.Errorf("update failed: %w", err)
@@ -295,10 +295,12 @@ func (p *NestedCUDProcessor) injectForeignKeys(data map[string]interface{}, mode
func (p *NestedCUDProcessor) processInsert(
ctx context.Context,
data map[string]interface{},
model interface{},
tableName string,
) (interface{}, error) {
logger.Debug("Inserting into %s with data: %+v", tableName, data)
reflection.RemoveNonWritableColumns(model, data)
query := p.db.NewInsert().Table(tableName)
for key, value := range data {
@@ -335,6 +337,7 @@ func (p *NestedCUDProcessor) processSelect(ctx context.Context, tableName string
func (p *NestedCUDProcessor) processUpdate(
ctx context.Context,
data map[string]interface{},
model interface{},
tableName string,
id interface{},
) (int64, error) {
@@ -345,6 +348,7 @@ func (p *NestedCUDProcessor) processUpdate(
logger.Debug("Updating %s with ID %v, data: %+v", tableName, id, data)
reflection.RemoveNonWritableColumns(model, data)
query := p.db.NewUpdate().Table(tableName).SetMap(data).Where(fmt.Sprintf("%s = ?", QuoteIdent(reflection.GetPrimaryKeyName(tableName))), id)
result, err := query.Exec(ctx)
+2
View File
@@ -99,6 +99,7 @@ func (m *mockInsertQuery) Value(column string, value interface{}) InsertQuery {
return m
}
func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m }
func (m *mockInsertQuery) ExcludeColumn(columns ...string) InsertQuery { return m }
func (m *mockInsertQuery) Returning(columns ...string) InsertQuery { return m }
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
m.db.insertCalls = append(m.db.insertCalls, m.values)
@@ -131,6 +132,7 @@ func (m *mockUpdateQuery) SetMap(values map[string]interface{}) UpdateQuery {
return m
}
func (m *mockUpdateQuery) Where(condition string, args ...interface{}) UpdateQuery { return m }
func (m *mockUpdateQuery) ExcludeColumn(columns ...string) UpdateQuery { return m }
func (m *mockUpdateQuery) Returning(columns ...string) UpdateQuery { return m }
func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) {
// Record the update call
+50 -3
View File
@@ -48,11 +48,59 @@ metrics.SetProvider(provider)
| `Namespace` | `string` | `""` | Prefix for all metric names |
| `HTTPRequestBuckets` | `[]float64` | See below | Histogram buckets for HTTP duration (seconds) |
| `DBQueryBuckets` | `[]float64` | See below | Histogram buckets for DB query duration (seconds) |
| `HTTPMaxPaths` | `int` | `1024` | Max distinct `path` label values; extras become `"other"` (negative disables) |
| `HTTPPathNormalizer` | `func(*http.Request) string` | `nil` | Custom request → `path` label mapping (return `""` to use the default) |
**HTTP `path` label:** the middleware uses, in order: `HTTPPathNormalizer`, the matched `http.ServeMux` pattern (`r.Pattern`, e.g. `/users/{id}`), then the raw path with numeric/UUID/hex/opaque-token segments replaced by `:id`. For routers other than `ServeMux`, supply `HTTPPathNormalizer` with your route template. The `HTTPMaxPaths` cap applies on top.
**Default HTTP Request Buckets:** `[0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10]`
**Default DB Query Buckets:** `[0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5]`
### Enabled flag and JSON pull
`Config.Enabled` is honoured: a disabled provider records nothing, `Middleware` passes requests straight through, `Handler()`/`JSONHandler()` answer 404, and push loops are not started (manual pushes return an error). Note a `&metrics.Config{}` literal has `Enabled: false`; use `DefaultConfig()` or set `Enabled: true`. `NewPrometheusProvider(nil)` is enabled.
`provider.JSONHandler()` serves the same JSON as the push `json` format on `GET`/`HEAD`:
```go
http.Handle("/metrics", provider.Handler()) // Prometheus text
http.Handle("/metrics.json", provider.JSONHandler()) // JSON
```
### Resetting Stats
- `provider.Reset()` clears counters, histograms and the cache-size gauge (live gauges such as in-flight requests are kept). Package-level `metrics.Reset()` does the same for the current provider if it implements `metrics.Resetter`.
- `provider.PushAndReset()` pushes to the Pushgateway and resets only if the push succeeded (errors if no Pushgateway is configured).
- `Config.PushgatewayResetOnPush: true` makes the automatic push loop do this on every tick.
- `provider.ResetHandler()` is a `POST`-only endpoint (`?push=true` to push first). It has no auth: mount it on an internal route.
```go
http.Handle("/metrics/reset", provider.ResetHandler())
```
Note: the normal `/metrics` scrape is read-only and never clears anything. Observations recorded between a push and its reset are lost. Prometheus handles the counter drop as a reset, but if you reset often, prefer `increase()`/`rate()` over raw counter values.
### Custom Push Endpoint (Optional)
POST metrics to your own server, optionally clearing local stats after a 2xx reply:
```go
provider := metrics.NewPrometheusProvider(&metrics.Config{
PushEndpointURL: "https://collector.example.com/metrics",
PushEndpointFormat: "json", // or "text" (Prometheus exposition, default)
PushEndpointHeaders: map[string]string{"Authorization": "Bearer token"},
PushEndpointInterval: 30, // seconds; 0 = manual only
PushEndpointTimeout: 10, // seconds (default 10)
PushEndpointResetOnSuccess: true, // clear local stats after a 2xx
})
err := provider.PushToEndpoint(ctx) // manual push; also honours ResetOnSuccess
provider.StopAutoPush() // stops the Pushgateway and endpoint loops
```
The `json` body is a list of `{name, help, type, metrics:[{labels, value | count, sum, buckets}]}`. Failures (non-2xx, network, timeout) are logged and never reset stats, so the next tick retries with the accumulated data. The payload covers everything in the default Prometheus registry, including Go runtime metrics.
### Pushgateway Configuration (Optional)
For batch jobs, cron tasks, or short-lived processes, you can push metrics to Prometheus Pushgateway:
@@ -457,10 +505,9 @@ scrape_configs:
- ✅ Good: `method`, `status_code`
- ❌ Bad: `user_id`, `timestamp`
2. **Path Normalization**: Normalize dynamic paths
2. **Path Normalization**: Done automatically for the `path` label (see Configuration Options)
```go
// Instead of /api/users/123
// Use /api/users/:id
// /api/users/123 is recorded as /api/users/:id
```
3. **Metric Naming**: Follow Prometheus conventions
+53
View File
@@ -1,5 +1,7 @@
package metrics
import "net/http"
// Config holds configuration for the metrics provider
type Config struct {
// Enabled determines whether metrics collection is enabled
@@ -19,6 +21,17 @@ type Config struct {
// Default: [0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5]
DBQueryBuckets []float64 `mapstructure:"db_query_buckets"`
// HTTPMaxPaths caps the number of distinct values of the "path" label on HTTP
// metrics. Paths beyond the cap are reported as "other". Paths are already
// normalized (route pattern, or dynamic segments replaced with ":id").
// Default: 1024. Set to a negative value to disable the cap.
HTTPMaxPaths int `mapstructure:"http_max_paths"`
// HTTPPathNormalizer optionally maps a request to its "path" label (e.g. the
// matched route template of your router). Return "" to fall back to the
// default behaviour (ServeMux pattern, then generic ID normalization).
HTTPPathNormalizer func(*http.Request) string `mapstructure:"-"`
// PushgatewayURL is the URL of the Prometheus Pushgateway (optional)
// If set, metrics will be pushed to this gateway instead of only being scraped
// Example: "http://pushgateway:9091"
@@ -32,6 +45,34 @@ type Config struct {
// Only used if PushgatewayURL is set. If 0, automatic pushing is disabled.
// Default: 0 (no automatic pushing)
PushgatewayInterval int `mapstructure:"pushgateway_interval"`
// PushEndpointURL is a custom HTTP endpoint that metrics are POSTed to
// (independent of Pushgateway). Example: "https://collector.example.com/metrics"
PushEndpointURL string `mapstructure:"push_endpoint_url"`
// PushEndpointFormat is the request body format: "text" (Prometheus text
// exposition, Content-Type text/plain; version=0.0.4) or "json".
// Default: "text"
PushEndpointFormat string `mapstructure:"push_endpoint_format"`
// PushEndpointHeaders are extra headers sent with each POST (e.g. Authorization).
PushEndpointHeaders map[string]string `mapstructure:"push_endpoint_headers"`
// PushEndpointInterval is the interval in seconds for automatic POSTs.
// If 0, automatic posting is disabled (PushToEndpoint can still be called manually).
PushEndpointInterval int `mapstructure:"push_endpoint_interval"`
// PushEndpointTimeout is the per-request timeout in seconds. Default: 10
PushEndpointTimeout int `mapstructure:"push_endpoint_timeout"`
// PushEndpointResetOnSuccess clears local counters and histograms after the
// endpoint answers with a 2xx status. Default: false.
PushEndpointResetOnSuccess bool `mapstructure:"push_endpoint_reset_on_success"`
// PushgatewayResetOnPush clears the local counters and histograms after each
// successful push (automatic or via PushAndReset), so each push carries only
// the activity since the previous one. Default: false.
PushgatewayResetOnPush bool `mapstructure:"pushgateway_reset_on_push"`
}
// DefaultConfig returns a Config with sensible defaults
@@ -43,6 +84,7 @@ func DefaultConfig() *Config {
HTTPRequestBuckets: []float64{0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10},
// DB queries are usually faster
DBQueryBuckets: []float64{0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5},
HTTPMaxPaths: defaultHTTPMaxPaths,
}
}
@@ -57,6 +99,17 @@ func (c *Config) ApplyDefaults() {
if len(c.DBQueryBuckets) == 0 {
c.DBQueryBuckets = []float64{0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5}
}
if c.PushEndpointURL != "" {
if c.PushEndpointFormat == "" {
c.PushEndpointFormat = "text"
}
if c.PushEndpointTimeout <= 0 {
c.PushEndpointTimeout = 10
}
}
if c.HTTPMaxPaths == 0 {
c.HTTPMaxPaths = defaultHTTPMaxPaths
}
// Set default job name if pushgateway is configured but job name is empty
if c.PushgatewayURL != "" && c.PushgatewayJobName == "" {
c.PushgatewayJobName = "resolvespec"
+189
View File
@@ -0,0 +1,189 @@
package metrics
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"sync"
"time"
"github.com/prometheus/client_golang/prometheus"
dto "github.com/prometheus/client_model/go"
"github.com/prometheus/common/expfmt"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
const textContentType = "text/plain; version=0.0.4; charset=utf-8"
// endpointPusher POSTs gathered metrics to a user-configured HTTP endpoint.
type endpointPusher struct {
url string
format string
headers map[string]string
client *http.Client
resetOnOK bool
provider *PrometheusProvider
gatherer prometheus.Gatherer
stopOnce sync.Once
stopCh chan struct{}
startedMu sync.Mutex
started bool
}
func newEndpointPusher(cfg *Config, p *PrometheusProvider) *endpointPusher {
return &endpointPusher{
url: cfg.PushEndpointURL,
format: cfg.PushEndpointFormat,
headers: cfg.PushEndpointHeaders,
client: &http.Client{Timeout: time.Duration(cfg.PushEndpointTimeout) * time.Second},
resetOnOK: cfg.PushEndpointResetOnSuccess,
provider: p,
gatherer: prometheus.DefaultGatherer,
stopCh: make(chan struct{}),
}
}
func (e *endpointPusher) start(interval time.Duration) {
e.startedMu.Lock()
defer e.startedMu.Unlock()
if e.started {
return
}
e.started = true
go func() {
t := time.NewTicker(interval)
defer t.Stop()
for {
select {
case <-t.C:
if err := e.push(context.Background()); err != nil {
logger.Warn("Failed to push metrics to endpoint %s: %v", e.url, err)
}
case <-e.stopCh:
return
}
}
}()
}
func (e *endpointPusher) stop() {
e.stopOnce.Do(func() { close(e.stopCh) })
}
func (e *endpointPusher) push(ctx context.Context) error {
mfs, err := e.gatherer.Gather()
if err != nil && len(mfs) == 0 {
return fmt.Errorf("gather metrics: %w", err)
}
body, contentType, err := encodeMetrics(mfs, e.format)
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, e.url, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", contentType)
for k, v := range e.headers {
req.Header.Set(k, v)
}
resp, err := e.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
snippet, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
if resp.StatusCode < 200 || resp.StatusCode > 299 {
return fmt.Errorf("endpoint returned %s: %s", resp.Status, bytes.TrimSpace(snippet))
}
if e.resetOnOK {
e.provider.Reset()
}
return nil
}
func encodeMetrics(mfs []*dto.MetricFamily, format string) (body []byte, contentType string, err error) {
switch format {
case "json":
b, err := json.Marshal(toJSONFamilies(mfs))
return b, "application/json", err
case "", "text":
var buf bytes.Buffer
enc := expfmt.NewEncoder(&buf, expfmt.NewFormat(expfmt.TypeTextPlain))
for _, mf := range mfs {
if err := enc.Encode(mf); err != nil {
return nil, "", err
}
}
return buf.Bytes(), textContentType, nil
default:
return nil, "", fmt.Errorf("unsupported push endpoint format %q", format)
}
}
type jsonFamily struct {
Name string `json:"name"`
Help string `json:"help,omitempty"`
Type string `json:"type"`
Metrics []jsonMetric `json:"metrics"`
}
type jsonMetric struct {
Labels map[string]string `json:"labels,omitempty"`
Value *float64 `json:"value,omitempty"`
Count *uint64 `json:"count,omitempty"`
Sum *float64 `json:"sum,omitempty"`
Buckets []jsonBucket `json:"buckets,omitempty"`
}
type jsonBucket struct {
UpperBound float64 `json:"le"`
Count uint64 `json:"count"`
}
func toJSONFamilies(mfs []*dto.MetricFamily) []jsonFamily {
out := make([]jsonFamily, 0, len(mfs))
for _, mf := range mfs {
f := jsonFamily{Name: mf.GetName(), Help: mf.GetHelp(), Type: mf.GetType().String()}
for _, m := range mf.GetMetric() {
jm := jsonMetric{}
if len(m.GetLabel()) > 0 {
jm.Labels = make(map[string]string, len(m.GetLabel()))
for _, l := range m.GetLabel() {
jm.Labels[l.GetName()] = l.GetValue()
}
}
switch {
case m.Counter != nil:
v := m.Counter.GetValue()
jm.Value = &v
case m.Gauge != nil:
v := m.Gauge.GetValue()
jm.Value = &v
case m.Untyped != nil:
v := m.Untyped.GetValue()
jm.Value = &v
case m.Histogram != nil:
c, s := m.Histogram.GetSampleCount(), m.Histogram.GetSampleSum()
jm.Count, jm.Sum = &c, &s
for _, b := range m.Histogram.GetBucket() {
jm.Buckets = append(jm.Buckets, jsonBucket{UpperBound: b.GetUpperBound(), Count: b.GetCumulativeCount()})
}
case m.Summary != nil:
c, s := m.Summary.GetSampleCount(), m.Summary.GetSampleSum()
jm.Count, jm.Sum = &c, &s
}
f.Metrics = append(f.Metrics, jm)
}
out = append(out, f)
}
return out
}
+15
View File
@@ -47,6 +47,21 @@ type Provider interface {
Handler() http.Handler
}
// Resetter is optionally implemented by providers that can clear their recorded stats.
type Resetter interface {
Reset()
}
// Reset clears the current provider's stats if it supports resetting.
// It returns false if the provider does not implement Resetter.
func Reset() bool {
if r, ok := GetProvider().(Resetter); ok {
r.Reset()
return true
}
return false
}
// globalProvider is the global metrics provider, protected by globalProviderMu.
var (
globalProviderMu sync.RWMutex
+175
View File
@@ -0,0 +1,175 @@
package metrics
import (
"net/http"
"strings"
"sync"
)
const (
// defaultHTTPMaxPaths is the default cap on distinct values of the "path" label.
defaultHTTPMaxPaths = 1024
// overflowPathLabel is used once the cap on distinct path labels is reached.
overflowPathLabel = "other"
)
// routeLabel returns the low-cardinality path label for a request, preferring
// (in order): the custom normalizer, the matched ServeMux pattern, and finally
// the generic normalization of the raw URL path.
func routeLabel(r *http.Request, custom func(*http.Request) string) string {
if custom != nil {
if p := custom(r); p != "" {
return p
}
}
if r.Pattern != "" {
return stripPatternMethod(r.Pattern)
}
return NormalizePath(r.URL.Path)
}
// stripPatternMethod removes the optional "METHOD " prefix (and host) from a
// Go 1.22+ ServeMux pattern, e.g. "GET /users/{id}" -> "/users/{id}".
func stripPatternMethod(pattern string) string {
if i := strings.IndexByte(pattern, ' '); i >= 0 {
pattern = strings.TrimLeft(pattern[i+1:], " ")
}
if i := strings.IndexByte(pattern, '/'); i > 0 {
pattern = pattern[i:] // drop host part
}
return pattern
}
// NormalizePath replaces dynamic-looking path segments (numeric IDs, UUIDs,
// long hex strings and other long opaque tokens) with ":id" so that
// /users/123 and /users/456 share one label value.
func NormalizePath(path string) string {
if path == "" {
return "/"
}
if !strings.Contains(path, "/") {
return path
}
segs := strings.Split(path, "/")
for i, s := range segs {
if isDynamicSegment(s) {
segs[i] = ":id"
}
}
return strings.Join(segs, "/")
}
func isDynamicSegment(s string) bool {
if s == "" {
return false
}
if allDigits(s) {
return true
}
if isUUID(s) {
return true
}
// Long hex strings (hashes, object IDs)
if len(s) >= 16 && allHex(s) {
return true
}
// Long opaque tokens containing digits (base64/ULID-like)
if len(s) >= 24 && hasDigit(s) && !strings.ContainsAny(s, ".") {
return true
}
return false
}
func allDigits(s string) bool {
for i := 0; i < len(s); i++ {
if s[i] < '0' || s[i] > '9' {
return false
}
}
return true
}
func hasDigit(s string) bool {
for i := 0; i < len(s); i++ {
if s[i] >= '0' && s[i] <= '9' {
return true
}
}
return false
}
func allHex(s string) bool {
for i := 0; i < len(s); i++ {
c := s[i]
if !isHexByte(c) {
return false
}
}
return true
}
func isUUID(s string) bool {
if len(s) != 36 {
return false
}
for i := 0; i < len(s); i++ {
c := s[i]
switch i {
case 8, 13, 18, 23:
if c != '-' {
return false
}
default:
if !isHexByte(c) {
return false
}
}
}
return true
}
// pathLimiter bounds the number of distinct path label values. Once the cap is
// reached, unseen paths are reported as "other".
type pathLimiter struct {
mu sync.RWMutex
max int // <= 0 disables the cap
seen map[string]struct{}
}
func newPathLimiter(limit int) *pathLimiter {
return &pathLimiter{max: limit, seen: make(map[string]struct{})}
}
func (l *pathLimiter) label(path string) string {
if l.max <= 0 {
return path
}
l.mu.RLock()
_, ok := l.seen[path]
l.mu.RUnlock()
if ok {
return path
}
l.mu.Lock()
defer l.mu.Unlock()
if _, ok := l.seen[path]; ok {
return path
}
if len(l.seen) >= l.max {
return overflowPathLabel
}
l.seen[path] = struct{}{}
return path
}
func (l *pathLimiter) reset() {
l.mu.Lock()
l.seen = make(map[string]struct{})
l.mu.Unlock()
}
func isHexByte(c byte) bool {
return c >= '0' && c <= '9' || c >= 'a' && c <= 'f' || c >= 'A' && c <= 'F'
}
+237
View File
@@ -0,0 +1,237 @@
package metrics
import (
"context"
"encoding/json"
"github.com/prometheus/client_golang/prometheus"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
)
func TestNormalizePath(t *testing.T) {
cases := map[string]string{
"": "/",
"/": "/",
"/users": "/users",
"/users/123": "/users/:id",
"/users/123/orders/9": "/users/:id/orders/:id",
"/x/550e8400-e29b-41d4-a716-446655440000": "/x/:id",
"/x/507f1f77bcf86cd799439011": "/x/:id",
"/api/public/users": "/api/public/users",
"/files/report.v2": "/files/report.v2",
}
for in, want := range cases {
if got := NormalizePath(in); got != want {
t.Errorf("NormalizePath(%q) = %q, want %q", in, got, want)
}
}
}
func TestRouteLabel(t *testing.T) {
r := httptest.NewRequest("GET", "/users/42", nil)
if got := routeLabel(r, nil); got != "/users/:id" {
t.Errorf("fallback = %q", got)
}
r.Pattern = "GET /users/{id}"
if got := routeLabel(r, nil); got != "/users/{id}" {
t.Errorf("pattern = %q", got)
}
got := routeLabel(r, func(*http.Request) string { return "/custom" })
if got != "/custom" {
t.Errorf("custom = %q", got)
}
}
func TestPathLimiter(t *testing.T) {
l := newPathLimiter(2)
for _, p := range []string{"/a", "/b", "/a"} {
if got := l.label(p); got != p {
t.Errorf("label(%q) = %q", p, got)
}
}
if got := l.label("/c"); got != overflowPathLabel {
t.Errorf("overflow = %q", got)
}
if got := newPathLimiter(-1).label("/z"); got != "/z" {
t.Errorf("disabled = %q", got)
}
}
func TestMiddlewareUsesPattern(t *testing.T) {
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "pathtest"})
mux := http.NewServeMux()
mux.HandleFunc("GET /users/{id}", func(w http.ResponseWriter, r *http.Request) {})
h := p.Middleware(mux)
for _, id := range []string{"1", "2", "abc"} {
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest("GET", "/users/"+id, nil))
}
if n := len(p.pathLimiter.seen); n != 1 {
t.Errorf("distinct paths = %d, want 1", n)
}
if _, ok := p.pathLimiter.seen["/users/{id}"]; !ok {
t.Errorf("seen = %v", p.pathLimiter.seen)
}
}
func TestResetAndHandler(t *testing.T) {
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "resettest"})
p.RecordHTTPRequest("GET", "/a/1", "200", 0)
p.RecordDBQuery("SELECT", "s", "e", "t", 0, nil)
p.IncRequestsInFlight()
count := func() int {
mfs, _ := prometheus.DefaultGatherer.Gather()
n := 0
for _, mf := range mfs {
if strings.HasPrefix(mf.GetName(), "resettest_") && mf.GetName() != "resettest_http_requests_in_flight" && mf.GetName() != "resettest_event_queue_size" {
n += len(mf.GetMetric())
}
}
return n
}
if count() == 0 {
t.Fatal("expected recorded series")
}
rec := httptest.NewRecorder()
p.ResetHandler().ServeHTTP(rec, httptest.NewRequest("GET", "/reset", nil))
if rec.Code != http.StatusMethodNotAllowed || count() == 0 {
t.Fatalf("GET should be rejected, code=%d", rec.Code)
}
rec = httptest.NewRecorder()
p.ResetHandler().ServeHTTP(rec, httptest.NewRequest("POST", "/reset", nil))
if rec.Code != http.StatusNoContent || count() != 0 {
t.Fatalf("reset failed, code=%d series=%d", rec.Code, count())
}
if len(p.pathLimiter.seen) != 0 {
t.Error("path limiter not reset")
}
// push=true without a pushgateway must fail and not be silent
rec = httptest.NewRecorder()
p.ResetHandler().ServeHTTP(rec, httptest.NewRequest("POST", "/reset?push=true", nil))
if rec.Code != http.StatusBadGateway {
t.Errorf("push without gateway code=%d", rec.Code)
}
}
func TestPushAndResetKeepsStatsOnFailure(t *testing.T) {
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "pushfail", PushgatewayURL: "http://127.0.0.1:1"})
p.RecordHTTPRequest("GET", "/a", "200", 0)
if err := p.PushAndReset(); err == nil {
t.Fatal("expected push error")
}
if len(p.pathLimiter.seen) != 1 {
t.Error("stats were reset despite failed push")
}
}
func TestPushToEndpoint(t *testing.T) {
for _, format := range []string{"text", "json"} {
var gotCT, gotAuth string
var gotBody []byte
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
t.Errorf("method = %s", r.Method)
}
gotCT, gotAuth = r.Header.Get("Content-Type"), r.Header.Get("Authorization")
gotBody, _ = io.ReadAll(r.Body)
}))
ns := "ep" + format
p := NewPrometheusProvider(&Config{
Enabled: true,
Namespace: ns,
PushEndpointURL: srv.URL,
PushEndpointFormat: format,
PushEndpointHeaders: map[string]string{"Authorization": "Bearer x"},
PushEndpointResetOnSuccess: true,
})
p.RecordHTTPRequest("GET", "/a", "200", time.Millisecond)
if err := p.PushToEndpoint(context.Background()); err != nil {
t.Fatalf("%s: %v", format, err)
}
srv.Close()
if gotAuth != "Bearer x" || !strings.Contains(string(gotBody), ns+"_http_requests_total") {
t.Errorf("%s: auth=%q body=%.200s", format, gotAuth, gotBody)
}
if format == "json" && gotCT != "application/json" || format == "text" && !strings.HasPrefix(gotCT, "text/plain") {
t.Errorf("%s: content-type %q", format, gotCT)
}
if len(p.pathLimiter.seen) != 0 {
t.Errorf("%s: stats not reset after success", format)
}
}
}
func TestPushToEndpointFailureKeepsStats(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "nope", http.StatusInternalServerError)
}))
defer srv.Close()
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "epfail", PushEndpointURL: srv.URL, PushEndpointResetOnSuccess: true})
p.RecordHTTPRequest("GET", "/a", "200", 0)
if err := p.PushToEndpoint(context.Background()); err == nil {
t.Fatal("expected error on 500")
}
if len(p.pathLimiter.seen) != 1 {
t.Error("stats reset despite failure")
}
if err := NewPrometheusProvider(&Config{Enabled: true, Namespace: "epnone"}).PushToEndpoint(context.Background()); err == nil {
t.Error("expected error without endpoint")
}
}
func TestDisabledProvider(t *testing.T) {
p := NewPrometheusProvider(&Config{Namespace: "disabled", PushEndpointURL: "http://127.0.0.1:1", PushEndpointInterval: 1})
p.RecordHTTPRequest("GET", "/a", "200", 0)
if len(p.pathLimiter.seen) != 0 {
t.Error("disabled provider recorded")
}
for name, h := range map[string]http.Handler{"handler": p.Handler(), "json": p.JSONHandler()} {
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest("GET", "/m", nil))
if rec.Code != http.StatusNotFound {
t.Errorf("%s code=%d", name, rec.Code)
}
}
if p.endpoint != nil || p.PushToEndpoint(context.Background()) == nil || p.Push() == nil {
t.Error("disabled provider must not push")
}
}
func TestJSONHandler(t *testing.T) {
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "jsonpull"})
p.RecordHTTPRequest("GET", "/a", "200", time.Millisecond)
rec := httptest.NewRecorder()
p.JSONHandler().ServeHTTP(rec, httptest.NewRequest("GET", "/m", nil))
if rec.Code != 200 || rec.Header().Get("Content-Type") != "application/json" {
t.Fatalf("code=%d ct=%q", rec.Code, rec.Header().Get("Content-Type"))
}
var fams []map[string]any
if err := json.Unmarshal(rec.Body.Bytes(), &fams); err != nil {
t.Fatal(err)
}
found := false
for _, f := range fams {
if f["name"] == "jsonpull_http_requests_total" {
found = true
}
}
if !found {
t.Error("metric family missing from JSON")
}
rec = httptest.NewRecorder()
p.JSONHandler().ServeHTTP(rec, httptest.NewRequest("POST", "/m", nil))
if rec.Code != http.StatusMethodNotAllowed {
t.Errorf("POST code=%d", rec.Code)
}
}
+200 -6
View File
@@ -1,6 +1,8 @@
package metrics
import (
"context"
"errors"
"net/http"
"strconv"
"time"
@@ -9,8 +11,12 @@ import (
"github.com/prometheus/client_golang/prometheus/promauto"
"github.com/prometheus/client_golang/prometheus/promhttp"
"github.com/prometheus/client_golang/prometheus/push"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
var errMetricsDisabled = errors.New("metrics: disabled")
// PrometheusProvider implements the Provider interface using Prometheus
type PrometheusProvider struct {
requestDuration *prometheus.HistogramVec
@@ -27,9 +33,16 @@ type PrometheusProvider struct {
eventQueueSize prometheus.Gauge
panicsTotal *prometheus.CounterVec
pathLimiter *pathLimiter
pathNormalizer func(*http.Request) string
enabled bool
endpoint *endpointPusher
// Pushgateway fields (optional)
pushgatewayURL string
pushgatewayJobName string
resetOnPush bool
pusher *push.Pusher
pushTicker *time.Ticker
pushStop chan bool
@@ -55,6 +68,7 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
}
p := &PrometheusProvider{
enabled: cfg.Enabled,
requestDuration: promauto.NewHistogramVec(
prometheus.HistogramOpts{
Name: metricName("http_request_duration_seconds"),
@@ -149,12 +163,17 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
[]string{"method"},
),
pathLimiter: newPathLimiter(cfg.HTTPMaxPaths),
pathNormalizer: cfg.HTTPPathNormalizer,
pushgatewayURL: cfg.PushgatewayURL,
pushgatewayJobName: cfg.PushgatewayJobName,
resetOnPush: cfg.PushgatewayResetOnPush,
}
// Initialize pushgateway if configured
if cfg.PushgatewayURL != "" {
// Pushing is never started for a disabled provider
if cfg.PushgatewayURL != "" && cfg.Enabled {
p.pusher = push.New(cfg.PushgatewayURL, cfg.PushgatewayJobName).
Gatherer(prometheus.DefaultGatherer)
@@ -166,6 +185,13 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
}
}
if cfg.PushEndpointURL != "" && cfg.Enabled {
p.endpoint = newEndpointPusher(cfg, p)
if cfg.PushEndpointInterval > 0 {
p.endpoint.start(time.Duration(cfg.PushEndpointInterval) * time.Second)
}
}
return p
}
@@ -188,23 +214,37 @@ func (rw *ResponseWriter) WriteHeader(code int) {
}
// RecordHTTPRequest implements Provider interface
// The path is normalized and capped to keep label cardinality bounded.
func (p *PrometheusProvider) RecordHTTPRequest(method, path, status string, duration time.Duration) {
if !p.enabled {
return
}
path = p.pathLimiter.label(NormalizePath(path))
p.requestDuration.WithLabelValues(method, path, status).Observe(duration.Seconds())
p.requestTotal.WithLabelValues(method, path, status).Inc()
}
// IncRequestsInFlight implements Provider interface
func (p *PrometheusProvider) IncRequestsInFlight() {
if !p.enabled {
return
}
p.requestsInFlight.Inc()
}
// DecRequestsInFlight implements Provider interface
func (p *PrometheusProvider) DecRequestsInFlight() {
if !p.enabled {
return
}
p.requestsInFlight.Dec()
}
// RecordDBQuery implements Provider interface
func (p *PrometheusProvider) RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error) {
if !p.enabled {
return
}
status := "success"
if err != nil {
status = "error"
@@ -215,47 +255,115 @@ func (p *PrometheusProvider) RecordDBQuery(operation, schema, entity, table stri
// RecordCacheHit implements Provider interface
func (p *PrometheusProvider) RecordCacheHit(provider string) {
if !p.enabled {
return
}
p.cacheHits.WithLabelValues(provider).Inc()
}
// RecordCacheMiss implements Provider interface
func (p *PrometheusProvider) RecordCacheMiss(provider string) {
if !p.enabled {
return
}
p.cacheMisses.WithLabelValues(provider).Inc()
}
// UpdateCacheSize implements Provider interface
func (p *PrometheusProvider) UpdateCacheSize(provider string, size int64) {
if !p.enabled {
return
}
p.cacheSize.WithLabelValues(provider).Set(float64(size))
}
// RecordEventPublished implements Provider interface
func (p *PrometheusProvider) RecordEventPublished(source, eventType string) {
if !p.enabled {
return
}
p.eventPublished.WithLabelValues(source, eventType).Inc()
}
// RecordEventProcessed implements Provider interface
func (p *PrometheusProvider) RecordEventProcessed(source, eventType, status string, duration time.Duration) {
if !p.enabled {
return
}
p.eventProcessed.WithLabelValues(source, eventType, status).Inc()
p.eventDuration.WithLabelValues(source, eventType).Observe(duration.Seconds())
}
// UpdateEventQueueSize implements Provider interface
func (p *PrometheusProvider) UpdateEventQueueSize(size int64) {
if !p.enabled {
return
}
p.eventQueueSize.Set(float64(size))
}
// RecordPanic implements the Provider interface
func (p *PrometheusProvider) RecordPanic(methodName string) {
if !p.enabled {
return
}
p.panicsTotal.WithLabelValues(methodName).Inc()
}
// Handler implements Provider interface
// It responds 404 when metrics are disabled.
func (p *PrometheusProvider) Handler() http.Handler {
if !p.enabled {
return disabledHandler()
}
return promhttp.Handler()
}
// JSONHandler returns an HTTP handler serving the current metrics as JSON
// (same shape as the "json" push endpoint format). Only GET and HEAD are
// accepted, and it responds 404 when metrics are disabled. It performs no
// authentication; mount it on an internal/protected route.
func (p *PrometheusProvider) JSONHandler() http.Handler {
if !p.enabled {
return disabledHandler()
}
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodHead {
w.Header().Set("Allow", "GET, HEAD")
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
mfs, err := prometheus.DefaultGatherer.Gather()
if err != nil && len(mfs) == 0 {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
body, contentType, err := encodeMetrics(mfs, "json")
if err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", contentType)
if r.Method == http.MethodGet {
if _, err := w.Write(body); err != nil {
logger.Warn("Failed to write metrics JSON: %v", err)
}
}
})
}
func disabledHandler() http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "metrics disabled", http.StatusNotFound)
})
}
// Middleware returns an HTTP middleware that collects metrics
// When metrics are disabled it returns next unchanged.
func (p *PrometheusProvider) Middleware(next http.Handler) http.Handler {
if !p.enabled {
return next
}
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
@@ -273,13 +381,17 @@ func (p *PrometheusProvider) Middleware(next http.Handler) http.Handler {
duration := time.Since(start)
status := strconv.Itoa(rw.statusCode)
p.RecordHTTPRequest(r.Method, r.URL.Path, status, duration)
// Read the label after next has run so the router has set r.Pattern.
p.RecordHTTPRequest(r.Method, routeLabel(r, p.pathNormalizer), status, duration)
})
}
// Push manually pushes metrics to the configured Pushgateway
// Returns an error if pushing fails or if Pushgateway is not configured
func (p *PrometheusProvider) Push() error {
if !p.enabled {
return errMetricsDisabled
}
if p.pusher == nil {
return nil // Pushgateway not configured, silently skip
}
@@ -291,10 +403,15 @@ func (p *PrometheusProvider) startAutoPush() {
for {
select {
case <-p.pushTicker.C:
if err := p.Push(); err != nil {
// Log error but continue pushing
// Note: In production, you might want to use a proper logger
_ = err
var err error
if p.resetOnPush {
err = p.PushAndReset()
} else {
err = p.Push()
}
if err != nil {
// Log and keep going; the next tick retries (and nothing was reset)
logger.Warn("Failed to push metrics to Pushgateway: %v", err)
}
case <-p.pushStop:
p.pushTicker.Stop()
@@ -303,10 +420,87 @@ func (p *PrometheusProvider) startAutoPush() {
}
}
// Reset clears all recorded counters, histograms and labelled gauges (cache size)
// and forgets the tracked HTTP path labels. Live gauges (requests in flight,
// event queue size) are left untouched since they reflect current state.
// Prometheus treats the drop in counters as a counter reset, so rate() and
// increase() keep working on the scraper side.
func (p *PrometheusProvider) Reset() {
p.requestDuration.Reset()
p.requestTotal.Reset()
p.dbQueryDuration.Reset()
p.dbQueryTotal.Reset()
p.cacheHits.Reset()
p.cacheMisses.Reset()
p.cacheSize.Reset()
p.eventPublished.Reset()
p.eventProcessed.Reset()
p.eventDuration.Reset()
p.panicsTotal.Reset()
p.pathLimiter.reset()
}
// PushAndReset pushes metrics to the Pushgateway and, only if the push
// succeeded, clears the local stats. Returns an error if Pushgateway is not
// configured, so stats are never discarded without being delivered. Observations
// recorded between the push and the reset are lost.
func (p *PrometheusProvider) PushAndReset() error {
if !p.enabled {
return errMetricsDisabled
}
if p.pusher == nil {
return errors.New("metrics: pushgateway not configured, refusing to reset")
}
if err := p.pusher.Push(); err != nil {
return err
}
p.Reset()
return nil
}
// PushToEndpoint POSTs the current metrics to the configured PushEndpointURL.
// If PushEndpointResetOnSuccess is set, local stats are cleared after a 2xx reply.
// Returns an error if no endpoint is configured.
func (p *PrometheusProvider) PushToEndpoint(ctx context.Context) error {
if !p.enabled {
return errMetricsDisabled
}
if p.endpoint == nil {
return errors.New("metrics: push endpoint not configured")
}
return p.endpoint.push(ctx)
}
// ResetHandler returns an HTTP handler that clears local stats on POST.
// With ?push=true it first pushes to the Pushgateway and only resets on success.
// The handler performs no authentication; mount it on an internal/protected route.
func (p *PrometheusProvider) ResetHandler() http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
w.Header().Set("Allow", http.MethodPost)
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
if r.URL.Query().Get("push") == "true" {
if err := p.PushAndReset(); err != nil {
http.Error(w, err.Error(), http.StatusBadGateway)
return
}
} else {
p.Reset()
}
w.WriteHeader(http.StatusNoContent)
})
}
// StopAutoPush stops the automatic push goroutine
// This should be called when shutting down the application
func (p *PrometheusProvider) StopAutoPush() {
if p.pushStop != nil {
close(p.pushStop)
p.pushStop = nil
}
if p.endpoint != nil {
p.endpoint.stop()
}
}
+5
View File
@@ -895,6 +895,9 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
// Insert record
query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
if generated := reflection.NonWritableColumns(hookCtx.Model); len(generated) > 0 {
query = query.ExcludeColumn(generated...)
}
if _, err := query.Exec(hookCtx.Context); err != nil {
return nil, fmt.Errorf("failed to create record: %w", err)
}
@@ -924,6 +927,8 @@ func (h *Handler) update(hookCtx *HookContext) error {
// the stored value unless disallowNulls is set, in which case null is skipped.
values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls)
reflection.RemoveNonWritableColumns(hookCtx.Model, values)
if len(values) > 0 {
query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID)
+65 -1
View File
@@ -656,7 +656,7 @@ func isColumnWritableInType(typ reflect.Type, columnName string) (found bool, wr
// Check bun tag for scanonly
bunTag := field.Tag.Get("bun")
if bunTag != "" {
if isBunFieldScanOnly(bunTag) {
if isBunFieldScanOnly(bunTag) || isBunFieldGenerated(bunTag) {
return true, false
}
}
@@ -689,6 +689,70 @@ func isBunFieldScanOnly(tag string) bool {
return false
}
// isBunFieldGenerated checks if a bun tag marks the column as database-generated
// (GENERATED ALWAYS AS ... STORED), which can be read but never written.
// Example: "email_normalized,generated" -> true
func isBunFieldGenerated(tag string) bool {
for _, part := range strings.Split(tag, ",") {
if strings.TrimSpace(part) == "generated" {
return true
}
}
return false
}
// RemoveNonWritableColumns deletes from values every key that maps to a
// non-writable model column (bun scanonly/generated, gorm read-only). Used
// before writing a read-merged record back with UPDATE ... SET.
func RemoveNonWritableColumns(model any, values map[string]interface{}) {
for key := range values {
if !IsColumnWritable(model, key) {
delete(values, key)
}
}
}
// NonWritableColumns returns the column names of the model that cannot be
// written (bun scanonly/generated, gorm read-only), including embedded structs.
func NonWritableColumns(model any) []string {
t := reflect.TypeOf(model)
for t != nil && (t.Kind() == reflect.Pointer || t.Kind() == reflect.Slice || t.Kind() == reflect.Array) {
t = t.Elem()
}
if t == nil || t.Kind() != reflect.Struct {
return nil
}
var cols []string
collectNonWritable(t, &cols)
return cols
}
func collectNonWritable(typ reflect.Type, cols *[]string) {
for i := 0; i < typ.NumField(); i++ {
field := typ.Field(i)
if field.Anonymous {
ft := field.Type
if ft.Kind() == reflect.Pointer {
ft = ft.Elem()
}
if ft.Kind() == reflect.Struct {
collectNonWritable(ft, cols)
continue
}
}
bunTag, gormTag := field.Tag.Get("bun"), field.Tag.Get("gorm")
if bunTag == "-" || gormTag == "-" {
continue
}
if (bunTag != "" && (isBunFieldScanOnly(bunTag) || isBunFieldGenerated(bunTag))) ||
(gormTag != "" && isGormFieldReadOnly(gormTag)) {
if name := getColumnNameFromField(field); name != "" {
*cols = append(*cols, name)
}
}
}
}
// isGormFieldReadOnly checks if a gorm tag indicates the field is read-only
// Examples:
// - "<-:false" -> true (no writes allowed)
+105 -34
View File
@@ -497,13 +497,13 @@ func TestIsColumnWritableWithEmbedded(t *testing.T) {
// Test models with relations for GetSQLModelColumns
type User struct {
ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"`
Email string `bun:"email" json:"email"`
ProfileData string `json:"profile_data"` // No bun/gorm tag
Posts []Post `bun:"rel:has-many,join:id=user_id" json:"posts"`
Profile *Profile `bun:"rel:has-one,join:id=user_id" json:"profile"`
RowNumber int64 `bun:",scanonly" json:"_rownumber"`
ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"`
Email string `bun:"email" json:"email"`
ProfileData string `json:"profile_data"` // No bun/gorm tag
Posts []Post `bun:"rel:has-many,join:id=user_id" json:"posts"`
Profile *Profile `bun:"rel:has-one,join:id=user_id" json:"profile"`
RowNumber int64 `bun:",scanonly" json:"_rownumber"`
}
type Post struct {
@@ -528,8 +528,8 @@ type Tag struct {
// Model with scan-only embedded struct
type EntityWithScanOnlyEmbedded struct {
ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"`
ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"`
AdhocBuffer `bun:",scanonly"` // Entire embedded struct is scan-only
}
@@ -1086,17 +1086,17 @@ func TestGetColumnTypeFromModel_SqlNullWrapper(t *testing.T) {
// Models for relation testing
type Author struct {
ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"`
Books []Book `bun:"rel:has-many,join:id=author_id" json:"books"`
ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"`
Books []Book `bun:"rel:has-many,join:id=author_id" json:"books"`
}
type Book struct {
ID int `bun:"id,pk" json:"id"`
Title string `bun:"title" json:"title"`
AuthorID int `bun:"author_id" json:"author_id"`
Author *Author `bun:"rel:belongs-to,join:author_id=id" json:"author"`
Publisher *Publisher `bun:"rel:has-one,join:id=book_id" json:"publisher"`
ID int `bun:"id,pk" json:"id"`
Title string `bun:"title" json:"title"`
AuthorID int `bun:"author_id" json:"author_id"`
Author *Author `bun:"rel:belongs-to,join:author_id=id" json:"author"`
Publisher *Publisher `bun:"rel:has-one,join:id=book_id" json:"publisher"`
}
type Publisher struct {
@@ -1106,9 +1106,9 @@ type Publisher struct {
}
type Student struct {
ID int `gorm:"column:id;primaryKey" json:"id"`
Name string `gorm:"column:name" json:"name"`
Courses []Course `gorm:"many2many:student_courses" json:"courses"`
ID int `gorm:"column:id;primaryKey" json:"id"`
Name string `gorm:"column:name" json:"name"`
Courses []Course `gorm:"many2many:student_courses" json:"courses"`
}
type Course struct {
@@ -1119,11 +1119,11 @@ type Course struct {
// Recursive relation model
type Category struct {
ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"`
ParentID *int `bun:"parent_id" json:"parent_id"`
Parent *Category `bun:"rel:belongs-to,join:parent_id=id" json:"parent"`
Children []Category `bun:"rel:has-many,join:id=parent_id" json:"children"`
ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"`
ParentID *int `bun:"parent_id" json:"parent_id"`
Parent *Category `bun:"rel:belongs-to,join:parent_id=id" json:"parent"`
Children []Category `bun:"rel:has-many,join:id=parent_id" json:"children"`
}
func TestGetRelationType(t *testing.T) {
@@ -1299,7 +1299,7 @@ func TestGetPrimaryKeyValue_EdgeCases(t *testing.T) {
expected: nil,
},
{
name: "model without primary key tags - fallback to ID field",
name: "model without primary key tags - fallback to ID field",
model: struct {
ID int
Name string
@@ -1307,7 +1307,7 @@ func TestGetPrimaryKeyValue_EdgeCases(t *testing.T) {
expected: 99,
},
{
name: "model without ID field",
name: "model without ID field",
model: struct {
Name string
}{Name: "Test"},
@@ -1508,10 +1508,10 @@ func TestGetSQLModelColumns_EdgeCases(t *testing.T) {
// Test models with table:, rel:, join: tags for ExtractColumnFromBunTag
type BunSpecialTagsModel struct {
Table string `bun:"table:users"`
Relation []Post `bun:"rel:has-many"`
Join string `bun:"join:id=user_id"`
NormalCol string `bun:"normal_col"`
Table string `bun:"table:users"`
Relation []Post `bun:"rel:has-many"`
Join string `bun:"join:id=user_id"`
NormalCol string `bun:"normal_col"`
}
func TestExtractColumnFromBunTag_SpecialTags(t *testing.T) {
@@ -1592,8 +1592,8 @@ func TestGetRelationType_GORMFallback(t *testing.T) {
func TestGetRelationType_AdditionalCases(t *testing.T) {
// Test model with GORM has-one (pointer without foreignKey or with references)
type Address struct {
ID int `gorm:"column:id;primaryKey"`
UserID int `gorm:"column:user_id"`
ID int `gorm:"column:id;primaryKey"`
UserID int `gorm:"column:user_id"`
}
type UserWithAddress struct {
@@ -1609,7 +1609,7 @@ func TestGetRelationType_AdditionalCases(t *testing.T) {
type Employee struct {
ID int
Company Company // Single struct (not pointer, not slice) - belongs-to
Company Company // Single struct (not pointer, not slice) - belongs-to
Coworkers []Employee // Slice without bun/gorm tags - has-many
}
@@ -1920,3 +1920,74 @@ func TestMapToStruct_Errors(t *testing.T) {
})
}
}
func TestRemoveNonWritableColumns_Generated(t *testing.T) {
type m struct {
ID int `bun:"id,pk"`
Email string `bun:"email"`
Norm string `bun:"email_normalized,generated"`
Scan string `bun:"scan_col,scanonly"`
}
vals := map[string]interface{}{"id": 1, "email": "A", "email_normalized": "a", "scan_col": "x", "dynamic": 1}
RemoveNonWritableColumns(&m{}, vals)
if _, ok := vals["email_normalized"]; ok {
t.Error("generated column not removed")
}
if _, ok := vals["scan_col"]; ok {
t.Error("scanonly column not removed")
}
if len(vals) != 3 {
t.Errorf("unexpected keys: %v", vals)
}
}
func TestNonWritableColumns(t *testing.T) {
type base struct {
Created string `bun:"created_at,scanonly"`
}
type m struct {
base
ID int `bun:"id,pk"`
Email string `bun:"email"`
Norm string `bun:"email_normalized,generated"`
Ro string `gorm:"column:ro;->"`
}
got := NonWritableColumns(&m{})
want := map[string]bool{"created_at": true, "email_normalized": true, "ro": true}
if len(got) != len(want) {
t.Fatalf("got %v", got)
}
for _, c := range got {
if !want[c] {
t.Errorf("unexpected %s", c)
}
}
}
func TestNonWritableColumns_EmbeddedScanOnlyBuffer(t *testing.T) {
type buffer struct {
CQL1 string `json:"cql1,omitempty" gorm:"->" bun:",scanonly"`
RowNumber int64 `json:"_rownumber,omitempty" gorm:"-" bun:",scanonly"`
}
type m struct {
ID int `json:"id" bun:"id,pk"`
Note string `json:"note" bun:"note,type:citext,"`
buffer `json:",omitempty" bun:",scanonly"`
}
got := NonWritableColumns(&m{})
has := map[string]bool{}
for _, c := range got {
has[c] = true
}
if !has["cql1"] {
t.Errorf("cql1 should be non-writable, got %v", got)
}
if has["id"] || has["note"] {
t.Errorf("writable columns reported as non-writable: %v", got)
}
vals := map[string]interface{}{"id": 1, "note": "x", "cql1": "y"}
RemoveNonWritableColumns(&m{}, vals)
if _, ok := vals["cql1"]; ok || len(vals) != 2 {
t.Errorf("unexpected values: %v", vals)
}
}
+2
View File
@@ -559,6 +559,7 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
if len(cols) == 0 {
return invalidArg("no writable fields in data")
}
reflection.RemoveNonWritableColumns(model, cols)
q := tx.NewInsert().Table(tableName)
for key, value := range cols {
q = q.Value(key, value)
@@ -726,6 +727,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
existingMap[key] = v
}
reflection.RemoveNonWritableColumns(model, setCols)
q := tx.NewUpdate().Table(tableName).SetMap(setCols).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
res, err := q.Exec(ctx)
+1
View File
@@ -194,6 +194,7 @@ func (h *Handler) executeWhere(ctx context.Context, req whereRequest) (_ *whereR
cond := fmt.Sprintf("%s IN (%s)", common.QuoteIdent(pkName), strings.Join(inList, ", "))
var affected int64
if req.op == "update" {
reflection.RemoveNonWritableColumns(model, setCols)
r, err := tx.NewUpdate().Table(tableName).SetMap(setCols).Where(cond, ids...).Exec(ctx)
if err != nil {
return fmt.Errorf("error updating records: %w", err)
+6
View File
@@ -824,6 +824,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
}
responseData = v
reflection.RemoveNonWritableColumns(model, v)
query := tx.NewInsert().Table(tableName)
for key, value := range v {
query = query.Value(key, common.ConvertSliceForBun(value))
@@ -971,6 +972,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
item = modifiedData
}
reflection.RemoveNonWritableColumns(model, item)
txQuery := tx.NewInsert().Table(tableName)
for key, value := range item {
txQuery = txQuery.Value(key, common.ConvertSliceForBun(value))
@@ -1127,6 +1129,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
itemMap = modifiedData
}
reflection.RemoveNonWritableColumns(model, itemMap)
txQuery := tx.NewInsert().Table(tableName)
for key, value := range itemMap {
txQuery = txQuery.Value(key, common.ConvertSliceForBun(value))
@@ -1322,6 +1325,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// Overwrite with every key present in the request (including "" and null unless disallowed)
common.MergeUpdateValues(existingMap, updates, h.disallowNulls)
reflection.RemoveNonWritableColumns(model, existingMap)
// Build update query with merged data
query := tx.NewUpdate().Table(tableName).SetMap(existingMap)
@@ -1507,6 +1511,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// Overwrite with every key present in the request (including "" and null unless disallowed)
common.MergeUpdateValues(existingMap, item, h.disallowNulls)
reflection.RemoveNonWritableColumns(model, existingMap)
txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if _, err := txQuery.Exec(ctx); err != nil {
@@ -1662,6 +1667,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// Overwrite with every key present in the request (including "" and null unless disallowed)
common.MergeUpdateValues(existingMap, itemMap, h.disallowNulls)
reflection.RemoveNonWritableColumns(model, existingMap)
txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if _, err := txQuery.Exec(ctx); err != nil {
+61
View File
@@ -0,0 +1,61 @@
package resolvespec
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/uptrace/bunrouter"
)
type wrapCtxKey struct{}
// The auth wrapper must hand the handler the middleware-enriched request
// without dropping the bunrouter route params.
func TestWrapBunRouterHandler_PreservesRouteParams(t *testing.T) {
var gotSchema, gotEntity, gotID string
var gotCtxVal any
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotSchema = req.Param("schema")
gotEntity = req.Param("entity")
gotID = req.Param("id")
gotCtxVal = req.Context().Value(wrapCtxKey{})
return nil
}
auth := func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), wrapCtxKey{}, "enriched")))
})
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, auth))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/public/users/42", nil))
if gotSchema != "public" || gotEntity != "users" || gotID != "42" {
t.Errorf("route params lost: schema=%q entity=%q id=%q", gotSchema, gotEntity, gotID)
}
if gotCtxVal != "enriched" {
t.Errorf("handler did not see middleware-enriched context, got %v", gotCtxVal)
}
}
func TestWrapBunRouterHandler_NilAuthPassesThrough(t *testing.T) {
var gotID string
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotID = req.Param("id")
return nil
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, nil))
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/public/users/7", nil))
if gotID != "7" {
t.Errorf("id = %q, want 7", gotID)
}
}
+15 -1
View File
@@ -417,6 +417,14 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
if id == "" {
options.SingleRecordAsObject = false
} else {
// The primary key is already filtered, so never return more than one
// record regardless of limit/offset/cursor headers or joins.
one := 1
options.Limit = &one
options.Offset = nil
options.CursorForward = ""
options.CursorBackward = ""
}
// Validate and unwrap model type to get base struct
@@ -726,7 +734,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
sanitizedOr = common.EnsureOuterParentheses(sanitizedOr)
}
if grouper, ok := query.(common.WhereGrouper); ok && sanitizedOr != "" && common.Hardening().SQLStrict {
if grouper, ok := query.(common.WhereGrouper); ok && sanitizedOr != "" {
query = grouper.WhereGroup(func(q common.SelectQuery) common.SelectQuery {
return applyUserConds(q).WhereOr(sanitizedOr)
})
@@ -1410,6 +1418,9 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
if provider, ok := modelValue.(common.TableNameProvider); !ok || provider.TableName() == "" {
query = query.Table(tableName)
}
if generated := reflection.NonWritableColumns(model); len(generated) > 0 {
query = query.ExcludeColumn(generated...)
}
fields := reflection.GetSQLModelColumns(model)
query = query.Returning(fields...)
@@ -1657,6 +1668,9 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
// Create update query using Model() to preserve custom types and driver.Valuer interfaces
query := tx.NewUpdate().Model(modelInstance)
if generated := reflection.NonWritableColumns(model); len(generated) > 0 {
query = query.ExcludeColumn(generated...)
}
query = query.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
// Execute BeforeScan hooks - pass query chain so hooks can modify it
+157
View File
@@ -0,0 +1,157 @@
package restheadspec
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"testing"
"github.com/DATA-DOG/go-sqlmock"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/pgdialect"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
)
// readCapturingSQL runs handleRead and returns every SELECT it issued.
func readCapturingSQL(t *testing.T, id string, options ExtendedRequestOptions) []string {
queries, _ := readCapturingSQLAndBody(t, id, options)
return queries
}
// readCapturingSQLAndBody is readCapturingSQL that also returns the response body.
// The mocked row carries the requested id so the body can be checked against it.
func readCapturingSQLAndBody(t *testing.T, id string, options ExtendedRequestOptions) ([]string, string) {
t.Helper()
resetTotalCache(t)
var queries []string
matcher := sqlmock.QueryMatcherFunc(func(_, actual string) error {
queries = append(queries, actual)
return nil
})
sqlDB, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(matcher))
if err != nil {
t.Fatal(err)
}
sqlDB.SetMaxOpenConns(1)
t.Cleanup(func() { _ = sqlDB.Close() })
h := NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry())
rowID, err := strconv.Atoi(id)
if err != nil {
rowID = 7
}
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(rowID, "a"))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectCommit()
rec := httptest.NewRecorder()
w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil))
h.handleRead(itemCtx(t), w, id, options)
if rec.Code != http.StatusOK {
t.Fatalf("status %d body %s", rec.Code, rec.Body)
}
return queries, rec.Body.String()
}
func TestReadByIDIgnoresLimitOffsetAndCursor(t *testing.T) {
limit, offset := 50, 10
queries := readCapturingSQL(t, "7", ExtendedRequestOptions{
RequestOptions: common.RequestOptions{
Limit: &limit,
Offset: &offset,
},
})
last := queries[len(queries)-1]
if !strings.Contains(last, "LIMIT 1") || strings.Contains(last, "OFFSET") {
t.Fatalf("read by id must be LIMIT 1 with no OFFSET: %s", last)
}
}
func TestReadWithoutIDKeepsRequestedLimit(t *testing.T) {
limit := 50
queries := readCapturingSQL(t, "", ExtendedRequestOptions{
RequestOptions: common.RequestOptions{Limit: &limit},
})
if last := queries[len(queries)-1]; !strings.Contains(last, "LIMIT 50") {
t.Fatalf("list read must keep its limit: %s", last)
}
}
// topLevelOr reports whether the WHERE clause has an OR outside any parentheses,
// i.e. one that would let rows bypass the AND-ed primary key condition.
func topLevelOr(sql string) bool {
where := sql[strings.Index(sql, "WHERE")+len("WHERE"):]
depth, inStr := 0, false
for i := 0; i < len(where); i++ {
switch c := where[i]; {
case c == '\'':
inStr = !inStr
case inStr:
case c == '(':
depth++
case c == ')':
depth--
case depth == 0 && strings.HasPrefix(where[i:], " OR "):
return true
}
}
return false
}
func TestReadByIDCustomSQLOrCannotEscapePrimaryKey(t *testing.T) {
queries := readCapturingSQL(t, "7", ExtendedRequestOptions{
RequestOptions: common.RequestOptions{
Filters: []common.FilterOption{{Column: "name", Operator: "eq", Value: "a"}},
},
CustomSQLOr: "name = 'x'",
})
last := queries[len(queries)-1]
if !strings.Contains(last, `"id" = '7'`) && !strings.Contains(last, `"id" = 7`) {
t.Fatalf("primary key filter missing: %s", last)
}
if topLevelOr(last) {
t.Fatalf("OR escapes the primary key filter: %s", last)
}
}
func TestReadByIDFiltersAndReturnsRequestedRecord(t *testing.T) {
queries, body := readCapturingSQLAndBody(t, "42", ExtendedRequestOptions{})
last := queries[len(queries)-1]
if !strings.Contains(last, `"items"."id" = '42'`) && !strings.Contains(last, `"items"."id" = 42`) {
t.Fatalf("query must filter the primary key to 42: %s", last)
}
if strings.Contains(last, "= 7") || strings.Contains(last, "= '7'") {
t.Fatalf("query filters a different id: %s", last)
}
// every query that touches rows (count and select) must carry the id filter
for _, q := range queries {
if strings.Contains(q, "FROM") && !strings.Contains(q, "42") {
t.Fatalf("query without the id filter: %s", q)
}
}
var rows []struct {
ID int `json:"id"`
}
data := body
if i := strings.Index(body, `"data"`); i >= 0 {
data = body[i+len(`"data"`):]
}
if i := strings.Index(data, "["); i >= 0 {
data = data[i:]
}
dec := json.NewDecoder(strings.NewReader(data))
if err := dec.Decode(&rows); err != nil {
t.Fatalf("decode %q: %v", body, err)
}
if len(rows) != 1 || rows[0].ID != 42 {
t.Fatalf("response must contain exactly the record with id 42: %s", body)
}
}
+61
View File
@@ -0,0 +1,61 @@
package restheadspec
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/uptrace/bunrouter"
)
type wrapCtxKey struct{}
// The auth wrapper must hand the handler the middleware-enriched request
// without dropping the bunrouter route params.
func TestWrapBunRouterHandler_PreservesRouteParams(t *testing.T) {
var gotSchema, gotEntity, gotID string
var gotCtxVal any
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotSchema = req.Param("schema")
gotEntity = req.Param("entity")
gotID = req.Param("id")
gotCtxVal = req.Context().Value(wrapCtxKey{})
return nil
}
auth := func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), wrapCtxKey{}, "enriched")))
})
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, auth))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/public/users/42", nil))
if gotSchema != "public" || gotEntity != "users" || gotID != "42" {
t.Errorf("route params lost: schema=%q entity=%q id=%q", gotSchema, gotEntity, gotID)
}
if gotCtxVal != "enriched" {
t.Errorf("handler did not see middleware-enriched context, got %v", gotCtxVal)
}
}
func TestWrapBunRouterHandler_NilAuthPassesThrough(t *testing.T) {
var gotID string
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotID = req.Param("id")
return nil
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, nil))
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/public/users/7", nil))
if gotID != "7" {
t.Errorf("id = %q, want 7", gotID)
}
}
+5
View File
@@ -758,6 +758,9 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
// Insert record
query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
if generated := reflection.NonWritableColumns(hookCtx.Model); len(generated) > 0 {
query = query.ExcludeColumn(generated...)
}
if _, err := query.Exec(hookCtx.Context); err != nil {
return nil, fmt.Errorf("failed to create record: %w", err)
}
@@ -786,6 +789,8 @@ func (h *Handler) update(hookCtx *HookContext) error {
// the stored value unless disallowNulls is set, in which case null is skipped.
values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls)
reflection.RemoveNonWritableColumns(hookCtx.Model, values)
if len(values) > 0 {
query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID)
+10
View File
@@ -226,6 +226,11 @@ func (m *MockInsertQuery) OnConflict(action string) common.InsertQuery {
return args.Get(0).(common.InsertQuery)
}
func (m *MockInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery {
args := m.Called(columns)
return args.Get(0).(common.InsertQuery)
}
func (m *MockInsertQuery) Returning(columns ...string) common.InsertQuery {
args := m.Called(columns)
return args.Get(0).(common.InsertQuery)
@@ -254,6 +259,11 @@ func (m *MockUpdateQuery) Model(model interface{}) common.UpdateQuery {
return args.Get(0).(common.UpdateQuery)
}
func (m *MockUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
args := m.Called(columns)
return args.Get(0).(common.UpdateQuery)
}
func (m *MockUpdateQuery) Table(table string) common.UpdateQuery {
args := m.Called(table)
return args.Get(0).(common.UpdateQuery)