mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-07 13:56:29 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
234aac9770 | ||
|
|
8cff3bde85 | ||
|
|
3e6224698c | ||
|
|
aec87a81e7 | ||
|
|
9235292586 | ||
|
|
23f10387c5 | ||
|
|
0d3ad9e4fd | ||
|
|
f5d232d971 |
@@ -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 }}
|
||||
@@ -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
|
||||
@@ -0,0 +1,5 @@
|
||||
# Changelog
|
||||
|
||||
## 0.1.0
|
||||
|
||||
- Initial release: ResolveSpec (JSON body) and FunctionSpec client.
|
||||
@@ -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,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,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,3 +1,3 @@
|
||||
module github.com/bitechdev/ResolveSpec/clients/resolvespec-go
|
||||
module git.warky.dev/wdevs/ResolveSpec/clients/resolvespec-go
|
||||
|
||||
go 1.22
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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'
|
||||
}
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user