mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-08 14:26:28 +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)
|
name: Create Go Release (Tag Versioning)
|
||||||
|
|
||||||
on:
|
on:
|
||||||
@@ -26,7 +23,9 @@ jobs:
|
|||||||
|
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout repository
|
- name: Checkout repository
|
||||||
uses: actions/checkout@v2
|
uses: actions/checkout@v4
|
||||||
|
with:
|
||||||
|
fetch-depth: 0
|
||||||
|
|
||||||
- name: Set up Git
|
- name: Set up Git
|
||||||
run: |
|
run: |
|
||||||
@@ -38,7 +37,7 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
git fetch --tags
|
git fetch --tags
|
||||||
latest_tag=$(git describe --tags `git rev-list --tags --max-count=1`)
|
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
|
- name: Determine new tag version
|
||||||
id: new_tag
|
id: new_tag
|
||||||
@@ -57,7 +56,7 @@ jobs:
|
|||||||
((minor++))
|
((minor++))
|
||||||
patch=0
|
patch=0
|
||||||
;;
|
;;
|
||||||
"release")
|
"major")
|
||||||
((major++))
|
((major++))
|
||||||
minor=0
|
minor=0
|
||||||
patch=0
|
patch=0
|
||||||
@@ -68,15 +67,11 @@ jobs:
|
|||||||
;;
|
;;
|
||||||
esac
|
esac
|
||||||
new_tag="v$major.$minor.$patch"
|
new_tag="v$major.$minor.$patch"
|
||||||
echo "::set-output name=tag::$new_tag"
|
echo "tag=${new_tag}" >> "${GITHUB_OUTPUT}"
|
||||||
|
|
||||||
- name: Create tag
|
- name: Create tag
|
||||||
run: |
|
run: |
|
||||||
git tag -a ${{ steps.new_tag.outputs.tag }} -m "Tagging ${{ steps.new_tag.outputs.tag }} for release"
|
git tag -a ${{ steps.new_tag.outputs.tag }} -m "Tagging ${{ steps.new_tag.outputs.tag }} for release"
|
||||||
|
|
||||||
- name: Push changes
|
- name: Push tag
|
||||||
uses: ad-m/github-push-action@master
|
run: git push origin ${{ steps.new_tag.outputs.tag }}
|
||||||
with:
|
|
||||||
github_token: ${{ secrets.BITECH_GITHUB_TOKEN }}
|
|
||||||
force: true
|
|
||||||
tags: true
|
|
||||||
@@ -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
|
name: Unit Tests
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v4
|
||||||
- name: Set up Go
|
- name: Set up Go
|
||||||
uses: actions/setup-go@v6
|
uses: actions/setup-go@v5
|
||||||
with:
|
with:
|
||||||
go-version: "1.24"
|
go-version: "1.24"
|
||||||
- name: Run unit tests
|
- name: Run unit tests
|
||||||
@@ -22,7 +22,7 @@ jobs:
|
|||||||
go test ./pkg/resolvespec ./pkg/restheadspec -coverprofile=coverage.out
|
go test ./pkg/resolvespec ./pkg/restheadspec -coverprofile=coverage.out
|
||||||
go tool cover -html=coverage.out -o coverage.html
|
go tool cover -html=coverage.out -o coverage.html
|
||||||
- name: Upload coverage
|
- name: Upload coverage
|
||||||
uses: actions/upload-artifact@v5
|
uses: actions/upload-artifact@v3
|
||||||
continue-on-error: true
|
continue-on-error: true
|
||||||
with:
|
with:
|
||||||
name: coverage-report
|
name: coverage-report
|
||||||
@@ -31,15 +31,16 @@ jobs:
|
|||||||
name: Race Detector
|
name: Race Detector
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v4
|
||||||
- name: Set up Go
|
- name: Set up Go
|
||||||
uses: actions/setup-go@v6
|
uses: actions/setup-go@v5
|
||||||
with:
|
with:
|
||||||
go-version: "1.24"
|
go-version: "1.24"
|
||||||
- name: Run unit tests with the race detector
|
- name: Run unit tests with the race detector
|
||||||
run: go test -race -count=1 ./pkg/...
|
run: go test -race -count=1 ./pkg/...
|
||||||
integration-tests:
|
integration-tests:
|
||||||
name: Integration Tests
|
name: Integration Tests
|
||||||
|
if: false # disabled for now
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
services:
|
services:
|
||||||
postgres:
|
postgres:
|
||||||
@@ -56,46 +57,51 @@ jobs:
|
|||||||
ports:
|
ports:
|
||||||
- 5432:5432
|
- 5432:5432
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v4
|
||||||
- name: Set up Go
|
- name: Set up Go
|
||||||
uses: actions/setup-go@v6
|
uses: actions/setup-go@v5
|
||||||
with:
|
with:
|
||||||
go-version: "1.24"
|
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
|
- name: Create test databases
|
||||||
env:
|
env:
|
||||||
PGPASSWORD: postgres
|
PGPASSWORD: postgres
|
||||||
run: |
|
run: |
|
||||||
psql -h localhost -U postgres -c "CREATE DATABASE resolvespec_test;"
|
psql -h postgres -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 restheadspec_test;"
|
||||||
- name: Run resolvespec integration tests
|
- name: Run resolvespec integration tests
|
||||||
continue-on-error: true
|
continue-on-error: true
|
||||||
env:
|
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
|
run: go test -tags=integration ./pkg/resolvespec -v -coverprofile=coverage-resolvespec-integration.out
|
||||||
- name: Run restheadspec integration tests
|
- name: Run restheadspec integration tests
|
||||||
continue-on-error: true
|
continue-on-error: true
|
||||||
env:
|
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
|
run: go test -tags=integration ./pkg/restheadspec -v -coverprofile=coverage-restheadspec-integration.out
|
||||||
- name: Generate integration coverage
|
- name: Generate integration coverage
|
||||||
continue-on-error: true
|
continue-on-error: true
|
||||||
env:
|
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: |
|
run: |
|
||||||
go tool cover -html=coverage-resolvespec-integration.out -o coverage-resolvespec-integration.html
|
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
|
go tool cover -html=coverage-restheadspec-integration.out -o coverage-restheadspec-integration.html
|
||||||
|
|
||||||
- name: Upload resolvespec integration coverage
|
- name: Upload resolvespec integration coverage
|
||||||
uses: actions/upload-artifact@v5
|
uses: actions/upload-artifact@v3
|
||||||
continue-on-error: true
|
continue-on-error: true
|
||||||
with:
|
with:
|
||||||
name: resolvespec-integration-coverage-report
|
name: resolvespec-integration-coverage-report
|
||||||
path: coverage-resolvespec-integration.html
|
path: coverage-resolvespec-integration.html
|
||||||
|
|
||||||
- name: Upload restheadspec integration coverage
|
- name: Upload restheadspec integration coverage
|
||||||
uses: actions/upload-artifact@v5
|
uses: actions/upload-artifact@v3
|
||||||
continue-on-error: true
|
continue-on-error: true
|
||||||
|
|
||||||
with:
|
with:
|
||||||
name: integration-coverage-restheadspec-report
|
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
|
name: resolvespec
|
||||||
description: Client for ResolveSpec (JSON body) and FunctionSpec endpoints.
|
description: Client for ResolveSpec (JSON body) and FunctionSpec endpoints.
|
||||||
version: 0.1.0
|
version: 0.1.0
|
||||||
|
repository: https://git.warky.dev/wdevs/ResolveSpec
|
||||||
publish_to: none
|
publish_to: none
|
||||||
|
|
||||||
environment:
|
environment:
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
# resolvespec-go
|
# 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
|
## 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
|
go 1.22
|
||||||
|
|||||||
@@ -109,8 +109,8 @@ require (
|
|||||||
github.com/pkg/errors v0.9.1 // indirect
|
github.com/pkg/errors v0.9.1 // indirect
|
||||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||||
github.com/power-devops/perfstat v0.0.0-20210106213030-5aafc221ea8c // indirect
|
github.com/power-devops/perfstat v0.0.0-20210106213030-5aafc221ea8c // indirect
|
||||||
github.com/prometheus/client_model v0.6.2 // indirect
|
github.com/prometheus/client_model v0.6.2
|
||||||
github.com/prometheus/common v0.67.5 // indirect
|
github.com/prometheus/common v0.67.5
|
||||||
github.com/prometheus/procfs v0.20.1 // indirect
|
github.com/prometheus/procfs v0.20.1 // indirect
|
||||||
github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect
|
github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/uptrace/bun"
|
"github.com/uptrace/bun"
|
||||||
|
"github.com/uptrace/bun/schema"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||||
@@ -1507,6 +1508,35 @@ func (b *BunInsertQuery) OnConflict(action string) common.InsertQuery {
|
|||||||
return b
|
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 {
|
func (b *BunInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||||
if len(columns) > 0 {
|
if len(columns) > 0 {
|
||||||
b.query = b.query.Returning(strings.Join(columns, ", "))
|
b.query = b.query.Returning(strings.Join(columns, ", "))
|
||||||
@@ -1619,6 +1649,13 @@ func (b *BunUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuer
|
|||||||
return b
|
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 {
|
func (b *BunUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery {
|
||||||
b.query = b.query.Where(query, args...)
|
b.query = b.query.Where(query, args...)
|
||||||
return b
|
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
|
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 {
|
func (g *GormInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||||
g.returningColumns = columns
|
g.returningColumns = columns
|
||||||
return g
|
return g
|
||||||
@@ -930,6 +937,13 @@ func (g *GormUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQue
|
|||||||
return g
|
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 {
|
func (g *GormUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery {
|
||||||
g.db = g.db.Where(query, args...)
|
g.db = g.db.Where(query, args...)
|
||||||
return g
|
return g
|
||||||
|
|||||||
@@ -691,6 +691,13 @@ func (p *PgSQLInsertQuery) OnConflict(action string) common.InsertQuery {
|
|||||||
return p
|
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 {
|
func (p *PgSQLInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||||
p.returning = columns
|
p.returning = columns
|
||||||
return p
|
return p
|
||||||
@@ -850,6 +857,13 @@ func (p *PgSQLUpdateQuery) Set(column string, value interface{}) common.UpdateQu
|
|||||||
return p
|
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 {
|
func (p *PgSQLUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuery {
|
||||||
pkName := ""
|
pkName := ""
|
||||||
if p.model != nil {
|
if p.model != nil {
|
||||||
|
|||||||
@@ -81,6 +81,8 @@ type InsertQuery interface {
|
|||||||
Table(table string) InsertQuery
|
Table(table string) InsertQuery
|
||||||
Value(column string, value interface{}) InsertQuery
|
Value(column string, value interface{}) InsertQuery
|
||||||
OnConflict(action string) 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
|
Returning(columns ...string) InsertQuery
|
||||||
|
|
||||||
// Execution
|
// Execution
|
||||||
@@ -94,6 +96,8 @@ type UpdateQuery interface {
|
|||||||
Table(table string) UpdateQuery
|
Table(table string) UpdateQuery
|
||||||
Set(column string, value interface{}) UpdateQuery
|
Set(column string, value interface{}) UpdateQuery
|
||||||
SetMap(values map[string]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
|
Where(query string, args ...interface{}) UpdateQuery
|
||||||
Returning(columns ...string) UpdateQuery
|
Returning(columns ...string) UpdateQuery
|
||||||
|
|
||||||
|
|||||||
@@ -116,7 +116,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
|||||||
case "insert", "create", "add":
|
case "insert", "create", "add":
|
||||||
// Only perform insert if we have data to insert
|
// Only perform insert if we have data to insert
|
||||||
if hasData {
|
if hasData {
|
||||||
id, err := p.processInsert(ctx, regularData, tableName)
|
id, err := p.processInsert(ctx, regularData, model, tableName)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("Insert failed for table=%s, data=%+v, error=%v", tableName, regularData, err)
|
logger.Error("Insert failed for table=%s, data=%+v, error=%v", tableName, regularData, err)
|
||||||
return nil, fmt.Errorf("insert failed: %w", err)
|
return nil, fmt.Errorf("insert failed: %w", err)
|
||||||
@@ -148,7 +148,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
|||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
if hasData {
|
if hasData {
|
||||||
rows, err := p.processUpdate(ctx, regularData, tableName, data[pkName])
|
rows, err := p.processUpdate(ctx, regularData, model, tableName, data[pkName])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("Update failed for table=%s, id=%v, data=%+v, error=%v", tableName, data[pkName], regularData, err)
|
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)
|
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(
|
func (p *NestedCUDProcessor) processInsert(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
data map[string]interface{},
|
data map[string]interface{},
|
||||||
|
model interface{},
|
||||||
tableName string,
|
tableName string,
|
||||||
) (interface{}, error) {
|
) (interface{}, error) {
|
||||||
logger.Debug("Inserting into %s with data: %+v", tableName, data)
|
logger.Debug("Inserting into %s with data: %+v", tableName, data)
|
||||||
|
|
||||||
|
reflection.RemoveNonWritableColumns(model, data)
|
||||||
query := p.db.NewInsert().Table(tableName)
|
query := p.db.NewInsert().Table(tableName)
|
||||||
|
|
||||||
for key, value := range data {
|
for key, value := range data {
|
||||||
@@ -335,6 +337,7 @@ func (p *NestedCUDProcessor) processSelect(ctx context.Context, tableName string
|
|||||||
func (p *NestedCUDProcessor) processUpdate(
|
func (p *NestedCUDProcessor) processUpdate(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
data map[string]interface{},
|
data map[string]interface{},
|
||||||
|
model interface{},
|
||||||
tableName string,
|
tableName string,
|
||||||
id interface{},
|
id interface{},
|
||||||
) (int64, error) {
|
) (int64, error) {
|
||||||
@@ -345,6 +348,7 @@ func (p *NestedCUDProcessor) processUpdate(
|
|||||||
|
|
||||||
logger.Debug("Updating %s with ID %v, data: %+v", tableName, id, data)
|
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)
|
query := p.db.NewUpdate().Table(tableName).SetMap(data).Where(fmt.Sprintf("%s = ?", QuoteIdent(reflection.GetPrimaryKeyName(tableName))), id)
|
||||||
|
|
||||||
result, err := query.Exec(ctx)
|
result, err := query.Exec(ctx)
|
||||||
|
|||||||
@@ -99,6 +99,7 @@ func (m *mockInsertQuery) Value(column string, value interface{}) InsertQuery {
|
|||||||
return m
|
return m
|
||||||
}
|
}
|
||||||
func (m *mockInsertQuery) OnConflict(action string) 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) Returning(columns ...string) InsertQuery { return m }
|
||||||
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
|
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
|
||||||
m.db.insertCalls = append(m.db.insertCalls, m.values)
|
m.db.insertCalls = append(m.db.insertCalls, m.values)
|
||||||
@@ -131,6 +132,7 @@ func (m *mockUpdateQuery) SetMap(values map[string]interface{}) UpdateQuery {
|
|||||||
return m
|
return m
|
||||||
}
|
}
|
||||||
func (m *mockUpdateQuery) Where(condition string, args ...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) Returning(columns ...string) UpdateQuery { return m }
|
||||||
func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) {
|
func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) {
|
||||||
// Record the update call
|
// Record the update call
|
||||||
|
|||||||
+50
-3
@@ -48,11 +48,59 @@ metrics.SetProvider(provider)
|
|||||||
| `Namespace` | `string` | `""` | Prefix for all metric names |
|
| `Namespace` | `string` | `""` | Prefix for all metric names |
|
||||||
| `HTTPRequestBuckets` | `[]float64` | See below | Histogram buckets for HTTP duration (seconds) |
|
| `HTTPRequestBuckets` | `[]float64` | See below | Histogram buckets for HTTP duration (seconds) |
|
||||||
| `DBQueryBuckets` | `[]float64` | See below | Histogram buckets for DB query 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 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]`
|
**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)
|
### Pushgateway Configuration (Optional)
|
||||||
|
|
||||||
For batch jobs, cron tasks, or short-lived processes, you can push metrics to Prometheus Pushgateway:
|
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`
|
- ✅ Good: `method`, `status_code`
|
||||||
- ❌ Bad: `user_id`, `timestamp`
|
- ❌ Bad: `user_id`, `timestamp`
|
||||||
|
|
||||||
2. **Path Normalization**: Normalize dynamic paths
|
2. **Path Normalization**: Done automatically for the `path` label (see Configuration Options)
|
||||||
```go
|
```go
|
||||||
// Instead of /api/users/123
|
// /api/users/123 is recorded as /api/users/:id
|
||||||
// Use /api/users/:id
|
|
||||||
```
|
```
|
||||||
|
|
||||||
3. **Metric Naming**: Follow Prometheus conventions
|
3. **Metric Naming**: Follow Prometheus conventions
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
package metrics
|
package metrics
|
||||||
|
|
||||||
|
import "net/http"
|
||||||
|
|
||||||
// Config holds configuration for the metrics provider
|
// Config holds configuration for the metrics provider
|
||||||
type Config struct {
|
type Config struct {
|
||||||
// Enabled determines whether metrics collection is enabled
|
// 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]
|
// 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"`
|
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)
|
// PushgatewayURL is the URL of the Prometheus Pushgateway (optional)
|
||||||
// If set, metrics will be pushed to this gateway instead of only being scraped
|
// If set, metrics will be pushed to this gateway instead of only being scraped
|
||||||
// Example: "http://pushgateway:9091"
|
// Example: "http://pushgateway:9091"
|
||||||
@@ -32,6 +45,34 @@ type Config struct {
|
|||||||
// Only used if PushgatewayURL is set. If 0, automatic pushing is disabled.
|
// Only used if PushgatewayURL is set. If 0, automatic pushing is disabled.
|
||||||
// Default: 0 (no automatic pushing)
|
// Default: 0 (no automatic pushing)
|
||||||
PushgatewayInterval int `mapstructure:"pushgateway_interval"`
|
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
|
// 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},
|
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
|
// 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},
|
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 {
|
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}
|
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
|
// Set default job name if pushgateway is configured but job name is empty
|
||||||
if c.PushgatewayURL != "" && c.PushgatewayJobName == "" {
|
if c.PushgatewayURL != "" && c.PushgatewayJobName == "" {
|
||||||
c.PushgatewayJobName = "resolvespec"
|
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
|
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.
|
// globalProvider is the global metrics provider, protected by globalProviderMu.
|
||||||
var (
|
var (
|
||||||
globalProviderMu sync.RWMutex
|
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
|
package metrics
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
@@ -9,8 +11,12 @@ import (
|
|||||||
"github.com/prometheus/client_golang/prometheus/promauto"
|
"github.com/prometheus/client_golang/prometheus/promauto"
|
||||||
"github.com/prometheus/client_golang/prometheus/promhttp"
|
"github.com/prometheus/client_golang/prometheus/promhttp"
|
||||||
"github.com/prometheus/client_golang/prometheus/push"
|
"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
|
// PrometheusProvider implements the Provider interface using Prometheus
|
||||||
type PrometheusProvider struct {
|
type PrometheusProvider struct {
|
||||||
requestDuration *prometheus.HistogramVec
|
requestDuration *prometheus.HistogramVec
|
||||||
@@ -27,9 +33,16 @@ type PrometheusProvider struct {
|
|||||||
eventQueueSize prometheus.Gauge
|
eventQueueSize prometheus.Gauge
|
||||||
panicsTotal *prometheus.CounterVec
|
panicsTotal *prometheus.CounterVec
|
||||||
|
|
||||||
|
pathLimiter *pathLimiter
|
||||||
|
pathNormalizer func(*http.Request) string
|
||||||
|
|
||||||
|
enabled bool
|
||||||
|
endpoint *endpointPusher
|
||||||
|
|
||||||
// Pushgateway fields (optional)
|
// Pushgateway fields (optional)
|
||||||
pushgatewayURL string
|
pushgatewayURL string
|
||||||
pushgatewayJobName string
|
pushgatewayJobName string
|
||||||
|
resetOnPush bool
|
||||||
pusher *push.Pusher
|
pusher *push.Pusher
|
||||||
pushTicker *time.Ticker
|
pushTicker *time.Ticker
|
||||||
pushStop chan bool
|
pushStop chan bool
|
||||||
@@ -55,6 +68,7 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
|
|||||||
}
|
}
|
||||||
|
|
||||||
p := &PrometheusProvider{
|
p := &PrometheusProvider{
|
||||||
|
enabled: cfg.Enabled,
|
||||||
requestDuration: promauto.NewHistogramVec(
|
requestDuration: promauto.NewHistogramVec(
|
||||||
prometheus.HistogramOpts{
|
prometheus.HistogramOpts{
|
||||||
Name: metricName("http_request_duration_seconds"),
|
Name: metricName("http_request_duration_seconds"),
|
||||||
@@ -149,12 +163,17 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
|
|||||||
[]string{"method"},
|
[]string{"method"},
|
||||||
),
|
),
|
||||||
|
|
||||||
|
pathLimiter: newPathLimiter(cfg.HTTPMaxPaths),
|
||||||
|
pathNormalizer: cfg.HTTPPathNormalizer,
|
||||||
|
|
||||||
pushgatewayURL: cfg.PushgatewayURL,
|
pushgatewayURL: cfg.PushgatewayURL,
|
||||||
pushgatewayJobName: cfg.PushgatewayJobName,
|
pushgatewayJobName: cfg.PushgatewayJobName,
|
||||||
|
resetOnPush: cfg.PushgatewayResetOnPush,
|
||||||
}
|
}
|
||||||
|
|
||||||
// Initialize pushgateway if configured
|
// 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).
|
p.pusher = push.New(cfg.PushgatewayURL, cfg.PushgatewayJobName).
|
||||||
Gatherer(prometheus.DefaultGatherer)
|
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
|
return p
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -188,23 +214,37 @@ func (rw *ResponseWriter) WriteHeader(code int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// RecordHTTPRequest implements Provider interface
|
// 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) {
|
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.requestDuration.WithLabelValues(method, path, status).Observe(duration.Seconds())
|
||||||
p.requestTotal.WithLabelValues(method, path, status).Inc()
|
p.requestTotal.WithLabelValues(method, path, status).Inc()
|
||||||
}
|
}
|
||||||
|
|
||||||
// IncRequestsInFlight implements Provider interface
|
// IncRequestsInFlight implements Provider interface
|
||||||
func (p *PrometheusProvider) IncRequestsInFlight() {
|
func (p *PrometheusProvider) IncRequestsInFlight() {
|
||||||
|
if !p.enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
p.requestsInFlight.Inc()
|
p.requestsInFlight.Inc()
|
||||||
}
|
}
|
||||||
|
|
||||||
// DecRequestsInFlight implements Provider interface
|
// DecRequestsInFlight implements Provider interface
|
||||||
func (p *PrometheusProvider) DecRequestsInFlight() {
|
func (p *PrometheusProvider) DecRequestsInFlight() {
|
||||||
|
if !p.enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
p.requestsInFlight.Dec()
|
p.requestsInFlight.Dec()
|
||||||
}
|
}
|
||||||
|
|
||||||
// RecordDBQuery implements Provider interface
|
// RecordDBQuery implements Provider interface
|
||||||
func (p *PrometheusProvider) RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error) {
|
func (p *PrometheusProvider) RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error) {
|
||||||
|
if !p.enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
status := "success"
|
status := "success"
|
||||||
if err != nil {
|
if err != nil {
|
||||||
status = "error"
|
status = "error"
|
||||||
@@ -215,47 +255,115 @@ func (p *PrometheusProvider) RecordDBQuery(operation, schema, entity, table stri
|
|||||||
|
|
||||||
// RecordCacheHit implements Provider interface
|
// RecordCacheHit implements Provider interface
|
||||||
func (p *PrometheusProvider) RecordCacheHit(provider string) {
|
func (p *PrometheusProvider) RecordCacheHit(provider string) {
|
||||||
|
if !p.enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
p.cacheHits.WithLabelValues(provider).Inc()
|
p.cacheHits.WithLabelValues(provider).Inc()
|
||||||
}
|
}
|
||||||
|
|
||||||
// RecordCacheMiss implements Provider interface
|
// RecordCacheMiss implements Provider interface
|
||||||
func (p *PrometheusProvider) RecordCacheMiss(provider string) {
|
func (p *PrometheusProvider) RecordCacheMiss(provider string) {
|
||||||
|
if !p.enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
p.cacheMisses.WithLabelValues(provider).Inc()
|
p.cacheMisses.WithLabelValues(provider).Inc()
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateCacheSize implements Provider interface
|
// UpdateCacheSize implements Provider interface
|
||||||
func (p *PrometheusProvider) UpdateCacheSize(provider string, size int64) {
|
func (p *PrometheusProvider) UpdateCacheSize(provider string, size int64) {
|
||||||
|
if !p.enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
p.cacheSize.WithLabelValues(provider).Set(float64(size))
|
p.cacheSize.WithLabelValues(provider).Set(float64(size))
|
||||||
}
|
}
|
||||||
|
|
||||||
// RecordEventPublished implements Provider interface
|
// RecordEventPublished implements Provider interface
|
||||||
func (p *PrometheusProvider) RecordEventPublished(source, eventType string) {
|
func (p *PrometheusProvider) RecordEventPublished(source, eventType string) {
|
||||||
|
if !p.enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
p.eventPublished.WithLabelValues(source, eventType).Inc()
|
p.eventPublished.WithLabelValues(source, eventType).Inc()
|
||||||
}
|
}
|
||||||
|
|
||||||
// RecordEventProcessed implements Provider interface
|
// RecordEventProcessed implements Provider interface
|
||||||
func (p *PrometheusProvider) RecordEventProcessed(source, eventType, status string, duration time.Duration) {
|
func (p *PrometheusProvider) RecordEventProcessed(source, eventType, status string, duration time.Duration) {
|
||||||
|
if !p.enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
p.eventProcessed.WithLabelValues(source, eventType, status).Inc()
|
p.eventProcessed.WithLabelValues(source, eventType, status).Inc()
|
||||||
p.eventDuration.WithLabelValues(source, eventType).Observe(duration.Seconds())
|
p.eventDuration.WithLabelValues(source, eventType).Observe(duration.Seconds())
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateEventQueueSize implements Provider interface
|
// UpdateEventQueueSize implements Provider interface
|
||||||
func (p *PrometheusProvider) UpdateEventQueueSize(size int64) {
|
func (p *PrometheusProvider) UpdateEventQueueSize(size int64) {
|
||||||
|
if !p.enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
p.eventQueueSize.Set(float64(size))
|
p.eventQueueSize.Set(float64(size))
|
||||||
}
|
}
|
||||||
|
|
||||||
// RecordPanic implements the Provider interface
|
// RecordPanic implements the Provider interface
|
||||||
func (p *PrometheusProvider) RecordPanic(methodName string) {
|
func (p *PrometheusProvider) RecordPanic(methodName string) {
|
||||||
|
if !p.enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
p.panicsTotal.WithLabelValues(methodName).Inc()
|
p.panicsTotal.WithLabelValues(methodName).Inc()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Handler implements Provider interface
|
// Handler implements Provider interface
|
||||||
|
// It responds 404 when metrics are disabled.
|
||||||
func (p *PrometheusProvider) Handler() http.Handler {
|
func (p *PrometheusProvider) Handler() http.Handler {
|
||||||
|
if !p.enabled {
|
||||||
|
return disabledHandler()
|
||||||
|
}
|
||||||
return promhttp.Handler()
|
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
|
// 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 {
|
func (p *PrometheusProvider) Middleware(next http.Handler) http.Handler {
|
||||||
|
if !p.enabled {
|
||||||
|
return next
|
||||||
|
}
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
|
|
||||||
@@ -273,13 +381,17 @@ func (p *PrometheusProvider) Middleware(next http.Handler) http.Handler {
|
|||||||
duration := time.Since(start)
|
duration := time.Since(start)
|
||||||
status := strconv.Itoa(rw.statusCode)
|
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
|
// Push manually pushes metrics to the configured Pushgateway
|
||||||
// Returns an error if pushing fails or if Pushgateway is not configured
|
// Returns an error if pushing fails or if Pushgateway is not configured
|
||||||
func (p *PrometheusProvider) Push() error {
|
func (p *PrometheusProvider) Push() error {
|
||||||
|
if !p.enabled {
|
||||||
|
return errMetricsDisabled
|
||||||
|
}
|
||||||
if p.pusher == nil {
|
if p.pusher == nil {
|
||||||
return nil // Pushgateway not configured, silently skip
|
return nil // Pushgateway not configured, silently skip
|
||||||
}
|
}
|
||||||
@@ -291,10 +403,15 @@ func (p *PrometheusProvider) startAutoPush() {
|
|||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-p.pushTicker.C:
|
case <-p.pushTicker.C:
|
||||||
if err := p.Push(); err != nil {
|
var err error
|
||||||
// Log error but continue pushing
|
if p.resetOnPush {
|
||||||
// Note: In production, you might want to use a proper logger
|
err = p.PushAndReset()
|
||||||
_ = err
|
} 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:
|
case <-p.pushStop:
|
||||||
p.pushTicker.Stop()
|
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
|
// StopAutoPush stops the automatic push goroutine
|
||||||
// This should be called when shutting down the application
|
// This should be called when shutting down the application
|
||||||
func (p *PrometheusProvider) StopAutoPush() {
|
func (p *PrometheusProvider) StopAutoPush() {
|
||||||
if p.pushStop != nil {
|
if p.pushStop != nil {
|
||||||
close(p.pushStop)
|
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
|
// Insert record
|
||||||
query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
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 {
|
if _, err := query.Exec(hookCtx.Context); err != nil {
|
||||||
return nil, fmt.Errorf("failed to create record: %w", err)
|
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.
|
// 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)
|
values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls)
|
||||||
|
|
||||||
|
reflection.RemoveNonWritableColumns(hookCtx.Model, values)
|
||||||
|
|
||||||
if len(values) > 0 {
|
if len(values) > 0 {
|
||||||
query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values).
|
query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values).
|
||||||
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID)
|
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
|
// Check bun tag for scanonly
|
||||||
bunTag := field.Tag.Get("bun")
|
bunTag := field.Tag.Get("bun")
|
||||||
if bunTag != "" {
|
if bunTag != "" {
|
||||||
if isBunFieldScanOnly(bunTag) {
|
if isBunFieldScanOnly(bunTag) || isBunFieldGenerated(bunTag) {
|
||||||
return true, false
|
return true, false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -689,6 +689,70 @@ func isBunFieldScanOnly(tag string) bool {
|
|||||||
return false
|
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
|
// isGormFieldReadOnly checks if a gorm tag indicates the field is read-only
|
||||||
// Examples:
|
// Examples:
|
||||||
// - "<-:false" -> true (no writes allowed)
|
// - "<-:false" -> true (no writes allowed)
|
||||||
|
|||||||
@@ -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 {
|
if len(cols) == 0 {
|
||||||
return invalidArg("no writable fields in data")
|
return invalidArg("no writable fields in data")
|
||||||
}
|
}
|
||||||
|
reflection.RemoveNonWritableColumns(model, cols)
|
||||||
q := tx.NewInsert().Table(tableName)
|
q := tx.NewInsert().Table(tableName)
|
||||||
for key, value := range cols {
|
for key, value := range cols {
|
||||||
q = q.Value(key, value)
|
q = q.Value(key, value)
|
||||||
@@ -726,6 +727,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
|
|||||||
existingMap[key] = v
|
existingMap[key] = v
|
||||||
}
|
}
|
||||||
|
|
||||||
|
reflection.RemoveNonWritableColumns(model, setCols)
|
||||||
q := tx.NewUpdate().Table(tableName).SetMap(setCols).
|
q := tx.NewUpdate().Table(tableName).SetMap(setCols).
|
||||||
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
|
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
|
||||||
res, err := q.Exec(ctx)
|
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, ", "))
|
cond := fmt.Sprintf("%s IN (%s)", common.QuoteIdent(pkName), strings.Join(inList, ", "))
|
||||||
var affected int64
|
var affected int64
|
||||||
if req.op == "update" {
|
if req.op == "update" {
|
||||||
|
reflection.RemoveNonWritableColumns(model, setCols)
|
||||||
r, err := tx.NewUpdate().Table(tableName).SetMap(setCols).Where(cond, ids...).Exec(ctx)
|
r, err := tx.NewUpdate().Table(tableName).SetMap(setCols).Where(cond, ids...).Exec(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("error updating records: %w", err)
|
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
|
responseData = v
|
||||||
|
|
||||||
|
reflection.RemoveNonWritableColumns(model, v)
|
||||||
query := tx.NewInsert().Table(tableName)
|
query := tx.NewInsert().Table(tableName)
|
||||||
for key, value := range v {
|
for key, value := range v {
|
||||||
query = query.Value(key, common.ConvertSliceForBun(value))
|
query = query.Value(key, common.ConvertSliceForBun(value))
|
||||||
@@ -971,6 +972,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
|||||||
item = modifiedData
|
item = modifiedData
|
||||||
}
|
}
|
||||||
|
|
||||||
|
reflection.RemoveNonWritableColumns(model, item)
|
||||||
txQuery := tx.NewInsert().Table(tableName)
|
txQuery := tx.NewInsert().Table(tableName)
|
||||||
for key, value := range item {
|
for key, value := range item {
|
||||||
txQuery = txQuery.Value(key, common.ConvertSliceForBun(value))
|
txQuery = txQuery.Value(key, common.ConvertSliceForBun(value))
|
||||||
@@ -1127,6 +1129,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
|||||||
itemMap = modifiedData
|
itemMap = modifiedData
|
||||||
}
|
}
|
||||||
|
|
||||||
|
reflection.RemoveNonWritableColumns(model, itemMap)
|
||||||
txQuery := tx.NewInsert().Table(tableName)
|
txQuery := tx.NewInsert().Table(tableName)
|
||||||
for key, value := range itemMap {
|
for key, value := range itemMap {
|
||||||
txQuery = txQuery.Value(key, common.ConvertSliceForBun(value))
|
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)
|
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
||||||
common.MergeUpdateValues(existingMap, updates, h.disallowNulls)
|
common.MergeUpdateValues(existingMap, updates, h.disallowNulls)
|
||||||
|
reflection.RemoveNonWritableColumns(model, existingMap)
|
||||||
|
|
||||||
// Build update query with merged data
|
// Build update query with merged data
|
||||||
query := tx.NewUpdate().Table(tableName).SetMap(existingMap)
|
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)
|
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
||||||
common.MergeUpdateValues(existingMap, item, h.disallowNulls)
|
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)
|
txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
|
||||||
if _, err := txQuery.Exec(ctx); err != nil {
|
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)
|
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
||||||
common.MergeUpdateValues(existingMap, itemMap, h.disallowNulls)
|
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)
|
txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
|
||||||
if _, err := txQuery.Exec(ctx); err != nil {
|
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 == "" {
|
if id == "" {
|
||||||
options.SingleRecordAsObject = false
|
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
|
// 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)
|
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 {
|
query = grouper.WhereGroup(func(q common.SelectQuery) common.SelectQuery {
|
||||||
return applyUserConds(q).WhereOr(sanitizedOr)
|
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() == "" {
|
if provider, ok := modelValue.(common.TableNameProvider); !ok || provider.TableName() == "" {
|
||||||
query = query.Table(tableName)
|
query = query.Table(tableName)
|
||||||
}
|
}
|
||||||
|
if generated := reflection.NonWritableColumns(model); len(generated) > 0 {
|
||||||
|
query = query.ExcludeColumn(generated...)
|
||||||
|
}
|
||||||
fields := reflection.GetSQLModelColumns(model)
|
fields := reflection.GetSQLModelColumns(model)
|
||||||
query = query.Returning(fields...)
|
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
|
// Create update query using Model() to preserve custom types and driver.Valuer interfaces
|
||||||
query := tx.NewUpdate().Model(modelInstance)
|
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)
|
query = query.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
|
||||||
|
|
||||||
// Execute BeforeScan hooks - pass query chain so hooks can modify it
|
// 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
|
// Insert record
|
||||||
query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
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 {
|
if _, err := query.Exec(hookCtx.Context); err != nil {
|
||||||
return nil, fmt.Errorf("failed to create record: %w", err)
|
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.
|
// 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)
|
values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls)
|
||||||
|
|
||||||
|
reflection.RemoveNonWritableColumns(hookCtx.Model, values)
|
||||||
|
|
||||||
if len(values) > 0 {
|
if len(values) > 0 {
|
||||||
query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values).
|
query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values).
|
||||||
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID)
|
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)
|
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 {
|
func (m *MockInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||||
args := m.Called(columns)
|
args := m.Called(columns)
|
||||||
return args.Get(0).(common.InsertQuery)
|
return args.Get(0).(common.InsertQuery)
|
||||||
@@ -254,6 +259,11 @@ func (m *MockUpdateQuery) Model(model interface{}) common.UpdateQuery {
|
|||||||
return args.Get(0).(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 {
|
func (m *MockUpdateQuery) Table(table string) common.UpdateQuery {
|
||||||
args := m.Called(table)
|
args := m.Called(table)
|
||||||
return args.Get(0).(common.UpdateQuery)
|
return args.Get(0).(common.UpdateQuery)
|
||||||
|
|||||||
Reference in New Issue
Block a user