Compare commits

...
31 Commits
Author SHA1 Message Date
Hein 5a3a1df3c8 feat(resolvemcp)!: make the server read-only by default
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 1m33s
Tests / Unit Tests (push) Successful in 2m1s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m36s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m38s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m49s
Tests / Race Detector (push) Successful in 4m31s
BREAKING CHANGE: Config.ReadOnly is now a *bool and unset means read-only.
Use ReadOnly: resolvemcp.Bool(false) to enable insert/update/delete,
annotations and function calls. Adds the Bool helper and updates docs.
2026-10-07 14:17:56 +02:00
Hein 4ed9506ad2 feat(resolvemcp): add read-only mode and function allowlist
- Config.ReadOnly disables insert/update/delete/annotation tools, reports only
  select in list_tables/describe_table and tells the agent it cannot write
- Config.AllowFunctionCalls keeps function tools on a read-only server
- Config.AllowedFunctions limits list_functions/call_function to named
  functions (empty allows all); others are reported as unknown
- reflect read-only mode in the usage guide and exported catalogue
2026-10-07 14:15:14 +02:00
Hein 431b674162 feat(resolvemcp): add model descriptions, usage guide and catalogue export
- modelregistry: ModelInfo (description, purpose, tags, column docs) with
  external JSON loader, Describer fallback and gorm/bun/comment tag support
- resolvemcp: surface descriptions in list_tables and describe_table, send a
  usage guide as MCP server instructions, add package docs
- add BuildCatalog/ExportCatalog to write a JSON or Markdown API catalogue
- document the descriptions map and catalogue in the README
2026-10-07 14:09:37 +02:00
Hein 234aac9770 feat(metrics): bound HTTP path labels, add reset, custom push endpoint and JSON pull
- normalize the HTTP path label (ServeMux pattern, custom normalizer, ID
  collapsing) and cap distinct values via HTTPMaxPaths (default 1024)
- add Reset, PushAndReset, ResetHandler and reset-on-push options
- add POST push to a custom endpoint (text or json) with optional reset
- add JSONHandler for JSON pull
- honour Config.Enabled; log Pushgateway push failures
2026-10-07 12:09:11 +02:00
Hein 8cff3bde85 test(wrap_bunrouter): add tests for route param preservation
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 1m28s
Tests / Unit Tests (push) Successful in 1m30s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m53s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m9s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m10s
Tests / Race Detector (push) Successful in 3m39s
2026-10-05 16:44:44 +02:00
Hein 3e6224698c fix(handler): enforce single record return for ID queries 2026-10-05 16:04:57 +02:00
Hein aec87a81e7 fix(bun): ignore scanonly columns in ExcludeColumn
Tests / Integration Tests (push) Skipped
Tests / Unit Tests (push) Successful in 1m37s
Tests / Race Detector (push) Successful in 3m52s
Build , Vet Test, and Lint / Build (push) Successful in 1m33s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m16s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m19s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m19s
Bun's ExcludeColumn errors with "can't find column" for scanonly fields
because they are not in the table's writable fields. Filter the exclude
list to writable bun fields so models with scanonly buffers can insert
and update again. Add tests for the adapter and reflection.
2026-10-05 14:10:52 +02:00
warkanum 9235292586 fix(crud): skip generated and read-only columns on insert and update
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 1m34s
Tests / Unit Tests (push) Successful in 1m42s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m11s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m25s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m28s
Tests / Race Detector (push) Successful in 4m6s
Read-merge-update wrote every model column back, so GENERATED ALWAYS
columns failed with SQLSTATE 428C9. Add a bun 'generated' tag option,
reflection.NonWritableColumns/RemoveNonWritableColumns, and apply them
in resolvespec, restheadspec, websocketspec, mqttspec, resolvemcp and
the nested CUD processor. Add ExcludeColumn to InsertQuery/UpdateQuery
for model-based writes.
2026-10-02 22:43:45 +02:00
warkanum 23f10387c5 ci(release): fix rust toolchain setup and dart publish validation warnings
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 3m15s
Tests / Unit Tests (push) Successful in 3m26s
Build , Vet Test, and Lint / Lint Code (push) Successful in 4m45s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 4m55s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 4m55s
Tests / Race Detector (push) Successful in 7m6s
2026-10-01 21:20:04 +02:00
warkanum 0d3ad9e4fd ci(tests): disable integration tests job and install psql client 2026-10-01 21:16:34 +02:00
warkanum f5d232d971 ci: move workflows to Gitea and add client release workflow
Tests / Integration Tests (push) Failing after 1m39s
Build , Vet Test, and Lint / Build (push) Successful in 2m3s
Tests / Unit Tests (push) Successful in 2m8s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m51s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m57s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m59s
Tests / Race Detector (push) Successful in 5m9s
Move .github/workflows to .gitea/workflows with Gitea-compatible action
versions, fix make_tag outputs/major bump, and add release_clients.yml to
build and publish all clients to the Gitea package registries. Rename the Go
client module to git.warky.dev/wdevs and add LICENSE/CHANGELOG for Dart.
2026-10-01 21:11:41 +02:00
Hein a4702161fb docs(readme): add breaking changes section for 2026-09-30 to 2026-10-01 2026-10-01 15:17:37 +02:00
Hein 640faeeeaf feat(security): full OAuth 2.1 / OpenID Connect server and OIDC relying-party client
Authorization server: consent and scopes, OIDC (nonce, auth_time, acr, sid,
at_hash, signed userinfo, RP-initiated and back-channel logout), managed
refresh tokens with rotation and reuse detection, RFC 9068 JWT access tokens,
DPoP, PAR, device grant, token exchange, private_key_jwt, RFC 7591/7592
registration, RFC 9207 iss, signing keyring with rotation.

State is DB-backed through a new lookup.OAuthGrantStore (procedure and direct
backends, four dialect DDLs, conformance cases).

Client side: WithOIDC discovery, PKCE, nonce, id_token validation, OAuth2LogoutURL.

PeekRefresh now returns already rotated tokens so RotateRefresh can detect reuse.

Docs: OAUTH2_SERVER.md, oauth2_full_example.go, breaking_changes.md step 8.
2026-10-01 14:42:12 +02:00
Hein f54b707040 feat(pgsql): add WhereGroup and a podman/docker hardening test
- PgSQLSelectQuery implements common.WhereGrouper so x-custom-sql-or
  is grouped with the client's own conditions on the pgx adapter too.
- Add a container test (opt-in via RESOLVESPEC_TEST_CONTAINERS=1) that
  starts PostgreSQL with podman or docker and checks the hardening
  against a real database: parenthesis escape, pg_sleep, catalog
  subquery, stacked statements, x-custom-sql-or grouping and the
  legacy behaviour when hardening is switched off.
2026-10-01 14:41:46 +02:00
Hein ca89cb8a73 fix(common): harden CORS, sort, raw-SQL WHERE and x-custom-sql-or
Add a `hardening` config section (RESOLVESPEC_HARDENING_*) so each
fix can be switched off to restore the previous behaviour:

- cors_strict_origins: only reflect origins listed in
  cors.allowed_origins / server URLs, with credentials; `*` never
  sends credentials; fix shared-slice append of expose headers.
- sort_strict: join aliases must match `alias.identifier` (empty alias
  no longer matches everything); sort expressions reject dangerous
  functions/catalogs; cql* columns must be identifier-safe.
- sql_strict: client raw-SQL fragments must have balanced parens and
  quotes, no comments/`;`/`$$`, DML keywords, dangerous functions or
  system catalogs; a rejected fragment now fails closed ("(1=0)")
  instead of dropping the filter. Subqueries stay allowed unless
  sql_block_subqueries is set.
- x-custom-sql-or is grouped together with the client's own
  conditions (new optional WhereGrouper, implemented for bun and gorm)
  so it can no longer OR past server-side filters.
2026-10-01 13:46:01 +02:00
Hein c1153522f2 docs(resolvemcp): rewrite README for meta tools, guard and limits; update plan status 2026-10-01 13:42:14 +02:00
Hein 155e04deea chore(resolvemcp): drop unused helpers, silence rangeValCopy 2026-10-01 13:40:21 +02:00
Hein e49c3a916e feat(resolvemcp): replace per-model tools with fixed meta tools, guarded filter writes and a function registry
Tools: list_tables, describe_table, select_table, insert_into_table, update_table,
delete_from_table, list_functions, call_function. Visibility follows the model rules.
Filter-based update/delete require filters (never dropped silently), cap the matched rows
(MaxWriteRows), support dry_run, and need a single-use confirm token bound to caller, table,
filters, data and the matched rows. RegisterFunction adds Go-callback and SQL-procedure
functions run in a transaction with BeforeCall/AfterCall hooks. Per-model tools and
resources are removed.

fix(pgsql): UPDATE with SET and a multi-placeholder WHERE renumbered the WHERE parameters
wrongly ($1, $2 became $3, $2); shift them in one pass.
2026-10-01 13:40:00 +02:00
Hein 276c3814d8 feat(resolvemcp): read/write limits, preload validation, query timeout, stable client error codes
Config gains DefaultLimit/MaxLimit/MaxOffset/MaxBatch/MaxPreloadDepth/MaxWriteRows/QueryTimeout/
ConfirmTTL. Reads are capped and the total COUNT is optional. Errors reach clients as
{code,message}; everything else is logged with a reference. Panics (handler and hooks) are
recovered without returning the panic value.
2026-10-01 13:35:00 +02:00
Hein 82f901a49c fix(resolvemcp): single transaction for create/update, hook registry mutex, uniform not-found, bounded SSE host cache 2026-10-01 13:33:11 +02:00
Hein ad2f54693f feat(resolvemcp): require authentication on MCP endpoints and enforce model rules on writes
Guard() rejects unauthenticated callers (no guest/optional mode); Setup*/New* helpers take a
SecurityList and have explicit *Unauthenticated variants. Model rules now reach the security
hooks, create checks CanCreate (security.CheckModelCreateAllowed), create/update validate keys
against the model's writable columns, update sets only given keys (NULL allowed), update and
delete go through row security via a new BeforeScan hook, and the annotation tool is opt-in
(Config.EnableAnnotations) and runs BeforeHandle.
2026-10-01 13:31:13 +02:00
Hein 7662d5055c test(security): seed expired key as UTC in direct auth test
Tests / Race Detector (push) Failing after 27s
Tests / Unit Tests (push) Failing after 29s
Tests / Integration Tests (push) Failing after 30s
Build , Vet Test, and Lint / Build (push) Successful in 1m14s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m32s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 1m42s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m40s
2026-10-01 13:24:38 +02:00
Hein ea6a2e705f test(security): add SQL Server container conformance test; bind timestamps as UTC in direct backend 2026-10-01 13:24:26 +02:00
Hein 2516fcb13d chore(security): apply golangci-lint fixes to lookup and security packages 2026-10-01 13:21:00 +02:00
Hein c9fa8c60f2 refactor(security): move all database access into pkg/security/lookup
pkg/security no longer contains SQL. Every provider calls a store interface
from lookup, implemented by a procedure backend (Postgres stored procedures,
the default there) and a direct backend (dialect-driven SQL for postgres,
sqlite, mysql and mssql with configurable table and column names).

- add sectypes, lookup, lookup/{dialect,procedure,direct,backends,ddl,conformance}
- split totp and providers sub packages out of the core package
- replace SQLNames/TableNames/QueryMode with lookup.Config (see breaking_changes.md)
- direct backend now covers column/row security and API-key login
- move txsettings SQL to lookup.ApplyTxSettings; remove password.go
- move schema scripts under lookup/, add reference DDL per dialect
- add a shared conformance suite; run it on sqlite, and on Postgres in a
  podman/docker container (RESOLVESPEC_TEST_CONTAINERS=1)
- fix procedure schema bugs found on real Postgres: duplicate p_data
  parameter, JSON null arrays, expires_at timezone casts, passkey list
  GROUP BY, missing resolvespec_passkey_login; accept zone-less timestamps
2026-10-01 13:19:44 +02:00
Hein 60bd0a6dd3 feat(security): add plan for pkg/security lookup sub package 2026-10-01 11:29:02 +02:00
Hein 982c90bfdd feat(websocketspec): fire BeforeDisconnect/AfterDisconnect hooks on close
Hooks fire from Connection.Close() once per connection. ConnectionManager
Shutdown now closes connections outside its lock so hooks can call back
into the manager. Update the single-transaction audit plan status.
2026-10-01 10:53:21 +02:00
Hein ae2b0a4ef4 fix(security): skip row security filter for insert queries
Insert queries have no Where clause and read no existing rows, so the
fail-closed check rejected every insert when a row security template
was defined.
2026-10-01 10:47:54 +02:00
Hein daeea241af fix(restheadspec): honour string URL ids on POST updates
POST with a non-numeric URL id (e.g. a string primary key) was treated as
having no id, so the body primary key was used to look up the existing row.
That broke primary key changes. Treat any non-empty, non-zero URL id as the
update target, and add an integration test through Handle for POST and PUT.
2026-10-01 10:23:23 +02:00
Hein 3bd3e46409 fix(restheadspec): handle primary key changes in updates
* Allow primary key changes when specified in the request body.
* Ensure correct record fetching after primary key updates.
* Add integration tests for primary key update scenarios.
2026-10-01 10:16:26 +02:00
Hein 247111c32e docs(README): update table of contents and feature descriptions 2026-10-01 09:46:17 +02:00
228 changed files with 27799 additions and 7937 deletions
@@ -1,6 +1,3 @@
# This workflow will build a golang project
# For more information see: https://docs.github.com/en/actions/automating-builds-and-tests/building-and-testing-go
name: Create Go Release (Tag Versioning) 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
+268
View File
@@ -0,0 +1,268 @@
name: Release Clients
on:
workflow_dispatch:
inputs:
version:
description: "Client version (e.g. 1.4.0)"
required: true
type: string
publish:
description: "Publish packages to Gitea (untick for a build/test dry run)"
required: true
default: true
type: boolean
env:
VERSION_INPUT: ${{ github.event.inputs.version }}
PUBLISH: ${{ github.event.inputs.publish }}
SERVER_URL: ${{ github.server_url }}
OWNER: ${{ github.repository_owner }}
REGISTRY_USER: ${{ secrets.PACKAGE_REGISTRY_USERNAME || vars.PACKAGE_REGISTRY_USERNAME }}
TOKEN: ${{ secrets.PACKAGE_REGISTRY_TOKEN || vars.PACKAGE_REGISTRY_TOKEN }}
jobs:
validate:
name: Validate version
runs-on: ubuntu-latest
outputs:
version: ${{ steps.v.outputs.version }}
steps:
- id: v
run: |
version="${VERSION_INPUT#v}"
if ! [[ "$version" =~ ^[0-9]+\.[0-9]+\.[0-9]+(-[0-9A-Za-z.-]+)?$ ]]; then
echo "Invalid version: $VERSION_INPUT" >&2
exit 1
fi
echo "version=${version}" >> "${GITHUB_OUTPUT}"
js:
name: JS (npm)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-js
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- uses: actions/setup-node@v4
with:
node-version: "22"
- name: Enable pnpm
run: corepack enable
- name: Install
run: pnpm install --frozen-lockfile
- name: Test
run: pnpm test
- name: Set version
run: npm version "$VERSION" --no-git-tag-version --allow-same-version
- name: Build
run: pnpm build
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
run: |
host="${SERVER_URL#*://}"
registry="${SERVER_URL}/api/packages/${OWNER}/npm/"
npm config set "@warkypublic:registry" "$registry"
npm config set "//${host}/api/packages/${OWNER}/npm/:_authToken" "$TOKEN"
npm publish --registry "$registry"
python:
name: Python (PyPI)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-python
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Install
run: pip install -e ".[dev]" build twine
- name: Test
run: pytest
- name: Set version
run: sed -i -E "s/^version = \".*\"/version = \"${VERSION}\"/" pyproject.toml
- name: Build
run: python -m build
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
run: |
twine upload \
--repository-url "${SERVER_URL}/api/packages/${OWNER}/pypi" \
-u "$REGISTRY_USER" -p "$TOKEN" \
dist/*
rust:
name: Rust (Cargo)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-rs
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- name: Set up Rust
uses: dtolnay/rust-toolchain@stable
- name: Test
run: cargo test
- name: Set version
run: sed -i -E '0,/^version = ".*"/s//version = "'"${VERSION}"'"/' Cargo.toml
- name: Package
run: cargo package --allow-dirty
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
env:
CARGO_REGISTRIES_GITEA_INDEX: sparse+${{ github.server_url }}/api/packages/${{ github.repository_owner }}/cargo/
run: |
export CARGO_REGISTRIES_GITEA_TOKEN="Bearer ${TOKEN}"
cargo publish --registry gitea --allow-dirty
dotnet:
name: C# (NuGet)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-cs
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- uses: actions/setup-dotnet@v4
with:
dotnet-version: "8.0.x"
- name: Test
run: dotnet test tests/ResolveSpec.Tests.csproj
- name: Pack
run: dotnet pack src/ResolveSpec.csproj -c Release -p:Version="$VERSION" -o out
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
run: |
dotnet nuget push out/*.nupkg \
--source "${SERVER_URL}/api/packages/${OWNER}/nuget/index.json" \
--api-key "$TOKEN" \
--skip-duplicate
go:
name: Go (Go registry)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-go
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- uses: actions/setup-go@v5
with:
go-version-file: clients/resolvespec-go/go.mod
- name: Test
run: go test ./...
- name: Build module zip
run: |
python3 - <<'PY'
import os, re, zipfile
version = "v" + os.environ["VERSION"]
module = re.search(r"^module\s+(\S+)", open("go.mod").read(), re.M).group(1)
prefix = f"{module}@{version}/"
with zipfile.ZipFile("../resolvespec-go.zip", "w", zipfile.ZIP_DEFLATED) as z:
for root, dirs, files in os.walk("."):
dirs[:] = [d for d in dirs if d != ".git"]
for f in files:
path = os.path.join(root, f)
z.write(path, prefix + os.path.relpath(path, "."))
PY
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
run: |
curl -f -X PUT \
--user "${REGISTRY_USER}:${TOKEN}" \
--upload-file ../resolvespec-go.zip \
"${SERVER_URL}/api/packages/${OWNER}/go/upload"
dart:
name: Dart (Pub)
needs: validate
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: clients/resolvespec-dart
env:
VERSION: ${{ needs.validate.outputs.version }}
steps:
- uses: actions/checkout@v4
- uses: dart-lang/setup-dart@v1
- name: Install
run: dart pub get
- name: Analyze
run: dart analyze
- name: Test
run: dart test
- name: Set version and registry
run: |
sed -i -E "s/^version: .*/version: ${VERSION}/" pubspec.yaml
sed -i -E "s#^publish_to: .*#publish_to: ${SERVER_URL}/api/packages/${OWNER}/pub#" pubspec.yaml
if ! grep -q "^## ${VERSION}\$" CHANGELOG.md; then
{ head -n 1 CHANGELOG.md; printf '\n## %s\n\n- Release %s.\n' "$VERSION" "$VERSION"; tail -n +2 CHANGELOG.md; } > CHANGELOG.tmp
mv CHANGELOG.tmp CHANGELOG.md
fi
# pub warns about a dirty git tree; commit the stamped files locally (never pushed)
git -c user.name=ci -c user.email=ci@localhost commit -q -am "ci: stamp dart version ${VERSION}"
- name: Dry run
if: ${{ env.PUBLISH != 'true' }}
run: dart pub publish --dry-run
- name: Publish
if: ${{ env.PUBLISH == 'true' }}
run: |
dart pub token add "${SERVER_URL}/api/packages/${OWNER}/pub" --env-var TOKEN
dart pub publish --force
@@ -9,9 +9,9 @@ jobs:
name: Unit Tests 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
+287 -212
View File
@@ -13,71 +13,70 @@ ResolveSpec is a flexible and powerful REST API specification and implementation
All share the same core architecture and provide dynamic data querying, relationship preloading, and complex filtering. All share the same core architecture and provide dynamic data querying, relationship preloading, and complex filtering.
## Table of Contents ## Table of Contents
* [Features](#features) - [Features](#features)
* [Installation](#installation) - [Installation](#installation)
* [Quick Start](#quick-start) - [Quick Start](#quick-start)
* [ResolveSpec (Body-Based API)](#resolvespec---body-based-api) - [ResolveSpec (Body-Based API)](#resolvespec---body-based-api)
* [RestHeadSpec (Header-Based API)](#restheadspec---header-based-api) - [RestHeadSpec (Header-Based API)](#restheadspec---header-based-api)
* [ResolveMCP (MCP Server)](#resolvemcp---mcp-server) - [ResolveMCP (MCP Server)](#resolvemcp---mcp-server)
* [Architecture](#architecture) - [Architecture](#architecture)
* [API Structure](#api-structure) - [API Structure](#api-structure)
* [RestHeadSpec Overview](#restheadspec-header-based-api) - [RestHeadSpec Overview](#restheadspec-header-based-api)
* [Example Usage](#example-usage) - [Example Usage](#example-usage)
* [Testing](#testing) - [Testing](#testing)
* [Additional Packages](#additional-packages) - [Additional Packages](#additional-packages)
* [Security Considerations](#security-considerations) - [Security Considerations](#security-considerations)
* [What's New](#whats-new) - [Breaking Changes](#breaking-changes)
- [What's New](#whats-new)
## Features ## Features
### Core Features ### Core Features
* **Dynamic Data Querying**: Select specific columns and relationships to return - **Dynamic Data Querying**: Select specific columns and relationships to return
* **Relationship Preloading**: Load related entities with custom column selection and filters - **Relationship Preloading**: Load related entities with custom column selection and filters
* **Complex Filtering**: Apply multiple filters with various operators - **Complex Filtering**: Apply multiple filters with various operators
* **Sorting**: Multi-column sort support - **Sorting**: Multi-column sort support
* **Pagination**: Built-in limit/offset and cursor-based pagination (both ResolveSpec and RestHeadSpec) - **Pagination**: Built-in limit/offset and cursor-based pagination (both ResolveSpec and RestHeadSpec)
* **Computed Columns**: Define virtual columns for complex calculations - **Computed Columns**: Define virtual columns for complex calculations
* **Custom Operators**: Add custom SQL conditions when needed - **Custom Operators**: Add custom SQL conditions when needed
* **🆕 One Transaction Per Request**: Every statement and DB-touching hook of a request runs on one transaction; `OnTxBegin` hook stamps transaction-local settings (RLS) first. See [pkg/common/TRANSACTIONS.md](pkg/common/TRANSACTIONS.md) - **🆕 One Transaction Per Request**: Every statement and DB-touching hook of a request runs on one transaction; `OnTxBegin` hook stamps transaction-local settings (RLS) first. See [pkg/common/TRANSACTIONS.md](pkg/common/TRANSACTIONS.md)
* **🆕 Recursive CRUD Handler**: Automatically handle nested object graphs with foreign key resolution and per-record operation control via `_request` field - **🆕 Recursive CRUD Handler**: Automatically handle nested object graphs with foreign key resolution and per-record operation control via `_request` field
### Architecture (v2.0+) ### Architecture (v2.0+)
* **🆕 Database Agnostic**: Works with GORM, Bun, or any database layer through adapters - **🆕 Database Agnostic**: Works with GORM, Bun, or any database layer through adapters
* **🆕 Router Flexible**: Integrates with Gorilla Mux, Gin, Echo, or custom routers - **🆕 Router Flexible**: Integrates with Gorilla Mux, Gin, Echo, or custom routers
* **🆕 Backward Compatible**: Existing code works without changes - **🆕 Backward Compatible**: Existing code works without changes
* **🆕 Better Testing**: Mockable interfaces for easy unit testing - **🆕 Better Testing**: Mockable interfaces for easy unit testing
### ResolveMCP (v3.2+) ### ResolveMCP (v3.2+)
* **🆕 MCP Server**: Expose any registered database model as Model Context Protocol tools and resources - **🆕 MCP Server**: Expose any registered database model as Model Context Protocol tools and resources
* **🆕 AI-Ready Descriptions**: Tool descriptions include the full column schema, primary key, nullable flags, and relations — giving AI models everything they need to query correctly without guessing - **🆕 AI-Ready Descriptions**: Tool descriptions include the full column schema, primary key, nullable flags, and relations — giving AI models everything they need to query correctly without guessing
* **🆕 Four Tools Per Model**: `read_`, `create_`, `update_`, `delete_` tools auto-registered per model - **🆕 Four Tools Per Model**: `read_`, `create_`, `update_`, `delete_` tools auto-registered per model
* **🆕 Full Query Support**: Filters, sort, limit/offset, cursor pagination, column selection, and relation preloading all available as tool parameters - **🆕 Full Query Support**: Filters, sort, limit/offset, cursor pagination, column selection, and relation preloading all available as tool parameters
* **🆕 HTTP/SSE Transport**: Standards-compliant SSE transport for use with Claude Desktop, Cursor, and any MCP-compatible client - **🆕 HTTP/SSE Transport**: Standards-compliant SSE transport for use with Claude Desktop, Cursor, and any MCP-compatible client
* **🆕 Lifecycle Hooks**: Same Before/After hook system as ResolveSpec for auth and side-effects - **🆕 Lifecycle Hooks**: Same Before/After hook system as ResolveSpec for auth and side-effects
### RestHeadSpec (v2.1+) ### RestHeadSpec (v2.1+)
* **🆕 Header-Based API**: All query options passed via HTTP headers instead of request body - **🆕 Header-Based API**: All query options passed via HTTP headers instead of request body
* **🆕 Lifecycle Hooks**: Before/after hooks for create, read, update, and delete operations - **🆕 Lifecycle Hooks**: Before/after hooks for create, read, update, and delete operations
* **🆕 Cursor Pagination**: Efficient cursor-based pagination with complex sort support - **🆕 Cursor Pagination**: Efficient cursor-based pagination with complex sort support
* **🆕 Multiple Response Formats**: Simple, detailed, and Syncfusion-compatible formats - **🆕 Multiple Response Formats**: Simple, detailed, and Syncfusion-compatible formats
* **🆕 Single Record as Object**: Automatically normalize single-element arrays to objects (enabled by default) - **🆕 Single Record as Object**: Automatically normalize single-element arrays to objects (enabled by default)
* **🆕 Advanced Filtering**: Field filters, search operators, AND/OR logic, and custom SQL - **🆕 Advanced Filtering**: Field filters, search operators, AND/OR logic, and custom SQL
* **🆕 Base64 Encoding**: Support for base64-encoded header values - **🆕 Base64 Encoding**: Support for base64-encoded header values
### Routing & CORS (v3.0+) ### Routing & CORS (v3.0+)
* **🆕 Explicit Route Registration**: Routes created per registered model instead of dynamic lookups - **🆕 Explicit Route Registration**: Routes created per registered model instead of dynamic lookups
* **🆕 OPTIONS Method Support**: Full OPTIONS method support returning model metadata - **🆕 OPTIONS Method Support**: Full OPTIONS method support returning model metadata
* **🆕 CORS Headers**: Comprehensive CORS support with all HeadSpec headers allowed - **🆕 CORS Headers**: Comprehensive CORS support with all HeadSpec headers allowed
* **🆕 Better Route Control**: Customize routes per model with more flexibility - **🆕 Better Route Control**: Customize routes per model with more flexibility
## API Structure ## API Structure
@@ -131,7 +130,6 @@ X-DetailApi: true
For complete documentation including setup, headers, lifecycle hooks, cursor pagination, and more, see [pkg/restheadspec/README.md](pkg/restheadspec/README.md). For complete documentation including setup, headers, lifecycle hooks, cursor pagination, and more, see [pkg/restheadspec/README.md](pkg/restheadspec/README.md).
## Example Usage ## Example Usage
For detailed examples of reading data, cursor pagination, recursive CRUD operations, filtering, sorting, and more, see [pkg/resolvespec/README.md](pkg/resolvespec/README.md). For detailed examples of reading data, cursor pagination, recursive CRUD operations, filtering, sorting, and more, see [pkg/resolvespec/README.md](pkg/resolvespec/README.md).
@@ -142,14 +140,14 @@ First-class support for PostGIS geometry/geography and pgvector columns in `reso
### Column types (`pkg/spectypes`) ### Column types (`pkg/spectypes`)
| Go type | SQL type | Wire / JSON | | Go type | SQL type | Wire / JSON |
|--------------------|--------------|--------------------------------------------------------| | ----------------- | -------------- | ------------------------------------------------------- |
| `SqlGeometry` | `geometry` | JSON in/out = **GeoJSON**; also accepts EWKT / hex-EWKB | | `SqlGeometry` | `geometry` | JSON in/out = **GeoJSON**; also accepts EWKT / hex-EWKB |
| `SqlGeography` | `geography` | same as `SqlGeometry` | | `SqlGeography` | `geography` | same as `SqlGeometry` |
| `SqlVector` | `vector` | `[]float32` ⇄ `[1,2,3]` | | `SqlVector` | `vector` | `[]float32` ⇄ `[1,2,3]` |
| `SqlHalfVector` | `halfvec` | `[]float32` ⇄ `[1,2,3]` | | `SqlHalfVector` | `halfvec` | `[]float32` ⇄ `[1,2,3]` |
| `SqlSparseVector` | `sparsevec` | `{"dim":8,"indices":[1,4],"values":[0.5,0.2]}` | | `SqlSparseVector` | `sparsevec` | `{"dim":8,"indices":[1,4],"values":[0.5,0.2]}` |
| `SqlBitVector` | `bit`/`varbit` | bool array or `"1011"` string | | `SqlBitVector` | `bit`/`varbit` | bool array or `"1011"` string |
- Geometry `Value()` emits `SRID=<n>;<WKT>` (PostGIS implicit text→geometry cast; no wrapper function needed). - Geometry `Value()` emits `SRID=<n>;<WKT>` (PostGIS implicit text→geometry cast; no wrapper function needed).
- Declare dimensioned types with a tag: `gorm:"type:vector(1536)"` — the tag wins over the canonical name in metadata/OpenAPI. - Declare dimensioned types with a tag: `gorm:"type:vector(1536)"` — the tag wins over the canonical name in metadata/OpenAPI.
@@ -159,32 +157,36 @@ First-class support for PostGIS geometry/geography and pgvector columns in `reso
`value` is a geometry (GeoJSON object, EWKT string, or hex-EWKB) unless noted. `value` is a geometry (GeoJSON object, EWKT string, or hex-EWKB) unless noted.
| Operator | Value shape | | Operator | Value shape |
|----------|-------------| | ----------------------------------------------------------------------------------------------------------------------------------------------- | --------------------------------------------------------- |
| `st_intersects`, `st_contains`, `st_within`, `st_covers`, `st_coveredby`, `st_overlaps`, `st_touches`, `st_crosses`, `st_equals`, `st_disjoint` | geometry | | `st_intersects`, `st_contains`, `st_within`, `st_covers`, `st_coveredby`, `st_overlaps`, `st_touches`, `st_crosses`, `st_equals`, `st_disjoint` | geometry |
| `st_dwithin` | `{"geom": <geometry>, "distance": <meters>}` | | `st_dwithin` | `{"geom": <geometry>, "distance": <meters>}` |
| `bbox` (alias `&&`) | geometry, or `{"bbox":[minx,miny,maxx,maxy],"srid":4326}` | | `bbox` (alias `&&`) | geometry, or `{"bbox":[minx,miny,maxx,maxy],"srid":4326}` |
### Vector similarity filter operators ### Vector similarity filter operators
| Operator | pgvector op | Value shape | | Operator | pgvector op | Value shape |
|----------|-------------|-------------| | -------------------------------- | ----------- | ----------------------------------------------------------------- |
| `l2_within` / `euclidean_within` | `<->` | `{"vector":[...], "distance": <n>}` | | `l2_within` / `euclidean_within` | `<->` | `{"vector":[...], "distance": <n>}` |
| `cosine_within` | `<=>` | same (also `"lt"`/`"lte"`/`"gt"`/`"gte"` instead of `"distance"`) | | `cosine_within` | `<=>` | same (also `"lt"`/`"lte"`/`"gt"`/`"gte"` instead of `"distance"`) |
| `ip_within` / `inner_within` | `<#>` | same | | `ip_within` / `inner_within` | `<#>` | same |
### KNN search (ordering + distance column) ### KNN search (ordering + distance column)
**resolvespec** — `options.vector_search`: **resolvespec** — `options.vector_search`:
```json ```json
{ "options": { "vector_search": { {
"column": "embedding", "options": {
"vector": [0.1, 0.2, 0.3], "vector_search": {
"metric": "cosine", "column": "embedding",
"as": "_distance", "vector": [0.1, 0.2, 0.3],
"direction": "asc" "metric": "cosine",
}}} "as": "_distance",
"direction": "asc"
}
}
}
``` ```
Orders rows by distance; when `as` is set, returns the distance as an extra column (all model columns are auto-selected). Orders rows by distance; when `as` is set, returns the distance as an extra column (all model columns are auto-selected).
@@ -276,32 +278,28 @@ ResolveMCP exposes registered models as Model Context Protocol tools so AI model
```go ```go
import "github.com/bitechdev/ResolveSpec/pkg/resolvemcp" import "github.com/bitechdev/ResolveSpec/pkg/resolvemcp"
// Create handler handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{BaseURL: "http://localhost:8080", BasePath: "/mcp"})
handler := resolvemcp.NewHandlerWithGORM(db)
securityList, _ := security.NewSecurityList(provider)
resolvemcp.RegisterSecurityHooks(handler, securityList)
// Register models — must be done BEFORE Build()
handler.RegisterModel("public", "users", &User{}) handler.RegisterModel("public", "users", &User{})
handler.RegisterModel("public", "posts", &Post{}) handler.RegisterModel("public", "posts", &Post{})
// Finalize: registers MCP tools and resources // Mount the guarded SSE transport (OAuth bearer, session token or API key required)
handler.Build()
// Mount SSE transport on your existing router
router := mux.NewRouter() router := mux.NewRouter()
resolvemcp.SetupMuxRoutes(router, handler, "http://localhost:8080") resolvemcp.SetupMuxRoutes(router, handler, securityList)
// MCP clients connect to: // MCP clients connect to:
// SSE stream: GET http://localhost:8080/mcp/sse // SSE stream: GET http://localhost:8080/mcp/sse
// Messages: POST http://localhost:8080/mcp/message // Messages: POST http://localhost:8080/mcp/message
// //
// Auto-registered tools per model: // Fixed meta tools (independent of the number of models):
// read_public_users — filter, sort, paginate, preload // list_tables, describe_table, select_table, insert_into_table,
// create_public_users — insert a new record // update_table, delete_from_table, list_functions, call_function
// update_public_users — update a record by ID
// delete_public_users — delete a record by ID
``` ```
For complete documentation, see [pkg/resolvemcp/README.md](pkg/resolvemcp/README.md) (if present) or the package source. For complete documentation, see [pkg/resolvemcp/README.md](pkg/resolvemcp/README.md) .
## Architecture ## Architecture
@@ -343,25 +341,25 @@ Your Application Code
### Supported Database Layers ### Supported Database Layers
* **GORM** - Full support for PostgreSQL, SQLite, MSSQL - **GORM** - Full support for PostgreSQL, SQLite, MSSQL
* **Bun** - Full support for PostgreSQL, SQLite, MSSQL - **Bun** - Full support for PostgreSQL, SQLite, MSSQL
* **Native SQL** - Standard library `*sql.DB` with all supported databases - **Native SQL** - Standard library `*sql.DB` with all supported databases
* **Custom ORMs** - Implement the `Database` interface - **Custom ORMs** - Implement the `Database` interface
### Supported Databases ### Supported Databases
* **PostgreSQL** - Full schema support - **PostgreSQL** - Full schema support
* **SQLite** - Automatic schema.table to schema_table translation - **SQLite** - Automatic schema.table to schema_table translation
* **Microsoft SQL Server** - Full schema support - **Microsoft SQL Server** - Full schema support
* **MongoDB** - NoSQL document database (via MQTTSpec and custom handlers) - **MongoDB** - NoSQL document database (via MQTTSpec and custom handlers)
### Supported Routers ### Supported Routers
* **Gorilla Mux** (built-in support with `SetupRoutes()`) - **Gorilla Mux** (built-in support with `SetupRoutes()`)
* **BunRouter** (built-in support with `SetupBunRouterWithResolveSpec()`) - **BunRouter** (built-in support with `SetupBunRouterWithResolveSpec()`)
* **Gin** (manual integration, see examples above) - **Gin** (manual integration, see examples above)
* **Echo** (manual integration, see examples above) - **Echo** (manual integration, see examples above)
* **Custom Routers** (implement request/response adapters) - **Custom Routers** (implement request/response adapters)
## Testing ## Testing
@@ -373,11 +371,11 @@ ResolveSpec is designed for testability with mockable interfaces. For testing ex
### Test Server (dbtrace, real PostgreSQL) ### Test Server (dbtrace, real PostgreSQL)
* `make testserver-up` / `make testserver-down`: testserver + PostgreSQL via compose (host networking) - `make testserver-up` / `make testserver-down`: testserver + PostgreSQL via compose (host networking)
* `make testserver-smoke`: create, read, update, delete, batch create/delete against the testserver - `make testserver-smoke`: create, read, update, delete, batch create/delete against the testserver
* Ports: testserver `8123`, PostgreSQL `8124` - Ports: testserver `8123`, PostgreSQL `8124`
* `dbtrace` logs per request `tx`, `tx_queries`, `pooled`, `raw`; `pooled=0` is the target - `dbtrace` logs per request `tx`, `tx_queries`, `pooled`, `raw`; `pooled=0` is the target
* Integration tests default to PostgreSQL on `localhost:8124` - Integration tests default to PostgreSQL on `localhost:8124`
## Continuous Integration ## Continuous Integration
@@ -387,10 +385,10 @@ ResolveSpec uses GitHub Actions for automated testing and quality checks. The CI
The project includes automated workflows that: The project includes automated workflows that:
* **Test**: Run all tests with race detection and code coverage - **Test**: Run all tests with race detection and code coverage
* **Lint**: Check code quality with golangci-lint - **Lint**: Check code quality with golangci-lint
* **Build**: Verify the project builds successfully - **Build**: Verify the project builds successfully
* **Multi-version**: Test against multiple Go versions (1.23.x, 1.24.x) - **Multi-version**: Test against multiple Go versions (1.23.x, 1.24.x)
### Running Tests Locally ### Running Tests Locally
@@ -412,9 +410,9 @@ golangci-lint run
The project includes comprehensive test coverage: The project includes comprehensive test coverage:
* **Unit Tests**: Individual component testing - **Unit Tests**: Individual component testing
* **Integration Tests**: End-to-end API testing - **Integration Tests**: End-to-end API testing
* **CRUD Tests**: Standalone tests for both ResolveSpec and RestHeadSpec APIs - **CRUD Tests**: Standalone tests for both ResolveSpec and RestHeadSpec APIs
To run only the CRUD standalone tests: To run only the CRUD standalone tests:
@@ -445,6 +443,7 @@ ResolveSpec includes several complementary packages that work together to provid
The core body-based REST API with GraphQL-like capabilities. The core body-based REST API with GraphQL-like capabilities.
**Key Features**: **Key Features**:
- JSON request body with operation and options - JSON request body with operation and options
- Recursive CRUD with nested object support - Recursive CRUD with nested object support
- Cursor and offset pagination - Cursor and offset pagination
@@ -458,6 +457,7 @@ For complete documentation, see [pkg/resolvespec/README.md](pkg/resolvespec/READ
Alternative REST API where query options are passed via HTTP headers. Alternative REST API where query options are passed via HTTP headers.
**Key Features**: **Key Features**:
- All query options via HTTP headers - All query options via HTTP headers
- Same capabilities as ResolveSpec - Same capabilities as ResolveSpec
- Cleaner separation of data and metadata - Cleaner separation of data and metadata
@@ -470,6 +470,7 @@ For complete documentation, see [pkg/restheadspec/README.md](pkg/restheadspec/RE
Expose any registered model as Model Context Protocol tools and resources consumable by AI models over HTTP/SSE. Expose any registered model as Model Context Protocol tools and resources consumable by AI models over HTTP/SSE.
**Key Features**: **Key Features**:
- Four tools per model: `read_`, `create_`, `update_`, `delete_` - Four tools per model: `read_`, `create_`, `update_`, `delete_`
- Rich AI-readable descriptions: column names, types, primary key, nullable flags, and preloadable relations - Rich AI-readable descriptions: column names, types, primary key, nullable flags, and preloadable relations
- Full query support: filters, sort, limit/offset, cursor pagination, column selection, preloads - Full query support: filters, sort, limit/offset, cursor pagination, column selection, preloads
@@ -483,6 +484,7 @@ For complete documentation, see [pkg/resolvemcp/](pkg/resolvemcp/).
Execute SQL functions and queries through a simple HTTP API with header-based parameters. Execute SQL functions and queries through a simple HTTP API with header-based parameters.
**Key Features**: **Key Features**:
- Direct SQL function invocation - Direct SQL function invocation
- Header-based parameter passing - Header-based parameter passing
- Automatic pagination and counting - Automatic pagination and counting
@@ -495,20 +497,21 @@ For complete documentation, see [pkg/funcspec/](pkg/funcspec/).
All clients are under [clients/](clients/README.md); wire behaviour is identical across them. All clients are under [clients/](clients/README.md); wire behaviour is identical across them.
| Client | Language | Specs | Docs | | Client | Language | Specs | Docs |
|---|---|---|---| | -------------------- | -------------- | ---------------------------------------------------- | ---------------------------------------------- |
| `resolvespec-js` | TypeScript | ResolveSpec, HeaderSpec, FunctionSpec, WebSocketSpec | [README](clients/resolvespec-js/README.md) | | `resolvespec-js` | TypeScript | ResolveSpec, HeaderSpec, FunctionSpec, WebSocketSpec | [README](clients/resolvespec-js/README.md) |
| `resolvespec-python` | Python >= 3.11 | ResolveSpec, HeaderSpec, FunctionSpec, WebSocketSpec | [README](clients/resolvespec-python/README.md) | | `resolvespec-python` | Python >= 3.11 | ResolveSpec, HeaderSpec, FunctionSpec, WebSocketSpec | [README](clients/resolvespec-python/README.md) |
| `resolvespec-go` | Go | ResolveSpec, FunctionSpec | [README](clients/resolvespec-go/README.md) | | `resolvespec-go` | Go | ResolveSpec, FunctionSpec | [README](clients/resolvespec-go/README.md) |
| `resolvespec-rs` | Rust | ResolveSpec, FunctionSpec | [README](clients/resolvespec-rs/README.md) | | `resolvespec-rs` | Rust | ResolveSpec, FunctionSpec | [README](clients/resolvespec-rs/README.md) |
| `resolvespec-cs` | C# (.NET 8) | ResolveSpec, FunctionSpec | [README](clients/resolvespec-cs/README.md) | | `resolvespec-cs` | C# (.NET 8) | ResolveSpec, FunctionSpec | [README](clients/resolvespec-cs/README.md) |
| `resolvespec-dart` | Dart / Flutter | ResolveSpec, FunctionSpec | [README](clients/resolvespec-dart/README.md) | | `resolvespec-dart` | Dart / Flutter | ResolveSpec, FunctionSpec | [README](clients/resolvespec-dart/README.md) |
#### ResolveSpec JS - TypeScript Client Library #### ResolveSpec JS - TypeScript Client Library
TypeScript/JavaScript client library supporting all three REST and WebSocket protocols. TypeScript/JavaScript client library supporting all three REST and WebSocket protocols.
**Clients**: **Clients**:
- Body-based REST client (`read`, `create`, `update`, `deleteEntity`) - Body-based REST client (`read`, `create`, `update`, `deleteEntity`)
- Header-based REST client (`HeaderSpecClient`) - Header-based REST client (`HeaderSpecClient`)
- WebSocket client (`WebSocketClient`) with CRUD, subscriptions, heartbeat, reconnect - WebSocket client (`WebSocketClient`) with CRUD, subscriptions, heartbeat, reconnect
@@ -522,6 +525,7 @@ For complete documentation, see [clients/resolvespec-js/README.md](clients/resol
Real-time bidirectional communication with full CRUD operations and subscriptions. Real-time bidirectional communication with full CRUD operations and subscriptions.
**Key Features**: **Key Features**:
- Persistent WebSocket connections - Persistent WebSocket connections
- Real-time subscriptions to entity changes - Real-time subscriptions to entity changes
- Automatic push notifications - Automatic push notifications
@@ -535,6 +539,7 @@ For complete documentation, see [pkg/websocketspec/README.md](pkg/websocketspec/
MQTT-based database operations ideal for IoT and mobile applications. MQTT-based database operations ideal for IoT and mobile applications.
**Key Features**: **Key Features**:
- Embedded or external MQTT broker support - Embedded or external MQTT broker support
- QoS 1 (at-least-once delivery) - QoS 1 (at-least-once delivery)
- Real-time subscriptions - Real-time subscriptions
@@ -550,6 +555,7 @@ For complete documentation, see [pkg/mqttspec/README.md](pkg/mqttspec/README.md)
Flexible, interface-driven static file server. Flexible, interface-driven static file server.
**Key Features**: **Key Features**:
- Router-agnostic with standard `http.Handler` - Router-agnostic with standard `http.Handler`
- Multiple filesystem backends (local, zip, embedded) - Multiple filesystem backends (local, zip, embedded)
- Pluggable cache, MIME, and fallback policies - Pluggable cache, MIME, and fallback policies
@@ -557,6 +563,7 @@ Flexible, interface-driven static file server.
- 140+ MIME types including modern formats - 140+ MIME types including modern formats
**Quick Example**: **Quick Example**:
```go ```go
import "github.com/bitechdev/ResolveSpec/pkg/server/staticweb" import "github.com/bitechdev/ResolveSpec/pkg/server/staticweb"
@@ -581,6 +588,7 @@ For complete documentation, see [pkg/server/staticweb/README.md](pkg/server/stat
Comprehensive event handling system for real-time event publishing and cross-instance communication. Comprehensive event handling system for real-time event publishing and cross-instance communication.
**Key Features**: **Key Features**:
- Multiple event sources (database, websockets, frontend, system) - Multiple event sources (database, websockets, frontend, system)
- Multiple providers (in-memory, Redis Streams, NATS, PostgreSQL) - Multiple providers (in-memory, Redis Streams, NATS, PostgreSQL)
- Pattern-based subscriptions - Pattern-based subscriptions
@@ -595,6 +603,7 @@ For complete documentation, see [pkg/eventbroker/README.md](pkg/eventbroker/READ
Centralized management of multiple database connections with support for PostgreSQL, SQLite, MSSQL, and MongoDB. Centralized management of multiple database connections with support for PostgreSQL, SQLite, MSSQL, and MongoDB.
**Key Features**: **Key Features**:
- Multiple named database connections - Multiple named database connections
- Multi-ORM access (Bun, GORM, Native SQL) sharing the same connection pool - Multi-ORM access (Bun, GORM, Native SQL) sharing the same connection pool
- Automatic SQLite schema translation (`schema.table` → `schema_table`) - Automatic SQLite schema translation (`schema.table` → `schema_table`)
@@ -634,9 +643,11 @@ For documentation, see [pkg/cache/README.md](pkg/cache/README.md).
#### Security #### Security
Authentication and authorization framework with hooks integration. Database-backed providers use PostgreSQL stored procedures by default, with a portable Direct mode (plain Go/SQL) for SQLite, MySQL, or Postgres without the procedures installed. Authentication and authorization framework with hooks integration. Database-backed providers use PostgreSQL stored procedures by default, with a direct SQL backend (SQLite, MySQL, SQL Server, or Postgres without the procedures) selected through `lookup.Config`.
For documentation, see [pkg/security/README.md](pkg/security/README.md) (see "Direct Mode" for the SQLite/portable-SQL path). For documentation, see [pkg/security/README.md](pkg/security/README.md) (see "Database access (lookup)" for the SQLite/portable-SQL path).
It includes a standards-based OAuth 2.1 / OpenID Connect authorization server (consent, rotating refresh tokens, JWT access tokens, DPoP, PAR, device grant, token exchange, logout) and an OIDC relying-party client; see [pkg/security/OAUTH2_SERVER.md](pkg/security/OAUTH2_SERVER.md).
#### Middleware #### Middleware
@@ -690,23 +701,23 @@ For documentation, see [pkg/dbtrace/README.md](pkg/dbtrace/README.md).
### Core Libraries ### Core Libraries
| Package | Purpose | | Package | Purpose |
|---|---| | ----------------------------------------- | --------------------------------------------------------------------------------------------------------------------------------------------------------- |
| [`pkg/common`](pkg/common/) | Shared interfaces (database, request/response adapters), validation, recursive CRUD, request transactions ([TRANSACTIONS.md](pkg/common/TRANSACTIONS.md)) | | [`pkg/common`](pkg/common/) | Shared interfaces (database, request/response adapters), validation, recursive CRUD, request transactions ([TRANSACTIONS.md](pkg/common/TRANSACTIONS.md)) |
| [`pkg/modelregistry`](pkg/modelregistry/) | Model registration by schema/entity and per-model access rules | | [`pkg/modelregistry`](pkg/modelregistry/) | Model registration by schema/entity and per-model access rules |
| [`pkg/reflection`](pkg/reflection/) | Model/struct reflection helpers (primary keys, columns, relations) | | [`pkg/reflection`](pkg/reflection/) | Model/struct reflection helpers (primary keys, columns, relations) |
| [`pkg/spectypes`](pkg/spectypes/) | SQL-aware types (nullable, JSONB, PostGIS, vector) | | [`pkg/spectypes`](pkg/spectypes/) | SQL-aware types (nullable, JSONB, PostGIS, vector) |
| [`pkg/logger`](pkg/logger/) | Logging used by all packages | | [`pkg/logger`](pkg/logger/) | Logging used by all packages |
| [`pkg/testmodels`](pkg/testmodels/) | Shared test models and data for tests and the testserver | | [`pkg/testmodels`](pkg/testmodels/) | Shared test models and data for tests and the testserver |
## Security Considerations ## Security Considerations
* Implement proper authentication and authorization - Implement proper authentication and authorization
* Validate all input parameters - Validate all input parameters
* Use prepared statements (handled by GORM/Bun/your ORM) - Use prepared statements (handled by GORM/Bun/your ORM)
* Implement rate limiting (`middleware.RateLimiter`) and per-client request queueing (`middleware.ClientQueue`) - Implement rate limiting (`middleware.RateLimiter`) and per-client request queueing (`middleware.ClientQueue`)
* Control access at schema/entity level - Control access at schema/entity level
* **New**: Database abstraction layer provides additional security through interface boundaries - **New**: Database abstraction layer provides additional security through interface boundaries
## Contributing ## Contributing
@@ -720,21 +731,73 @@ For documentation, see [pkg/dbtrace/README.md](pkg/dbtrace/README.md).
This project is licensed under the MIT License - see the [LICENSE](LICENSE) file for details. This project is licensed under the MIT License - see the [LICENSE](LICENSE) file for details.
## Breaking Changes
Changes from 2026-09-30 to 2026-10-01 that require action when upgrading. The full `pkg/security` migration tables are in [pkg/security/breaking_changes.md](pkg/security/breaking_changes.md).
### Security package (`pkg/security`)
- **No SQL in `pkg/security`**: all database access moved to `pkg/security/lookup`. Removed `SQLNames`, `TableNames`, `KeyStoreSQLNames`, `KeyStoreTableNames`, `QueryMode` (`ModeAuto`/`ModeProcedure`/`ModeDirect`), `ErrDirectModeUnsupported` and the `SQLNames`/`TableNames`/`QueryMode` option fields and `WithQueryMode`/`WithTableNames` builders. Use `lookup.Config` (`Dialect`, `Mode`, `Overrides`, `Procs`, `Schema`) via the `Lookup` option field or `WithLookup`/`WithLookupProvider`. The variadic `names ...*SQLNames` argument was dropped from `NewJWTAuthenticator`, `NewDatabaseColumnSecurityProvider`, `NewDatabaseRowSecurityProvider` and `NewDatabaseTwoFactorProvider`.
- **Default lookup mode is per dialect**: stored procedures on Postgres, direct SQL elsewhere. `ModeAuto` is opt-in.
- **Packages moved (no aliases)**: TOTP types to `pkg/security/totp`; header, config key store and config column/row providers to `pkg/security/providers`.
- **SQL schema files moved** from `pkg/security/` to `pkg/security/lookup/`.
- **Password handling**: bcrypt is verified in Direct mode and in the shipped procedures; passwords are hashed on register and reset; client-supplied roles/level are ignored at registration; legacy cleartext passwords need an explicit opt-in to upgrade. Row security templates bind the user as a parameter, and a filter that cannot be attached now fails the request.
- **OAuth2 / OIDC server**: existing databases need new columns (`oauth_clients.metadata`, `oauth_codes.extra`) and new tables (`oauth_consents`, `oauth_refresh_tokens`, `oauth_device_codes`, `oauth_par_requests`, `oauth_jti`); Postgres procedure mode needs the schema reapplied. `/oauth/introspect` and `/oauth/revoke` now require client authentication (`AllowAnonymousIntrospection` restores the old behaviour). Only PKCE `S256` is accepted. Authorization errors redirect to the client after `redirect_uri` validation.
- **Row security** now applies to update and delete queries (skipped for inserts); hidden/masked columns are excluded from create and update payloads.
### Request hardening (`pkg/common`)
Enabled by default under the new `hardening` config section (`RESOLVESPEC_HARDENING_*`); each can be switched off to restore the old behaviour.
- **`cors_strict_origins`**: only origins listed in `cors.allowed_origins` or the server URLs are reflected; `*` never sends credentials.
- **`sort_strict`**: sort expressions and join aliases must be identifier-safe; dangerous functions and catalogs are rejected.
- **`sql_strict`**: client raw-SQL fragments with unbalanced quotes/parens, comments, `;`, DML keywords or system catalogs are rejected, and a rejected fragment now fails closed (`1=0`) instead of dropping the filter.
- **`x-custom-sql-or`** is grouped with the client's own conditions so it can no longer OR past server-side filters. `common.Database` query builders gain an optional `WhereGrouper` (implemented for bun and gorm).
### ResolveMCP (`pkg/resolvemcp`)
- **Authentication required**: `SetupMuxRoutes`, `SetupBunRouterRoutes`, `SetupMuxStreamableHTTPRoutes`, `SetupBunRouterStreamableHTTPRoutes`, `NewSSEServer` and `NewStreamableHTTPHandler` now take a `*security.SecurityList`. Use the `*Unauthenticated` variants to keep the old open behaviour.
- **Per-model tools removed**: the `read_`/`create_`/`update_`/`delete_` tools and per-model resources are replaced by fixed meta tools: `list_tables`, `describe_table`, `select_table`, `insert_into_table`, `update_table`, `delete_from_table`, `list_functions`, `call_function`.
- **Guarded writes**: filter-based update/delete require filters, are capped by `MaxWriteRows`, and need a single-use confirm token (or `dry_run`).
- **Limits and errors**: reads are capped (`DefaultLimit`, `MaxLimit`, `MaxOffset`, `MaxBatch`, `MaxPreloadDepth`, `QueryTimeout`); the total `COUNT` is optional; errors reach clients as `{code, message}` only.
- **Writes**: create checks `CanCreate`; create/update reject keys outside the model's writable columns; update sets only the given keys; the annotation tool is opt-in via `Config.EnableAnnotations`.
### Update semantics
- **`""` and `null` now overwrite** stored values on update in resolvespec and restheadspec (previously skipped). Use `Handler.SetDisallowNulls` to skip nulls.
- **websocketspec / mqttspec** update only the keys present in the payload (`SetMap`) instead of writing the whole zeroed model.
### Transactions and hooks
- **One transaction per request** in every spec. Hooks must use `hookCtx.Tx`, not the pool; `BeforeHandle` runs before any transaction and must not touch the DB. `OnTxBegin` fires first in every transaction.
- **`AfterDelete` failure now rolls the delete back.**
- **Create/update re-fetch, `BeforeScan` and post-commit hooks** (`AfterCreate`, `AfterUpdate`, restheadspec `AfterRead`, funcspec `BeforeResponse`) run on a second, short transaction after the first commits.
- **websocketspec / mqttspec / funcspec** begin or commit failures answer `transaction_error`.
- **`BeforeDisconnect` / `AfterDisconnect`** hooks in websocketspec now fire on close.
### Other
- **Clients moved**: the JS and Python clients now live under `clients/`.
- **Test server ports**: `8123` (testserver) and `8124` (PostgreSQL), previously `8080` and `5434`.
- **pgsql**: subquery preload errors are returned instead of being logged and skipped; a `SET` plus multi-placeholder `WHERE` update previously renumbered `WHERE` parameters wrongly (fixed).
- **Cache**: `Clear()` on the Redis and Memcache providers requires `AllowFlush`; a missing key returns `ErrNotFound`; Memcache keys are hashed and namespaced, so existing entries are not found.
- **Config**: `NewManager` no longer replaces the global manager (use `SetConfigManager`); saved configs are written `0600`; `PathsConfig.Join` is confined to its base path.
## What's New ## What's New
### Unreleased ### Unreleased
**Single transaction per request**: **Single transaction per request**:
* **One tx per request**: hooks get the transaction in `hookCtx.Tx`, never the pool (`BeforeHandle` runs before any tx and must not touch the DB) - **One tx per request**: hooks get the transaction in `hookCtx.Tx`, never the pool (`BeforeHandle` runs before any tx and must not touch the DB)
* **`OnTxBegin` hook**: all specs (mqttspec re-exports websocketspec's); fires once, first, in every tx; error or abort rolls back with no detail to the client - **`OnTxBegin` hook**: all specs (mqttspec re-exports websocketspec's); fires once, first, in every tx; error or abort rolls back with no detail to the client
* **Second short tx**: create/update re-fetch, `BeforeScan` and post-commit hooks (`AfterCreate`, `AfterUpdate`, restheadspec `AfterRead`, funcspec `BeforeResponse`) run on a new tx after the first commits - **Second short tx**: create/update re-fetch, `BeforeScan` and post-commit hooks (`AfterCreate`, `AfterUpdate`, restheadspec `AfterRead`, funcspec `BeforeResponse`) run on a new tx after the first commits
* **Delete**: single and batch delete, hooks included, in one tx - **Delete**: single and batch delete, hooks included, in one tx
* **websocketspec / mqttspec**: one tx per message; begin/commit failures answer `transaction_error` - **websocketspec / mqttspec**: one tx per message; begin/commit failures answer `transaction_error`
* **resolvemcp**: read, create, update, delete transactional - **resolvemcp**: read, create, update, delete transactional
* **RLS stamping**: `SecurityList.SetTxSettings(fn)`; `RegisterSecurityHooks` of every spec stamps `set_config(name, value, true)` on `OnTxBegin`; fails closed - **RLS stamping**: `SecurityList.SetTxSettings(fn)`; `RegisterSecurityHooks` of every spec stamps `set_config(name, value, true)` on `OnTxBegin`; fails closed
* **New**: `common.RunRequestTx`, `common.TxContext`, `common.TxHookName` - **New**: `common.RunRequestTx`, `common.TxContext`, `common.TxHookName`
* **Behavior changes**: `AfterDelete` failure now rolls the delete back; funcspec begin/commit failure answers 500 `transaction_error` - **Behavior changes**: `AfterDelete` failure now rolls the delete back; funcspec begin/commit failure answers 500 `transaction_error`
**Clients**: Go, Rust, C# and Dart clients for ResolveSpec and FunctionSpec under `clients/`. **Clients**: Go, Rust, C# and Dart clients for ResolveSpec and FunctionSpec under `clients/`.
@@ -744,121 +807,133 @@ This project is licensed under the MIT License - see the [LICENSE](LICENSE) file
**ResolveMCP - Model Context Protocol Server (🆕)**: **ResolveMCP - Model Context Protocol Server (🆕)**:
* **MCP Tools**: Four tools auto-registered per model (`read_`, `create_`, `update_`, `delete_`) over HTTP/SSE transport - **MCP Tools**: Four tools auto-registered per model (`read_`, `create_`, `update_`, `delete_`) over HTTP/SSE transport
* **AI-Ready Descriptions**: Full column schema, primary key, nullable flags, and relation names surfaced in tool descriptions so AI models can query without guessing - **AI-Ready Descriptions**: Full column schema, primary key, nullable flags, and relation names surfaced in tool descriptions so AI models can query without guessing
* **Full Query Support**: Filters, sort, limit/offset, cursor pagination, column selection, and relation preloading all available as tool parameters - **Full Query Support**: Filters, sort, limit/offset, cursor pagination, column selection, and relation preloading all available as tool parameters
* **HTTP/SSE Transport**: Standards-compliant transport compatible with Claude Desktop, Cursor, and any MCP 2024-11-05 client - **HTTP/SSE Transport**: Standards-compliant transport compatible with Claude Desktop, Cursor, and any MCP 2024-11-05 client
* **Lifecycle Hooks**: Same Before/After hook system as ResolveSpec for auth, auditing, and side-effects - **Lifecycle Hooks**: Same Before/After hook system as ResolveSpec for auth, auditing, and side-effects
* **MCP Resources**: Each model also exposed as a named resource for direct data access by AI clients - **MCP Resources**: Each model also exposed as a named resource for direct data access by AI clients
### v3.1 (February 2026) ### v3.1 (February 2026)
**SQLite Schema Translation (🆕)**: **SQLite Schema Translation (🆕)**:
* **Automatic Schema Translation**: SQLite support with automatic `schema.table` to `schema_table` conversion - **Automatic Schema Translation**: SQLite support with automatic `schema.table` to `schema_table` conversion
* **Database Agnostic Models**: Write models once, use across PostgreSQL, SQLite, and MSSQL - **Database Agnostic Models**: Write models once, use across PostgreSQL, SQLite, and MSSQL
* **Transparent Handling**: Translation occurs automatically in all operations (SELECT, INSERT, UPDATE, DELETE, preloads) - **Transparent Handling**: Translation occurs automatically in all operations (SELECT, INSERT, UPDATE, DELETE, preloads)
* **All ORMs Supported**: Works with Bun, GORM, and Native SQL adapters - **All ORMs Supported**: Works with Bun, GORM, and Native SQL adapters
### v3.0 (December 2025) ### v3.0 (December 2025)
**Explicit Route Registration (🆕)**: **Explicit Route Registration (🆕)**:
* **Breaking Change**: Routes are now created explicitly for each registered model - **Breaking Change**: Routes are now created explicitly for each registered model
* **Better Control**: Customize routes per model with more flexibility - **Better Control**: Customize routes per model with more flexibility
* **Registration Order**: Models must be registered BEFORE calling SetupMuxRoutes/SetupBunRouterRoutes - **Registration Order**: Models must be registered BEFORE calling SetupMuxRoutes/SetupBunRouterRoutes
* **Benefits**: More flexible routing, easier to add custom routes per model, better performance - **Benefits**: More flexible routing, easier to add custom routes per model, better performance
**OPTIONS Method & CORS Support (🆕)**: **OPTIONS Method & CORS Support (🆕)**:
* **OPTIONS Endpoint**: Full OPTIONS method support for CORS preflight requests - **OPTIONS Endpoint**: Full OPTIONS method support for CORS preflight requests
* **Metadata Response**: OPTIONS returns model metadata (same as GET /metadata) - **Metadata Response**: OPTIONS returns model metadata (same as GET /metadata)
* **CORS Headers**: Comprehensive CORS headers on all responses - **CORS Headers**: Comprehensive CORS headers on all responses
* **Header Support**: All HeadSpec custom headers (`X-Select-Fields`, `X-FieldFilter-*`, etc.) allowed - **Header Support**: All HeadSpec custom headers (`X-Select-Fields`, `X-FieldFilter-*`, etc.) allowed
* **No Auth on OPTIONS**: CORS preflight requests don't require authentication - **No Auth on OPTIONS**: CORS preflight requests don't require authentication
* **Configurable**: Customize CORS settings via `common.CORSConfig` - **Configurable**: Customize CORS settings via `common.CORSConfig`
### v2.1 ### v2.1
**Cursor Pagination for ResolveSpec (🆕 Dec 9, 2025)**: **Cursor Pagination for ResolveSpec (🆕 Dec 9, 2025)**:
* **Cursor-Based Pagination**: Efficient cursor pagination now available in ResolveSpec (body-based API) - **Cursor-Based Pagination**: Efficient cursor pagination now available in ResolveSpec (body-based API)
* **Consistent with RestHeadSpec**: Both APIs now support cursor pagination for feature parity - **Consistent with RestHeadSpec**: Both APIs now support cursor pagination for feature parity
* **Multi-Column Sort Support**: Works seamlessly with complex sorting requirements - **Multi-Column Sort Support**: Works seamlessly with complex sorting requirements
* **Better Performance**: Improved performance for large datasets compared to offset pagination - **Better Performance**: Improved performance for large datasets compared to offset pagination
* **SQL Safety**: Proper SQL sanitization for cursor values - **SQL Safety**: Proper SQL sanitization for cursor values
**Recursive CRUD Handler (🆕 Nov 11, 2025)**: **Recursive CRUD Handler (🆕 Nov 11, 2025)**:
* **Nested Object Graphs**: Automatically handle complex object hierarchies with parent-child relationships - **Nested Object Graphs**: Automatically handle complex object hierarchies with parent-child relationships
* **Foreign Key Resolution**: Automatic propagation of parent IDs to child records - **Foreign Key Resolution**: Automatic propagation of parent IDs to child records
* **Per-Record Operations**: Control create/update/delete operations per record via `_request` field - **Per-Record Operations**: Control create/update/delete operations per record via `_request` field
* **Transaction Safety**: All nested operations execute atomically within database transactions - **Transaction Safety**: All nested operations execute atomically within database transactions
* **Relationship Detection**: Automatic detection of belongsTo, hasMany, hasOne, and many2many relationships - **Relationship Detection**: Automatic detection of belongsTo, hasMany, hasOne, and many2many relationships
* **Deep Nesting Support**: Handle relationships at any depth level - **Deep Nesting Support**: Handle relationships at any depth level
* **Mixed Operations**: Combine insert, update, and delete operations in a single request - **Mixed Operations**: Combine insert, update, and delete operations in a single request
**Primary Key Improvements (Nov 11, 2025)**: **Primary Key Improvements (Nov 11, 2025)**:
* **GetPrimaryKeyName**: Enhanced primary key detection for better preload and ID field handling - **GetPrimaryKeyName**: Enhanced primary key detection for better preload and ID field handling
* **Better GORM/Bun Support**: Improved compatibility with both ORMs for primary key operations - **Better GORM/Bun Support**: Improved compatibility with both ORMs for primary key operations
* **Computed Column Support**: Fixed computed columns functionality across handlers - **Computed Column Support**: Fixed computed columns functionality across handlers
**Database Adapter Enhancements (Nov 11, 2025)**: **Database Adapter Enhancements (Nov 11, 2025)**:
* **Bun ORM Relations**: Using Scan model method for better has-many and many-to-many relationship handling - **Bun ORM Relations**: Using Scan model method for better has-many and many-to-many relationship handling
* **Model Method Support**: Enhanced query building with proper model registration - **Model Method Support**: Enhanced query building with proper model registration
* **Improved Type Safety**: Better handling of relationship queries with type-aware scanning - **Improved Type Safety**: Better handling of relationship queries with type-aware scanning
**RestHeadSpec - Header-Based REST API**: **RestHeadSpec - Header-Based REST API**:
* **Header-Based Querying**: All query options via HTTP headers instead of request body - **Header-Based Querying**: All query options via HTTP headers instead of request body
* **Lifecycle Hooks**: Before/after hooks for create, read, update, delete operations - **Lifecycle Hooks**: Before/after hooks for create, read, update, delete operations
* **Cursor Pagination**: Efficient cursor-based pagination with complex sorting - **Cursor Pagination**: Efficient cursor-based pagination with complex sorting
* **Advanced Filtering**: Field filters, search operators, AND/OR logic - **Advanced Filtering**: Field filters, search operators, AND/OR logic
* **Multiple Response Formats**: Simple, detailed, and Syncfusion-compatible responses - **Multiple Response Formats**: Simple, detailed, and Syncfusion-compatible responses
* **Single Record as Object**: Automatically return single-element arrays as objects (default, toggleable via header) - **Single Record as Object**: Automatically return single-element arrays as objects (default, toggleable via header)
* **Base64 Support**: Base64-encoded header values for complex queries - **Base64 Support**: Base64-encoded header values for complex queries
* **Type-Aware Filtering**: Automatic type detection and conversion for filters - **Type-Aware Filtering**: Automatic type detection and conversion for filters
**Core Improvements**: **Core Improvements**:
* Better model registry with schema.table format support - Better model registry with schema.table format support
* Enhanced validation and error handling - Enhanced validation and error handling
* Improved reflection safety - Improved reflection safety
* Fixed COUNT query issues with table aliasing - Fixed COUNT query issues with table aliasing
* Better pointer handling throughout the codebase - Better pointer handling throughout the codebase
* **Comprehensive Test Coverage**: Added standalone CRUD tests for both ResolveSpec and RestHeadSpec - **Comprehensive Test Coverage**: Added standalone CRUD tests for both ResolveSpec and RestHeadSpec
### v2.0 ### v2.0
**Breaking Changes**: **Breaking Changes**:
* **None!** Full backward compatibility maintained - **None!** Full backward compatibility maintained
**New Features**: **New Features**:
* **Database Abstraction**: Support for GORM, Bun, and custom ORMs - **Database Abstraction**: Support for GORM, Bun, and custom ORMs
* **Router Flexibility**: Works with any HTTP router through adapters - **Router Flexibility**: Works with any HTTP router through adapters
* **BunRouter Integration**: Built-in support for uptrace/bunrouter - **BunRouter Integration**: Built-in support for uptrace/bunrouter
* **Better Architecture**: Clean separation of concerns with interfaces - **Better Architecture**: Clean separation of concerns with interfaces
* **Enhanced Testing**: Mockable interfaces for comprehensive testing - **Enhanced Testing**: Mockable interfaces for comprehensive testing
**Performance Improvements**: **Performance Improvements**:
* More efficient query building through interface design - More efficient query building through interface design
* Reduced coupling between components - Reduced coupling between components
* Better memory management with interface boundaries - Better memory management with interface boundaries
# Security Policy
## Reporting a vulnerability
Please do not open a public issue for security problems.
Report privately through GitHub: Security → Report a vulnerability
(https://github.com/bitechdev/ResolveSpec/security/advisories/new),
or email hein@bitechsystems.co.za / hein@warky.dev
You'll get an acknowledgement within 7 days. We aim to release a fix within
90 days and will credit reporters in the advisory unless they prefer otherwise.
## Acknowledgments ## Acknowledgments
* Inspired by REST, OData, and GraphQL's flexibility - Inspired by REST, OData, and GraphQL's flexibility
* **Header-based approach**: Inspired by REST best practices and clean API design - **Header-based approach**: Inspired by REST best practices and clean API design
* **Database Support**: [GORM](https://gorm.io) and [Bun](https://bun.uptrace.dev/) - **Database Support**: [GORM](https://gorm.io) and [Bun](https://bun.uptrace.dev/)
* **Router Support**: Gorilla Mux (built-in), BunRouter, Gin, Echo, and others through adapters - **Router Support**: Gorilla Mux (built-in), BunRouter, Gin, Echo, and others through adapters
* Slogan generated using DALL-E - Slogan generated using DALL-E
* AI used for documentation checking and correction - AI used for documentation checking and correction
* Community feedback and contributions that made v2.0 and v2.1 possible - Community feedback and contributions that made v2.0 and v2.1 possible
![1.00](./generated_slogan.webp) ![1.00](./generated_slogan.webp)
+2 -2
View File
@@ -1,6 +1,6 @@
# resolvemcp rewrite plan # resolvemcp rewrite plan
Source: `audit/pkg/resolvemcp.audit.md`. Status: plan only, no code changed. Source: `audit/pkg/resolvemcp.audit.md`. Status: items 1-8 and 10 implemented; item 9 (tests) mostly done, see git log.
## Goal ## Goal
@@ -72,7 +72,7 @@ Same rules as resolvespec CRUD, plus guardrails.
### 1. API key login (`pkg/security`) ### 1. API key login (`pkg/security`)
- Existing: keystore has `ValidateKey` and `KeyStoreAuthenticator`; `Login` needs a password; no key-to-session path. - Existing: keystore has `ValidateKey` and `KeyStoreAuthenticator`; `Login` needs a password; no key-to-session path.
- Add `resolvespec_login_api_key` to `SQLNames` (default + override) and a SQL script beside the existing procedures. Contract: `p_success, p_error, p_data`, input raw key; hashes, validates active/non-expired key, creates session for the key's user. - Add `resolvespec_login_api_key` to `lookup.ProcNames` (default + override; was `SQLNames` before the lookup refactor) and a SQL script beside the existing procedures. Contract: `p_success, p_error, p_data`, input raw key; hashes, validates active/non-expired key, creates session for the key's user.
- Add `DatabaseAuthenticator.LoginWithAPIKey(ctx, rawKey)`; procedure first, direct-SQL fallback via `ShouldUseProcedure`. - Add `DatabaseAuthenticator.LoginWithAPIKey(ctx, rawKey)`; procedure first, direct-SQL fallback via `ShouldUseProcedure`.
- Hashed lookup; same generic error for unknown, expired or inactive key; no key material in logs. - Hashed lookup; same generic error for unknown, expired or inactive key; no key material in logs.
- Expose through the chain/composite authenticators so the middleware can accept it. - Expose through the chain/composite authenticators so the middleware can accept it.
+1
View File
@@ -6,6 +6,7 @@
| **Files** | `handler.go` (901), `tools.go` (720), `cursor.go`, `oauth2.go`, `oauth2_server.go`, `annotation.go`, `hooks.go`, `security_hooks.go`, `context.go`, `resolvemcp.go` | | **Files** | `handler.go` (901), `tools.go` (720), `cursor.go`, `oauth2.go`, `oauth2_server.go`, `annotation.go`, `hooks.go`, `security_hooks.go`, `context.go`, `resolvemcp.go` |
| **Tests** | `tools_test.go` (34), `tx_test.go` (207); `go test` passes. No hostile-input tests, no `-race` | | **Tests** | `tools_test.go` (34), `tx_test.go` (207); `go test` passes. No hostile-input tests, no `-race` |
| **Audit date** | 2026-09-30 | | **Audit date** | 2026-09-30 |
| **Status** | Rewrite implemented (meta tools, guard, limits, guardrails, function registry); see `audit/mcp_plan.md` |
| **Axes** | thread locking/waiting, slowness, security, panic handling & logging, agent usability | | **Axes** | thread locking/waiting, slowness, security, panic handling & logging, agent usability |
| **Threat model** | hostile or confused MCP client (LLM agent, possibly prompt-injected); tool arguments are attacker-controlled | | **Threat model** | hostile or confused MCP client (LLM agent, possibly prompt-injected); tool arguments are attacker-controlled |
| **Depth** | targeted (request path, security wiring; verified against source) | | **Depth** | targeted (request path, security wiring; verified against source) |
+1
View File
@@ -9,6 +9,7 @@
| **Docs** | `README.md`, `SECURITY_FEATURES.md`, `QUICK_REFERENCE.md`, `OAUTH2.md`, `OAUTH2_REFRESH_*.md`, `PASSKEY_QUICK_REFERENCE.md`, `KEYSTORE.md` | | **Docs** | `README.md`, `SECURITY_FEATURES.md`, `QUICK_REFERENCE.md`, `OAUTH2.md`, `OAUTH2_REFRESH_*.md`, `PASSKEY_QUICK_REFERENCE.md`, `KEYSTORE.md` |
| **Tests** | 6 359 lines across 13 `_test.go` files | | **Tests** | 6 359 lines across 13 `_test.go` files |
| **Audit date** | 2026-09-29 | | **Audit date** | 2026-09-29 |
| **Note** | Point-in-time snapshot. File names and line numbers refer to the code as audited. Since then all SQL moved out of `pkg/security` into `pkg/security/lookup`: `providers_direct.go`, `sql_names.go`, `table_names.go`, `query_mode.go` and `password.go` are gone, `SQLNames` / `TableNames` / `QueryMode` became `lookup.Config`, and the SQL files moved to `pkg/security/lookup/`. See `pkg/security/breaking_changes.md` for the mapping. |
| **Axes** | thread locking/waiting, slowness, security, panic handling & logging | | **Axes** | thread locking/waiting, slowness, security, panic handling & logging |
| **Threat model** | hostile internet client; request bodies, headers, query params, schema/table/column names and filter expressions all attacker-controlled | | **Threat model** | hostile internet client; request bodies, headers, query params, schema/table/column names and filter expressions all attacker-controlled |
| **Depth** | deep | | **Depth** | deep |
+346
View File
@@ -0,0 +1,346 @@
# pkg/security lookup sub package plan
Status: implemented (steps 0-7). What shipped and every API change is recorded in `pkg/security/breaking_changes.md`;
usage is documented in `pkg/security/README.md` ("Database access (lookup)"). This file is kept as the design record.
Deviations from the plan below: only `totp` and `providers` were split out of `pkg/security` (no `oauth` package, the
OAuth server and passkey provider stay in `security`), the `Database*` constructors stay in `security`, and
`ddl/postgres.sql` is tables only and cannot be combined with the procedure schema.
Related: `audit/mcp_plan.md` (work item 1, API key login).
## Problem
- `pkg/security` mixes two data-access styles: stored procedures (`resolvespec_*`) and ~1,800 lines of hand-written
"direct" SQL (`*_direct.go`) selected per call by `QueryMode` / `ShouldUseProcedure`.
- Direct SQL is written once with `?` placeholders and only rewritten for Postgres. It assumes one fixed schema
(table and column names, JSON stored as TEXT, bool/time handling), and is only really exercised on SQLite.
- Table names are configurable (`TableNames`), column names are not. Procedure names are configurable (`SQLNames`).
- Postgres cannot be run "tables only" in a first-class way, and other databases have no defined support.
- Rule going forward: **`pkg/security` itself contains no SQL.** All lookups go through one sub package.
## Goal
New sub package `pkg/security/lookup` that owns every database read/write the security package needs.
| Requirement | Decision |
|---|---|
| Postgres default | Existing stored procedures, existing names, unchanged behaviour out of the box |
| Postgres direct | Optional: work on tables directly with no procs installed |
| SQLite | First-class: configurable tables and columns |
| Other DBs | MySQL/MariaDB and MSSQL via dialects; adding more = adding a dialect |
| Config | Procedure names, table names, column names, mode (procedure / direct / auto), per backend |
| `pkg/security` | Calls lookup interfaces only; no `SELECT`/`INSERT`/`UPDATE`/`DELETE`, no `pg_proc` probing |
## Design
### Package layout
```
pkg/security/ # core: behaviour interfaces, SecurityList, middleware, chain, composite, hooks, write security, tx settings, type aliases
pkg/security/sectypes/ # shared data types (no deps, no SQL, no logic beyond small helpers)
pkg/security/providers/ # concrete authenticators + security providers (see Package split)
pkg/security/oauth/ # OAuth2 client login + OAuth2 authorization server
pkg/security/totp/ # two-factor: generator, providers, TwoFactorAuthenticator
pkg/security/passkey/ # WebAuthn passkey provider + passkey login flow
pkg/security/lookup/
lookup.go # store interfaces + record types + Config + New(db, cfg)
schema.go # Schema: table + column names per entity, defaults, merge, validate
mode.go # Mode (Auto/Procedure/Direct), per-operation resolution, proc probe (pg only)
dialect/ # Dialect interface + postgres, sqlite, mysql, mssql
procedure/ # procedure backend (current SQLNames, p_success/p_error/p_data contract)
direct/ # dialect-driven SQL backend (no fixed SQL strings per dialect)
ddl/ # reference schemas per dialect (replaces database_schema*.sql variants)
```
### Shared types: `pkg/security/sectypes`
`lookup` cannot import `pkg/security` (cycle: security -> lookup -> security), so the plain data types move to a
dependency-free sub package that both import.
- Moves to `sectypes`: `UserContext`, `LoginRequest`, `LoginResponse`, `RegisterRequest`, `LogoutRequest`,
`PasswordResetRequest/Response/CompleteRequest`, `KeyType`, `UserKey`, `CreateKeyRequest/Response`,
`PasskeyCredential` (+ passkey request/option structs the stores return), `TwoFactorSecret`, OAuth server
client/code/token-info structs (`OAuthServerClient`, `OAuthCode`, `OAuthTokenInfo`), `ColumnSecurity`, `RowSecurity`.
- Stays in `pkg/security`: all behaviour interfaces (`Authenticator`, `SecurityProvider`, `Registrable`, ...),
authenticators, middleware, `OAuthServer`, TOTP generator, hooks. They reference `sectypes` types.
- Compatibility: `pkg/security` re-exports each moved type as an alias (`type UserContext = sectypes.UserContext`) and
the `KeyType*` constants, so `security.UserContext` etc. keep compiling and are identical types. In-repo users
(eventbroker, funcspec, mqttspec, resolvemcp, resolvespec, restheadspec, websocketspec, ~17 files) need no edits.
- `lookup` stores take and return `sectypes` types directly; no separate record types and no conversion layer.
- Rules: `sectypes` imports only the standard library (and `oauth2` types only if unavoidable, else a local struct);
JSON tags unchanged so wire formats and procedure `p_data` payloads stay identical.
### Package split
Dependency direction (no cycles): `sectypes` <- `lookup` <- `providers`/`oauth`/`totp`/`passkey` -> `security` (core) -> `sectypes`.
Core `security` never imports the sub packages. Sub packages do not import each other (see Dependency rules).
### Dependency rules (how cycles are avoided)
1. **Layers, imports only point down.** L0 `sectypes` (stdlib only) -> L1 `lookup`, core `security` interfaces ->
L2 `providers`, `oauth`, `totp`, `passkey`. A package may import lower layers, never its own layer or above.
2. **Define interfaces where they are consumed, not where they are implemented** (Go idiom). E.g. `oauth` declares
the small `SessionCreator` it needs; `providers.DatabaseAuthenticator` satisfies it without `oauth` importing `providers`.
3. **Shared data goes down, not sideways.** If two L2 packages need the same struct, it moves to `sectypes`
(or a tiny `internal/` package), never "A imports B for one type".
4. **No L2 -> L2 imports.** Composition happens in the application (or an optional top-level `security/setup`
package that imports everything and is imported by nobody in `pkg/security`).
5. **Dependency injection by constructor**, passing interfaces/stores (`lookup.Provider`, `security.Authenticator`);
no package-level registries that need a back-import; use functional options for optional collaborators.
6. **Core never imports concrete implementations**; where core needs behaviour it calls an interface it owns
(hooks, `SecurityContext`, `Authenticator`).
7. **Tests:** external test packages (`package foo_test`) for cross-package integration tests, so test-only
imports cannot create cycles; shared fixtures in an `internal/testutil` package.
8. **Guard in CI:** `go list -deps` / a small test that asserts the layer rules (e.g. `sectypes` imports only stdlib,
`lookup` does not import `security`, no L2 package imports another L2 package). `go build` already rejects true cycles.
| Package | Contents (from today's files) |
|---|---|
| `security` (core) | `Authenticator`, `SecurityProvider`, `Registrable`, `Refreshable`, `APIKeyLoginable`, ... interfaces; `SecurityList`; `SecurityContext`; middleware + cookie options; `ChainAuthenticator`; `CompositeSecurityProvider`; hooks; `WriteDataContext`; `TxSettings`; type aliases to `sectypes` |
| `providers` | `DatabaseAuthenticator`, `JWTAuthenticator`, `HeaderAuthenticator`, `KeyStoreAuthenticator`, `ConfigKeyStore`, `DatabaseKeyStore`, `DatabaseColumnSecurityProvider`, `DatabaseRowSecurityProvider`, `Config*SecurityProvider` |
| `oauth` | `OAuth2Config`, `OAuth2Provider`, Google/GitHub/Microsoft/Facebook/multi-provider constructors, OAuth2 refresh, `OAuthServer` + `OAuthServerConfig`, oauth server persistence (via `lookup.OAuthClientStore`) |
| `passkey` | `PasskeyProvider` impl (`DatabasePasskeyProvider`), registration/authentication flows, passkey request/option types that are not shared (shared ones stay in `sectypes`) |
| `totp` | `TwoFactorAuthProvider`, `TwoFactorConfig`, `TOTPGenerator`, `MemoryTwoFactorProvider`, `DatabaseTwoFactorProvider`, `TwoFactorAuthenticator` |
Consequences to design for:
- **Methods cannot span packages.** Today OAuth2 and passkey logic are methods on `DatabaseAuthenticator`
(`oauth2_methods*.go`, `oauth_server_db*.go`, passkey methods) and `NewOAuthServer` takes `*DatabaseAuthenticator`.
They become standalone types in `oauth` / `providers` that depend on `lookup` stores and on small interfaces
(e.g. `oauth.SessionCreator`) instead of the concrete authenticator. `NewGoogleAuthenticator(...)` etc. return a
`providers.DatabaseAuthenticator` configured with an `oauth.Provider`, or an `oauth.Authenticator` that implements
`security.Authenticator`; pick one in step 5 (see Open).
- **Constructors cannot be re-exported from core `security`** (it would import the sub packages = cycle). Types that
move to `sectypes` keep aliases; constructors and concrete types do not. This is a breaking import change.
In-repo callers affected (outside `pkg/security`): `pkg/resolvemcp` (`oauth2.go`, `oauth2_server.go`, `handler.go`),
`pkg/middleware/clientqueue.go`, docs and examples. Provide a mechanical migration table
(`security.NewDatabaseAuthenticator` -> `providers.NewDatabaseAuthenticator`, `security.OAuthServer` -> `oauth.Server`, ...).
- **Interfaces core needs from sub packages** (e.g. 2FA hook points) are defined in core or `sectypes`, implemented in
`totp`; core never imports `totp`.
- `examples*.go` / `oauth2_examples.go` / `passkey_examples.go` move next to the package they exemplify (or to
`_example_test.go` files) so core has no dependency on them.
- Tests move with their code; shared helpers (sqlite test DB, `authenticatedRequest`) go to an internal test helper package.
### Store interfaces (one per domain, mirrors current procs)
| Store | Operations (current proc in brackets) |
|---|---|
| `AuthStore` | `Login` [login], `Register` [register], `Logout` [logout], `Session` [session], `TouchSession` [session_update], `Refresh` [refresh_token], `LoginAPIKey` [login_api_key], `JWTLogin`, `JWTLogout`, `ResetRequest`, `ResetComplete` |
| `KeyStore` | `Create`, `List`, `Delete`, `Validate` [keystore_*] |
| `OAuthClientStore` | register client, get client, save code, exchange code, introspect, revoke |
| `OAuthUserStore` | get-or-create user, create session, get/update refresh token, get user |
| `PasskeyStore` | store, get, update counter, list, delete, rename, get by username, login |
| `TOTPStore` | enable, disable, status, secret, regenerate backup codes, validate backup code |
| `PolicyStore` | column security, row security: procedure backend (default) + direct backend over the `sec_*` table layout below |
Each store has a procedure implementation and a direct implementation. A `Provider` bundles them; `security`
constructors take a `lookup.Provider` (or build one from `db` + `lookup.Config`, so existing constructors keep working).
### Config
```go
type Config struct {
Dialect string // "postgres" | "sqlite" | "mysql" | "mssql"; empty = detect from driver
Mode Mode // Auto | Procedure | Direct; default: Procedure for postgres, Direct otherwise
Overrides map[Op]Mode // optional per-operation mode, e.g. direct for Session, procedure for Login
Procs ProcNames // = today's SQLNames (+ LoginAPIKey), defaults unchanged
Schema Schema // tables + columns, see below
}
```
- `Schema` = per entity `{Table string; Columns map[Column]string}` with typed column keys covering every column
(`users.id`, `users.username`, `users.password`, `users.is_active`, ...). Defaults reproduce the current schema,
so zero config behaves as today. Optional `Schema` name per entity for `schema.table` qualification.
- Merge + validation as today: non-empty override wins; every identifier checked against `^[a-zA-Z_][a-zA-Z0-9_]*$`
(plus optional single `schema.` prefix). Identifiers are quoted by the dialect, never interpolated raw.
- No back-compat for config: `SQLNames`, `KeyStoreSQLNames`, `TableNames`, `KeyStoreTableNames` and `QueryMode` are
removed from `pkg/security`; `lookup.Config` replaces them (decision 8).
### Dialect interface (the per-database adaptor)
One adaptor per database type (`dialect/postgres`, `sqlite`, `mysql`, `mssql`), selected by `Config.Dialect` or
detected from the driver, registered through `dialect.Register(name, factory)` so more databases can be added later
without touching core code. Each adaptor supplies only the things that differ; the direct backend builds queries from it.
| Concern | Dialect method |
|---|---|
| Placeholders | `Placeholder(n)` (`$n`, `?`, `@pn`) |
| Identifier quoting | `Quote(ident)` (`"x"`, `` `x` ``, `[x]`) |
| Booleans | `Bool(v)` / scan helper (bool vs 0/1) |
| Time | `Now()` expr or Go-side `time.Now()`; scan helper for drivers returning strings |
| Insert returning id | `InsertReturningID(table, cols, idCol)` returns the SQL + scan strategy: postgres `... RETURNING id` (QueryRow), sqlite/mysql `LastInsertId`, mssql `... OUTPUT INSERTED.id` (QueryRow). The only dialect-specific write construct (decision 12 / confirmed) |
| Get-or-create | none; standard SQL select-then-insert inside a tx (no upsert) |
| Limit/top | only if a query needs it |
| JSON columns (scopes/meta/roles) | `EncodeJSON` / `DecodeJSON` (native jsonb vs TEXT) |
| Random / hashing | done in Go (token generation, SHA-256 key hash, bcrypt) so direct mode needs no `pgcrypto` and no DB functions |
| Driver detection | `Detect(*sql.DB)` from driver type (replaces `driverIsPostgres` / `driverIsPortableOnly`) |
Queries are assembled by a small internal builder (select/insert/update/delete with named columns from `Schema`),
not string-concatenated per dialect and not via an ORM, to keep `pkg/security` free of bun/gorm.
### PolicyStore table layout (column / row security, direct backend)
Approved layout. Both the procedure backend (`resolvespec_column_security` / `resolvespec_row_security`, rewritten
in `database_schema.sql`) and the direct backend read these tables; the former external schema is no longer
referenced anywhere in the repo. All table and column names are configurable via `Schema`, defaults shown.
`sec_group_members` (optional; omit to use direct user rules only)
| Column | Type | Notes |
|---|---|---|
| `group_id` | int, not null | group a user belongs to |
| `user_id` | int, not null | FK users.id; PK (`group_id`, `user_id`) |
`sec_column_rules`
| Column | Type | Notes |
|---|---|---|
| `id` | int PK | |
| `user_id` | int null | rule for one user |
| `group_id` | int null | rule for every member of the group; exactly one of `user_id` / `group_id` set (check constraint) |
| `schema_name` | text, not null | matched case-insensitively |
| `table_name` | text, not null | matched case-insensitively |
| `column_path` | text, not null | dot path under the table (`col` or `col.sub.field`) = `ColumnSecurity.Path` joined by `.` |
| `access_type` | text, not null | `ColumnSecurity.Accesstype` (e.g. `mask`, `hide`, `read`) |
| `mask_start`, `mask_end` | int null | default 0 |
| `mask_invert` | bool null | default false |
| `mask_char` | text null | default `*` |
| `extra_filters` | text/JSON null | `ExtraFilters` map, JSON-encoded via dialect `EncodeJSON` |
| `is_active` | bool, not null | default true |
`sec_row_rules`
| Column | Type | Notes |
|---|---|---|
| `id` | int PK | |
| `user_id` | int null / `group_id` int null | as above, exactly one set |
| `schema_name`, `table_name` | text, not null | case-insensitive match |
| `template` | text null | SQL fragment with the existing placeholders (`RowSecurity.Template`) |
| `has_block` | bool, not null | default false; true = no rows visible (`RowSecurity.HasBlock`) |
| `is_active` | bool, not null | default true |
Resolution rules (direct backend; also the contract the conformance tests assert):
- Applicable rules = active rules where `user_id` = caller, plus rules of every group the caller belongs to.
- Column security: all applicable rules for the exact schema + table, returned as `[]ColumnSecurity` (union).
Exact table match, not a prefix match (a prefix would match `users_archive` for `users`).
- Row security: any applicable `has_block` wins; otherwise templates of all applicable rules are combined with
`AND` (each wrapped in parentheses); no rule = `RowSecurity{}` with `ErrNoRowSecurity` semantics unchanged.
- Templates are still validated/substituted by the existing safe-identifier code in core; the store only loads text.
- Both loaders keep the guarantees of today: user reference reduced to a scalar, failures are errors (fail closed),
no rule is "no rules".
### Mode resolution
- `Procedure`: always call the proc; missing proc = error (no silent fallback).
- `Direct`: always use tables via the dialect builder.
- `Auto`: Postgres probes `pg_proc` once per proc (cached, as today); other dialects resolve to `Direct`.
Probe lives in `lookup` and is the only place that queries the catalog.
- Postgres default stays `Procedure` so current installs do not change behaviour.
- Roles/user-level safety rules already enforced in direct mode (Register ignores client-supplied level/roles, bcrypt
hash, opt-in password upgrade) become backend-agnostic tests that both backends must pass.
## Work items
### 0. Extract shared types (prerequisite, no behaviour change)
- Create `pkg/security/sectypes`, move the types listed above, add aliases in `pkg/security`.
- Verify `go build ./...` and `go test ./pkg/...` unchanged; `go vet` for alias/import cycles.
- Do this first and on its own so the diff is a pure move.
### 0b. Package split (after 0, before `lookup` wiring)
- Create `providers`, `oauth`, `totp`; move files per the Package split table, one package per commit:
`totp` and `passkey` (self-contained) -> `providers` (key stores, authenticators, policy providers) -> `oauth` (needs de-methoding
from `DatabaseAuthenticator`).
- Break the method-on-`DatabaseAuthenticator` coupling for OAuth2 and passkey first (extract interfaces), then move.
- Update in-repo callers and docs; add migration table to `pkg/security/README.md`.
- Behaviour unchanged; at this point stores are still the old direct/proc code, only relocated.
### 1. Skeleton and contracts
- Create `lookup` package: records, store interfaces, `Config`, `Schema` (defaults, merge, validate), `Mode`.
- No behaviour change yet; compile-only.
### 2. Dialects
- Implement `postgres`, `sqlite`, `mysql`, `mssql` against the dialect interface; driver detection.
- Unit tests per dialect: placeholders, quoting, bool/time round trip, insert-returning-id.
### 3. Procedure backend
- Move existing proc calls out of `pkg/security` into `lookup/procedure` using `ProcNames` (current defaults).
- Keep the `p_success, p_error, p_data` contracts and reconnect-on-closed-DB helper.
- Include `resolvespec_login_api_key` (added in mcp_plan item 1) with the generic error behaviour.
### 4. Direct backend
- Port each `*_direct.go` to `lookup/direct` using `Schema` + dialect builder, one store at a time:
AuthStore -> KeyStore -> OAuth stores -> Passkey -> TOTP -> PolicyStore (column/row security tables).
- Add direct `LoginAPIKey` here (select by key hash, active, unexpired, key type in header_api/api, user active;
one generic error), since SQL is now allowed only inside `lookup`.
- Transactions: multi-step writes (login = session insert + last_login; register; reset complete) run in one tx.
### 5. Wire `pkg/security`
- Constructors accept `lookup.Provider` / `lookup.Config`; old options map onto it (deprecated).
- Replace every `*_direct.go` call and `ShouldUseProcedure` branch with a store call.
- Delete `*_direct.go`, `query_mode.go` probe/placeholder code, direct `TableNames` use; keep only aliases.
- Check no non-test code in `pkg/security` contains SQL keywords (CI grep guard).
### 5b. `pkg/security/breaking_changes.md`
- Create at step 0 and append as each step lands: moved types (aliased, no action), moved constructors/types
with rename table (old -> new import path and symbol), removed config types (`SQLNames`, `TableNames`,
`KeyStore*Names`, `QueryMode`) with the `lookup.Config` replacement, removed `database_schema_sqlite.sql`,
`ModeAuto` behaviour change, API key login procedure.
### 6. Schemas and docs
- `lookup/ddl`: reference DDL for postgres (tables only, no procs), sqlite, mysql, mssql; existing proc scripts
stay beside the procedure backend. Replace `database_schema_sqlite.sql`.
- Document: default (procs), Postgres tables-only, SQLite, custom column mapping, adding a dialect.
- Update `pkg/security` README/QUICK_REFERENCE; note API key procedure in security docs (mcp_plan item 10).
### 7. Tests
- Shared conformance suite run against every backend/dialect: login/register/logout/session/refresh, reset,
API keys (valid/expired/inactive/unknown/wrong type), keystore, OAuth server, passkey, TOTP, privilege rules.
- Backends covered: sqlite (in-memory, direct), postgres direct and postgres procedure (needs a Postgres instance,
skipped without `RESOLVESPEC_TEST_PG_DSN`), mysql/mssql behind env DSNs. Dialect unit tests need no DB.
- Procedure backend unit tests with sqlmock for the call/contract shape.
- Check for existing test data before creating any; ask before generating.
- Run with `-race`; migrate current `direct_mode_test.go` / `query_mode_test.go` cases into the suite.
## Order
0. Extract `sectypes` types + aliases (0)
0b. Package split: `totp`, `passkey` -> `providers` -> `oauth` (0b), callers + docs updated
1. Skeleton + Schema/Config (1)
2. Dialects (2)
3. Procedure backend extraction, `pkg/security` wired to it for procs only (3, part of 5) - zero behaviour change
4. Direct backend per store, then remove old `*_direct.go` as each store lands (4, 5)
5. DDL + docs + conformance suite (6, 7)
6. Resume `audit/mcp_plan.md` step 2 (guard) on top of `lookup`
## Breaking changes
- `QueryMode`, `SQLNames`, `TableNames`, `KeyStoreSQLNames`, `KeyStoreTableNames` removed; replaced by `lookup.Config`.
- `ModeAuto` no longer silently falls back from procedure to SQL on non-Postgres drivers without telling: resolution
is explicit and logged once per op.
- `database_schema_sqlite.sql` replaced by `lookup/ddl`.
- Concrete types and constructors move to `providers`, `oauth`, `totp` (e.g. `security.NewDatabaseAuthenticator`
-> `providers.NewDatabaseAuthenticator`, `security.NewOAuthServer` -> `oauth.NewServer`,
`security.NewTOTPGenerator` -> `totp.NewGenerator`). No aliases possible (import cycle); import paths must change.
- Type identity is preserved via aliases; code that used reflection on the package path of these types
(`security.UserContext` -> `sectypes.UserContext`) would see the new path (none found in-repo; re-check).
- Anything outside `lookup` that relied on SQL living in `pkg/security` (none found in-repo) must use the stores.
## Decisions
| # | Topic | Decision |
|---|---|---|
| 1 | Shared package name | `sectypes` |
| 2 | Type aliases in `security` | Kept permanently (public API) |
| 3 | `ColumnSecurity` / `RowSecurity` | Move to `sectypes` (types and their helper logic that has no outside deps) |
| 4 | Moved type names | Renamed for new paths (`oauth.Server`, `totp.Generator`, ...) |
| 5 | OAuth2 client login | Dedicated `oauth.Authenticator` type; `New{Google,GitHub,Microsoft,Facebook}Authenticator` return it |
| 6 | Migration | Rename table only; recorded in new `pkg/security/breaking_changes.md` |
| 7 | Dialects v1 | postgres, sqlite, mysql/mariadb, mssql; more later via the dialect interface |
| 8 | Deprecated config (`SQLNames`, `TableNames`, `KeyStore*Names`, `QueryMode`) | Removed, no aliases; recorded in `breaking_changes.md` |
| 9 | Column mapping | Every column of every entity configurable |
| 10 | DB handle | `*sql.DB` |
| 11 | Column / row security | Procedures (default) **and** a table layout for direct mode |
| 12 | Postgres direct SQL | Standard SQL only: no `ON CONFLICT` / `MERGE`; get-or-create = select then insert in a tx |
| 13 | Token format | Keep `sess_<hex>_<unix>`, generated in Go for direct mode |
## Open
- None.
+16 -6
View File
@@ -10,7 +10,7 @@
- Each un-transacted call takes its own pool connection → bursts with a small pool (see `dbtrace`). - Each un-transacted call takes its own pool connection → bursts with a small pool (see `dbtrace`).
- Already fixed: read/create hooks in `resolvespec` + `restheadspec` (commit `47708fc`, tag >= v1.1.28). Consumers on older tags still show the bug. - Already fixed: read/create hooks in `resolvespec` + `restheadspec` (commit `47708fc`, tag >= v1.1.28). Consumers on older tags still show the bug.
## Current state (verified by reading code; not yet by `dbtrace`) ## Current state — BEFORE this work (historical baseline; everything below is now fixed, see Progress and Status)
| Spec | Read | Create | Update | Delete | | Spec | Read | Create | Update | Delete |
|---|---|---|---|---| |---|---|---|---|---|
| restheadspec | tx; `AfterRead` post-commit on pool | tx; `AfterCreate` post-commit on pool | tx; re-fetch + `BeforeScan` post-commit on pool (`:1667-1674`) | **single: no tx, hook + select + delete on pool (`:1945-1994`)**; batch: tx, per-item `BeforeDelete` inside | | restheadspec | tx; `AfterRead` post-commit on pool | tx; `AfterCreate` post-commit on pool | tx; re-fetch + `BeforeScan` post-commit on pool (`:1667-1674`) | **single: no tx, hook + select + delete on pool (`:1945-1994`)**; batch: tx, per-item `BeforeDelete` inside |
@@ -65,8 +65,18 @@
- Update re-fetch is a plain SELECT in that second tx. No `RETURNING`. - Update re-fetch is a plain SELECT in that second tx. No `RETURNING`.
- `OnTxBegin` failure aborts the whole request, rolls back, returns an error with no detail leaked to the client. - `OnTxBegin` failure aborts the whole request, rolls back, returns an error with no detail leaked to the client.
## Open ## Status summary
- Consumer's ResolveSpec version: confirm it is >= v1.1.28 (read/create already in tx). Not blocking. **Done**
- P0-P7 all DONE (baseline, delete in one tx, `OnTxBegin` + `runInTx`, second short tx, websocketspec/mqttspec, resolvemcp, funcspec, security stamping).
- Regression tests in all six specs, plus source guard `pkg/common/tx_guard_test.go`.
- `AfterRead` decided and done (restheadspec: second short tx; websocketspec/mqttspec: inside the read tx).
- websocketspec `BeforeDisconnect`/`AfterDisconnect` wired (see Progress); `unwiredHooks` allowlist is now empty.
**Not done**
- Real-Postgres `dbtrace` measurement for websocketspec, mqttspec, resolvemcp, restheadspec, funcspec (only resolvespec measured: `pooled=0` on every op). "Done when" bullet 1 is proven for resolvespec only.
- resolvespec batch delete: per-item `BeforeDelete` (one hook per request today). Deferred on purpose: behavior change.
- Confirm the consumer's ResolveSpec version is >= v1.1.28 (read/create already in tx). Not blocking; needs the consumer.
- Known, pre-existing, not ours: `pkg/security` `TestDatabaseAuthenticator` fails with `-count=2` (use `-count=1`); mqttspec integration tests need a DB.
## Phases ## Phases
| # | Status | Change | Files | Notes | | # | Status | Change | Files | Notes |
@@ -88,7 +98,7 @@
- DONE P2 (resolvespec + restheadspec): `common.TxHookName`, `common.TxContext` (`SetTx` only; no abort/context accessors needed since `Execute` already returns an error on abort), `common.RunRequestTx`; per-spec `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Every `RunInTransaction` in both handlers now goes through it. Tests: `pkg/*/on_tx_begin_test.go` (once, first, on tx, failure rolls back). Not yet: the post-commit second tx (P3) and the security stamping registration (P7). - DONE P2 (resolvespec + restheadspec): `common.TxHookName`, `common.TxContext` (`SetTx` only; no abort/context accessors needed since `Execute` already returns an error on abort), `common.RunRequestTx`; per-spec `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Every `RunInTransaction` in both handlers now goes through it. Tests: `pkg/*/on_tx_begin_test.go` (once, first, on tx, failure rolls back). Not yet: the post-commit second tx (P3) and the security stamping registration (P7).
- DONE P3: restheadspec update re-fetch + `BeforeScan` + `AfterUpdate` and `AfterCreate` run in a second short `runInTx`; resolvespec update re-fetches (single, both batch paths) run in a second short `runInTx`. Fixed the pool reads inside the first tx (resolvespec single/batch update existing-record select, restheadspec update existence select) to use `tx`. Tests: `pkg/*/update_tx_test.go` (restheadspec uses the bun adapter; the pgsql adapter cannot build model-based updates). - DONE P3: restheadspec update re-fetch + `BeforeScan` + `AfterUpdate` and `AfterCreate` run in a second short `runInTx`; resolvespec update re-fetches (single, both batch paths) run in a second short `runInTx`. Fixed the pool reads inside the first tx (resolvespec single/batch update existing-record select, restheadspec update existence select) to use `tx`. Tests: `pkg/*/update_tx_test.go` (restheadspec uses the bun adapter; the pgsql adapter cannot build model-based updates).
- NOTE: resolvespec fires no `AfterCreate`/`AfterRead`/`AfterUpdate`-post-commit hooks other than `AfterUpdate` inside the tx; nothing more to move there. - NOTE: resolvespec fires no `AfterCreate`/`AfterRead`/`AfterUpdate`-post-commit hooks other than `AfterUpdate` inside the tx; nothing more to move there.
- OPEN: restheadspec `AfterRead` still runs post-commit with `Tx = h.db` (`:1004`); decision says read has no second tx. Needs a call: run it inside the read tx, or in a short second tx. - RESOLVED: restheadspec `AfterRead` question (see DONE AfterRead below).
- DONE P4: websocketspec + mqttspec. `OnTxBegin` (mqttspec re-exports the websocketspec constant), `HookContext.SetTx`, per-handler `runInTx`/`sendTxError`. Per message: read = 1 tx (Before/After hooks + queries); delete = 1 tx (Before, delete, After); create/update = tx 1 (Before + write) then tx 2 (re-fetch + `BeforeScan` + After). `create()`/`update()` no longer re-fetch; `read*`/`create`/`update`/`delete` use `hookCtx.Tx`. websocketspec `FetchRowNumber` keeps its public signature and delegates to a new tx-aware `fetchRowNumber`. A failure in begin/`OnTxBegin`/commit answers `transaction_error` with no detail. Tests: `pkg/websocketspec/tx_test.go` (sqlmock), `pkg/mqttspec/tx_test.go` (sqlite); mqttspec `update` tests now pass `Tx`. - DONE P4: websocketspec + mqttspec. `OnTxBegin` (mqttspec re-exports the websocketspec constant), `HookContext.SetTx`, per-handler `runInTx`/`sendTxError`. Per message: read = 1 tx (Before/After hooks + queries); delete = 1 tx (Before, delete, After); create/update = tx 1 (Before + write) then tx 2 (re-fetch + `BeforeScan` + After). `create()`/`update()` no longer re-fetch; `read*`/`create`/`update`/`delete` use `hookCtx.Tx`. websocketspec `FetchRowNumber` keeps its public signature and delegates to a new tx-aware `fetchRowNumber`. A failure in begin/`OnTxBegin`/commit answers `transaction_error` with no detail. Tests: `pkg/websocketspec/tx_test.go` (sqlmock), `pkg/mqttspec/tx_test.go` (sqlite); mqttspec `update` tests now pass `Tx`.
- DECIDED in P4 (follow `AfterRead` question above): websocketspec/mqttspec run `AfterRead` inside the read tx (keeps "read has no second tx"). - DECIDED in P4 (follow `AfterRead` question above): websocketspec/mqttspec run `AfterRead` inside the read tx (keeps "read has no second tx").
- DONE P5: resolvemcp. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Read = 1 tx (`BeforeRead`, count, scan, `AfterRead`; `readInTx`). Delete = 1 tx (`BeforeDelete` moved inside, after `OnTxBegin`). Create (single and batch, unified) = tx 1 (`BeforeCreate` + inserts) then tx 2 (re-fetch + `AfterCreate`); the old single-record pool insert/re-fetch is gone. Update = tx 1 (select, `BeforeUpdate`, update, `AfterUpdate`) then tx 2 (re-fetch). `BeforeHandle` still runs before any tx with `Tx = h.db`. Tests: `pkg/resolvemcp/tx_test.go` (sqlmock). - DONE P5: resolvemcp. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Read = 1 tx (`BeforeRead`, count, scan, `AfterRead`; `readInTx`). Delete = 1 tx (`BeforeDelete` moved inside, after `OnTxBegin`). Create (single and batch, unified) = tx 1 (`BeforeCreate` + inserts) then tx 2 (re-fetch + `AfterCreate`); the old single-record pool insert/re-fetch is gone. Update = tx 1 (select, `BeforeUpdate`, update, `AfterUpdate`) then tx 2 (re-fetch). `BeforeHandle` still runs before any tx with `Tx = h.db`. Tests: `pkg/resolvemcp/tx_test.go` (sqlmock).
@@ -103,7 +113,7 @@
## Tests ## Tests
- Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests.
- Done: delete tx tests (`pkg/*/delete_tx_test.go`, sqlmock, 1-conn pool detects pool use). Missing: same for read/create/update, `OnTxBegin`, other specs. - Done: delete tx tests (`pkg/*/delete_tx_test.go`, sqlmock, 1-conn pool detects pool use), and read/create/update/`OnTxBegin`/other-spec tests (see "DONE regression tests" in Progress). Still missing: `dbtrace` `pooled == 0` on real Postgres for all specs except resolvespec.
- Add per spec/op: hook `Tx` is not the pool; `OnTxBegin` fires once per tx, before other hooks; single-ID delete = 1 tx; `dbtrace` `pooled == 0` on the request path. - Add per spec/op: hook `Tx` is not the pool; `OnTxBegin` fires once per tx, before other hooks; single-ID delete = 1 tx; `dbtrace` `pooled == 0` on the request path.
- Test data: reuse `pkg/testmodels`; **ask before generating new data** (per project rule). - Test data: reuse `pkg/testmodels`; **ask before generating new data** (per project rule).
- Regression: full `go test -race` for security, dbmanager, common, restheadspec, resolvespec, websocketspec, mqttspec, resolvemcp, funcspec. Known pre-existing failures: mqttspec integration (no DB). - Regression: full `go test -race` for security, dbmanager, common, restheadspec, resolvespec, websocketspec, mqttspec, resolvemcp, funcspec. Known pre-existing failures: mqttspec integration (no DB).
@@ -118,5 +128,5 @@
- `dbtrace` shows `pooled=0` for every handler op on a hooked model. - `dbtrace` shows `pooled=0` for every handler op on a hooked model.
- RLS GUC set in `OnTxBegin` is visible to read, create, update, delete queries and hooks. - RLS GUC set in `OnTxBegin` is visible to read, create, update, delete queries and hooks.
- No `Tx: h.db` / `hookCtx.Tx = h.db` left in spec handlers. - No `Tx: h.db` / `hookCtx.Tx = h.db` left in spec handlers.
- OPEN: websocketspec `BeforeDisconnect`/`AfterDisconnect` are defined but never executed (connection lifecycle, not DB). Allowlisted in `TestEveryDefinedHookHasACallSite`; wire them to remove the entry. - DONE: websocketspec `BeforeDisconnect`/`AfterDisconnect` fire from `Connection.Close()` (single close path, `closedOnce`, so exactly once per registered connection however it closes: read error, write error, slow-consumer eviction, shutdown). Connection lifecycle, not DB: `Tx` is not set. The hook context is detached from the connection cancel (`context.WithoutCancel`) so `AfterDisconnect` still has a live context. Errors are logged and never block the close. A connection rejected by `BeforeConnect` gets no disconnect hooks. `ConnectionManager.Shutdown` now closes connections outside its lock (a hook calling `Count()` would have deadlocked). Allowlist in `TestEveryDefinedHookHasACallSite` is now empty. Tests: `pkg/websocketspec/connection_test.go`. mqttspec already fired these in `Handler.Shutdown`; its per-client disconnect is unchanged.
- DONE: column-level hide/mask columns are dropped from create/update payloads (`security.ApplyWriteColumnSecurity`); rules preloaded in `BeforeHandle` for create/update. resolvemcp update now runs `BeforeHandle`. - DONE: column-level hide/mask columns are dropped from create/update payloads (`security.ApplyWriteColumnSecurity`); rules preloaded in `BeforeHandle` for create/update. resolvemcp update now runs `BeforeHandle`.
+5
View File
@@ -0,0 +1,5 @@
# Changelog
## 0.1.0
- Initial release: ResolveSpec (JSON body) and FunctionSpec client.
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2026 Hein
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+1
View File
@@ -1,6 +1,7 @@
name: resolvespec 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 -1
View File
@@ -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 -1
View File
@@ -1,3 +1,3 @@
module github.com/bitechdev/ResolveSpec/clients/resolvespec-go module git.warky.dev/wdevs/ResolveSpec/clients/resolvespec-go
go 1.22 go 1.22
+2 -2
View File
@@ -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
+58
View File
@@ -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"
@@ -275,6 +276,14 @@ func (b *BunAdapter) GetUnderlyingDB() interface{} {
return b.getDB() return b.getDB()
} }
// SQLDB implements common.SQLDBProvider.
func (b *BunAdapter) SQLDB() *sql.DB {
if db := b.getDB(); db != nil {
return db.DB
}
return nil
}
func (b *BunAdapter) DriverName() string { func (b *BunAdapter) DriverName() string {
// Normalize Bun's dialect name to match the project's canonical vocabulary. // Normalize Bun's dialect name to match the project's canonical vocabulary.
// Bun returns "pg" for PostgreSQL; the rest of the project uses "postgres". // Bun returns "pg" for PostgreSQL; the rest of the project uses "postgres".
@@ -529,6 +538,19 @@ func (b *BunSelectQuery) WhereOr(query string, args ...interface{}) common.Selec
return b return b
} }
// WhereGroup wraps the conditions added by fn in one parenthesised group ANDed with the rest.
func (b *BunSelectQuery) WhereGroup(fn func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
b.query = b.query.WhereGroup(" AND ", func(q *bun.SelectQuery) *bun.SelectQuery {
inner := *b
inner.query = q
if res, ok := fn(&inner).(*BunSelectQuery); ok {
return res.query
}
return q
})
return b
}
func (b *BunSelectQuery) Join(query string, args ...interface{}) common.SelectQuery { func (b *BunSelectQuery) Join(query string, args ...interface{}) common.SelectQuery {
// Extract optional prefix from args // Extract optional prefix from args
// If the last arg is a string that looks like a table prefix, use it // If the last arg is a string that looks like a table prefix, use it
@@ -1486,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, ", "))
@@ -1598,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
}
+39
View File
@@ -2,6 +2,7 @@ package database
import ( import (
"context" "context"
"database/sql"
"fmt" "fmt"
"reflect" "reflect"
"strings" "strings"
@@ -227,6 +228,20 @@ func (g *GormAdapter) GetUnderlyingDB() interface{} {
return g.getDB() return g.getDB()
} }
// SQLDB implements common.SQLDBProvider. It returns nil when GORM has no *sql.DB
// (for example a ConnPool that is not database/sql).
func (g *GormAdapter) SQLDB() *sql.DB {
db := g.getDB()
if db == nil {
return nil
}
sqlDB, err := db.DB()
if err != nil {
return nil
}
return sqlDB
}
func (g *GormAdapter) DriverName() string { func (g *GormAdapter) DriverName() string {
return normalizeGormDriverName(g.getDB()) return normalizeGormDriverName(g.getDB())
} }
@@ -362,6 +377,16 @@ func (g *GormSelectQuery) WhereOr(query string, args ...interface{}) common.Sele
return g return g
} }
// WhereGroup wraps the conditions added by fn in one parenthesised group ANDed with the rest.
func (g *GormSelectQuery) WhereGroup(fn func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
inner := *g
inner.db = g.db.Session(&gorm.Session{NewDB: true})
if res, ok := fn(&inner).(*GormSelectQuery); ok {
g.db = g.db.Where(res.db)
}
return g
}
func (g *GormSelectQuery) Join(query string, args ...interface{}) common.SelectQuery { func (g *GormSelectQuery) Join(query string, args ...interface{}) common.SelectQuery {
// Extract optional prefix from args // Extract optional prefix from args
// If the last arg is a string that looks like a table prefix, use it // If the last arg is a string that looks like a table prefix, use it
@@ -726,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
@@ -905,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
@@ -0,0 +1,230 @@
package database
import (
"bytes"
"context"
"database/sql"
"fmt"
"net"
"os"
"os/exec"
"strings"
"testing"
"time"
_ "github.com/jackc/pgx/v5/stdlib"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/pgdialect"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/config"
)
// These tests start a throwaway PostgreSQL server with podman or docker (whichever is
// installed, podman first) and run the client-SQL hardening against a real database. They
// pull an image, so they only run when RESOLVESPEC_TEST_CONTAINERS=1 and not with -short.
// The container is removed when the test ends.
const hardeningPGPassword = "Resolve_Spec_1"
type hardeningItem struct {
bun.BaseModel `bun:"table:items,alias:items"`
ID int `bun:"id"`
Tenant int `bun:"tenant"`
Name string `bun:"name"`
}
func hardeningRuntime(t *testing.T) string {
t.Helper()
if testing.Short() {
t.Skip("container tests are skipped with -short")
}
if os.Getenv("RESOLVESPEC_TEST_CONTAINERS") != "1" {
t.Skip("set RESOLVESPEC_TEST_CONTAINERS=1 to run tests that start a podman/docker container")
}
for _, rt := range []string{"podman", "docker"} {
if p, err := exec.LookPath(rt); err == nil {
return p
}
}
t.Skip("neither podman nor docker found in PATH")
return ""
}
func hardeningRun(t *testing.T, timeout time.Duration, name string, args ...string) string {
t.Helper()
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
var out, errb bytes.Buffer
cmd := exec.CommandContext(ctx, name, args...)
cmd.Stdout, cmd.Stderr = &out, &errb
if err := cmd.Run(); err != nil {
t.Fatalf("%s %s: %v\n%s", name, strings.Join(args, " "), err, errb.String())
}
return strings.TrimSpace(out.String())
}
// startHardeningPostgres runs postgres on a random localhost port and returns a ready *sql.DB.
func startHardeningPostgres(t *testing.T, rt string) *sql.DB {
t.Helper()
id := hardeningRun(t, 10*time.Minute, rt, "run", "-d", "--rm", "-p", "127.0.0.1::5432",
"-e", "POSTGRES_PASSWORD="+hardeningPGPassword, "docker.io/library/postgres:16-alpine") // first run may pull
t.Cleanup(func() { _ = exec.Command(rt, "rm", "-f", id).Run() })
out := hardeningRun(t, 30*time.Second, rt, "port", id, "5432")
line := strings.Fields(out)[len(strings.Fields(out))-1]
for _, l := range strings.Split(out, "\n") {
if strings.Contains(l, "127.0.0.1:") {
line = l[strings.LastIndex(l, " ")+1:]
break
}
}
_, port, err := net.SplitHostPort(line)
if err != nil {
t.Fatalf("cannot parse published port %q: %v", out, err)
}
dsn := fmt.Sprintf("postgres://postgres:%s@127.0.0.1:%s/postgres?sslmode=disable", hardeningPGPassword, port)
// The official image restarts once during init: wait, pause, wait again.
wait := func(d time.Duration) *sql.DB {
deadline := time.Now().Add(d)
var last error
for time.Now().Before(deadline) {
db, err := sql.Open("pgx", dsn)
if err == nil {
if last = db.Ping(); last == nil {
return db
}
_ = db.Close()
} else {
last = err
}
time.Sleep(time.Second)
}
t.Fatalf("postgres not ready within %s: %v", d, last)
return nil
}
_ = wait(90 * time.Second).Close()
time.Sleep(2 * time.Second)
db := wait(60 * time.Second)
t.Cleanup(func() { _ = db.Close() })
return db
}
// setHardeningConfig overrides the global hardening switches for the test.
func setHardeningConfig(t *testing.T, h config.HardeningConfig) {
t.Helper()
m := config.GetConfigManager()
cfg, err := m.GetConfig()
if err != nil {
t.Fatal(err)
}
old := cfg.Hardening
cfg.Hardening = h
if err := m.SetConfig(cfg); err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
cfg.Hardening = old
_ = m.SetConfig(cfg)
})
}
// clientWhere mirrors the restheadspec handler pipeline for x-custom-sql-w.
func clientWhere(raw string) string {
w := common.AddTablePrefixToColumns(raw, "items")
w = common.SanitizeWhereClause(w, "items")
return common.EnsureOuterParentheses(w)
}
func TestHardeningAgainstPostgresContainer(t *testing.T) {
rt := hardeningRuntime(t)
sqldb := startHardeningPostgres(t, rt)
if _, err := sqldb.Exec(`
CREATE TABLE items (id int PRIMARY KEY, tenant int, name text);
INSERT INTO items VALUES (1,5,'mine-a'),(2,5,'mine-b'),(3,6,'other-a'),(4,6,'awaiting update approval');`); err != nil {
t.Fatal(err)
}
bdb := bun.NewDB(sqldb, pgdialect.New())
adapter := NewBunAdapter(bdb)
ctx := context.Background()
// list runs the handler-shaped query: client x-custom-sql-w, then the server tenant filter.
list := func(t *testing.T, where string) []hardeningItem {
t.Helper()
var rows []hardeningItem
q := adapter.NewSelect().Model(&rows)
if w := clientWhere(where); w != "" {
q = q.Where(w)
}
q = q.Where("items.tenant = ?", 5)
if err := q.Scan(ctx, &rows); err != nil {
t.Fatalf("query failed for %q: %v", where, err)
}
return rows
}
t.Run("strict", func(t *testing.T) {
setHardeningConfig(t, config.HardeningConfig{CORSStrictOrigins: true, SortStrict: true, SQLStrict: true})
t.Run("legitimate filters keep working", func(t *testing.T) {
if got := list(t, "name = 'mine-a'"); len(got) != 1 || got[0].ID != 1 {
t.Errorf("simple filter: %v", got)
}
if got := list(t, "id in (select id from items where tenant = 5)"); len(got) != 2 {
t.Errorf("subquery filter: %v", got)
}
})
t.Run("parenthesis escape cannot leave the tenant", func(t *testing.T) {
if got := list(t, "1=1)) OR ((1=1"); len(got) != 0 {
t.Errorf("escape not rejected closed, got rows: %v", got)
}
})
t.Run("hostile fragments fail closed", func(t *testing.T) {
for _, w := range []string{
"id = 1 and pg_sleep(10) is not null",
"id = 1 or (select count(*) from pg_shadow) > 0",
"id = 1; delete/**/from items",
} {
start := time.Now()
if got := list(t, w); len(got) != 0 {
t.Errorf("%q returned rows: %v", w, got)
}
if time.Since(start) > 5*time.Second {
t.Errorf("%q was executed (took %s)", w, time.Since(start))
}
}
var n int
if err := sqldb.QueryRow("SELECT count(*) FROM items").Scan(&n); err != nil || n != 4 {
t.Errorf("items table modified: count=%d err=%v", n, err)
}
})
t.Run("x-custom-sql-or stays inside the tenant", func(t *testing.T) {
var rows []hardeningItem
q := adapter.NewSelect().Model(&rows)
orClause := common.EnsureOuterParentheses(common.SanitizeWhereClause("items.name = 'other-a'", "items"))
q = q.(common.WhereGrouper).WhereGroup(func(g common.SelectQuery) common.SelectQuery {
return g.Where("items.name = ?", "mine-a").WhereOr(orClause)
})
q = q.Where("items.tenant = ?", 5)
if err := q.Scan(ctx, &rows); err != nil {
t.Fatal(err)
}
if len(rows) != 1 || rows[0].ID != 1 {
t.Errorf("OR clause leaked outside tenant filter: %v", rows)
}
})
})
t.Run("switch off restores legacy behaviour", func(t *testing.T) {
setHardeningConfig(t, config.HardeningConfig{})
// Proves the strict assertions above are meaningful: without hardening the same
// escape returns the other tenant's rows.
if got := list(t, "1=1)) OR ((1=1"); len(got) < 3 {
t.Errorf("expected the legacy escape to leak rows, got %v", got)
}
})
}
+58 -14
View File
@@ -5,7 +5,9 @@ import (
"database/sql" "database/sql"
"fmt" "fmt"
"reflect" "reflect"
"regexp"
"sort" "sort"
"strconv"
"strings" "strings"
"sync" "sync"
"time" "time"
@@ -223,6 +225,11 @@ func (p *PgSQLAdapter) GetUnderlyingDB() interface{} {
return p.db return p.db
} }
// SQLDB implements common.SQLDBProvider.
func (p *PgSQLAdapter) SQLDB() *sql.DB {
return p.db
}
func (p *PgSQLAdapter) DriverName() string { func (p *PgSQLAdapter) DriverName() string {
return p.driverName return p.driverName
} }
@@ -318,6 +325,32 @@ func (p *PgSQLSelectQuery) WhereOr(query string, args ...interface{}) common.Sel
return p return p
} }
// WhereGroup wraps the conditions added by fn (Where = AND, WhereOr = OR) in one
// parenthesised group ANDed with the rest of the query. Inside the group the
// semantics match Bun: `w1 AND w2 OR o1 OR o2`.
func (p *PgSQLSelectQuery) WhereGroup(fn func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
sub := &PgSQLSelectQuery{driverName: p.driverName, paramCounter: p.paramCounter, args: make([]interface{}, 0)}
res, ok := fn(sub).(*PgSQLSelectQuery)
if !ok {
res = sub
}
var group string
switch {
case len(res.whereClauses) > 0 && len(res.orClauses) > 0:
group = "(" + strings.Join(res.whereClauses, " AND ") + ") OR " + strings.Join(res.orClauses, " OR ")
case len(res.whereClauses) > 0:
group = strings.Join(res.whereClauses, " AND ")
case len(res.orClauses) > 0:
group = strings.Join(res.orClauses, " OR ")
default:
return p
}
p.whereClauses = append(p.whereClauses, "("+group+")")
p.args = append(p.args, res.args...)
p.paramCounter = res.paramCounter
return p
}
func (p *PgSQLSelectQuery) Join(query string, args ...interface{}) common.SelectQuery { func (p *PgSQLSelectQuery) Join(query string, args ...interface{}) common.SelectQuery {
query = p.replacePlaceholders(query, len(args)) query = p.replacePlaceholders(query, len(args))
p.joins = append(p.joins, "JOIN "+query) p.joins = append(p.joins, "JOIN "+query)
@@ -658,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
@@ -762,6 +802,9 @@ func (p *PgSQLInsertQuery) Scan(ctx context.Context, dest interface{}) (err erro
return nil return nil
} }
// placeholderRe matches a numbered SQL parameter such as $12.
var placeholderRe = regexp.MustCompile(`\$\d+`)
// PgSQLUpdateQuery implements UpdateQuery for PostgreSQL // PgSQLUpdateQuery implements UpdateQuery for PostgreSQL
type PgSQLUpdateQuery struct { type PgSQLUpdateQuery struct {
db *sql.DB db *sql.DB
@@ -814,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 {
@@ -892,23 +942,17 @@ func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err err
p.tableName, p.tableName,
strings.Join(setClauses, ", ")) strings.Join(setClauses, ", "))
// Update WHERE clause parameter numbers to continue after SET parameters // WHERE placeholders were numbered from $1 as the clauses were added; shift every one past
// the SET parameters in a single pass (replacing one number at a time would rewrite a
// number it had just produced, e.g. "$1, $2" -> "$3, $2").
if len(p.whereClauses) > 0 { if len(p.whereClauses) > 0 {
shift := len(setArgs)
updatedWhereClauses := make([]string, 0, len(p.whereClauses)) updatedWhereClauses := make([]string, 0, len(p.whereClauses))
for _, whereClause := range p.whereClauses { for _, whereClause := range p.whereClauses {
// Find and replace parameter placeholders updatedWhereClauses = append(updatedWhereClauses, placeholderRe.ReplaceAllStringFunc(whereClause, func(m string) string {
updatedClause := whereClause n, _ := strconv.Atoi(m[1:])
paramNum := i return fmt.Sprintf("$%d", n+shift)
// Count how many parameters are in this WHERE clause }))
placeholderCount := strings.Count(whereClause, "$")
for j := 0; j < placeholderCount; j++ {
oldParam := fmt.Sprintf("$%d", j+1)
newParam := fmt.Sprintf("$%d", paramNum)
updatedClause = strings.Replace(updatedClause, oldParam, newParam, 1)
paramNum++
}
updatedWhereClauses = append(updatedWhereClauses, updatedClause)
i = paramNum
} }
p.whereClauses = updatedWhereClauses p.whereClauses = updatedWhereClauses
} }
@@ -627,3 +627,23 @@ func TestRawSQL(t *testing.T) {
assert.NoError(t, mock.ExpectationsWereMet()) assert.NoError(t, mock.ExpectationsWereMet())
} }
// WHERE placeholders must be shifted past the SET parameters without rewriting numbers the
// shift itself produced: "a = ? AND b = ?" must stay in order.
func TestPgSQLUpdateQuery_WherePlaceholdersAfterSet(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer db.Close()
mock.ExpectExec(`UPDATE users SET name = \$1 WHERE a = \$2 AND b = \$3 AND "id" IN \(\$4, \$5\)`).
WithArgs("n", 10, 20, 7, 8).
WillReturnResult(sqlmock.NewResult(0, 2))
adapter := NewPgSQLAdapter(db)
_, err = adapter.NewUpdate().Table("users").SetMap(map[string]interface{}{"name": "n"}).
Where("a = ? AND b = ?", 10, 20).
Where(`"id" IN (?, ?)`, 7, 8).
Exec(context.Background())
require.NoError(t, err)
require.NoError(t, mock.ExpectationsWereMet())
}
@@ -0,0 +1,50 @@
package database
import (
"database/sql"
"strings"
"testing"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/pgdialect"
"github.com/bitechdev/ResolveSpec/pkg/common"
)
// The client's OR must widen only the client's own conditions, never the
// server-side condition ANDed after it.
func TestBunWhereGroupConfinesOr(t *testing.T) {
db := bun.NewDB(&sql.DB{}, pgdialect.New())
q := &BunSelectQuery{query: db.NewSelect().TableExpr("items"), db: db, driverName: "postgres"}
var sq common.SelectQuery = q.Where("a = 1")
sq = sq.(common.WhereGrouper).WhereGroup(func(g common.SelectQuery) common.SelectQuery {
return g.Where("b = 2").WhereOr("(c = 3)")
})
sq = sq.Where("tenant = 5")
got := sq.(*BunSelectQuery).query.String()
want := `WHERE (a = 1) AND ((b = 2) OR ((c = 3))) AND (tenant = 5)`
if !strings.Contains(got, want) {
t.Fatalf("unexpected SQL:\n got: %s\nwant to contain: %s", got, want)
}
}
func TestPgSQLWhereGroupConfinesOr(t *testing.T) {
var q common.SelectQuery = &PgSQLSelectQuery{driverName: "postgres", tableName: "items", columns: []string{"*"}, args: []interface{}{}}
q = q.Where("a = ?", 1)
q = q.(common.WhereGrouper).WhereGroup(func(g common.SelectQuery) common.SelectQuery {
return g.Where("b = ?", 2).WhereOr("(c = 3)")
})
q = q.Where("tenant = ?", 5)
pq := q.(*PgSQLSelectQuery)
got := pq.buildSQL()
want := `WHERE (a = $1 AND ((b = $2) OR (c = 3)) AND tenant = $3)`
if !strings.Contains(got, want) {
t.Fatalf("unexpected SQL:\n got: %s\nwant to contain: %s", got, want)
}
if len(pq.args) != 3 || pq.args[0] != 1 || pq.args[1] != 2 || pq.args[2] != 5 {
t.Fatalf("args out of order: %v", pq.args)
}
}
+85 -12
View File
@@ -20,7 +20,15 @@ func DefaultCORSConfig() CORSConfig {
configManager := config.GetConfigManager() configManager := config.GetConfigManager()
cfg, _ := configManager.GetConfig() cfg, _ := configManager.GetConfig()
hosts := make([]string, 0) hosts := make([]string, 0)
// hosts = append(hosts, "*") if cfg == nil {
return CORSConfig{
AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"},
AllowedHeaders: GetHeadSpecHeaders(),
MaxAge: 86400,
}
}
// Explicitly configured origins (cors.allowed_origins); "*" allows any origin without credentials
hosts = append(hosts, cfg.CORS.AllowedOrigins...)
_, _, ipsList := config.GetIPs() _, _, ipsList := config.GetIPs()
@@ -113,25 +121,63 @@ func GetHeadSpecHeaders() []string {
} }
} }
// SetCORSHeaders sets CORS headers on a response writer // originAllowed reports whether origin matches config.AllowedOrigins exactly
// (case-insensitive, trailing slash ignored). wildcard is true when the list
// contains "*".
func originAllowed(origin string, allowed []string) (ok bool, wildcard bool) {
norm := strings.ToLower(strings.TrimRight(strings.TrimSpace(origin), "/"))
for _, a := range allowed {
a = strings.ToLower(strings.TrimRight(strings.TrimSpace(a), "/"))
if a == "*" {
wildcard = true
continue
}
if a != "" && a == norm {
return true, wildcard
}
}
return false, wildcard
}
// SetCORSHeaders sets CORS headers on a response writer.
//
// The request Origin is only reflected (and credentials only allowed) when it
// is listed in config.AllowedOrigins. A "*" entry allows any origin but never
// with credentials. Unlisted origins get no CORS headers, so browsers block
// the cross-origin read.
func SetCORSHeaders(w ResponseWriter, r Request, config CORSConfig) { func SetCORSHeaders(w ResponseWriter, r Request, config CORSConfig) {
// Reflect the request origin; fall back to wildcard only when no origin is present if !Hardening().CORSStrictOrigins {
setCORSHeadersLegacy(w, r, config)
return
}
origin := r.Header("Origin") origin := r.Header("Origin")
if origin == "" { if origin == "" {
origin = "*" // Not a cross-origin browser request; nothing to protect.
w.SetHeader("Access-Control-Allow-Origin", "*")
} else { } else {
// Vary must be set so caches don't serve one origin's response to another // Vary must be set so caches don't serve one origin's response to another
httpW := w.UnderlyingResponseWriter() w.UnderlyingResponseWriter().Header().Set("Vary", "Origin")
httpW.Header().Set("Vary", "Origin")
ok, wildcard := originAllowed(origin, config.AllowedOrigins)
switch {
case ok:
w.SetHeader("Access-Control-Allow-Origin", origin)
w.SetHeader("Access-Control-Allow-Credentials", "true")
case wildcard:
w.SetHeader("Access-Control-Allow-Origin", "*")
default:
return
}
} }
w.SetHeader("Access-Control-Allow-Origin", origin)
// Set allowed methods // Set allowed methods
if len(config.AllowedMethods) > 0 { if len(config.AllowedMethods) > 0 {
w.SetHeader("Access-Control-Allow-Methods", strings.Join(config.AllowedMethods, ", ")) w.SetHeader("Access-Control-Allow-Methods", strings.Join(config.AllowedMethods, ", "))
} }
// Reflect the preflight request headers when present; otherwise use the explicit config list // The origin is trusted at this point, so reflecting the preflight request
// headers is safe (the config list contains "X-Foo-*" patterns that browsers
// cannot match literally).
requestedHeaders := r.Header("Access-Control-Request-Headers") requestedHeaders := r.Header("Access-Control-Request-Headers")
if requestedHeaders != "" { if requestedHeaders != "" {
w.SetHeader("Access-Control-Allow-Headers", requestedHeaders) w.SetHeader("Access-Control-Allow-Headers", requestedHeaders)
@@ -144,13 +190,40 @@ func SetCORSHeaders(w ResponseWriter, r Request, config CORSConfig) {
w.SetHeader("Access-Control-Max-Age", fmt.Sprintf("%d", config.MaxAge)) w.SetHeader("Access-Control-Max-Age", fmt.Sprintf("%d", config.MaxAge))
} }
// Allow credentials only when a specific origin is reflected (not wildcard) // Expose headers that clients can read (fresh slice: avoid appending into
// config.AllowedHeaders' backing array)
exposeHeaders := make([]string, 0, len(config.AllowedHeaders)+3)
exposeHeaders = append(exposeHeaders, config.AllowedHeaders...)
exposeHeaders = append(exposeHeaders, "Content-Range", "X-Api-Range-Total", "X-Api-Range-Size")
w.SetHeader("Access-Control-Expose-Headers", strings.Join(exposeHeaders, ", "))
}
// setCORSHeadersLegacy is the pre-hardening behaviour (reflect any origin with
// credentials). Used only when hardening.cors_strict_origins is false.
func setCORSHeadersLegacy(w ResponseWriter, r Request, config CORSConfig) {
origin := r.Header("Origin")
if origin == "" {
origin = "*"
} else {
w.UnderlyingResponseWriter().Header().Set("Vary", "Origin")
}
w.SetHeader("Access-Control-Allow-Origin", origin)
if len(config.AllowedMethods) > 0 {
w.SetHeader("Access-Control-Allow-Methods", strings.Join(config.AllowedMethods, ", "))
}
if requested := r.Header("Access-Control-Request-Headers"); requested != "" {
w.SetHeader("Access-Control-Allow-Headers", requested)
} else if len(config.AllowedHeaders) > 0 {
w.SetHeader("Access-Control-Allow-Headers", strings.Join(config.AllowedHeaders, ", "))
}
if config.MaxAge > 0 {
w.SetHeader("Access-Control-Max-Age", fmt.Sprintf("%d", config.MaxAge))
}
if origin != "*" { if origin != "*" {
w.SetHeader("Access-Control-Allow-Credentials", "true") w.SetHeader("Access-Control-Allow-Credentials", "true")
} }
exposeHeaders := make([]string, 0, len(config.AllowedHeaders)+3)
// Expose headers that clients can read exposeHeaders = append(exposeHeaders, config.AllowedHeaders...)
exposeHeaders := config.AllowedHeaders
exposeHeaders = append(exposeHeaders, "Content-Range", "X-Api-Range-Total", "X-Api-Range-Size") exposeHeaders = append(exposeHeaders, "Content-Range", "X-Api-Range-Total", "X-Api-Range-Size")
w.SetHeader("Access-Control-Expose-Headers", strings.Join(exposeHeaders, ", ")) w.SetHeader("Access-Control-Expose-Headers", strings.Join(exposeHeaders, ", "))
} }
+16
View File
@@ -0,0 +1,16 @@
package common
import "github.com/bitechdev/ResolveSpec/pkg/config"
// hardeningProvider returns the active hardening switches. Tests may replace it.
var hardeningProvider = func() config.HardeningConfig {
cfg, err := config.GetConfigManager().GetConfig()
if err != nil || cfg == nil {
// Fail secure: hardening on when config is unavailable.
return config.HardeningConfig{CORSStrictOrigins: true, SortStrict: true, SQLStrict: true}
}
return cfg.Hardening
}
// Hardening returns the security hardening toggles (config section "hardening").
func Hardening() config.HardeningConfig { return hardeningProvider() }
+121
View File
@@ -0,0 +1,121 @@
package common
import (
"testing"
"github.com/bitechdev/ResolveSpec/pkg/config"
)
func setHardening(t *testing.T, h config.HardeningConfig) {
t.Helper()
old := hardeningProvider
hardeningProvider = func() config.HardeningConfig { return h }
t.Cleanup(func() { hardeningProvider = old })
}
var (
hardOn = config.HardeningConfig{CORSStrictOrigins: true, SortStrict: true, SQLStrict: true}
hardOff = config.HardeningConfig{}
)
func TestSanitizeWhereClause_Strict(t *testing.T) {
setHardening(t, hardOn)
allowed := []string{
"status = 'awaiting update approval'",
"last_update > '2020-01-01'",
"name = 'it''s; fine'",
"(a = 1 or b = 2) and ifblnk(c) = 'x'",
}
for _, w := range allowed {
if got := SanitizeWhereClause(w, ""); got == "" || got == "(1=0)" {
t.Errorf("legitimate clause %q rejected: %q", w, got)
}
}
hostile := []string{
"1=1)) OR ((1=1",
"id = 1 or (select count(*) from pg_shadow) > 0",
"id in (select oid from pg_catalog.pg_class)",
"id = 1 and pg_sleep(5) is not null",
"id = 1; delete/**/from items",
"id = 1 -- x",
"name = 'unterminated",
}
for _, w := range hostile {
if got := SanitizeWhereClause(w, "t"); got != "(1=0)" {
t.Errorf("hostile clause %q not rejected closed: %q", w, got)
}
}
}
func TestSanitizeWhereClause_Subqueries(t *testing.T) {
q := "id in (select l.id from other l where l.x = 1)"
setHardening(t, hardOn)
if got := SanitizeWhereClause(q, "t"); got == "(1=0)" || got == "" {
t.Errorf("subquery should be allowed by default: %q", got)
}
h := hardOn
h.SQLBlockSubqueries = true
setHardening(t, h)
if got := SanitizeWhereClause(q, "t"); got != "(1=0)" {
t.Errorf("subquery should be blocked with sql_block_subqueries: %q", got)
}
}
func TestSanitizeWhereClause_StrictOff_LegacyBehaviour(t *testing.T) {
setHardening(t, hardOff)
if got := SanitizeWhereClause("a = 1 and drop table x", "t"); got != "" {
t.Errorf("legacy fail-open expected empty, got %q", got)
}
}
func TestSortStrict(t *testing.T) {
v := NewColumnValidator(TestModel{})
opts := RequestOptions{
JoinAliases: []string{""},
Sort: []SortOption{{Column: "x.id, (select pg_sleep(10))"}, {Column: "(select pg_sleep(1))"}},
}
setHardening(t, hardOn)
if got := v.FilterRequestOptions(opts).Sort; len(got) != 0 {
t.Errorf("strict: expected no sorts, got %v", got)
}
opts.JoinAliases = []string{"j"}
opts.Sort = []SortOption{{Column: "j.id"}, {Column: "j.id, (select 1)"}, {Column: "(select max(age) from users)"}}
if got := v.FilterRequestOptions(opts).Sort; len(got) != 2 {
t.Errorf("strict: join column and plain subquery sort must be allowed, got %v", got)
}
opts.Sort = []SortOption{{Column: "j.id"}, {Column: "j.id, (select 1)"}}
if got := v.FilterRequestOptions(opts).Sort; len(got) != 1 || got[0].Column != "j.id" {
t.Errorf("strict: expected only j.id, got %v", got)
}
setHardening(t, hardOff)
opts.JoinAliases = []string{""}
opts.Sort = []SortOption{{Column: "x.id, (select pg_sleep(10))"}}
if got := v.FilterRequestOptions(opts).Sort; len(got) != 1 {
t.Errorf("legacy: expected sort kept, got %v", got)
}
}
func TestCQLColumn(t *testing.T) {
v := NewColumnValidator(TestModel{})
setHardening(t, hardOn)
if !v.IsValidColumn("cqlComputed1") || v.IsValidColumn("cql1); drop") {
t.Error("strict cql validation wrong")
}
setHardening(t, hardOff)
if !v.IsValidColumn("cql1); drop") {
t.Error("legacy cql behaviour should be permissive")
}
}
func TestOriginAllowed(t *testing.T) {
list := []string{"https://app.example.com/", "http://localhost:8080"}
if ok, _ := originAllowed("https://APP.example.com", list); !ok {
t.Error("listed origin should match case-insensitively, ignoring trailing slash")
}
if ok, _ := originAllowed("https://evil.example", list); ok {
t.Error("unlisted origin must not match")
}
if ok, wc := originAllowed("https://evil.example", []string{"*"}); ok || !wc {
t.Error("wildcard must be reported separately and never as an exact match")
}
}
+22
View File
@@ -2,6 +2,7 @@ package common
import ( import (
"context" "context"
"database/sql"
"encoding/json" "encoding/json"
"io" "io"
"net/http" "net/http"
@@ -38,6 +39,15 @@ type Database interface {
DriverName() string DriverName() string
} }
// SQLDBProvider is implemented by adapters that wrap a *sql.DB (directly or through an ORM).
// It lets packages that work on database/sql, such as pkg/security/lookup, reuse the
// connection an application already configured. Transaction adapters do not implement it.
type SQLDBProvider interface {
// SQLDB returns the current underlying *sql.DB. After an adapter reconnects it
// returns the new handle, so do not cache it across reconnects.
SQLDB() *sql.DB
}
// SelectQuery interface for building SELECT queries (compatible with both GORM and Bun) // SelectQuery interface for building SELECT queries (compatible with both GORM and Bun)
type SelectQuery interface { type SelectQuery interface {
Model(model interface{}) SelectQuery Model(model interface{}) SelectQuery
@@ -71,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
@@ -84,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
@@ -309,3 +323,11 @@ type QueryHandler interface {
SpecHandler SpecHandler
// Methods are defined in funcspec package due to different function signature requirements // Methods are defined in funcspec package due to different function signature requirements
} }
// WhereGrouper is implemented by query builders that can wrap a set of
// conditions (including WhereOr) in one parenthesised group that is ANDed with
// the rest of the query. It is optional so existing SelectQuery implementations
// keep compiling.
type WhereGrouper interface {
WhereGroup(fn func(SelectQuery) SelectQuery) SelectQuery
}
+6 -2
View File
@@ -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)
+2
View File
@@ -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
+95 -1
View File
@@ -148,6 +148,93 @@ func validateWhereClauseSecurity(where string) error {
return nil return nil
} }
var (
reStrictDML = regexp.MustCompile(`(?i)\b(delete|update|truncate|drop|alter|create|insert|grant|revoke|exec|execute|copy|call|do|merge|vacuum|listen|notify|set|returning|into)\b`)
// reStrictSubquery is only enforced when hardening.sql_block_subqueries is on.
reStrictSubquery = regexp.MustCompile(`(?i)\b(select|union|lateral|with)\b`)
// reStrictDangerousFunc matches functions/schemas that enable DoS or data
// exfiltration through a WHERE fragment (sleep, file/large-object access,
// dblink, config access, catalogs).
reStrictDangerousFunc = regexp.MustCompile(`(?i)\b(pg_[a-z0-9_]*|lo_[a-z0-9_]*|dblink[a-z0-9_]*|set_config|current_setting|query_to_xml[a-z_]*|xpath[a-z_]*|generate_series|repeat|crypt|information_schema|sleep|benchmark)\b`)
)
// stripSQLLiterals blanks out single-quoted literals (honouring ”) and
// double-quoted identifiers so structural checks only see SQL syntax.
// ok is false when a quote is left unterminated.
func stripSQLLiterals(s string) (out string, ok bool) {
var b strings.Builder
for i := 0; i < len(s); i++ {
ch := s[i]
if ch != '\'' && ch != '"' {
b.WriteByte(ch)
continue
}
q := ch
closed := false
for i++; i < len(s); i++ {
if s[i] == q {
if i+1 < len(s) && s[i+1] == q { // escaped quote
i++
continue
}
closed = true
break
}
}
if !closed {
return "", false
}
b.WriteString("''")
}
return b.String(), true
}
// validateWhereClauseStrict is the hardened check for client raw-SQL fragments
// (hardening.sql_strict). It inspects syntax outside string literals: quotes
// and parentheses must be balanced, and comments, statement separators,
// dollar-quoting, DML keywords, dangerous functions and system catalogs are
// rejected. Subqueries and ordinary functions stay allowed unless
// hardening.sql_block_subqueries is set. isJoin (custom joins) still gets all
// checks except the subquery block, since joins legitimately use subqueries.
func validateWhereClauseStrict(where string, isJoin bool) error {
stripped, ok := stripSQLLiterals(where)
if !ok {
return fmt.Errorf("unterminated quote")
}
for _, bad := range []string{"--", "/*", "*/", ";", "$$", "\\"} {
if strings.Contains(stripped, bad) {
return fmt.Errorf("forbidden token %q", bad)
}
}
depth := 0
for i := 0; i < len(stripped); i++ {
switch stripped[i] {
case '(':
depth++
case ')':
depth--
if depth < 0 {
return fmt.Errorf("unbalanced parentheses")
}
}
}
if depth != 0 {
return fmt.Errorf("unbalanced parentheses")
}
if m := reStrictDML.FindString(stripped); m != "" {
return fmt.Errorf("forbidden keyword %q", strings.ToLower(m))
}
if m := reStrictDangerousFunc.FindString(stripped); m != "" {
return fmt.Errorf("forbidden function or schema %q", strings.ToLower(m))
}
if !isJoin && Hardening().SQLBlockSubqueries {
if m := reStrictSubquery.FindString(stripped); m != "" {
return fmt.Errorf("subqueries not allowed (%q)", strings.ToLower(m))
}
}
return nil
}
// SanitizeWhereClause removes trivial conditions and fixes incorrect table prefixes // SanitizeWhereClause removes trivial conditions and fixes incorrect table prefixes
// This function should be used everywhere a WHERE statement is sent to ensure clean, efficient SQL // This function should be used everywhere a WHERE statement is sent to ensure clean, efficient SQL
// //
@@ -174,7 +261,14 @@ func SanitizeWhereClause(where string, tableName string, options ...*RequestOpti
where = strings.TrimSpace(where) where = strings.TrimSpace(where)
// Validate that the WHERE clause doesn't contain dangerous SQL statements // Validate that the WHERE clause doesn't contain dangerous SQL statements
if err := validateWhereClauseSecurity(where); err != nil { if Hardening().SQLStrict {
// Strict mode: fail closed. A rejected client fragment must not turn into
// "no filter", so substitute a clause that matches no rows.
if err := validateWhereClauseStrict(where, tableName == ""); err != nil {
logger.Warn("Rejected client SQL fragment (%v): %s", err, where)
return "(1=0)"
}
} else if err := validateWhereClauseSecurity(where); err != nil {
logger.Debug("Security validation failed for WHERE clause: %v", err) logger.Debug("Security validation failed for WHERE clause: %v", err)
return "" return ""
} }
+4 -4
View File
@@ -102,25 +102,25 @@ func TestSanitizeWhereClause(t *testing.T) {
name: "dangerous DELETE keyword - blocked", name: "dangerous DELETE keyword - blocked",
where: "status = 'active'; DELETE FROM users", where: "status = 'active'; DELETE FROM users",
tableName: "users", tableName: "users",
expected: "", expected: "(1=0)", // fail closed,
}, },
{ {
name: "dangerous UPDATE keyword - blocked", name: "dangerous UPDATE keyword - blocked",
where: "1=1; UPDATE users SET admin = true", where: "1=1; UPDATE users SET admin = true",
tableName: "users", tableName: "users",
expected: "", expected: "(1=0)", // fail closed,
}, },
{ {
name: "dangerous TRUNCATE keyword - blocked", name: "dangerous TRUNCATE keyword - blocked",
where: "status = 'active' OR TRUNCATE TABLE users", where: "status = 'active' OR TRUNCATE TABLE users",
tableName: "users", tableName: "users",
expected: "", expected: "(1=0)", // fail closed,
}, },
{ {
name: "dangerous DROP keyword - blocked", name: "dangerous DROP keyword - blocked",
where: "status = 'active'; DROP TABLE users", where: "status = 'active'; DROP TABLE users",
tableName: "users", tableName: "users",
expected: "", expected: "(1=0)", // fail closed,
}, },
{ {
name: "subquery with table alias should not be modified", name: "subquery with table alias should not be modified",
+5 -7
View File
@@ -27,12 +27,13 @@ var allowedPoolHookTx = map[string]int{
"resolvespec/handler.go": 1, "resolvespec/handler.go": 1,
"websocketspec/handler.go": 1, "websocketspec/handler.go": 1,
"resolvemcp/handler.go": 4, "resolvemcp/handler.go": 4,
"resolvemcp/annotation.go": 1, // BeforeHandle context of the annotate tool
"resolvemcp/writewhere.go": 1, // BeforeHandle context of filter writes
"resolvemcp/functions.go": 1, // BeforeHandle context of call_function
} }
// allowedPoolQuery: statements outside the request path. // allowedPoolQuery: statements outside the request path.
var allowedPoolQuery = map[string]int{ var allowedPoolQuery = map[string]int{}
"resolvemcp/annotation.go": 2, // tool annotations, not a data request
}
func guardedFiles(t *testing.T) map[string][]string { func guardedFiles(t *testing.T) map[string][]string {
t.Helper() t.Helper()
@@ -102,10 +103,7 @@ func TestSpecHandlersDoNotQueryThePoolDirectly(t *testing.T) {
// Anything else defined in a spec's hooks.go must have an Execute call site: an // Anything else defined in a spec's hooks.go must have an Execute call site: an
// unwired hook silently disables whatever is registered on it (resolvespec's // unwired hook silently disables whatever is registered on it (resolvespec's
// AfterRead skipped column-level security masking until it was wired). // AfterRead skipped column-level security masking until it was wired).
var unwiredHooks = map[string]string{ var unwiredHooks = map[string]string{}
"websocketspec/BeforeDisconnect": "connection close is not hooked yet",
"websocketspec/AfterDisconnect": "connection close is not hooked yet",
}
var hookConstRE = regexp.MustCompile(`(?m)^\s*([A-Z][A-Za-z0-9]*)\s+HookType\s*=`) var hookConstRE = regexp.MustCompile(`(?m)^\s*([A-Z][A-Za-z0-9]*)\s+HookType\s*=`)
+60 -3
View File
@@ -3,6 +3,7 @@ package common
import ( import (
"fmt" "fmt"
"reflect" "reflect"
"regexp"
"sort" "sort"
"strings" "strings"
@@ -105,7 +106,7 @@ func (v *ColumnValidator) ValidateColumn(column string) error {
} }
// Allow columns prefixed with "cql" (case insensitive) for computed columns // Allow columns prefixed with "cql" (case insensitive) for computed columns
if strings.HasPrefix(strings.ToLower(column), "cql") { if lc := strings.ToLower(column); strings.HasPrefix(lc, "cql") && (!Hardening().SortStrict || reCQLColumn.MatchString(lc)) {
return nil return nil
} }
@@ -275,8 +276,14 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
validSorts = append(validSorts, sort) validSorts = append(validSorts, sort)
} else { } else {
foundJoin := false foundJoin := false
strictSort := Hardening().SortStrict
for _, j := range options.JoinAliases { for _, j := range options.JoinAliases {
if strings.Contains(sort.Column, j) { if strictSort {
if isJoinAliasColumn(sort.Column, j) {
foundJoin = true
break
}
} else if strings.Contains(sort.Column, j) {
foundJoin = true foundJoin = true
break break
} }
@@ -287,7 +294,7 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
} }
if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") { if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
// Allow sort by expression/subquery, but validate for security // Allow sort by expression/subquery, but validate for security
if IsSafeSortExpression(sort.Column) { if IsSafeSortExpression(sort.Column) && (!strictSort || isSortExpressionRestricted(sort.Column)) {
validSorts = append(validSorts, sort) validSorts = append(validSorts, sort)
} else { } else {
logger.Warn("Unsafe sort expression '%s' removed", sort.Column) logger.Warn("Unsafe sort expression '%s' removed", sort.Column)
@@ -376,6 +383,56 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
return filtered return filtered
} }
var reCQLColumn = regexp.MustCompile(`^cql[a-z0-9_]*$`)
var (
reJoinColumnIdent = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
)
// isJoinAliasColumn reports whether col is exactly "<alias>.<identifier>".
// An empty alias never matches.
func isJoinAliasColumn(col, alias string) bool {
if alias == "" || !strings.HasPrefix(col, alias+".") {
return false
}
return reJoinColumnIdent.MatchString(col[len(alias)+1:])
}
// isSortExpressionRestricted applies the hardened checks to a client sort
// expression. Subqueries and ordinary functions are allowed; dangerous
// functions/catalogs (pg_sleep, pg_*, dblink, ...) and unbalanced parentheses are
// rejected. Subqueries are blocked only when hardening.sql_block_subqueries is on.
func isSortExpressionRestricted(expr string) bool {
stripped, ok := stripSQLLiterals(expr)
if !ok {
return false
}
depth := 0
for i := 0; i < len(stripped); i++ {
switch stripped[i] {
case '(':
depth++
case ')':
depth--
if depth < 0 {
return false
}
}
}
if depth != 0 {
return false
}
if m := reStrictDangerousFunc.FindString(stripped); m != "" {
logger.Warn("Forbidden function '%s' in sort expression: %s", m, expr)
return false
}
if Hardening().SQLBlockSubqueries && reStrictSubquery.MatchString(stripped) {
logger.Warn("Subquery in sort expression rejected: %s", expr)
return false
}
return true
}
// IsSafeSortExpression validates that a sort expression (enclosed in brackets) is safe // IsSafeSortExpression validates that a sort expression (enclosed in brackets) is safe
// and doesn't contain SQL injection attempts or dangerous commands // and doesn't contain SQL injection attempts or dangerous commands
func IsSafeSortExpression(expr string) bool { func IsSafeSortExpression(expr string) bool {
+9 -9
View File
@@ -435,19 +435,19 @@ func TestFilterRequestOptions_WithSortExpressions(t *testing.T) {
options := RequestOptions{ options := RequestOptions{
Sort: []SortOption{ Sort: []SortOption{
{Column: "id", Direction: "ASC"}, // Valid column {Column: "id", Direction: "ASC"}, // Valid column
{Column: "(SELECT MAX(age) FROM users)", Direction: "DESC"}, // Safe expression {Column: "(SELECT MAX(age) FROM users)", Direction: "DESC"}, // Safe expression
{Column: "name", Direction: "ASC"}, // Valid column {Column: "name", Direction: "ASC"}, // Valid column
{Column: "(id); DROP TABLE users; --", Direction: "DESC"}, // Dangerous expression {Column: "(id); DROP TABLE users; --", Direction: "DESC"}, // Dangerous expression
{Column: "invalid_col", Direction: "ASC"}, // Invalid column {Column: "invalid_col", Direction: "ASC"}, // Invalid column
{Column: "(CASE WHEN age > 18 THEN 1 ELSE 0 END)", Direction: "ASC"}, // Safe expression {Column: "(CASE WHEN age > 18 THEN 1 ELSE 0 END)", Direction: "ASC"}, // Safe expression
}, },
} }
filtered := validator.FilterRequestOptions(options) filtered := validator.FilterRequestOptions(options)
// Should keep: id, safe expression, name, another safe expression // Keeps: id, subquery expression, name, CASE expression
// Should remove: dangerous expression, invalid column // Removes: dangerous expression, invalid column
expectedCount := 4 expectedCount := 4
if len(filtered.Sort) != expectedCount { if len(filtered.Sort) != expectedCount {
t.Errorf("Expected %d sort options, got %d", expectedCount, len(filtered.Sort)) t.Errorf("Expected %d sort options, got %d", expectedCount, len(filtered.Sort))
@@ -474,8 +474,8 @@ type RelatedModel struct {
// PreloadParentModel has a has-one relation to RelatedModel. The json tag on // PreloadParentModel has a has-one relation to RelatedModel. The json tag on
// the relation field is the name used in x-preload headers. // the relation field is the name used in x-preload headers.
type PreloadParentModel struct { type PreloadParentModel struct {
ID int64 `bun:"id,pk"` ID int64 `bun:"id,pk"`
Name string `bun:"name"` Name string `bun:"name"`
RELATED *RelatedModel `json:"RELATED" bun:"rel:has-one,join:id=related_id"` RELATED *RelatedModel `json:"RELATED" bun:"rel:has-one,join:id=related_id"`
} }
+20
View File
@@ -17,6 +17,7 @@ type Config struct {
EventBroker EventBrokerConfig `mapstructure:"event_broker"` EventBroker EventBrokerConfig `mapstructure:"event_broker"`
DBManager DBManagerConfig `mapstructure:"dbmanager"` DBManager DBManagerConfig `mapstructure:"dbmanager"`
DBTrace DBTraceConfig `mapstructure:"db_trace"` DBTrace DBTraceConfig `mapstructure:"db_trace"`
Hardening HardeningConfig `mapstructure:"hardening"`
Paths PathsConfig `mapstructure:"paths"` Paths PathsConfig `mapstructure:"paths"`
Extensions map[string]interface{} `mapstructure:"extensions"` Extensions map[string]interface{} `mapstructure:"extensions"`
} }
@@ -143,6 +144,25 @@ type CORSConfig struct {
MaxAge int `mapstructure:"max_age"` MaxAge int `mapstructure:"max_age"`
} }
// HardeningConfig toggles security hardening that may reject requests which
// older clients relied on. All switches default to true; set one to false to
// restore the previous (permissive) behaviour.
// Env: RESOLVESPEC_HARDENING_CORS_STRICT_ORIGINS, _SORT_STRICT, _SQL_STRICT, _SQL_BLOCK_SUBQUERIES.
type HardeningConfig struct {
// CORSStrictOrigins only reflects origins listed in cors.allowed_origins (and the
// server URLs); credentials are never sent for unlisted or wildcard origins.
CORSStrictOrigins bool `mapstructure:"cors_strict_origins"`
// SortStrict stops empty/substring join aliases from admitting arbitrary sort strings.
SortStrict bool `mapstructure:"sort_strict"`
// SQLStrict hardens client raw-SQL fragments (x-custom-sql-*, preload where, cursor):
// balanced parentheses, no subqueries/functions/comments, and rejection instead of
// silently dropping the filter.
SQLStrict bool `mapstructure:"sql_strict"`
// SQLBlockSubqueries additionally rejects subqueries (select/union/with) in client
// WHERE fragments (not custom joins). Off by default: existing clients use them.
SQLBlockSubqueries bool `mapstructure:"sql_block_subqueries"`
}
// DBTraceConfig controls database usage logging (off by default). // DBTraceConfig controls database usage logging (off by default).
// Env: RESOLVESPEC_DB_TRACE_ENABLED, _MIN_CALLS, _MIN_DURATION, _POOL_LOG. // Env: RESOLVESPEC_DB_TRACE_ENABLED, _MIN_CALLS, _MIN_DURATION, _POOL_LOG.
type DBTraceConfig struct { type DBTraceConfig struct {
+7
View File
@@ -168,6 +168,7 @@ func (m *Manager) SetConfig(cfg *Config) error {
m.v.Set("event_broker", cfg.EventBroker) m.v.Set("event_broker", cfg.EventBroker)
m.v.Set("dbmanager", cfg.DBManager) m.v.Set("dbmanager", cfg.DBManager)
m.v.Set("db_trace", cfg.DBTrace) m.v.Set("db_trace", cfg.DBTrace)
m.v.Set("hardening", cfg.Hardening)
m.v.Set("paths", cfg.Paths) m.v.Set("paths", cfg.Paths)
m.v.Set("extensions", cfg.Extensions) m.v.Set("extensions", cfg.Extensions)
@@ -279,6 +280,12 @@ func setDefaults(v *viper.Viper) {
v.SetDefault("cors.allowed_headers", []string{"*"}) v.SetDefault("cors.allowed_headers", []string{"*"})
v.SetDefault("cors.max_age", 3600) v.SetDefault("cors.max_age", 3600)
// Security hardening defaults (on)
v.SetDefault("hardening.cors_strict_origins", true)
v.SetDefault("hardening.sort_strict", true)
v.SetDefault("hardening.sql_strict", true)
v.SetDefault("hardening.sql_block_subqueries", false)
// Database defaults // Database defaults
v.SetDefault("database.url", "") v.SetDefault("database.url", "")
+1 -1
View File
@@ -15,7 +15,7 @@ Wire: `dbtrace.Configure(dbtrace.FromConfig(cfg.DBTrace))` and wrap handlers wit
## Log fields ## Log fields
- `tx` transactions begun · `tx_queries` adapter queries inside `RunInTransaction` (share the tx connection) - `tx` transactions begun · `tx_queries` adapter queries inside `RunInTransaction` (share the tx connection)
- `pooled` adapter queries outside a tx (each takes a pool connection) - `pooled` adapter queries outside a tx (each takes a pool connection)
- `raw` direct `*sql.DB` calls, with kinds: `auth.session`, `auth.activity`, `security.column`, `security.row`, `probe.pg_proc`, `keystore.validate` - `raw` direct `*sql.DB` calls, with kinds: `auth.session`, `auth.activity`, `security.column`, `security.row`, `probe.pg_proc` (lookup `ModeAuto` only), `keystore.validate`
- Connections used ≈ `tx + pooled + raw` - Connections used ≈ `tx + pooled + raw`
## Pool log ## Pool log
+50 -3
View File
@@ -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
+53
View File
@@ -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"
+189
View File
@@ -0,0 +1,189 @@
package metrics
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"sync"
"time"
"github.com/prometheus/client_golang/prometheus"
dto "github.com/prometheus/client_model/go"
"github.com/prometheus/common/expfmt"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
const textContentType = "text/plain; version=0.0.4; charset=utf-8"
// endpointPusher POSTs gathered metrics to a user-configured HTTP endpoint.
type endpointPusher struct {
url string
format string
headers map[string]string
client *http.Client
resetOnOK bool
provider *PrometheusProvider
gatherer prometheus.Gatherer
stopOnce sync.Once
stopCh chan struct{}
startedMu sync.Mutex
started bool
}
func newEndpointPusher(cfg *Config, p *PrometheusProvider) *endpointPusher {
return &endpointPusher{
url: cfg.PushEndpointURL,
format: cfg.PushEndpointFormat,
headers: cfg.PushEndpointHeaders,
client: &http.Client{Timeout: time.Duration(cfg.PushEndpointTimeout) * time.Second},
resetOnOK: cfg.PushEndpointResetOnSuccess,
provider: p,
gatherer: prometheus.DefaultGatherer,
stopCh: make(chan struct{}),
}
}
func (e *endpointPusher) start(interval time.Duration) {
e.startedMu.Lock()
defer e.startedMu.Unlock()
if e.started {
return
}
e.started = true
go func() {
t := time.NewTicker(interval)
defer t.Stop()
for {
select {
case <-t.C:
if err := e.push(context.Background()); err != nil {
logger.Warn("Failed to push metrics to endpoint %s: %v", e.url, err)
}
case <-e.stopCh:
return
}
}
}()
}
func (e *endpointPusher) stop() {
e.stopOnce.Do(func() { close(e.stopCh) })
}
func (e *endpointPusher) push(ctx context.Context) error {
mfs, err := e.gatherer.Gather()
if err != nil && len(mfs) == 0 {
return fmt.Errorf("gather metrics: %w", err)
}
body, contentType, err := encodeMetrics(mfs, e.format)
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, e.url, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", contentType)
for k, v := range e.headers {
req.Header.Set(k, v)
}
resp, err := e.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
snippet, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
if resp.StatusCode < 200 || resp.StatusCode > 299 {
return fmt.Errorf("endpoint returned %s: %s", resp.Status, bytes.TrimSpace(snippet))
}
if e.resetOnOK {
e.provider.Reset()
}
return nil
}
func encodeMetrics(mfs []*dto.MetricFamily, format string) (body []byte, contentType string, err error) {
switch format {
case "json":
b, err := json.Marshal(toJSONFamilies(mfs))
return b, "application/json", err
case "", "text":
var buf bytes.Buffer
enc := expfmt.NewEncoder(&buf, expfmt.NewFormat(expfmt.TypeTextPlain))
for _, mf := range mfs {
if err := enc.Encode(mf); err != nil {
return nil, "", err
}
}
return buf.Bytes(), textContentType, nil
default:
return nil, "", fmt.Errorf("unsupported push endpoint format %q", format)
}
}
type jsonFamily struct {
Name string `json:"name"`
Help string `json:"help,omitempty"`
Type string `json:"type"`
Metrics []jsonMetric `json:"metrics"`
}
type jsonMetric struct {
Labels map[string]string `json:"labels,omitempty"`
Value *float64 `json:"value,omitempty"`
Count *uint64 `json:"count,omitempty"`
Sum *float64 `json:"sum,omitempty"`
Buckets []jsonBucket `json:"buckets,omitempty"`
}
type jsonBucket struct {
UpperBound float64 `json:"le"`
Count uint64 `json:"count"`
}
func toJSONFamilies(mfs []*dto.MetricFamily) []jsonFamily {
out := make([]jsonFamily, 0, len(mfs))
for _, mf := range mfs {
f := jsonFamily{Name: mf.GetName(), Help: mf.GetHelp(), Type: mf.GetType().String()}
for _, m := range mf.GetMetric() {
jm := jsonMetric{}
if len(m.GetLabel()) > 0 {
jm.Labels = make(map[string]string, len(m.GetLabel()))
for _, l := range m.GetLabel() {
jm.Labels[l.GetName()] = l.GetValue()
}
}
switch {
case m.Counter != nil:
v := m.Counter.GetValue()
jm.Value = &v
case m.Gauge != nil:
v := m.Gauge.GetValue()
jm.Value = &v
case m.Untyped != nil:
v := m.Untyped.GetValue()
jm.Value = &v
case m.Histogram != nil:
c, s := m.Histogram.GetSampleCount(), m.Histogram.GetSampleSum()
jm.Count, jm.Sum = &c, &s
for _, b := range m.Histogram.GetBucket() {
jm.Buckets = append(jm.Buckets, jsonBucket{UpperBound: b.GetUpperBound(), Count: b.GetCumulativeCount()})
}
case m.Summary != nil:
c, s := m.Summary.GetSampleCount(), m.Summary.GetSampleSum()
jm.Count, jm.Sum = &c, &s
}
f.Metrics = append(f.Metrics, jm)
}
out = append(out, f)
}
return out
}
+15
View File
@@ -47,6 +47,21 @@ type Provider interface {
Handler() http.Handler 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
+175
View File
@@ -0,0 +1,175 @@
package metrics
import (
"net/http"
"strings"
"sync"
)
const (
// defaultHTTPMaxPaths is the default cap on distinct values of the "path" label.
defaultHTTPMaxPaths = 1024
// overflowPathLabel is used once the cap on distinct path labels is reached.
overflowPathLabel = "other"
)
// routeLabel returns the low-cardinality path label for a request, preferring
// (in order): the custom normalizer, the matched ServeMux pattern, and finally
// the generic normalization of the raw URL path.
func routeLabel(r *http.Request, custom func(*http.Request) string) string {
if custom != nil {
if p := custom(r); p != "" {
return p
}
}
if r.Pattern != "" {
return stripPatternMethod(r.Pattern)
}
return NormalizePath(r.URL.Path)
}
// stripPatternMethod removes the optional "METHOD " prefix (and host) from a
// Go 1.22+ ServeMux pattern, e.g. "GET /users/{id}" -> "/users/{id}".
func stripPatternMethod(pattern string) string {
if i := strings.IndexByte(pattern, ' '); i >= 0 {
pattern = strings.TrimLeft(pattern[i+1:], " ")
}
if i := strings.IndexByte(pattern, '/'); i > 0 {
pattern = pattern[i:] // drop host part
}
return pattern
}
// NormalizePath replaces dynamic-looking path segments (numeric IDs, UUIDs,
// long hex strings and other long opaque tokens) with ":id" so that
// /users/123 and /users/456 share one label value.
func NormalizePath(path string) string {
if path == "" {
return "/"
}
if !strings.Contains(path, "/") {
return path
}
segs := strings.Split(path, "/")
for i, s := range segs {
if isDynamicSegment(s) {
segs[i] = ":id"
}
}
return strings.Join(segs, "/")
}
func isDynamicSegment(s string) bool {
if s == "" {
return false
}
if allDigits(s) {
return true
}
if isUUID(s) {
return true
}
// Long hex strings (hashes, object IDs)
if len(s) >= 16 && allHex(s) {
return true
}
// Long opaque tokens containing digits (base64/ULID-like)
if len(s) >= 24 && hasDigit(s) && !strings.ContainsAny(s, ".") {
return true
}
return false
}
func allDigits(s string) bool {
for i := 0; i < len(s); i++ {
if s[i] < '0' || s[i] > '9' {
return false
}
}
return true
}
func hasDigit(s string) bool {
for i := 0; i < len(s); i++ {
if s[i] >= '0' && s[i] <= '9' {
return true
}
}
return false
}
func allHex(s string) bool {
for i := 0; i < len(s); i++ {
c := s[i]
if !isHexByte(c) {
return false
}
}
return true
}
func isUUID(s string) bool {
if len(s) != 36 {
return false
}
for i := 0; i < len(s); i++ {
c := s[i]
switch i {
case 8, 13, 18, 23:
if c != '-' {
return false
}
default:
if !isHexByte(c) {
return false
}
}
}
return true
}
// pathLimiter bounds the number of distinct path label values. Once the cap is
// reached, unseen paths are reported as "other".
type pathLimiter struct {
mu sync.RWMutex
max int // <= 0 disables the cap
seen map[string]struct{}
}
func newPathLimiter(limit int) *pathLimiter {
return &pathLimiter{max: limit, seen: make(map[string]struct{})}
}
func (l *pathLimiter) label(path string) string {
if l.max <= 0 {
return path
}
l.mu.RLock()
_, ok := l.seen[path]
l.mu.RUnlock()
if ok {
return path
}
l.mu.Lock()
defer l.mu.Unlock()
if _, ok := l.seen[path]; ok {
return path
}
if len(l.seen) >= l.max {
return overflowPathLabel
}
l.seen[path] = struct{}{}
return path
}
func (l *pathLimiter) reset() {
l.mu.Lock()
l.seen = make(map[string]struct{})
l.mu.Unlock()
}
func isHexByte(c byte) bool {
return c >= '0' && c <= '9' || c >= 'a' && c <= 'f' || c >= 'A' && c <= 'F'
}
+237
View File
@@ -0,0 +1,237 @@
package metrics
import (
"context"
"encoding/json"
"github.com/prometheus/client_golang/prometheus"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
)
func TestNormalizePath(t *testing.T) {
cases := map[string]string{
"": "/",
"/": "/",
"/users": "/users",
"/users/123": "/users/:id",
"/users/123/orders/9": "/users/:id/orders/:id",
"/x/550e8400-e29b-41d4-a716-446655440000": "/x/:id",
"/x/507f1f77bcf86cd799439011": "/x/:id",
"/api/public/users": "/api/public/users",
"/files/report.v2": "/files/report.v2",
}
for in, want := range cases {
if got := NormalizePath(in); got != want {
t.Errorf("NormalizePath(%q) = %q, want %q", in, got, want)
}
}
}
func TestRouteLabel(t *testing.T) {
r := httptest.NewRequest("GET", "/users/42", nil)
if got := routeLabel(r, nil); got != "/users/:id" {
t.Errorf("fallback = %q", got)
}
r.Pattern = "GET /users/{id}"
if got := routeLabel(r, nil); got != "/users/{id}" {
t.Errorf("pattern = %q", got)
}
got := routeLabel(r, func(*http.Request) string { return "/custom" })
if got != "/custom" {
t.Errorf("custom = %q", got)
}
}
func TestPathLimiter(t *testing.T) {
l := newPathLimiter(2)
for _, p := range []string{"/a", "/b", "/a"} {
if got := l.label(p); got != p {
t.Errorf("label(%q) = %q", p, got)
}
}
if got := l.label("/c"); got != overflowPathLabel {
t.Errorf("overflow = %q", got)
}
if got := newPathLimiter(-1).label("/z"); got != "/z" {
t.Errorf("disabled = %q", got)
}
}
func TestMiddlewareUsesPattern(t *testing.T) {
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "pathtest"})
mux := http.NewServeMux()
mux.HandleFunc("GET /users/{id}", func(w http.ResponseWriter, r *http.Request) {})
h := p.Middleware(mux)
for _, id := range []string{"1", "2", "abc"} {
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest("GET", "/users/"+id, nil))
}
if n := len(p.pathLimiter.seen); n != 1 {
t.Errorf("distinct paths = %d, want 1", n)
}
if _, ok := p.pathLimiter.seen["/users/{id}"]; !ok {
t.Errorf("seen = %v", p.pathLimiter.seen)
}
}
func TestResetAndHandler(t *testing.T) {
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "resettest"})
p.RecordHTTPRequest("GET", "/a/1", "200", 0)
p.RecordDBQuery("SELECT", "s", "e", "t", 0, nil)
p.IncRequestsInFlight()
count := func() int {
mfs, _ := prometheus.DefaultGatherer.Gather()
n := 0
for _, mf := range mfs {
if strings.HasPrefix(mf.GetName(), "resettest_") && mf.GetName() != "resettest_http_requests_in_flight" && mf.GetName() != "resettest_event_queue_size" {
n += len(mf.GetMetric())
}
}
return n
}
if count() == 0 {
t.Fatal("expected recorded series")
}
rec := httptest.NewRecorder()
p.ResetHandler().ServeHTTP(rec, httptest.NewRequest("GET", "/reset", nil))
if rec.Code != http.StatusMethodNotAllowed || count() == 0 {
t.Fatalf("GET should be rejected, code=%d", rec.Code)
}
rec = httptest.NewRecorder()
p.ResetHandler().ServeHTTP(rec, httptest.NewRequest("POST", "/reset", nil))
if rec.Code != http.StatusNoContent || count() != 0 {
t.Fatalf("reset failed, code=%d series=%d", rec.Code, count())
}
if len(p.pathLimiter.seen) != 0 {
t.Error("path limiter not reset")
}
// push=true without a pushgateway must fail and not be silent
rec = httptest.NewRecorder()
p.ResetHandler().ServeHTTP(rec, httptest.NewRequest("POST", "/reset?push=true", nil))
if rec.Code != http.StatusBadGateway {
t.Errorf("push without gateway code=%d", rec.Code)
}
}
func TestPushAndResetKeepsStatsOnFailure(t *testing.T) {
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "pushfail", PushgatewayURL: "http://127.0.0.1:1"})
p.RecordHTTPRequest("GET", "/a", "200", 0)
if err := p.PushAndReset(); err == nil {
t.Fatal("expected push error")
}
if len(p.pathLimiter.seen) != 1 {
t.Error("stats were reset despite failed push")
}
}
func TestPushToEndpoint(t *testing.T) {
for _, format := range []string{"text", "json"} {
var gotCT, gotAuth string
var gotBody []byte
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
t.Errorf("method = %s", r.Method)
}
gotCT, gotAuth = r.Header.Get("Content-Type"), r.Header.Get("Authorization")
gotBody, _ = io.ReadAll(r.Body)
}))
ns := "ep" + format
p := NewPrometheusProvider(&Config{
Enabled: true,
Namespace: ns,
PushEndpointURL: srv.URL,
PushEndpointFormat: format,
PushEndpointHeaders: map[string]string{"Authorization": "Bearer x"},
PushEndpointResetOnSuccess: true,
})
p.RecordHTTPRequest("GET", "/a", "200", time.Millisecond)
if err := p.PushToEndpoint(context.Background()); err != nil {
t.Fatalf("%s: %v", format, err)
}
srv.Close()
if gotAuth != "Bearer x" || !strings.Contains(string(gotBody), ns+"_http_requests_total") {
t.Errorf("%s: auth=%q body=%.200s", format, gotAuth, gotBody)
}
if format == "json" && gotCT != "application/json" || format == "text" && !strings.HasPrefix(gotCT, "text/plain") {
t.Errorf("%s: content-type %q", format, gotCT)
}
if len(p.pathLimiter.seen) != 0 {
t.Errorf("%s: stats not reset after success", format)
}
}
}
func TestPushToEndpointFailureKeepsStats(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "nope", http.StatusInternalServerError)
}))
defer srv.Close()
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "epfail", PushEndpointURL: srv.URL, PushEndpointResetOnSuccess: true})
p.RecordHTTPRequest("GET", "/a", "200", 0)
if err := p.PushToEndpoint(context.Background()); err == nil {
t.Fatal("expected error on 500")
}
if len(p.pathLimiter.seen) != 1 {
t.Error("stats reset despite failure")
}
if err := NewPrometheusProvider(&Config{Enabled: true, Namespace: "epnone"}).PushToEndpoint(context.Background()); err == nil {
t.Error("expected error without endpoint")
}
}
func TestDisabledProvider(t *testing.T) {
p := NewPrometheusProvider(&Config{Namespace: "disabled", PushEndpointURL: "http://127.0.0.1:1", PushEndpointInterval: 1})
p.RecordHTTPRequest("GET", "/a", "200", 0)
if len(p.pathLimiter.seen) != 0 {
t.Error("disabled provider recorded")
}
for name, h := range map[string]http.Handler{"handler": p.Handler(), "json": p.JSONHandler()} {
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest("GET", "/m", nil))
if rec.Code != http.StatusNotFound {
t.Errorf("%s code=%d", name, rec.Code)
}
}
if p.endpoint != nil || p.PushToEndpoint(context.Background()) == nil || p.Push() == nil {
t.Error("disabled provider must not push")
}
}
func TestJSONHandler(t *testing.T) {
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "jsonpull"})
p.RecordHTTPRequest("GET", "/a", "200", time.Millisecond)
rec := httptest.NewRecorder()
p.JSONHandler().ServeHTTP(rec, httptest.NewRequest("GET", "/m", nil))
if rec.Code != 200 || rec.Header().Get("Content-Type") != "application/json" {
t.Fatalf("code=%d ct=%q", rec.Code, rec.Header().Get("Content-Type"))
}
var fams []map[string]any
if err := json.Unmarshal(rec.Body.Bytes(), &fams); err != nil {
t.Fatal(err)
}
found := false
for _, f := range fams {
if f["name"] == "jsonpull_http_requests_total" {
found = true
}
}
if !found {
t.Error("metric family missing from JSON")
}
rec = httptest.NewRecorder()
p.JSONHandler().ServeHTTP(rec, httptest.NewRequest("POST", "/m", nil))
if rec.Code != http.StatusMethodNotAllowed {
t.Errorf("POST code=%d", rec.Code)
}
}
+200 -6
View File
@@ -1,6 +1,8 @@
package metrics 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()
} }
} }
+32
View File
@@ -0,0 +1,32 @@
// Package modelregistry is the shared catalogue of the Go model structs that the
// ResolveSpec front ends (resolvespec, restheadspec, websocketspec, mqttspec,
// resolvemcp, ...) expose as database entities.
//
// A registry maps a model name ("schema.entity") to a struct type and holds:
// - ModelRules: which operations (read/create/update/delete, public or not)
// are allowed, and whether security checks are disabled.
// - ModelInfo: optional documentation (description, purpose, tags, per-column
// descriptions) meant for humans and AI agents. It never affects queries or
// permissions.
//
// Register models on a registry created with NewModelRegistry, or through the
// package-level functions that use the default registry:
//
// reg := modelregistry.NewModelRegistry()
// _ = reg.RegisterModelWithRules("public.users", User{}, modelregistry.DefaultModelRules())
// reg.SetModelInfo("public.users", modelregistry.ModelInfo{
// Description: "Application accounts",
// Purpose: "Look up who a person is; never store credentials here",
// Columns: map[string]string{"email": "Login address, unique"},
// })
//
// Descriptions come from, in priority order:
// 1. ModelInfo set with SetModelInfo or loaded from an external JSON map with
// LoadModelInfoFile (the map can be maintained outside the Go code).
// 2. The model's Describer (ModelDescription() string) for the description.
// 3. Struct tags, read per column by FieldComment: comment, note, desc or
// description tags, then "comment:" inside the gorm or bun tag.
//
// Models must be non-pointer structs; pointers, slices and arrays of structs are
// unwrapped on registration. All registry methods are safe for concurrent use.
package modelregistry
+172
View File
@@ -0,0 +1,172 @@
package modelregistry
import (
"encoding/json"
"fmt"
"io"
"os"
"path/filepath"
"reflect"
"strings"
)
// ModelInfo is human/AI-facing documentation for a registered model: what it
// is for and what its columns mean. It is optional and has no effect on
// permissions or queries.
type ModelInfo struct {
// Description says what the model/table holds.
Description string `json:"description,omitempty"`
// Purpose says why it exists / when an agent should use it.
Purpose string `json:"purpose,omitempty"`
// Tags are free-form labels (e.g. "billing", "pii").
Tags []string `json:"tags,omitempty"`
// Columns maps a JSON column name to its description.
Columns map[string]string `json:"columns,omitempty"`
}
// IsZero reports whether the info carries no documentation.
func (i ModelInfo) IsZero() bool {
return i.Description == "" && i.Purpose == "" && len(i.Tags) == 0 && len(i.Columns) == 0
}
// Describer can be implemented by a model to document itself. It is the
// fallback used when no ModelInfo description was registered or loaded.
// (The method is not called Description so models may keep a Description field.)
type Describer interface {
ModelDescription() string
}
// commentTagKeys are the standalone struct tags read as a column description,
// in priority order.
var commentTagKeys = []string{"comment", "note", "desc", "description"}
// FieldComment returns the description of a struct field from its tags. Order:
// standalone comment/note/desc/description tags, then a "comment:" entry inside
// the gorm tag (semicolon separated), then inside the bun tag (comma separated).
// It returns "" when the field carries none.
func FieldComment(sf reflect.StructField) string {
for _, key := range commentTagKeys {
if v := strings.TrimSpace(sf.Tag.Get(key)); v != "" {
return v
}
}
if v := tagOption(sf.Tag.Get("gorm"), ';', "comment:"); v != "" {
return v
}
return tagOption(sf.Tag.Get("bun"), ',', "comment:")
}
func tagOption(tag string, sep byte, key string) string {
for _, part := range strings.Split(tag, string(sep)) {
part = strings.TrimSpace(part)
if len(part) >= len(key) && strings.EqualFold(part[:len(key)], key) {
return strings.Trim(strings.TrimSpace(part[len(key):]), `'"`)
}
}
return ""
}
// SetModelInfo stores documentation for a model name ("schema.entity"). The
// model does not have to be registered yet, so descriptions can be loaded
// before or after registration. Any previous info for the name is replaced.
func (r *DefaultModelRegistry) SetModelInfo(name string, info ModelInfo) {
r.mutex.Lock()
defer r.mutex.Unlock()
if r.info == nil {
r.info = make(map[string]ModelInfo)
}
r.info[name] = cloneInfo(info)
}
// GetModelInfo returns the documentation stored with SetModelInfo (or loaded
// from a descriptions file), without any fallback.
func (r *DefaultModelRegistry) GetModelInfo(name string) (ModelInfo, bool) {
r.mutex.RLock()
defer r.mutex.RUnlock()
info, ok := r.info[name]
return cloneInfo(info), ok
}
// RegisterModelWithInfo registers a model together with its documentation.
func (r *DefaultModelRegistry) RegisterModelWithInfo(name string, model interface{}, info ModelInfo) error {
if err := r.RegisterModel(name, model); err != nil {
return err
}
r.SetModelInfo(name, info)
return nil
}
// ResolveModelInfo returns the effective documentation for a registered model.
// Stored/loaded info wins; an empty Description falls back to the model's
// Describer. Column descriptions are not resolved here: use the stored map and
// fall back to FieldComment per field.
func (r *DefaultModelRegistry) ResolveModelInfo(name string) ModelInfo {
info, _ := r.GetModelInfo(name)
if info.Description == "" {
if model, err := r.GetModel(name); err == nil {
info.Description = describerText(model)
}
}
return info
}
func describerText(model interface{}) (text string) {
defer func() {
if recover() != nil {
text = ""
}
}()
if d, ok := model.(Describer); ok {
return strings.TrimSpace(d.ModelDescription())
}
if t := reflect.TypeOf(model); t != nil && t.Kind() != reflect.Pointer {
if d, ok := reflect.New(t).Interface().(Describer); ok {
return strings.TrimSpace(d.ModelDescription())
}
}
return ""
}
// LoadModelInfo reads an external descriptions map from r and applies it. The
// JSON is an object keyed by model name:
//
// {"public.users": {"description": "...", "purpose": "...", "tags": ["x"],
// "columns": {"email": "Login address"}}}
//
// Entries replace any existing info for the same name and take precedence over
// the model's Describer and its struct-tag comments. It returns the number of
// models loaded.
func (r *DefaultModelRegistry) LoadModelInfo(src io.Reader) (int, error) {
var m map[string]ModelInfo
dec := json.NewDecoder(src)
dec.DisallowUnknownFields()
if err := dec.Decode(&m); err != nil {
return 0, fmt.Errorf("modelregistry: decode model info: %w", err)
}
for name, info := range m {
r.SetModelInfo(name, info)
}
return len(m), nil
}
// LoadModelInfoFile is LoadModelInfo reading from a JSON file.
func (r *DefaultModelRegistry) LoadModelInfoFile(path string) (int, error) {
f, err := os.Open(filepath.Clean(path)) //nolint:gosec // operator-supplied descriptions file
if err != nil {
return 0, fmt.Errorf("modelregistry: %w", err)
}
defer f.Close()
return r.LoadModelInfo(f)
}
func cloneInfo(in ModelInfo) ModelInfo {
out := in
out.Tags = append([]string(nil), in.Tags...)
if in.Columns != nil {
out.Columns = make(map[string]string, len(in.Columns))
for k, v := range in.Columns {
out.Columns[k] = v
}
}
return out
}
+68
View File
@@ -0,0 +1,68 @@
package modelregistry
import (
"reflect"
"strings"
"testing"
)
type infoModel struct {
ID int `json:"id" gorm:"primaryKey;comment:Row id"`
Email string `json:"email" bun:"email,comment:Login address"`
Name string `json:"name" note:"Display name"`
Plain string `json:"plain"`
}
func (infoModel) ModelDescription() string { return " From the model " }
func TestFieldComment(t *testing.T) {
typ := reflect.TypeOf(infoModel{})
want := map[string]string{"ID": "Row id", "Email": "Login address", "Name": "Display name", "Plain": ""}
for field, exp := range want {
sf, _ := typ.FieldByName(field)
if got := FieldComment(sf); got != exp {
t.Errorf("%s = %q, want %q", field, got, exp)
}
}
}
func TestModelInfoPrecedence(t *testing.T) {
r := NewModelRegistry()
if err := r.RegisterModel("public.items", infoModel{}); err != nil {
t.Fatal(err)
}
if got := r.ResolveModelInfo("public.items").Description; got != "From the model" {
t.Errorf("describer fallback = %q", got)
}
n, err := r.LoadModelInfo(strings.NewReader(
`{"public.items":{"description":"From file","tags":["a"],"columns":{"email":"Mail"}},"public.later":{"purpose":"p"}}`))
if err != nil || n != 2 {
t.Fatalf("load n=%d err=%v", n, err)
}
info := r.ResolveModelInfo("public.items")
if info.Description != "From file" || info.Columns["email"] != "Mail" || len(info.Tags) != 1 {
t.Errorf("file info = %+v", info)
}
if _, ok := r.GetModelInfo("public.later"); !ok {
t.Error("info for a not-yet-registered model must be kept")
}
// returned info is a copy
info.Columns["email"] = "changed"
if got, _ := r.GetModelInfo("public.items"); got.Columns["email"] != "Mail" {
t.Error("GetModelInfo leaked internal map")
}
}
func TestLoadModelInfoRejectsBadInput(t *testing.T) {
r := NewModelRegistry()
for _, in := range []string{`not json`, `{"a":{"descripton":"typo"}}`} {
if _, err := r.LoadModelInfo(strings.NewReader(in)); err == nil {
t.Errorf("expected error for %q", in)
}
}
if _, err := r.LoadModelInfoFile("/nonexistent/x.json"); err == nil {
t.Error("expected error for missing file")
}
}
+1
View File
@@ -41,6 +41,7 @@ func DefaultModelRules() ModelRules {
type DefaultModelRegistry struct { type DefaultModelRegistry struct {
models map[string]interface{} models map[string]interface{}
rules map[string]ModelRules rules map[string]ModelRules
info map[string]ModelInfo
mutex sync.RWMutex mutex sync.RWMutex
} }
+5
View File
@@ -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)
+65 -1
View File
@@ -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)
+105 -34
View File
@@ -497,13 +497,13 @@ func TestIsColumnWritableWithEmbedded(t *testing.T) {
// Test models with relations for GetSQLModelColumns // Test models with relations for GetSQLModelColumns
type User struct { type User struct {
ID int `bun:"id,pk" json:"id"` ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"` Name string `bun:"name" json:"name"`
Email string `bun:"email" json:"email"` Email string `bun:"email" json:"email"`
ProfileData string `json:"profile_data"` // No bun/gorm tag ProfileData string `json:"profile_data"` // No bun/gorm tag
Posts []Post `bun:"rel:has-many,join:id=user_id" json:"posts"` Posts []Post `bun:"rel:has-many,join:id=user_id" json:"posts"`
Profile *Profile `bun:"rel:has-one,join:id=user_id" json:"profile"` Profile *Profile `bun:"rel:has-one,join:id=user_id" json:"profile"`
RowNumber int64 `bun:",scanonly" json:"_rownumber"` RowNumber int64 `bun:",scanonly" json:"_rownumber"`
} }
type Post struct { type Post struct {
@@ -528,8 +528,8 @@ type Tag struct {
// Model with scan-only embedded struct // Model with scan-only embedded struct
type EntityWithScanOnlyEmbedded struct { type EntityWithScanOnlyEmbedded struct {
ID int `bun:"id,pk" json:"id"` ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"` Name string `bun:"name" json:"name"`
AdhocBuffer `bun:",scanonly"` // Entire embedded struct is scan-only AdhocBuffer `bun:",scanonly"` // Entire embedded struct is scan-only
} }
@@ -1086,17 +1086,17 @@ func TestGetColumnTypeFromModel_SqlNullWrapper(t *testing.T) {
// Models for relation testing // Models for relation testing
type Author struct { type Author struct {
ID int `bun:"id,pk" json:"id"` ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"` Name string `bun:"name" json:"name"`
Books []Book `bun:"rel:has-many,join:id=author_id" json:"books"` Books []Book `bun:"rel:has-many,join:id=author_id" json:"books"`
} }
type Book struct { type Book struct {
ID int `bun:"id,pk" json:"id"` ID int `bun:"id,pk" json:"id"`
Title string `bun:"title" json:"title"` Title string `bun:"title" json:"title"`
AuthorID int `bun:"author_id" json:"author_id"` AuthorID int `bun:"author_id" json:"author_id"`
Author *Author `bun:"rel:belongs-to,join:author_id=id" json:"author"` Author *Author `bun:"rel:belongs-to,join:author_id=id" json:"author"`
Publisher *Publisher `bun:"rel:has-one,join:id=book_id" json:"publisher"` Publisher *Publisher `bun:"rel:has-one,join:id=book_id" json:"publisher"`
} }
type Publisher struct { type Publisher struct {
@@ -1106,9 +1106,9 @@ type Publisher struct {
} }
type Student struct { type Student struct {
ID int `gorm:"column:id;primaryKey" json:"id"` ID int `gorm:"column:id;primaryKey" json:"id"`
Name string `gorm:"column:name" json:"name"` Name string `gorm:"column:name" json:"name"`
Courses []Course `gorm:"many2many:student_courses" json:"courses"` Courses []Course `gorm:"many2many:student_courses" json:"courses"`
} }
type Course struct { type Course struct {
@@ -1119,11 +1119,11 @@ type Course struct {
// Recursive relation model // Recursive relation model
type Category struct { type Category struct {
ID int `bun:"id,pk" json:"id"` ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"` Name string `bun:"name" json:"name"`
ParentID *int `bun:"parent_id" json:"parent_id"` ParentID *int `bun:"parent_id" json:"parent_id"`
Parent *Category `bun:"rel:belongs-to,join:parent_id=id" json:"parent"` Parent *Category `bun:"rel:belongs-to,join:parent_id=id" json:"parent"`
Children []Category `bun:"rel:has-many,join:id=parent_id" json:"children"` Children []Category `bun:"rel:has-many,join:id=parent_id" json:"children"`
} }
func TestGetRelationType(t *testing.T) { func TestGetRelationType(t *testing.T) {
@@ -1299,7 +1299,7 @@ func TestGetPrimaryKeyValue_EdgeCases(t *testing.T) {
expected: nil, expected: nil,
}, },
{ {
name: "model without primary key tags - fallback to ID field", name: "model without primary key tags - fallback to ID field",
model: struct { model: struct {
ID int ID int
Name string Name string
@@ -1307,7 +1307,7 @@ func TestGetPrimaryKeyValue_EdgeCases(t *testing.T) {
expected: 99, expected: 99,
}, },
{ {
name: "model without ID field", name: "model without ID field",
model: struct { model: struct {
Name string Name string
}{Name: "Test"}, }{Name: "Test"},
@@ -1508,10 +1508,10 @@ func TestGetSQLModelColumns_EdgeCases(t *testing.T) {
// Test models with table:, rel:, join: tags for ExtractColumnFromBunTag // Test models with table:, rel:, join: tags for ExtractColumnFromBunTag
type BunSpecialTagsModel struct { type BunSpecialTagsModel struct {
Table string `bun:"table:users"` Table string `bun:"table:users"`
Relation []Post `bun:"rel:has-many"` Relation []Post `bun:"rel:has-many"`
Join string `bun:"join:id=user_id"` Join string `bun:"join:id=user_id"`
NormalCol string `bun:"normal_col"` NormalCol string `bun:"normal_col"`
} }
func TestExtractColumnFromBunTag_SpecialTags(t *testing.T) { func TestExtractColumnFromBunTag_SpecialTags(t *testing.T) {
@@ -1592,8 +1592,8 @@ func TestGetRelationType_GORMFallback(t *testing.T) {
func TestGetRelationType_AdditionalCases(t *testing.T) { func TestGetRelationType_AdditionalCases(t *testing.T) {
// Test model with GORM has-one (pointer without foreignKey or with references) // Test model with GORM has-one (pointer without foreignKey or with references)
type Address struct { type Address struct {
ID int `gorm:"column:id;primaryKey"` ID int `gorm:"column:id;primaryKey"`
UserID int `gorm:"column:user_id"` UserID int `gorm:"column:user_id"`
} }
type UserWithAddress struct { type UserWithAddress struct {
@@ -1609,7 +1609,7 @@ func TestGetRelationType_AdditionalCases(t *testing.T) {
type Employee struct { type Employee struct {
ID int ID int
Company Company // Single struct (not pointer, not slice) - belongs-to Company Company // Single struct (not pointer, not slice) - belongs-to
Coworkers []Employee // Slice without bun/gorm tags - has-many Coworkers []Employee // Slice without bun/gorm tags - has-many
} }
@@ -1920,3 +1920,74 @@ func TestMapToStruct_Errors(t *testing.T) {
}) })
} }
} }
func TestRemoveNonWritableColumns_Generated(t *testing.T) {
type m struct {
ID int `bun:"id,pk"`
Email string `bun:"email"`
Norm string `bun:"email_normalized,generated"`
Scan string `bun:"scan_col,scanonly"`
}
vals := map[string]interface{}{"id": 1, "email": "A", "email_normalized": "a", "scan_col": "x", "dynamic": 1}
RemoveNonWritableColumns(&m{}, vals)
if _, ok := vals["email_normalized"]; ok {
t.Error("generated column not removed")
}
if _, ok := vals["scan_col"]; ok {
t.Error("scanonly column not removed")
}
if len(vals) != 3 {
t.Errorf("unexpected keys: %v", vals)
}
}
func TestNonWritableColumns(t *testing.T) {
type base struct {
Created string `bun:"created_at,scanonly"`
}
type m struct {
base
ID int `bun:"id,pk"`
Email string `bun:"email"`
Norm string `bun:"email_normalized,generated"`
Ro string `gorm:"column:ro;->"`
}
got := NonWritableColumns(&m{})
want := map[string]bool{"created_at": true, "email_normalized": true, "ro": true}
if len(got) != len(want) {
t.Fatalf("got %v", got)
}
for _, c := range got {
if !want[c] {
t.Errorf("unexpected %s", c)
}
}
}
func TestNonWritableColumns_EmbeddedScanOnlyBuffer(t *testing.T) {
type buffer struct {
CQL1 string `json:"cql1,omitempty" gorm:"->" bun:",scanonly"`
RowNumber int64 `json:"_rownumber,omitempty" gorm:"-" bun:",scanonly"`
}
type m struct {
ID int `json:"id" bun:"id,pk"`
Note string `json:"note" bun:"note,type:citext,"`
buffer `json:",omitempty" bun:",scanonly"`
}
got := NonWritableColumns(&m{})
has := map[string]bool{}
for _, c := range got {
has[c] = true
}
if !has["cql1"] {
t.Errorf("cql1 should be non-writable, got %v", got)
}
if has["id"] || has["note"] {
t.Errorf("writable columns reported as non-writable: %v", got)
}
vals := map[string]interface{}{"id": 1, "note": "x", "cql1": "y"}
RemoveNonWritableColumns(&m{}, vals)
if _, ok := vals["cql1"]; ok || len(vals) != 2 {
t.Errorf("unexpected values: %v", vals)
}
}
+196 -126
View File
@@ -1,46 +1,55 @@
# resolvemcp # resolvemcp
Package `resolvemcp` exposes registered database models as **Model Context Protocol (MCP) tools and resources** over HTTP/SSE transport. It mirrors the `resolvespec` package patterns — same model registration API, same filter/sort/pagination/preload options, same lifecycle hook system. Package `resolvemcp` exposes registered database models to AI clients through a **fixed set of Model Context Protocol (MCP) meta tools** over SSE or Streamable HTTP. The tool count does not grow with the number of models. It mirrors the `resolvespec` package — same model registration, filter/sort/pagination/preload options, hook system and security rules.
Every endpoint **requires authentication**; tools run as the authenticated caller.
## Quick Start ## Quick Start
```go ```go
import ( import (
"github.com/bitechdev/ResolveSpec/pkg/resolvemcp" "github.com/bitechdev/ResolveSpec/pkg/resolvemcp"
"github.com/bitechdev/ResolveSpec/pkg/security"
"github.com/gorilla/mux" "github.com/gorilla/mux"
) )
// 1. Create a handler
handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{ handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{
BaseURL: "http://localhost:8080", BaseURL: "http://localhost:8080",
BasePath: "/mcp",
// Read-only by default; uncomment to allow writes:
// ReadOnly: resolvemcp.Bool(false),
}) })
// 2. Register models securityList, _ := security.NewSecurityList(provider)
resolvemcp.RegisterSecurityHooks(handler, securityList)
handler.RegisterModel("public", "users", &User{}) handler.RegisterModel("public", "users", &User{})
handler.RegisterModel("public", "orders", &Order{}) handler.RegisterModel("public", "orders", &Order{})
// 3. Mount routes
r := mux.NewRouter() r := mux.NewRouter()
resolvemcp.SetupMuxRoutes(r, handler) resolvemcp.SetupMuxRoutes(r, handler, securityList) // guarded
``` ```
--- ---
## Config ## Config
```go | Field | Default | Purpose |
type Config struct { |---|---|---|
// BaseURL is the public-facing base URL of the server (e.g. "http://localhost:8080"). | `BaseURL` | request-detected | Public base URL sent to SSE clients |
// Sent to MCP clients during the SSE handshake so they know where to POST messages. | `BasePath` | request-detected | Mount path (e.g. `/mcp`) |
// If empty, it is detected from each incoming request using the Host header and | `DefaultLimit` | 50 | Page size when a read gives no limit |
// TLS state (X-Forwarded-Proto is honoured for reverse-proxy deployments). | `MaxLimit` | 1000 | Larger limits are clamped |
BaseURL string | `MaxOffset` | 100000 | Larger offsets are rejected |
| `MaxBatch` | 100 | Rows in one batch insert |
| `MaxPreloadDepth` | 2 | Depth of a preload path (`a.b.c`) |
| `MaxWriteRows` | 100 | Rows a filter-based update/delete may touch |
| `QueryTimeout` | 30s | One tool call, hooks and queries included |
| `ConfirmTTL` | 5m | Lifetime of a filter-write confirm token |
| `AllowedHosts` | any | Host allowlist for SSE when `BaseURL` is empty (prefer setting `BaseURL`) |
| `EnableAnnotations` | false | Registers `resolvespec_annotate` (opt-in) |
// BasePath is the URL path prefix where MCP endpoints are mounted (e.g. "/mcp"). ---
// Required.
BasePath string
}
```
## Handler Creation ## Handler Creation
@@ -63,14 +72,36 @@ handler.RegisterModel(schema, entity string, model interface{}) error
- `entity` — table/entity name (e.g. `"users"`). - `entity` — table/entity name (e.g. `"users"`).
- `model` — a pointer to a struct (e.g. `&User{}`). - `model` — a pointer to a struct (e.g. `&User{}`).
Each call immediately creates four MCP **tools** and one MCP **resource** for the model. `RegisterModel` only adds the model to the registry; it creates no tools. All registered models are visible to `list_tables`; per-entity rules (see [Security](#security)) restrict the operations.
### Functions
```go
// Go callback
handler.RegisterFunction(resolvemcp.Function{
Name: "recalc_totals",
Description: "Recalculate order totals",
Params: []resolvemcp.FunctionParam{{Name: "order_id", Type: resolvemcp.ParamNumber, Required: true}},
Handler: func(ctx context.Context, tx common.Database, args map[string]any) (any, error) { return nil, nil },
Authorize: func(ctx context.Context) error { return nil }, // optional per-caller gate
})
// SQL procedure: SELECT * FROM public.my_proc($1, $2::jsonb)
handler.RegisterFunction(resolvemcp.Function{
Name: "my_proc", Procedure: "public.my_proc",
Params: []resolvemcp.FunctionParam{{Name: "a", Type: resolvemcp.ParamString, Required: true}},
})
```
Only registered functions are callable. Arguments are validated against `Params`; calls run in a transaction (`OnTxBegin` fired). A function the caller is not authorized for looks identical to an unknown one.
--- ---
## HTTP Transports ## HTTP Transports
`Config.BasePath` is required and used for all route registration. `Config.BasePath` is used for route registration. `Config.BaseURL` is optional — when empty it is detected from each request.
`Config.BaseURL` is optional — when empty it is detected from each request.
All `Setup*`/`New*` helpers wrap the endpoint in `Guard(securityList)`: a valid OAuth bearer token, session token or API key is required, there is no guest/optional mode, and it fails closed. `handler.SSEServer()` / `handler.StreamableHTTPServer()` and the `*Unauthenticated` variants serve **without** a guard and log a warning; use them only behind your own authentication.
Two transports are supported: **SSE** (legacy, two-endpoint) and **Streamable HTTP** (recommended, single-endpoint). Two transports are supported: **SSE** (legacy, two-endpoint) and **Streamable HTTP** (recommended, single-endpoint).
@@ -83,7 +114,7 @@ Two endpoints: `GET {BasePath}/sse` (subscribe) + `POST {BasePath}/message` (sen
#### Gorilla Mux #### Gorilla Mux
```go ```go
resolvemcp.SetupMuxRoutes(r, handler) resolvemcp.SetupMuxRoutes(r, handler, securityList)
``` ```
| Route | Method | Description | | Route | Method | Description |
@@ -94,13 +125,13 @@ resolvemcp.SetupMuxRoutes(r, handler)
#### bunrouter #### bunrouter
```go ```go
resolvemcp.SetupBunRouterRoutes(router, handler) resolvemcp.SetupBunRouterRoutes(router, handler, securityList)
``` ```
#### Gin / net/http / Echo #### Gin / net/http / Echo
```go ```go
sse := handler.SSEServer() sse := resolvemcp.NewSSEServer(handler, securityList) // guarded
engine.Any("/mcp/*path", gin.WrapH(sse)) // Gin engine.Any("/mcp/*path", gin.WrapH(sse)) // Gin
http.Handle("/mcp/", sse) // net/http http.Handle("/mcp/", sse) // net/http
@@ -116,7 +147,7 @@ Single endpoint at `{BasePath}`. Handles POST (client→server) and GET (server
#### Gorilla Mux #### Gorilla Mux
```go ```go
resolvemcp.SetupMuxStreamableHTTPRoutes(r, handler) resolvemcp.SetupMuxStreamableHTTPRoutes(r, handler, securityList)
``` ```
Mounts the handler at `{BasePath}` (all methods). Mounts the handler at `{BasePath}` (all methods).
@@ -124,7 +155,7 @@ Mounts the handler at `{BasePath}` (all methods).
#### bunrouter #### bunrouter
```go ```go
resolvemcp.SetupBunRouterStreamableHTTPRoutes(router, handler) resolvemcp.SetupBunRouterStreamableHTTPRoutes(router, handler, securityList)
``` ```
Registers GET, POST, DELETE on `{BasePath}`. Registers GET, POST, DELETE on `{BasePath}`.
@@ -132,8 +163,7 @@ Registers GET, POST, DELETE on `{BasePath}`.
#### Gin / net/http / Echo #### Gin / net/http / Echo
```go ```go
h := handler.StreamableHTTPServer() h := resolvemcp.NewStreamableHTTPHandler(handler, securityList) // guarded
// or: h := resolvemcp.NewStreamableHTTPHandler(handler)
engine.Any("/mcp", gin.WrapH(h)) // Gin engine.Any("/mcp", gin.WrapH(h)) // Gin
http.Handle("/mcp", h) // net/http http.Handle("/mcp", h) // net/http
@@ -151,6 +181,8 @@ It can operate as:
- **An OAuth2 federation layer** — delegates to external providers (Google, GitHub, Microsoft, etc.) - **An OAuth2 federation layer** — delegates to external providers (Google, GitHub, Microsoft, etc.)
- **Both simultaneously** - **Both simultaneously**
> The underlying `security.OAuthServer` also supports consent, OpenID Connect, rotating refresh tokens, JWT access tokens, DPoP, PAR, the device grant and token exchange; they are opt-in `OAuthServerConfig` options described in [pkg/security/OAUTH2_SERVER.md](../security/OAUTH2_SERVER.md). The options of `resolvemcp.OAuth2Config` are unchanged.
### Standard endpoints served ### Standard endpoints served
| Path | Spec | Purpose | | Path | Spec | Purpose |
@@ -187,7 +219,7 @@ handler.EnableOAuthServer(security.OAuthServerConfig{
provider, _ := security.NewCompositeSecurityProvider(auth, colSec, rowSec) provider, _ := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
securityList, _ := security.NewSecurityList(provider) securityList, _ := security.NewSecurityList(provider)
security.RegisterSecurityHooks(handler, securityList) resolvemcp.RegisterSecurityHooks(handler, securityList)
http.ListenAndServe(":8080", handler.HTTPHandler(securityList)) http.ListenAndServe(":8080", handler.HTTPHandler(securityList))
``` ```
@@ -286,7 +318,10 @@ resolvemcp.SetupMuxRoutesWithAuth(r, handler, securityList)
```go ```go
import "github.com/bitechdev/ResolveSpec/pkg/security" import "github.com/bitechdev/ResolveSpec/pkg/security"
securityList := security.NewSecurityList(mySecurityProvider) securityList, err := security.NewSecurityList(mySecurityProvider)
if err != nil {
log.Fatal(err)
}
resolvemcp.RegisterSecurityHooks(handler, securityList) resolvemcp.RegisterSecurityHooks(handler, securityList)
``` ```
@@ -294,12 +329,17 @@ Call `RegisterSecurityHooks` **once**, after creating the handler and before reg
| Hook | Effect | | Hook | Effect |
|---|---| |---|---|
| `BeforeHandle` | Enforces per-entity operation rules (see below) | | `OnTxBegin` | Stamps transaction-local settings (RLS GUCs) set with `SecurityList.SetTxSettings` |
| `BeforeHandle` | Enforces per-entity operation rules (see below); preloads column rules for writes |
| `BeforeRead` | Loads RLS/CLS rules, then injects a user-scoped WHERE clause | | `BeforeRead` | Loads RLS/CLS rules, then injects a user-scoped WHERE clause |
| `BeforeScan` | Applies row security to the row an update or delete targets; a row the user cannot see is "not found" |
| `AfterRead` | Masks/hides columns per column-security rules; writes audit log | | `AfterRead` | Masks/hides columns per column-security rules; writes audit log |
| `BeforeUpdate` | Blocks update if `CanUpdate` is false | | `BeforeCreate` | Blocks create if `CanCreate` is false; drops hidden/masked columns from the payload |
| `BeforeUpdate` | Blocks update if `CanUpdate` is false; drops hidden/masked columns from the payload |
| `BeforeDelete` | Blocks delete if `CanDelete` is false | | `BeforeDelete` | Blocks delete if `CanDelete` is false |
Additional hooks: `BeforeScan` (row pre-read of update/delete/filter writes), `BeforeCall`/`AfterCall` (functions) and `OnTxBegin`. Hooks are mutex-protected and panics in hooks are recovered.
### Per-entity operation rules ### Per-entity operation rules
Use `RegisterModelWithRules` instead of `RegisterModel` to set access rules at registration time: Use `RegisterModelWithRules` instead of `RegisterModel` to set access rules at registration time:
@@ -354,123 +394,139 @@ handler.SetModelRules("public", "users", modelregistry.ModelRules{
--- ---
## MCP Tools ## Describing the API for agents
### Tool Naming Give agents context about what each table is for:
``` ```go
{operation}_{schema}_{entity} // e.g. read_public_users // 1. Explicitly, in code
{operation}_{entity} // e.g. read_users (when schema is empty) handler.SetModelDescription("public", "users", modelregistry.ModelInfo{
Description: "Application accounts",
Purpose: "Look up who a person is",
Tags: []string{"identity"},
Columns: map[string]string{"email": "Login address, unique"},
})
// 2. From an external JSON map (keyed by "schema.entity"); entries here win
n, err := handler.LoadModelDescriptions("docs/model-descriptions.json")
``` ```
Operations: `read`, `create`, `update`, `delete`. Example `docs/model-descriptions.json` (every key is optional; unknown keys are rejected):
### Read Tool — `read_{schema}_{entity}`
Fetch one or many records.
| Argument | Type | Description |
|---|---|---|
| `id` | string | Primary key value. Omit to return multiple records. |
| `limit` | number | Max records per page (recommended: 10–100). |
| `offset` | number | Records to skip (offset-based pagination). |
| `cursor_forward` | string | PK of the **last** record on the current page (next-page cursor). |
| `cursor_backward` | string | PK of the **first** record on the current page (prev-page cursor). |
| `columns` | array | Column names to include. Omit for all columns. |
| `omit_columns` | array | Column names to exclude. |
| `filters` | array | Filter objects (see [Filtering](#filtering)). |
| `sort` | array | Sort objects (see [Sorting](#sorting)). |
| `preloads` | array | Relation preload objects (see [Preloading](#preloading)). |
**Response:**
```json ```json
{ {
"success": true, "public.users": {
"data": [...], "description": "Application accounts, one row per person who can sign in.",
"metadata": { "purpose": "Look up who someone is. Use public.orders for what they bought.",
"total": 100, "tags": ["identity", "pii"],
"filtered": 100, "columns": {
"count": 10, "id": "Internal account id",
"limit": 10, "email": "Login address, unique and lower-cased",
"offset": 0 "created_at": "When the account was created (UTC)"
}
},
"public.orders": {
"description": "Customer orders.",
"columns": {
"status": "One of: pending, paid, shipped, cancelled"
}
} }
} }
``` ```
### Create Tool — `create_{schema}_{entity}` Keys are `schema.entity` names as registered with `RegisterModel`. Column keys are the JSON column names shown by `describe_table`. An entry replaces any info set earlier for that table, and columns it leaves out still fall back to field tags.
Insert one or more records. Fallbacks when nothing is set for a table or column, in order: the model's `ModelDescription() string` method (table), then field tags (column): `comment`, `note`, `desc` or `description` tags, then `comment:` inside the `gorm` or `bun` tag.
The text appears in `list_tables` and `describe_table`. The server also sends a short usage guide as MCP `instructions` on connect.
### Catalogue file
`handler.ExportCatalog(path)` writes the usage guide, tools, limits and every table (columns, types, keys, relations, allowed operations, descriptions) to disk, JSON for a `.json` path and Markdown otherwise. The file is replaced atomically. It lists every table with at least one allowed operation, regardless of caller, so keep it out of public directories. Call it after registering models (for example at startup, or from a `go generate` step).
## Read-only mode
The server is **read-only unless you enable writes**: `Config.ReadOnly` is a `*bool` and an unset (nil) value means on. To allow inserts, updates and deletes:
```go
handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{ReadOnly: resolvemcp.Bool(false)})
```
While read-only is on:
- The insert, update, delete and annotation tools are not registered, so the agent never sees them. `list_functions`/`call_function` are off too, because a registered function may change data, unless you set `AllowFunctionCalls` (below).
- `list_tables` and `describe_table` report only `select`; `describe_table` also sets `read_only: true` and lists no writable columns.
- The MCP server instructions (and the exported catalogue) say the server is read-only and tell the agent not to attempt writes.
- A write that reaches a handler anyway is refused with a `forbidden` error ("this server is read-only: writes are disabled").
### Function calls and the allowlist
```go
// Read-only server that may still run two named functions
resolvemcp.Config{
// ReadOnly is on by default
AllowFunctionCalls: true, // keep list_functions / call_function on a read-only server
AllowedFunctions: []string{"report_totals", "search_customers"},
}
```
- `AllowFunctionCalls` only matters while read-only is on; with writes enabled (`ReadOnly: resolvemcp.Bool(false)`), functions are always available. Set it only for functions that do not change data.
- `AllowedFunctions` works in either mode. When empty, every registered function is allowed. When set, only the named functions are listed and callable; any other is reported as `unknown function`, so its existence is not revealed. Per-function `Authorize` still applies on top.
## MCP Tools
Fixed set, independent of the models. `table` is `schema.entity`. Errors return `{"success":false,"error":{"code","message"}}` with codes `invalid_argument`, `not_found`, `forbidden`, `limit_exceeded`, `internal` (internal details are logged, the client gets a reference id).
| Tool | Purpose |
|---|---|
| `list_tables` | Tables the caller may use and the allowed operations |
| `describe_table` | Columns, PK, relations, writable columns, operations, limits |
| `select_table` | Read rows (filters, sort, columns, preloads, paging) |
| `insert_into_table` | Insert one row or a capped batch |
| `update_table` | Update by `id` or `filters` |
| `delete_from_table` | Delete by `id` or `filters` |
| `list_functions` | Registered functions the caller may call, with parameters |
| `call_function` | Call a registered function |
| `resolvespec_annotate` | Only with `EnableAnnotations` |
### `select_table`
| Argument | Type | Description | | Argument | Type | Description |
|---|---|---| |---|---|---|
| `data` | object \| array | Single object or array of objects to insert. | | `table` | string (required) | `schema.entity` |
| `id` | string | Primary key of one row |
| `filters`, `sort` | array | See [Filtering](#filtering), [Sorting](#sorting) |
| `columns`, `omit_columns` | array | Column selection |
| `preloads` | array | Relations (validated against the model, max depth `MaxPreloadDepth`) |
| `limit`, `offset` | number | Clamped to `MaxLimit` / rejected above `MaxOffset` |
| `cursor_forward`, `cursor_backward` | string | PK cursor, requires `sort` |
| `include_count` | boolean | Also compute totals (slower); otherwise `total`/`filtered` are 0 |
Array input runs inside a single transaction — all succeed or all fail. Response: `{"success":true,"data":[...],"metadata":{"total","filtered","count","limit","offset"}}`
**Response:** ### `insert_into_table`
```json
{ "success": true, "data": { ... } }
```
### Update Tool — `update_{schema}_{entity}` `data` is an object or an array (one transaction, max `MaxBatch`). Unknown, duplicate or read-only keys are rejected; keys are resolved to columns from the model.
Partially update an existing record. Only non-null, non-empty fields in `data` are applied; existing values are preserved for omitted fields. ### `update_table` / `delete_from_table`
| Argument | Type | Description | Either `id` or `filters` is required.
|---|---|---|
| `id` | string | Primary key of the record. Can also be included inside `data`. |
| `data` | object (required) | Fields to update. |
**Response:** | Mode | Behaviour |
```json |---|---|
{ "success": true, "data": { ...merged record... } } | `id` | One row, applied immediately. The row is locked and row security applies; an invisible row is "not found". |
``` | `filters` | Matching rows are found inside the transaction (row security applied, max `MaxWriteRows`). The first call returns a preview and a `confirm_token`; repeat the identical call with `confirm_token` to apply. |
| `dry_run` | Report match count and preview ids; change nothing. |
### Delete Tool — `delete_{schema}_{entity}` The token is single-use, expires after `ConfirmTTL`, and is bound to user, table, operation and a hash of filters, data and matched ids; it is held in memory (lost on restart, single instance). Update changes only the keys in `data`; `null` sets NULL. Filters are strictly parsed (never silently dropped), columns validated, and only the documented operators are accepted.
Delete a record by primary key. **Irreversible.** ### `call_function`
| Argument | Type | Description | `name` and `arguments` (object). See [Functions](#functions).
|---|---|---|
| `id` | string (required) | Primary key of the record to delete. |
**Response:** ### `resolvespec_annotate`
```json
{ "success": true, "data": { ...deleted record... } }
```
### Annotation Tool — `resolvespec_annotate` Opt-in (`Config.EnableAnnotations`). Stores/retrieves freeform annotations through `resolvespec_set_annotation` / `resolvespec_get_annotation`; runs `BeforeHandle` hooks (`annotate_set` / `annotate_get`) and a transaction.
Store or retrieve freeform annotation records for any tool, model, or entity. Registered automatically on every handler.
| Argument | Type | Description |
|---|---|---|
| `tool_name` | string (required) | Key to annotate — an MCP tool name (e.g. `read_public_users`), a model name (e.g. `public.users`), or any other identifier. |
| `annotations` | object | Annotation data to persist. Omit to retrieve existing annotations instead. |
**Set annotations** (calls `resolvespec_set_annotation(tool_name, annotations)`):
```json
{ "tool_name": "read_public_users", "annotations": { "description": "Returns active users", "owner": "platform-team" } }
```
**Response:**
```json
{ "success": true, "tool_name": "read_public_users", "action": "set" }
```
**Get annotations** (calls `resolvespec_get_annotation(tool_name)`):
```json
{ "tool_name": "read_public_users" }
```
**Response:**
```json
{ "success": true, "tool_name": "read_public_users", "action": "get", "annotations": { ... } }
```
---
### Resource — `{schema}.{entity}`
Each model is also registered as an MCP resource with URI `schema.entity` (or just `entity` when schema is empty). Reading the resource returns up to 100 records as `application/json`.
--- ---
@@ -570,6 +626,9 @@ Hooks let you intercept and modify CRUD operations at well-defined lifecycle poi
| `BeforeCreate` / `AfterCreate` | Around insert | | `BeforeCreate` / `AfterCreate` | Around insert |
| `BeforeUpdate` / `AfterUpdate` | Around update | | `BeforeUpdate` / `AfterUpdate` | Around update |
| `BeforeDelete` / `AfterDelete` | Around delete | | `BeforeDelete` / `AfterDelete` | Around delete |
| `BeforeScan` | Row pre-read for update/delete/filter writes |
| `BeforeCall` / `AfterCall` | Around `call_function` |
| `OnTxBegin` | Start of every transaction |
### Registering Hooks ### Registering Hooks
@@ -599,7 +658,7 @@ handler.Hooks().RegisterMultiple(
| `Entity` | `string` | Entity/table name | | `Entity` | `string` | Entity/table name |
| `Model` | `interface{}` | Registered model instance | | `Model` | `interface{}` | Registered model instance |
| `Options` | `common.RequestOptions` | Parsed request options (read operations) | | `Options` | `common.RequestOptions` | Parsed request options (read operations) |
| `Operation` | `string` | `"read"`, `"create"`, `"update"`, or `"delete"` | | `Operation` | `string` | `"read"`, `"create"`, `"update"`, `"delete"`, `"call"`, `"annotate_set"` or `"annotate_get"` |
| `ID` | `string` | Primary key from request (read/update/delete) | | `ID` | `string` | Primary key from request (read/update/delete) |
| `Data` | `interface{}` | Input data (create/update — modifiable) | | `Data` | `interface{}` | Input data (create/update — modifiable) |
| `Result` | `interface{}` | Output data (set by After hooks) | | `Result` | `interface{}` | Output data (set by After hooks) |
@@ -633,7 +692,7 @@ registry.ClearAll() // remove all hooks
## Context Helpers ## Context Helpers
Request metadata is threaded through `context.Context` during handler execution. Hooks and custom tools can read it: The caller's `security.UserContext` reaches every tool call through the request context. Request metadata is threaded through `context.Context` during handler execution. Hooks and custom tools can read it:
```go ```go
schema := resolvemcp.GetSchema(ctx) schema := resolvemcp.GetSchema(ctx)
@@ -653,7 +712,7 @@ ctx = resolvemcp.WithSchema(ctx, "tenant_a")
## Adding Custom MCP Tools ## Adding Custom MCP Tools
Access the underlying `*server.MCPServer` to register additional tools: Access the underlying `*server.MCPServer` to register additional tools (they sit behind the same guard). Prefer `RegisterFunction` for database-backed actions:
```go ```go
mcpServer := handler.MCPServer() mcpServer := handler.MCPServer()
@@ -669,3 +728,14 @@ The handler resolves table names in priority order:
1. `TableNameProvider` interface — `TableName() string` (can return `"schema.table"`) 1. `TableNameProvider` interface — `TableName() string` (can return `"schema.table"`)
2. `SchemaProvider` interface — `SchemaName() string` (combined with entity name) 2. `SchemaProvider` interface — `SchemaName() string` (combined with entity name)
3. Fallback: `schema.entity` (or `schema_entity` for SQLite) 3. Fallback: `schema.entity` (or `schema_entity` for SQLite)
---
## Breaking changes
- The server is read-only by default. Writes (insert/update/delete), annotations and function calls need `Config{ReadOnly: resolvemcp.Bool(false)}` (function calls can also be kept on a read-only server with `AllowFunctionCalls`).
- Per-model tools (`read_/create_/update_/delete_{schema}_{entity}`) and per-model resources are gone; use the meta tools.
- `Setup*` / `NewSSEServer` / `NewStreamableHTTPHandler` take a `*security.SecurityList` and require authentication. `OptionalAuth*` helpers were removed; `*Unauthenticated` variants exist for explicit opt-out.
- `resolvespec_annotate` is opt-in via `Config.EnableAnnotations`.
- `Handler.Build()` is not needed.
- Update is now a partial update by validated keys; reads are capped by the configured limits.
+39 -2
View File
@@ -6,6 +6,8 @@ import (
"fmt" "fmt"
"github.com/mark3labs/mcp-go/mcp" "github.com/mark3labs/mcp-go/mcp"
"github.com/bitechdev/ResolveSpec/pkg/common"
) )
const annotationToolName = "resolvespec_annotate" const annotationToolName = "resolvespec_annotate"
@@ -50,13 +52,42 @@ func registerAnnotationTool(h *Handler) {
}) })
} }
// maxAnnotationKey caps tool_name so the key space cannot be abused as storage.
const maxAnnotationKey = 200
// annotationGate runs the BeforeHandle hooks for an annotation call and returns the hook
// context. A hook error (e.g. authentication required) is returned to the caller.
func annotationGate(ctx context.Context, h *Handler, operation, toolName string) (*HookContext, error) {
if len(toolName) > maxAnnotationKey {
return nil, fmt.Errorf("tool_name too long")
}
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Entity: toolName,
Operation: operation,
Tx: h.db,
}
if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil {
return nil, err
}
return hookCtx, nil
}
func executeSetAnnotation(ctx context.Context, h *Handler, toolName string, annotations interface{}) (*mcp.CallToolResult, error) { func executeSetAnnotation(ctx context.Context, h *Handler, toolName string, annotations interface{}) (*mcp.CallToolResult, error) {
hookCtx, err := annotationGate(ctx, h, "annotate_set", toolName)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
jsonBytes, err := json.Marshal(annotations) jsonBytes, err := json.Marshal(annotations)
if err != nil { if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("failed to marshal annotations: %v", err)), nil return mcp.NewToolResultError(fmt.Sprintf("failed to marshal annotations: %v", err)), nil
} }
_, err = h.db.Exec(ctx, "SELECT resolvespec_set_annotation($1, $2)", toolName, string(jsonBytes)) err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
_, err := tx.Exec(ctx, "SELECT resolvespec_set_annotation($1, $2)", toolName, string(jsonBytes))
return err
})
if err != nil { if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("failed to set annotation: %v", err)), nil return mcp.NewToolResultError(fmt.Sprintf("failed to set annotation: %v", err)), nil
} }
@@ -69,8 +100,14 @@ func executeSetAnnotation(ctx context.Context, h *Handler, toolName string, anno
} }
func executeGetAnnotation(ctx context.Context, h *Handler, toolName string) (*mcp.CallToolResult, error) { func executeGetAnnotation(ctx context.Context, h *Handler, toolName string) (*mcp.CallToolResult, error) {
hookCtx, err := annotationGate(ctx, h, "annotate_get", toolName)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
var rows []map[string]interface{} var rows []map[string]interface{}
err := h.db.Query(ctx, &rows, "SELECT resolvespec_get_annotation($1)", toolName) err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
return tx.Query(ctx, &rows, "SELECT resolvespec_get_annotation($1)", toolName)
})
if err != nil { if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("failed to get annotation: %v", err)), nil return mcp.NewToolResultError(fmt.Sprintf("failed to get annotation: %v", err)), nil
} }
+329
View File
@@ -0,0 +1,329 @@
package resolvemcp
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"reflect"
"sort"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
)
// usageGuide is the short agent-facing guide sent as the MCP server instructions and
// embedded in the exported catalogue. It is generic: it never mentions concrete models.
const usageGuide = `This server exposes database tables through a fixed set of tools.
1. Call list_tables to see the tables you may use, what they hold and the operations allowed.
2. Call describe_table for a table before using it: columns, types, primary key, relations (preloadable), writable columns and limits.
3. Read with select_table (filters, sort, columns, preloads). Results are paged; use limit/offset or cursors, and include_count only when you need a total.
4. Write with insert_into_table, update_table, delete_from_table. Address one row by id, or several by filters. A filter-based write first returns a preview; repeat the call with the confirm_token to apply it (dry_run only previews).
5. Use list_functions / call_function for registered functions.
Read the error message when a call fails: it says which argument was wrong.`
// readOnlyGuide replaces usageGuide on a read-only server.
const readOnlyGuide = `This server exposes database tables through a fixed set of tools. It is READ-ONLY: you cannot insert, update or delete data or write annotations, and no tool for that exists. Do not attempt a write; tell the user it is not possible through this server.
1. Call list_tables to see the tables you may read and what they hold.
2. Call describe_table for a table before using it: columns, types, primary key, relations (preloadable) and limits.
3. Read with select_table (filters, sort, columns, preloads). Results are paged; use limit/offset or cursors, and include_count only when you need a total.
Read the error message when a call fails: it says which argument was wrong.`
// readOnlyFunctionsGuide is the extra step of a read-only server that still allows functions.
const readOnlyFunctionsGuide = `
4. Use list_functions / call_function for the registered functions. Only call functions that fit a read-only server; the server decides what is allowed.`
// guideFor returns the usage guide for the server mode.
func guideFor(readOnly, functions bool) string {
if !readOnly {
return usageGuide
}
if functions {
return readOnlyGuide + readOnlyFunctionsGuide
}
return readOnlyGuide
}
// Catalog is a snapshot of what the server offers: the usage guide, the tools, the limits
// and every table with its columns, relations, allowed operations and descriptions.
type Catalog struct {
GeneratedAt time.Time `json:"generated_at"`
Server string `json:"server"`
Version string `json:"version"`
ReadOnly bool `json:"read_only"`
Guide string `json:"guide"`
Limits CatalogLimits `json:"limits"`
Tools []CatalogTool `json:"tools"`
Tables []CatalogTable `json:"tables"`
}
// CatalogLimits mirrors the configured server limits.
type CatalogLimits struct {
DefaultLimit int `json:"default_limit"`
MaxLimit int `json:"max_limit"`
MaxOffset int `json:"max_offset"`
MaxBatch int `json:"max_batch"`
MaxPreloadDepth int `json:"max_preload_depth"`
MaxWriteRows int `json:"max_write_rows"`
}
// CatalogTool is one MCP tool.
type CatalogTool struct {
Name string `json:"name"`
Description string `json:"description"`
}
// CatalogTable is one table in the catalogue.
type CatalogTable struct {
Table string `json:"table"`
Description string `json:"description,omitempty"`
Purpose string `json:"purpose,omitempty"`
Tags []string `json:"tags,omitempty"`
Operations []string `json:"operations"`
PrimaryKey string `json:"primary_key,omitempty"`
Columns []CatalogColumn `json:"columns"`
Relations []string `json:"relations,omitempty"`
}
// CatalogColumn is one column of a table.
type CatalogColumn struct {
Name string `json:"name"`
Type string `json:"type,omitempty"`
Nullable bool `json:"nullable"`
PrimaryKey bool `json:"primary_key,omitempty"`
Unique bool `json:"unique,omitempty"`
Writable bool `json:"writable"`
Description string `json:"description,omitempty"`
}
// modelDocs returns the effective documentation of a table. Registry info (including a
// loaded descriptions file) wins, then the model's ModelDescription().
func (h *Handler) modelDocs(schema, entity string) modelregistry.ModelInfo {
if reg, ok := h.registry.(*modelregistry.DefaultModelRegistry); ok {
return reg.ResolveModelInfo(buildModelName(schema, entity))
}
return modelregistry.ModelInfo{}
}
// columnDescription picks a column's description: the registry/file map first, then the
// struct-tag comment.
func columnDescription(docs modelregistry.ModelInfo, c columnInfo) string {
if d := docs.Columns[c.jsonName]; d != "" {
return d
}
return c.comment
}
// SetModelDescription stores documentation for a registered or soon-to-be-registered table.
// It returns an error when the handler's registry does not keep model info.
func (h *Handler) SetModelDescription(schema, entity string, info modelregistry.ModelInfo) error {
reg, ok := h.registry.(*modelregistry.DefaultModelRegistry)
if !ok {
return fmt.Errorf("resolvemcp: registry does not support model descriptions (use NewHandlerWithGORM/Bun/DB)")
}
reg.SetModelInfo(buildModelName(schema, entity), info)
return nil
}
// LoadModelDescriptions loads an external JSON map of descriptions keyed by "schema.entity"
// (see modelregistry.LoadModelInfo for the format) and returns how many tables it covered.
// Loaded entries override the model's own comments.
func (h *Handler) LoadModelDescriptions(path string) (int, error) {
reg, ok := h.registry.(*modelregistry.DefaultModelRegistry)
if !ok {
return 0, fmt.Errorf("resolvemcp: registry does not support model descriptions (use NewHandlerWithGORM/Bun/DB)")
}
return reg.LoadModelInfoFile(path)
}
// BuildCatalog snapshots the server's tools and tables. Tables with no allowed operation are
// left out, exactly as list_tables does.
func (h *Handler) BuildCatalog() Catalog {
cat := Catalog{
GeneratedAt: time.Now().UTC(),
Server: h.name,
Version: h.version,
ReadOnly: h.config.readOnly,
Guide: guideFor(h.config.readOnly, h.config.AllowFunctionCalls),
Limits: CatalogLimits{
DefaultLimit: h.config.DefaultLimit,
MaxLimit: h.config.MaxLimit,
MaxOffset: h.config.MaxOffset,
MaxBatch: h.config.MaxBatch,
MaxPreloadDepth: h.config.MaxPreloadDepth,
MaxWriteRows: h.config.MaxWriteRows,
},
Tools: []CatalogTool{},
Tables: []CatalogTable{},
}
for name, tool := range h.mcpServer.ListTools() {
cat.Tools = append(cat.Tools, CatalogTool{Name: name, Description: tool.Tool.Description})
}
sort.Slice(cat.Tools, func(i, j int) bool { return cat.Tools[i].Name < cat.Tools[j].Name })
for name, model := range h.registry.GetAllModels() {
schema, entity, _ := splitTable(name)
rules := h.modelRules(schema, entity)
ops := h.opsFor(rules)
if len(ops) == 0 {
continue
}
info := buildModelInfo(schema, entity, model)
docs := h.modelDocs(schema, entity)
writable := map[string]bool{}
mt := reflect.TypeOf(model)
for mt != nil && (mt.Kind() == reflect.Pointer || mt.Kind() == reflect.Slice) {
mt = mt.Elem()
}
if !h.config.readOnly && mt != nil && mt.Kind() == reflect.Struct {
for k := range reflectionJSONColumns(mt) {
writable[k] = true
}
}
t := CatalogTable{
Table: info.fullName,
Description: docs.Description,
Purpose: docs.Purpose,
Tags: docs.Tags,
Operations: ops,
PrimaryKey: info.pkName,
Relations: info.relationNames,
Columns: make([]CatalogColumn, 0, len(info.columns)),
}
for _, c := range info.columns {
typ := c.sqlType
if typ == "" {
typ = c.goType
}
t.Columns = append(t.Columns, CatalogColumn{
Name: c.jsonName, Type: typ, Nullable: c.nullable, PrimaryKey: c.isPrimary,
Unique: c.isUnique, Writable: writable[c.jsonName], Description: columnDescription(docs, c),
})
}
cat.Tables = append(cat.Tables, t)
}
sort.Slice(cat.Tables, func(i, j int) bool { return cat.Tables[i].Table < cat.Tables[j].Table })
return cat
}
// ExportCatalog writes the catalogue to path. A ".json" extension writes JSON; anything
// else writes Markdown. The file is replaced atomically (written to a temp file in the same
// directory, then renamed) and created with mode 0600. It lists every table the registry
// allows any operation on, regardless of caller, so keep it out of public directories.
func (h *Handler) ExportCatalog(path string) error {
cat := h.BuildCatalog()
var data []byte
if strings.EqualFold(filepath.Ext(path), ".json") {
b, err := json.MarshalIndent(cat, "", " ")
if err != nil {
return err
}
data = b
data = append(data, '\n')
} else {
data = []byte(cat.Markdown())
}
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0o750); err != nil {
return fmt.Errorf("resolvemcp: export catalog: %w", err)
}
tmp, err := os.CreateTemp(dir, ".catalog-*")
if err != nil {
return fmt.Errorf("resolvemcp: export catalog: %w", err)
}
tmpName := tmp.Name()
_, werr := tmp.Write(data)
cerr := tmp.Close()
if werr == nil {
werr = cerr
}
if werr == nil {
werr = os.Rename(tmpName, path)
}
if werr != nil {
_ = os.Remove(tmpName)
return fmt.Errorf("resolvemcp: export catalog: %w", werr)
}
return nil
}
// Markdown renders the catalogue as a Markdown document.
func (c Catalog) Markdown() string {
var sb strings.Builder
fmt.Fprintf(&sb, "# %s API catalogue\n\nGenerated %s.\n\n", c.Server, c.GeneratedAt.Format(time.RFC3339))
if c.ReadOnly {
sb.WriteString("**This server is read-only.**\n\n")
}
sb.WriteString("## How to use\n\n" + c.Guide + "\n\n")
fmt.Fprintf(&sb, "## Limits\n\ndefault limit %d, max limit %d, max offset %d, max batch %d, max preload depth %d, max rows per filter write %d.\n\n",
c.Limits.DefaultLimit, c.Limits.MaxLimit, c.Limits.MaxOffset, c.Limits.MaxBatch, c.Limits.MaxPreloadDepth, c.Limits.MaxWriteRows)
sb.WriteString("## Tools\n\n")
for _, t := range c.Tools {
fmt.Fprintf(&sb, "- `%s`: %s\n", t.Name, oneLine(t.Description))
}
sb.WriteString("\n## Tables\n\n")
if len(c.Tables) == 0 {
sb.WriteString("No tables are registered.\n")
}
for i := range c.Tables {
t := &c.Tables[i]
fmt.Fprintf(&sb, "### %s\n\n", t.Table)
if t.Description != "" {
sb.WriteString(t.Description + "\n\n")
}
if t.Purpose != "" {
sb.WriteString("Purpose: " + t.Purpose + "\n\n")
}
if len(t.Tags) > 0 {
sb.WriteString("Tags: " + strings.Join(t.Tags, ", ") + "\n\n")
}
fmt.Fprintf(&sb, "Operations: %s", strings.Join(t.Operations, ", "))
if t.PrimaryKey != "" {
fmt.Fprintf(&sb, " · Primary key: `%s`", t.PrimaryKey)
}
sb.WriteString("\n\n| Column | Type | Flags | Description |\n|---|---|---|---|\n")
for _, col := range t.Columns {
var flags []string
if col.PrimaryKey {
flags = append(flags, "pk")
}
if col.Unique {
flags = append(flags, "unique")
}
if col.Nullable {
flags = append(flags, "nullable")
}
if !col.Writable {
flags = append(flags, "read-only")
}
fmt.Fprintf(&sb, "| `%s` | %s | %s | %s |\n", col.Name, mdCell(col.Type), strings.Join(flags, ", "), mdCell(col.Description))
}
if len(t.Relations) > 0 {
sb.WriteString("\nRelations (preloadable): " + strings.Join(t.Relations, ", ") + "\n")
}
sb.WriteString("\n")
}
return sb.String()
}
func oneLine(s string) string {
return strings.Join(strings.Fields(s), " ")
}
func mdCell(s string) string {
return strings.ReplaceAll(oneLine(s), "|", `\|`)
}
// reflectionJSONColumns returns the JSON names of the columns a write may set.
func reflectionJSONColumns(t reflect.Type) map[string]string {
return reflection.BuildJSONToDBColumnMap(t)
}
+126
View File
@@ -0,0 +1,126 @@
package resolvemcp
import (
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
)
type docItem struct {
ID int `json:"id" bun:"id,pk"`
Email string `json:"email" bun:"email,comment:Tag comment"`
Name string `json:"name" bun:"name" note:"Name from tag"`
}
func (docItem) ModelDescription() string { return "Model-level fallback" }
func newDocHandler(t *testing.T) *Handler {
t.Helper()
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(), Config{})
if err := h.RegisterModel("public", "items", &docItem{}); err != nil {
t.Fatal(err)
}
hidden := modelregistry.ModelRules{} // no operation allowed
if err := h.RegisterModelWithRules("public", "secret", &docItem{}, hidden); err != nil {
t.Fatal(err)
}
return h
}
func TestCatalogDescriptionsPrecedence(t *testing.T) {
h := newDocHandler(t)
cat := h.BuildCatalog()
if len(cat.Tables) != 1 || cat.Tables[0].Table != "public.items" {
t.Fatalf("tables = %+v (hidden table must be left out)", cat.Tables)
}
tb := cat.Tables[0]
if tb.Description != "Model-level fallback" {
t.Errorf("fallback description = %q", tb.Description)
}
col := map[string]string{}
for _, c := range tb.Columns {
col[c.Name] = c.Description
}
if col["email"] != "Tag comment" || col["name"] != "Name from tag" {
t.Errorf("tag comments = %v", col)
}
path := filepath.Join(t.TempDir(), "desc.json")
if err := os.WriteFile(path, []byte(`{"public.items":{"description":"From file","columns":{"email":"File email"}}}`), 0o600); err != nil {
t.Fatal(err)
}
if n, err := h.LoadModelDescriptions(path); err != nil || n != 1 {
t.Fatalf("load n=%d err=%v", n, err)
}
tb = h.BuildCatalog().Tables[0]
if tb.Description != "From file" {
t.Errorf("file must win, got %q", tb.Description)
}
for _, c := range tb.Columns {
switch c.Name {
case "email":
if c.Description != "File email" {
t.Errorf("email = %q", c.Description)
}
case "name":
if c.Description != "Name from tag" {
t.Errorf("name must fall back to tag, got %q", c.Description)
}
}
}
}
func TestExportCatalogFiles(t *testing.T) {
h := newDocHandler(t)
dir := t.TempDir()
md := filepath.Join(dir, "sub", "catalog.md")
if err := h.ExportCatalog(md); err != nil {
t.Fatal(err)
}
b, _ := os.ReadFile(md)
for _, want := range []string{"# resolvemcp API catalogue", "### public.items", "Model-level fallback", "`list_tables`", "Tag comment"} {
if !strings.Contains(string(b), want) {
t.Errorf("markdown missing %q", want)
}
}
if strings.Contains(string(b), "public.secret") {
t.Error("table without operations leaked into the catalogue")
}
js := filepath.Join(dir, "catalog.json")
if err := h.ExportCatalog(js); err != nil {
t.Fatal(err)
}
var cat Catalog
b, _ = os.ReadFile(js)
if err := json.Unmarshal(b, &cat); err != nil || len(cat.Tables) != 1 || len(cat.Tools) == 0 {
t.Fatalf("json catalog bad: err=%v %+v", err, cat)
}
entries, _ := os.ReadDir(dir)
for _, e := range entries {
if strings.HasPrefix(e.Name(), ".catalog-") {
t.Errorf("temp file left behind: %s", e.Name())
}
}
}
func TestDescribeAndListIncludeDescriptions(t *testing.T) {
h := newDocHandler(t)
res, _ := h.handleListTables(nil, callReq(nil))
tables, _ := payload(t, res)["tables"].([]any)
if len(tables) != 1 || tables[0].(map[string]any)["description"] != "Model-level fallback" {
t.Errorf("list_tables = %v", tables)
}
res, _ = h.handleDescribeTable(nil, callReq(map[string]any{"table": "public.items"}))
p := payload(t, res)
if p["description"] != "Model-level fallback" {
t.Errorf("describe_table description = %v", p["description"])
}
}
+83
View File
@@ -0,0 +1,83 @@
package resolvemcp
import (
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"sync"
"time"
)
// confirmStore holds the single-use confirmation tokens for filter-based writes. It is
// in-memory: tokens are lost on restart and are not shared between instances, which only costs
// the client one more preview call.
type confirmStore struct {
mu sync.Mutex
tokens map[string]confirmEntry
now func() time.Time
maxLive int
}
type confirmEntry struct {
user, table, op, binding string
expires time.Time
}
func newConfirmStore() *confirmStore {
return &confirmStore{tokens: map[string]confirmEntry{}, now: time.Now, maxLive: 10000}
}
var errConfirmInvalid = NewClientError(CodeInvalidArgument, "confirm_token is invalid or expired; repeat the call without it to get a new preview")
// issue returns a token bound to the caller, table, operation and binding (a hash of the
// filters, data and matched rows the preview showed).
func (c *confirmStore) issue(user, table, op, binding string, ttl time.Duration) (string, error) {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
return "", err
}
tok := hex.EncodeToString(b[:])
c.mu.Lock()
defer c.mu.Unlock()
now := c.now()
for k, e := range c.tokens {
if now.After(e.expires) {
delete(c.tokens, k)
}
}
if len(c.tokens) >= c.maxLive {
return "", NewClientError(CodeLimitExceeded, "too many pending confirmations; try again later")
}
c.tokens[tok] = confirmEntry{user: user, table: table, op: op, binding: binding, expires: now.Add(ttl)}
return tok, nil
}
// consume validates and removes a token. Any mismatch (other user, table, operation or
// changed binding) is the same error, and the token is spent either way.
func (c *confirmStore) consume(tok, user, table, op, binding string) error {
c.mu.Lock()
e, ok := c.tokens[tok]
delete(c.tokens, tok)
now := c.now()
c.mu.Unlock()
if !ok || now.After(e.expires) || e.user != user || e.table != table || e.op != op || e.binding != binding {
return errConfirmInvalid
}
return nil
}
// bindingHash fingerprints the parts of a write a confirmation covers.
func bindingHash(parts ...any) (string, error) {
h := sha256.New()
for _, p := range parts {
b, err := json.Marshal(p)
if err != nil {
return "", errors.New("cannot fingerprint request")
}
h.Write(b)
h.Write([]byte{0})
}
return hex.EncodeToString(h.Sum(nil)), nil
}
+21 -1
View File
@@ -1,6 +1,11 @@
package resolvemcp package resolvemcp
import "context" import (
"context"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/security"
)
type contextKey string type contextKey string
@@ -69,3 +74,18 @@ func withRequestData(ctx context.Context, schema, entity, tableName string, mode
ctx = WithModelPtr(ctx, modelPtr) ctx = WithModelPtr(ctx, modelPtr)
return ctx return ctx
} }
// withModelRules puts the handler registry's rules for the model into the context, where the
// security hooks look them up first. The handler registry is private, so without this the
// hooks would not see rules set by RegisterModelWithRules / SetModelRules.
func (h *Handler) withModelRules(ctx context.Context, schema, entity string) context.Context {
reg, ok := h.registry.(*modelregistry.DefaultModelRegistry)
if !ok {
return ctx
}
rules, err := reg.GetModelRules(buildModelName(schema, entity))
if err != nil {
return ctx
}
return context.WithValue(ctx, security.ModelRulesKey, rules)
}
+50
View File
@@ -0,0 +1,50 @@
// Package resolvemcp exposes registered database models as Model Context Protocol (MCP)
// tools over HTTP (SSE and streamable HTTP), so an AI agent can discover, read and
// change data without a tool per table.
//
// # How an agent uses it
//
// The tool set is fixed and does not grow with the models:
//
// list_tables tables the caller may use, their description and allowed operations
// describe_table columns, types, keys, relations, writable fields, limits
// select_table read rows: filters, sort, columns, preloads, paging, cursors
// insert_into_table / update_table / delete_from_table
// writes; filter-based writes are previewed (dry_run) and need the
// confirm_token from the preview
// list_functions / call_function registered stored functions
// resolvespec_annotate optional free-text notes (Config.EnableAnnotations)
//
// The same guide is sent to MCP clients as the server instructions.
//
// The server is read-only by default (Config.ReadOnly nil means on); set
// ReadOnly: resolvemcp.Bool(false) to enable the write tools.
//
// # Setting it up
//
// handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{BaseURL: "http://localhost:8080"})
// handler.RegisterModel("public", "users", &User{})
//
// r := mux.NewRouter()
// resolvemcp.SetupMuxRoutes(r, handler, securityList) // requires an authenticated caller
//
// # Describing the API for agents
//
// Models are documented through the model registry (see package modelregistry):
// SetModelDescription / LoadModelDescriptions on the handler, a ModelDescription()
// method on the model, or comment tags on its fields (gorm/bun "comment:" or
// comment/note/desc tags). The descriptions show up in list_tables and
// describe_table.
//
// ExportCatalog writes the whole picture (guide, tools, limits, tables with columns,
// relations, operations and descriptions) to a JSON or Markdown file on disk, so
// agents and developers can learn the API without connecting:
//
// handler.ExportCatalog("docs/mcp-catalog.md")
//
// # Security
//
// Routes must be mounted behind Guard(securityList); the *Unauthenticated setup
// functions exist only for use behind another trusted layer. Per-entity rules come from
// modelregistry.ModelRules, and BeforeHandle/AfterHandle hooks can veto or audit any call.
package resolvemcp
+81
View File
@@ -0,0 +1,81 @@
package resolvemcp
import (
"crypto/rand"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"github.com/mark3labs/mcp-go/mcp"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// Stable error codes returned to MCP clients.
const (
CodeInvalidArgument = "invalid_argument"
CodeNotFound = "not_found"
CodeForbidden = "forbidden"
CodeLimitExceeded = "limit_exceeded"
CodeInternal = "internal"
)
// ClientError is an error whose code and message are safe to show to an MCP client.
// Hooks may return one (NewClientError) to give the client a specific reason; every other error
// is reported as an opaque internal error with a reference that matches the server log.
type ClientError struct {
Code string
Message string
}
func (e *ClientError) Error() string { return e.Message }
// NewClientError returns an error that reaches the client as {code, message}.
func NewClientError(code, message string) error {
return &ClientError{Code: code, Message: message}
}
func invalidArg(format string, a ...any) error {
return NewClientError(CodeInvalidArgument, fmt.Sprintf(format, a...))
}
// errInternal is what recovered panics return: the details are logged, not sent.
var errInternal = NewClientError(CodeInternal, "internal error")
// clientFacing maps err to the code and message the client sees. Anything that is not a
// ClientError or a not-found is logged in full with a short reference and reported as an
// opaque internal error carrying that reference.
func clientFacing(op string, err error) (code, message string) {
var ce *ClientError
switch {
case errors.As(err, &ce):
if ce.Code == CodeInternal {
ref := newRef()
logger.Error("[resolvemcp] %s: internal error ref=%s: %v", op, ref, err)
return CodeInternal, "internal error (ref " + ref + ")"
}
return ce.Code, ce.Message
case errors.Is(err, errRecordNotFound):
return CodeNotFound, "record not found"
}
ref := newRef()
logger.Error("[resolvemcp] %s failed ref=%s: %v", op, ref, err)
return CodeInternal, "internal error (ref " + ref + ")"
}
func newRef() string {
var b [4]byte
_, _ = rand.Read(b[:])
return hex.EncodeToString(b[:])
}
// toolError builds the error result for a tool call.
func toolError(op string, err error) *mcp.CallToolResult {
code, msg := clientFacing(op, err)
b, _ := json.Marshal(map[string]any{
"success": false,
"error": map[string]string{"code": code, "message": msg},
})
return mcp.NewToolResultError(string(b))
}
+292
View File
@@ -0,0 +1,292 @@
package resolvemcp
import (
"context"
"encoding/json"
"fmt"
"math"
"regexp"
"sort"
"strings"
"sync"
"github.com/bitechdev/ResolveSpec/pkg/common"
)
// Parameter types a function may declare.
const (
ParamString = "string"
ParamInteger = "integer"
ParamNumber = "number"
ParamBoolean = "boolean"
ParamObject = "object"
ParamArray = "array"
)
// FunctionParam declares one argument of a registered function.
type FunctionParam struct {
Name string
Type string // one of the Param* constants
Description string
Required bool
}
// FunctionFunc is a Go function callable through call_function. tx is the transaction the call
// runs in (OnTxBegin already fired on it); args are validated against the declared params.
type FunctionFunc func(ctx context.Context, tx common.Database, args map[string]any) (any, error)
// Function is a callable registered with Handler.RegisterFunction. Exactly one of Handler (a Go
// callback) and Procedure (a SQL function called by name) is set.
type Function struct {
Name string
Description string
Params []FunctionParam
// Handler is the Go callback.
Handler FunctionFunc
// Procedure is the SQL function to call, optionally schema-qualified. Declared params are
// passed positionally in declaration order; an omitted optional param is NULL. Rows
// returned are the result.
Procedure string
// Authorize, when set, decides per caller whether the function is listed and callable.
// Return nil to allow.
Authorize func(ctx context.Context) error
}
var (
functionNameRe = regexp.MustCompile(`^[A-Za-z][A-Za-z0-9_]{0,63}$`)
procedureRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)?$`)
paramNameRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]{0,63}$`)
)
type functionRegistry struct {
mu sync.RWMutex
funcs map[string]Function
}
// RegisterFunction makes f callable through the call_function meta tool. Only registered
// functions are callable; there is no way to reach an arbitrary SQL function.
func (h *Handler) RegisterFunction(f Function) error {
if !functionNameRe.MatchString(f.Name) {
return fmt.Errorf("resolvemcp: invalid function name %q", f.Name)
}
if (f.Handler == nil) == (f.Procedure == "") {
return fmt.Errorf("resolvemcp: function %q needs exactly one of Handler and Procedure", f.Name)
}
if f.Procedure != "" && !procedureRe.MatchString(f.Procedure) {
return fmt.Errorf("resolvemcp: invalid procedure name %q", f.Procedure)
}
seen := map[string]bool{}
for _, p := range f.Params {
if !paramNameRe.MatchString(p.Name) || seen[p.Name] {
return fmt.Errorf("resolvemcp: function %q: invalid or duplicate param %q", f.Name, p.Name)
}
seen[p.Name] = true
switch p.Type {
case ParamString, ParamInteger, ParamNumber, ParamBoolean, ParamObject, ParamArray:
default:
return fmt.Errorf("resolvemcp: function %q param %q: unknown type %q", f.Name, p.Name, p.Type)
}
}
h.functions.mu.Lock()
defer h.functions.mu.Unlock()
if h.functions.funcs == nil {
h.functions.funcs = map[string]Function{}
}
if _, dup := h.functions.funcs[f.Name]; dup {
return fmt.Errorf("resolvemcp: function %q already registered", f.Name)
}
h.functions.funcs[f.Name] = f
return nil
}
func (h *Handler) function(name string) (Function, bool) {
h.functions.mu.RLock()
defer h.functions.mu.RUnlock()
f, ok := h.functions.funcs[name]
return f, ok
}
// functionAllowed reports whether Config.AllowedFunctions lets the function through.
func (h *Handler) functionAllowed(name string) bool {
if h.allowedFns == nil {
return true
}
_, ok := h.allowedFns[name]
return ok
}
// visibleFunctions returns the functions the caller may call, sorted by name.
func (h *Handler) visibleFunctions(ctx context.Context) []Function {
h.functions.mu.RLock()
out := make([]Function, 0, len(h.functions.funcs))
for _, f := range h.functions.funcs {
out = append(out, f)
}
h.functions.mu.RUnlock()
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
visible := out[:0]
for _, f := range out {
if !h.functionAllowed(f.Name) {
continue
}
if f.Authorize == nil || f.Authorize(ctx) == nil {
visible = append(visible, f)
}
}
return visible
}
// validateArgs checks args against the declared params: required present, no undeclared names,
// types match. The returned map is what the function receives.
func validateArgs(params []FunctionParam, args map[string]any) (map[string]any, error) {
decl := make(map[string]FunctionParam, len(params))
for _, p := range params {
decl[p.Name] = p
}
for name := range args {
if _, ok := decl[name]; !ok {
return nil, invalidArg("unknown argument %q", truncate(name))
}
}
out := make(map[string]any, len(args))
for _, p := range params {
v, ok := args[p.Name]
if !ok || v == nil {
if p.Required {
return nil, invalidArg("missing required argument %q", p.Name)
}
continue
}
if !argMatches(p.Type, v) {
return nil, invalidArg("argument %q must be %s", p.Name, p.Type)
}
out[p.Name] = v
}
return out, nil
}
func truncate(s string) string {
if len(s) > maxKeyEcho {
return s[:maxKeyEcho] + "..."
}
return s
}
func argMatches(typ string, v any) bool {
switch typ {
case ParamString:
_, ok := v.(string)
return ok
case ParamBoolean:
_, ok := v.(bool)
return ok
case ParamNumber:
return isNumber(v)
case ParamInteger:
switch n := v.(type) {
case float64:
return n == math.Trunc(n) && !math.IsInf(n, 0)
case int, int64:
return true
}
return false
case ParamObject:
_, ok := v.(map[string]any)
return ok
case ParamArray:
_, ok := v.([]any)
return ok
}
return false
}
func isNumber(v any) bool {
switch v.(type) {
case float64, int, int64:
return true
}
return false
}
// callProcedure runs f.Procedure on tx.
func callProcedure(ctx context.Context, tx common.Database, f Function, args map[string]any) (any, error) {
placeholders := make([]string, len(f.Params))
values := make([]any, len(f.Params))
for i, p := range f.Params {
ph := fmt.Sprintf("$%d", i+1)
if p.Type == ParamObject || p.Type == ParamArray {
ph += "::jsonb"
}
if v, ok := args[p.Name]; !ok {
values[i] = nil
} else if p.Type == ParamObject || p.Type == ParamArray {
b, err := json.Marshal(v)
if err != nil {
return nil, invalidArg("argument %q cannot be encoded", p.Name)
}
values[i] = string(b)
} else {
values[i] = v
}
placeholders[i] = ph
}
var rows []map[string]any
query := fmt.Sprintf("SELECT * FROM %s(%s)", f.Procedure, strings.Join(placeholders, ", "))
if err := tx.Query(ctx, &rows, query, values...); err != nil {
return nil, err
}
return rows, nil
}
// executeCall validates and runs a registered function in a transaction: BeforeHandle (auth and
// rules), Authorize, then OnTxBegin, BeforeCall, the function, AfterCall.
func (h *Handler) executeCall(ctx context.Context, name string, rawArgs map[string]any) (_ any, retErr error) {
defer recoverPanic(&retErr)
ctx, cancel := h.callContext(ctx)
defer cancel()
f, ok := h.function(name)
if !ok || !h.functionAllowed(name) {
return nil, invalidArg("unknown function %q", truncate(name))
}
hookCtx := &HookContext{Context: ctx, Handler: h, Entity: name, Operation: "call_function", Tx: h.db}
if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil {
return nil, err
}
// Same answer for "not authorized" and "does not exist" so names cannot be probed.
if f.Authorize != nil && f.Authorize(ctx) != nil {
return nil, invalidArg("unknown function %q", truncate(name))
}
args, err := validateArgs(f.Params, rawArgs)
if err != nil {
return nil, err
}
hookCtx.Data = args
var result any
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
if err := h.hooks.Execute(BeforeCall, hookCtx); err != nil {
return err
}
if m, ok := hookCtx.Data.(map[string]any); ok {
args = m
}
var err error
if f.Handler != nil {
result, err = f.Handler(ctx, tx, args)
} else {
result, err = callProcedure(ctx, tx, f, args)
}
if err != nil {
return err
}
hookCtx.Result = result
return h.hooks.Execute(AfterCall, hookCtx)
})
if err != nil {
return nil, err
}
return result, nil
}
+52
View File
@@ -0,0 +1,52 @@
package resolvemcp
import (
"net/http"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/security"
)
// Guard returns middleware that requires an authenticated caller on every request.
//
// The security list's provider decides which credentials are accepted: build it from a
// security.ChainAuthenticator over an OAuth bearer token, a session token (header or cookie)
// and an API key authenticator. The authenticated security.UserContext is placed in the request
// context, which the MCP transports pass on to every tool call, so rules, row security and
// OnTxBegin apply to that caller.
//
// Unlike security.NewAuthMiddleware this guard has no guest or optional mode: it ignores
// security.SkipAuth / security.OptionalAuth markers on the request context, and fails closed
// (500) when no provider is configured.
func Guard(securityList *security.SecurityList) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
authed := security.NewAuthHandler(securityList, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if uc, ok := security.GetUserContext(r.Context()); !ok || uc == nil {
http.Error(w, "Authentication failed", http.StatusUnauthorized)
return
}
next.ServeHTTP(w, r)
}))
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if securityList == nil {
http.Error(w, "Security provider not configured", http.StatusInternalServerError)
return
}
authed.ServeHTTP(w, r)
})
}
}
// requireGuard reports whether securityList can guard a route. Setup helpers use it to refuse
// to mount an endpoint rather than serve it unauthenticated by mistake.
func requireGuard(fn string, securityList *security.SecurityList) bool {
if securityList == nil || securityList.Provider() == nil {
logger.Error("resolvemcp.%s: no security provider configured; MCP endpoint NOT mounted", fn)
return false
}
return true
}
func warnUnauthenticated(fn string) {
logger.Warn("resolvemcp.%s: serving the MCP endpoint WITHOUT authentication; every caller can read and write all registered models", fn)
}
+97
View File
@@ -0,0 +1,97 @@
package resolvemcp
import (
"errors"
"net/http"
"net/http/httptest"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/security"
"github.com/bitechdev/ResolveSpec/pkg/security/providers"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// tokenAuth accepts the bearer token "good" and nothing else.
type tokenAuth struct{ security.Authenticator }
func (tokenAuth) Authenticate(r *http.Request) (*security.UserContext, error) {
if r.Header.Get("Authorization") != "Bearer good" {
return nil, errors.New("bad credentials")
}
return &security.UserContext{UserID: 7, UserName: "kim"}, nil
}
func newTestSecurityList(t *testing.T) *security.SecurityList {
t.Helper()
p, err := security.NewCompositeSecurityProvider(tokenAuth{},
providers.NewConfigColumnSecurityProvider(map[string][]sectypes.ColumnSecurity{}),
providers.NewConfigRowSecurityProvider(nil, nil))
if err != nil {
t.Fatal(err)
}
sl, err := security.NewSecurityList(p)
if err != nil {
t.Fatal(err)
}
return sl
}
func serve(h http.Handler, auth string, mark func(*http.Request) *http.Request) *httptest.ResponseRecorder {
r := httptest.NewRequest(http.MethodGet, "/mcp", nil)
if auth != "" {
r.Header.Set("Authorization", auth)
}
if mark != nil {
r = mark(r)
}
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
return w
}
func TestGuardRejectsUnauthenticated(t *testing.T) {
var gotUser int
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
uc, ok := security.GetUserContext(r.Context())
if !ok {
t.Error("user context missing downstream")
return
}
gotUser = uc.UserID
})
g := Guard(newTestSecurityList(t))(next)
for name, auth := range map[string]string{"none": "", "wrong": "Bearer bad"} {
if w := serve(g, auth, nil); w.Code != http.StatusUnauthorized {
t.Errorf("%s: status %d, want 401", name, w.Code)
}
}
if w := serve(g, "Bearer good", nil); w.Code != http.StatusOK || gotUser != 7 {
t.Errorf("good: status %d user %d, want 200 / 7", w.Code, gotUser)
}
}
// Skip/optional markers on the request context must not open the MCP endpoint.
func TestGuardIgnoresSkipAndOptionalMarkers(t *testing.T) {
called := false
g := Guard(newTestSecurityList(t))(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { called = true }))
for name, mark := range map[string]func(*http.Request) *http.Request{
"skip": func(r *http.Request) *http.Request { return r.WithContext(security.SkipAuth(r.Context())) },
"optional": func(r *http.Request) *http.Request { return r.WithContext(security.OptionalAuth(r.Context())) },
} {
if w := serve(g, "", mark); w.Code != http.StatusUnauthorized || called {
t.Errorf("%s: status %d called=%v, want 401 and not called", name, w.Code, called)
}
}
}
func TestGuardFailsClosedWithoutProvider(t *testing.T) {
called := false
g := Guard(nil)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { called = true }))
if w := serve(g, "Bearer good", nil); w.Code != http.StatusInternalServerError || called {
t.Errorf("status %d called=%v, want 500 and not called", w.Code, called)
}
if requireGuard("test", nil) {
t.Error("requireGuard(nil) must be false")
}
}
+232 -96
View File
@@ -4,9 +4,12 @@ import (
"context" "context"
"database/sql" "database/sql"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"net/http" "net/http"
"reflect" "reflect"
"regexp"
"runtime/debug"
"strings" "strings"
"sync" "sync"
@@ -21,6 +24,7 @@ import (
// Handler exposes registered database models as MCP tools and resources. // Handler exposes registered database models as MCP tools and resources.
type Handler struct { type Handler struct {
allowedFns map[string]struct{} // nil: every function is allowed
db common.Database db common.Database
registry common.ModelRegistry registry common.ModelRegistry
hooks *HookRegistry hooks *HookRegistry
@@ -30,6 +34,8 @@ type Handler struct {
version string version string
oauth2Regs []oauth2Registration oauth2Regs []oauth2Registration
oauthSrv *security.OAuthServer oauthSrv *security.OAuthServer
functions functionRegistry
confirms *confirmStore
} }
// NewHandler creates a Handler with the given database, model registry, and config. // NewHandler creates a Handler with the given database, model registry, and config.
@@ -38,12 +44,22 @@ func NewHandler(db common.Database, registry common.ModelRegistry, cfg Config) *
db: db, db: db,
registry: registry, registry: registry,
hooks: NewHookRegistry(), hooks: NewHookRegistry(),
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"), mcpServer: server.NewMCPServer("resolvemcp", "1.0.0", server.WithInstructions(guideFor(cfg.withDefaults().readOnly, cfg.AllowFunctionCalls))),
config: cfg, config: cfg.withDefaults(),
confirms: newConfirmStore(),
name: "resolvemcp", name: "resolvemcp",
version: "1.0.0", version: "1.0.0",
} }
registerAnnotationTool(h) if len(cfg.AllowedFunctions) > 0 {
h.allowedFns = make(map[string]struct{}, len(cfg.AllowedFunctions))
for _, n := range cfg.AllowedFunctions {
h.allowedFns[n] = struct{}{}
}
}
registerMetaTools(h)
if cfg.EnableAnnotations && !h.config.readOnly {
registerAnnotationTool(h)
}
return h return h
} }
@@ -97,7 +113,20 @@ type dynamicSSEHandler struct {
pool map[string]*server.SSEServer pool map[string]*server.SSEServer
} }
// maxSSEPool bounds the per-base-URL server cache; Host and X-Forwarded-Proto are client
// controlled, so without a bound a client could grow it forever.
const maxSSEPool = 32
func (d *dynamicSSEHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { func (d *dynamicSSEHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
if !d.h.hostAllowed(r.Host) {
http.Error(w, "host not allowed", http.StatusBadRequest)
return
}
proto := r.Header.Get("X-Forwarded-Proto")
if proto != "" && proto != "http" && proto != "https" {
http.Error(w, "invalid forwarded protocol", http.StatusBadRequest)
return
}
baseURL := requestBaseURL(r) baseURL := requestBaseURL(r)
d.mu.Lock() d.mu.Lock()
@@ -106,6 +135,12 @@ func (d *dynamicSSEHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
} }
s, ok := d.pool[baseURL] s, ok := d.pool[baseURL]
if !ok { if !ok {
if len(d.pool) >= maxSSEPool {
d.mu.Unlock()
logger.Warn("resolvemcp: SSE base URL cache full; set Config.BaseURL or Config.AllowedHosts")
http.Error(w, "too many hosts", http.StatusServiceUnavailable)
return
}
s = d.h.newSSEServer(baseURL, d.h.config.BasePath) s = d.h.newSSEServer(baseURL, d.h.config.BasePath)
d.pool[baseURL] = s d.pool[baseURL] = s
} }
@@ -114,6 +149,20 @@ func (d *dynamicSSEHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
s.ServeHTTP(w, r) s.ServeHTTP(w, r)
} }
// hostAllowed reports whether host may be used to build the SSE message URL. With no
// Config.AllowedHosts every host is accepted (the pool cap still applies).
func (h *Handler) hostAllowed(host string) bool {
if len(h.config.AllowedHosts) == 0 {
return true
}
for _, a := range h.config.AllowedHosts {
if strings.EqualFold(a, host) {
return true
}
}
return false
}
// requestBaseURL builds the base URL from an incoming request. // requestBaseURL builds the base URL from an incoming request.
// It honours the X-Forwarded-Proto header for deployments behind a proxy. // It honours the X-Forwarded-Proto header for deployments behind a proxy.
func requestBaseURL(r *http.Request) string { func requestBaseURL(r *http.Request) string {
@@ -127,13 +176,13 @@ func requestBaseURL(r *http.Request) string {
return scheme + "://" + r.Host return scheme + "://" + r.Host
} }
// RegisterModel registers a model and immediately exposes it as MCP tools and a resource. // RegisterModel registers a model. It becomes visible to the fixed meta tools (list_tables,
// select_table, ...); no per-model tools are created.
func (h *Handler) RegisterModel(schema, entity string, model interface{}) error { func (h *Handler) RegisterModel(schema, entity string, model interface{}) error {
fullName := buildModelName(schema, entity) fullName := buildModelName(schema, entity)
if err := h.registry.RegisterModel(fullName, model); err != nil { if err := h.registry.RegisterModel(fullName, model); err != nil {
return err return err
} }
registerModelTools(h, schema, entity, model)
return nil return nil
} }
@@ -149,7 +198,6 @@ func (h *Handler) RegisterModelWithRules(schema, entity string, model interface{
if err := reg.RegisterModelWithRules(fullName, model, rules); err != nil { if err := reg.RegisterModelWithRules(fullName, model, rules); err != nil {
return err return err
} }
registerModelTools(h, schema, entity, model)
return nil return nil
} }
@@ -200,22 +248,36 @@ func (h *Handler) getSchemaAndTable(defaultSchema, entity string, model interfac
return defaultSchema, entity return defaultSchema, entity
} }
// errRecordNotFound is the one error update and delete return for a row that does not exist,
// is hidden by row security, or vanished mid-write, so ids cannot be enumerated by error text.
var errRecordNotFound = errors.New("record not found")
// recoverPanic catches a panic from the current goroutine and returns it as an error. // recoverPanic catches a panic from the current goroutine and returns it as an error.
// Usage: defer recoverPanic(&returnedErr) // Usage: defer recoverPanic(&returnedErr)
func recoverPanic(err *error) { func recoverPanic(err *error) {
if r := recover(); r != nil { if r := recover(); r != nil {
msg := fmt.Sprintf("%v", r) logger.Error("[resolvemcp] panic recovered: %v\n%s", r, debug.Stack())
logger.Error("[resolvemcp] panic recovered: %s", msg) *err = errInternal
*err = fmt.Errorf("internal error: %s", msg)
} }
} }
// executeRead reads records from the database and returns raw data + metadata. // executeRead reads records from the database and returns raw data + metadata.
func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, options common.RequestOptions) (_ interface{}, _ *common.Metadata, retErr error) { func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, options common.RequestOptions) (interface{}, *common.Metadata, error) {
return h.executeReadCounted(ctx, schema, entity, id, options, true)
}
// executeReadCounted is executeRead with control over the total-row COUNT, which costs a full
// scan of the filtered set and is only run when the caller asks for it.
func (h *Handler) executeReadCounted(ctx context.Context, schema, entity, id string, options common.RequestOptions, count bool) (_ interface{}, _ *common.Metadata, retErr error) {
defer recoverPanic(&retErr) defer recoverPanic(&retErr)
ctx, cancel := h.callContext(ctx)
defer cancel()
if err := h.checkReadLimits(&options); err != nil {
return nil, nil, err
}
model, err := h.registry.GetModelByEntity(schema, entity) model, err := h.registry.GetModelByEntity(schema, entity)
if err != nil { if err != nil {
return nil, nil, fmt.Errorf("model not found: %w", err) return nil, nil, invalidArg("model not found: %s", buildModelName(schema, entity))
} }
unwrapped, err := common.ValidateAndUnwrapModel(model) unwrapped, err := common.ValidateAndUnwrapModel(model)
@@ -226,7 +288,7 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
model = unwrapped.Model model = unwrapped.Model
modelType := unwrapped.ModelType modelType := unwrapped.ModelType
tableName := h.getTableName(schema, entity, model) tableName := h.getTableName(schema, entity, model)
ctx = withRequestData(ctx, schema, entity, tableName, model, unwrapped.ModelPtr) ctx = withRequestData(h.withModelRules(ctx, schema, entity), schema, entity, tableName, model, unwrapped.ModelPtr)
validator := common.NewColumnValidator(model) validator := common.NewColumnValidator(model)
options = validator.FilterRequestOptions(options) options = validator.FilterRequestOptions(options)
@@ -253,7 +315,7 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
var metadata *common.Metadata var metadata *common.Metadata
err = h.runInTx(ctx, hookCtx, func(common.Database) error { err = h.runInTx(ctx, hookCtx, func(common.Database) error {
var err error var err error
data, metadata, err = h.readInTx(ctx, hookCtx, model, modelType, tableName, id, options) data, metadata, err = h.readInTx(ctx, hookCtx, model, modelType, tableName, id, options, count)
return err return err
}) })
if err != nil { if err != nil {
@@ -263,7 +325,7 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
} }
// readInTx runs the read hooks and queries on hookCtx.Tx. // readInTx runs the read hooks and queries on hookCtx.Tx.
func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model interface{}, modelType reflect.Type, tableName, id string, options common.RequestOptions) (interface{}, *common.Metadata, error) { func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model interface{}, modelType reflect.Type, tableName, id string, options common.RequestOptions, count bool) (interface{}, *common.Metadata, error) {
sliceType := reflect.SliceOf(reflect.PointerTo(modelType)) sliceType := reflect.SliceOf(reflect.PointerTo(modelType))
modelPtr := reflect.New(sliceType).Interface() modelPtr := reflect.New(sliceType).Interface()
@@ -314,7 +376,7 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
// expandJoins is empty for resolvemcp — no custom SQL join support yet // expandJoins is empty for resolvemcp — no custom SQL join support yet
cursorFilter, err := getCursorFilter(tableName, pkName, modelColumns, options, nil) cursorFilter, err := getCursorFilter(tableName, pkName, modelColumns, options, nil)
if err != nil { if err != nil {
return nil, nil, fmt.Errorf("cursor error: %w", err) return nil, nil, invalidArg("invalid cursor")
} }
if cursorFilter != "" { if cursorFilter != "" {
@@ -327,9 +389,13 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
} }
// Count — must happen before preloads are applied; Bun panics when counting with relations. // Count — must happen before preloads are applied; Bun panics when counting with relations.
total, err := query.Count(ctx) total := 0
if err != nil { if count {
return nil, nil, fmt.Errorf("error counting records: %w", err) var err error
total, err = query.Count(ctx)
if err != nil {
return nil, nil, fmt.Errorf("error counting records: %w", err)
}
} }
// Pagination // Pagination
@@ -342,6 +408,9 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
// Preloads — applied after count to avoid Bun panic when counting with relations. // Preloads — applied after count to avoid Bun panic when counting with relations.
if len(options.Preload) > 0 { if len(options.Preload) > 0 {
if err := h.validatePreloads(model, options.Preload); err != nil {
return nil, nil, err
}
var preloadErr error var preloadErr error
query, preloadErr = h.applyPreloads(model, query, options.Preload) query, preloadErr = h.applyPreloads(model, query, options.Preload)
if preloadErr != nil { if preloadErr != nil {
@@ -363,7 +432,7 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
// a destination when the query preloads a has-many relation. // a destination when the query preloads a has-many relation.
if err := query.ScanModel(ctx); err != nil { if err := query.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
return nil, nil, fmt.Errorf("record not found") return nil, nil, errRecordNotFound
} }
return nil, nil, fmt.Errorf("query error: %w", err) return nil, nil, fmt.Errorf("query error: %w", err)
} }
@@ -372,7 +441,7 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
// for both collection and single-record reads. Extract its one result. // for both collection and single-record reads. Extract its one result.
scannedResults := reflect.ValueOf(modelPtr).Elem() scannedResults := reflect.ValueOf(modelPtr).Elem()
if scannedResults.Len() == 0 { if scannedResults.Len() == 0 {
return nil, nil, fmt.Errorf("record not found") return nil, nil, errRecordNotFound
} }
data = scannedResults.Index(0).Interface() data = scannedResults.Index(0).Interface()
} else { } else {
@@ -421,9 +490,14 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
// executeCreate inserts one or more records. // executeCreate inserts one or more records.
func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data interface{}) (_ interface{}, retErr error) { func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data interface{}) (_ interface{}, retErr error) {
defer recoverPanic(&retErr) defer recoverPanic(&retErr)
ctx, cancel := h.callContext(ctx)
defer cancel()
if items, ok := data.([]interface{}); ok && len(items) > h.config.MaxBatch {
return nil, NewClientError(CodeLimitExceeded, fmt.Sprintf("batch of %d exceeds the maximum of %d", len(items), h.config.MaxBatch))
}
model, err := h.registry.GetModelByEntity(schema, entity) model, err := h.registry.GetModelByEntity(schema, entity)
if err != nil { if err != nil {
return nil, fmt.Errorf("model not found: %w", err) return nil, invalidArg("model not found: %s", buildModelName(schema, entity))
} }
result, err := common.ValidateAndUnwrapModel(model) result, err := common.ValidateAndUnwrapModel(model)
@@ -433,7 +507,7 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
model = result.Model model = result.Model
tableName := h.getTableName(schema, entity, model) tableName := h.getTableName(schema, entity, model)
ctx = withRequestData(ctx, schema, entity, tableName, model, result.ModelPtr) ctx = withRequestData(h.withModelRules(ctx, schema, entity), schema, entity, tableName, model, result.ModelPtr)
hookCtx := &HookContext{ hookCtx := &HookContext{
Context: ctx, Context: ctx,
@@ -455,11 +529,11 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
modelType = modelType.Elem() modelType = modelType.Elem()
} }
// Transaction 1: BeforeCreate + inserts.
var ( var (
single bool single bool
originals []map[string]interface{} originals []map[string]interface{}
insertedIDs []interface{} insertedIDs []interface{}
results []interface{}
) )
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil { if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil {
@@ -475,18 +549,26 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
for _, item := range v { for _, item := range v {
itemMap, ok := item.(map[string]interface{}) itemMap, ok := item.(map[string]interface{})
if !ok { if !ok {
return fmt.Errorf("each item must be an object") return invalidArg("each item must be an object")
} }
originals = append(originals, itemMap) originals = append(originals, itemMap)
} }
default: default:
return fmt.Errorf("data must be an object or array of objects") return invalidArg("data must be an object or array of objects")
} }
insertedIDs = make([]interface{}, 0, len(originals)) insertedIDs = make([]interface{}, 0, len(originals))
for _, itemMap := range originals { for _, itemMap := range originals {
cols, err := writeColumns(model, itemMap)
if err != nil {
return err
}
if len(cols) == 0 {
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 itemMap { for key, value := range cols {
q = q.Value(key, value) q = q.Value(key, value)
} }
if pkName == "" { if pkName == "" {
@@ -502,22 +584,11 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
} }
insertedIDs = append(insertedIDs, returnedID) insertedIDs = append(insertedIDs, returnedID)
} }
return nil
})
if err != nil {
if single {
return nil, fmt.Errorf("create error: %w", err)
}
if _, ok := hookCtx.Data.([]interface{}); ok {
return nil, fmt.Errorf("batch create error: %w", err)
}
return nil, err
}
// Transaction 2: re-fetch to capture DB-generated defaults/triggers, then AfterCreate. // Re-fetch inside the same transaction to capture DB-generated defaults/triggers, then
results := make([]interface{}, 0, len(insertedIDs)) // AfterCreate: the write is only committed when the whole sequence succeeds, so a
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { // failure here cannot leave a committed insert behind an error the client may retry.
results = results[:0] results = make([]interface{}, 0, len(insertedIDs))
for i, pkVal := range insertedIDs { for i, pkVal := range insertedIDs {
if pkVal == nil { if pkVal == nil {
results = append(results, originals[i]) results = append(results, originals[i])
@@ -544,6 +615,12 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
return nil return nil
}) })
if err != nil { if err != nil {
if single {
return nil, fmt.Errorf("create error: %w", err)
}
if _, ok := hookCtx.Data.([]interface{}); ok {
return nil, fmt.Errorf("batch create error: %w", err)
}
return nil, err return nil, err
} }
if single { if single {
@@ -555,9 +632,11 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
// executeUpdate updates a record by ID. // executeUpdate updates a record by ID.
func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, data interface{}) (_ interface{}, retErr error) { func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, data interface{}) (_ interface{}, retErr error) {
defer recoverPanic(&retErr) defer recoverPanic(&retErr)
ctx, cancel := h.callContext(ctx)
defer cancel()
model, err := h.registry.GetModelByEntity(schema, entity) model, err := h.registry.GetModelByEntity(schema, entity)
if err != nil { if err != nil {
return nil, fmt.Errorf("model not found: %w", err) return nil, invalidArg("model not found: %s", buildModelName(schema, entity))
} }
result, err := common.ValidateAndUnwrapModel(model) result, err := common.ValidateAndUnwrapModel(model)
@@ -567,11 +646,11 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
model = result.Model model = result.Model
tableName := h.getTableName(schema, entity, model) tableName := h.getTableName(schema, entity, model)
ctx = withRequestData(ctx, schema, entity, tableName, model, result.ModelPtr) ctx = withRequestData(h.withModelRules(ctx, schema, entity), schema, entity, tableName, model, result.ModelPtr)
updates, ok := data.(map[string]interface{}) updates, ok := data.(map[string]interface{})
if !ok { if !ok {
return nil, fmt.Errorf("data must be an object") return nil, invalidArg("data must be an object")
} }
if id == "" { if id == "" {
@@ -580,7 +659,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
} }
} }
if id == "" { if id == "" {
return nil, fmt.Errorf("update requires an ID") return nil, invalidArg("update requires an id")
} }
pkName := reflection.GetPrimaryKeyName(model) pkName := reflection.GetPrimaryKeyName(model)
@@ -602,23 +681,47 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
var updateResult interface{} var updateResult interface{}
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
// Read existing record if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
return err
}
if modifiedData, ok := hookCtx.Data.(map[string]interface{}); ok {
updates = modifiedData
}
// SET only the validated incoming keys; the primary key addresses the row, it is not
// rewritten. nil and "" are real values (NULL / empty string).
setCols, err := writeColumns(model, updates)
if err != nil {
return err
}
for col := range setCols {
if strings.EqualFold(col, pkName) {
delete(setCols, col)
}
}
if len(setCols) == 0 {
return invalidArg("no updatable fields in data")
}
// Load the target through the BeforeScan hooks (row security) so a row the caller
// cannot see is reported as not found and never written.
modelType := reflect.TypeOf(model) modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Pointer { if modelType.Kind() == reflect.Pointer {
modelType = modelType.Elem() modelType = modelType.Elem()
} }
existingRecord := reflect.New(modelType).Interface() existingRecord := reflect.New(modelType).Interface()
selectQuery := tx.NewSelect().Model(existingRecord).Column("*"). hookCtx.Query = tx.NewSelect().Model(existingRecord).Column("*").
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
if err := h.hooks.Execute(BeforeScan, hookCtx); err != nil {
if err := selectQuery.ScanModel(ctx); err != nil { return err
}
if err := hookCtx.Query.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
return fmt.Errorf("no records found to update") return errRecordNotFound
} }
return fmt.Errorf("error fetching existing record: %w", err) return fmt.Errorf("error fetching existing record: %w", err)
} }
// Convert to map
existingMap := make(map[string]interface{}) existingMap := make(map[string]interface{})
jsonData, err := json.Marshal(existingRecord) jsonData, err := json.Marshal(existingRecord)
if err != nil { if err != nil {
@@ -627,81 +730,58 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
if err := json.Unmarshal(jsonData, &existingMap); err != nil { if err := json.Unmarshal(jsonData, &existingMap); err != nil {
return fmt.Errorf("error unmarshaling existing record: %w", err) return fmt.Errorf("error unmarshaling existing record: %w", err)
} }
for key, v := range updates {
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil { existingMap[key] = v
return err
}
if modifiedData, ok := hookCtx.Data.(map[string]interface{}); ok {
updates = modifiedData
} }
// Merge non-nil, non-empty values reflection.RemoveNonWritableColumns(model, setCols)
for key, newValue := range updates { q := tx.NewUpdate().Table(tableName).SetMap(setCols).
if newValue == nil {
continue
}
if strVal, ok := newValue.(string); ok && strVal == "" {
continue
}
existingMap[key] = newValue
}
q := tx.NewUpdate().Table(tableName).SetMap(existingMap).
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)
if err != nil { if err != nil {
return fmt.Errorf("error updating record: %w", err) return fmt.Errorf("error updating record: %w", err)
} }
if res.RowsAffected() == 0 { if res.RowsAffected() == 0 {
return fmt.Errorf("no records found to update") return errRecordNotFound
} }
updateResult = existingMap hookCtx.Result = existingMap
hookCtx.Result = updateResult
return h.hooks.Execute(AfterUpdate, hookCtx)
})
if err != nil { // Re-fetch inside the same transaction to capture DB-generated changes, then
return nil, err // AfterUpdate; see executeCreate.
}
// Transaction 2: re-fetch to capture DB-generated changes.
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Pointer {
modelType = modelType.Elem()
}
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
fetchedRecord := reflect.New(modelType).Interface() fetchedRecord := reflect.New(modelType).Interface()
if err := tx.NewSelect().Model(fetchedRecord). if err := tx.NewSelect().Model(fetchedRecord).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id).
ScanModel(ctx); err == nil { ScanModel(ctx); err == nil {
jsonData, marshalErr := json.Marshal(fetchedRecord) if jsonData, marshalErr := json.Marshal(fetchedRecord); marshalErr == nil {
if marshalErr == nil {
var fetchedMap map[string]interface{} var fetchedMap map[string]interface{}
if json.Unmarshal(jsonData, &fetchedMap) == nil { if json.Unmarshal(jsonData, &fetchedMap) == nil {
updateResult = fetchedMap existingMap = fetchedMap
hookCtx.Result = fetchedMap
} }
} }
} }
return nil updateResult = existingMap
return h.hooks.Execute(AfterUpdate, hookCtx)
}) })
if err != nil { if err != nil {
return nil, err return nil, err
} }
return updateResult, nil return updateResult, nil
} }
// executeDelete deletes a record by ID. // executeDelete deletes a record by ID.
func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string) (_ interface{}, retErr error) { func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string) (_ interface{}, retErr error) {
defer recoverPanic(&retErr) defer recoverPanic(&retErr)
ctx, cancel := h.callContext(ctx)
defer cancel()
if id == "" { if id == "" {
return nil, fmt.Errorf("delete requires an ID") return nil, invalidArg("delete requires an id")
} }
model, err := h.registry.GetModelByEntity(schema, entity) model, err := h.registry.GetModelByEntity(schema, entity)
if err != nil { if err != nil {
return nil, fmt.Errorf("model not found: %w", err) return nil, invalidArg("model not found: %s", buildModelName(schema, entity))
} }
result, err := common.ValidateAndUnwrapModel(model) result, err := common.ValidateAndUnwrapModel(model)
@@ -711,7 +791,7 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
model = result.Model model = result.Model
tableName := h.getTableName(schema, entity, model) tableName := h.getTableName(schema, entity, model)
ctx = withRequestData(ctx, schema, entity, tableName, model, result.ModelPtr) ctx = withRequestData(h.withModelRules(ctx, schema, entity), schema, entity, tableName, model, result.ModelPtr)
pkName := reflection.GetPrimaryKeyName(model) pkName := reflection.GetPrimaryKeyName(model)
@@ -741,11 +821,14 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
return err return err
} }
record := reflect.New(modelType).Interface() record := reflect.New(modelType).Interface()
selectQuery := tx.NewSelect().Model(record). hookCtx.Query = tx.NewSelect().Model(record).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
if err := selectQuery.ScanModel(ctx); err != nil { if err := h.hooks.Execute(BeforeScan, hookCtx); err != nil {
return err
}
if err := hookCtx.Query.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
return fmt.Errorf("record not found") return errRecordNotFound
} }
return fmt.Errorf("error fetching record: %w", err) return fmt.Errorf("error fetching record: %w", err)
} }
@@ -757,7 +840,7 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
return fmt.Errorf("delete error: %w", err) return fmt.Errorf("delete error: %w", err)
} }
if res.RowsAffected() == 0 { if res.RowsAffected() == 0 {
return fmt.Errorf("record not found or already deleted") return errRecordNotFound
} }
recordToDelete = record recordToDelete = record
@@ -902,3 +985,56 @@ func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(t
return h.hooks.Execute(OnTxBegin, hookCtx) return h.hooks.Execute(OnTxBegin, hookCtx)
}, body) }, body)
} }
// callContext bounds one tool call by Config.QueryTimeout.
func (h *Handler) callContext(ctx context.Context) (context.Context, context.CancelFunc) {
return context.WithTimeout(ctx, h.config.QueryTimeout)
}
// checkReadLimits applies the paging caps to options in place: a missing limit takes
// DefaultLimit, a larger one is clamped to MaxLimit, and an offset above MaxOffset is rejected.
func (h *Handler) checkReadLimits(options *common.RequestOptions) error {
limit := h.config.DefaultLimit
if options.Limit != nil && *options.Limit > 0 {
limit = *options.Limit
}
if limit > h.config.MaxLimit {
limit = h.config.MaxLimit
}
options.Limit = &limit
if options.Offset != nil {
if *options.Offset < 0 {
return invalidArg("offset must not be negative")
}
if *options.Offset > h.config.MaxOffset {
return NewClientError(CodeLimitExceeded, fmt.Sprintf("offset exceeds the maximum of %d; use cursor paging", h.config.MaxOffset))
}
}
return nil
}
var preloadSegmentRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
// validatePreloads checks preload paths against the model: the first segment must be one of
// the model's relations and the path may not be deeper than Config.MaxPreloadDepth.
func (h *Handler) validatePreloads(model interface{}, preloads []common.PreloadOption) error {
relations := map[string]bool{}
for _, name := range buildModelInfo("", "", model).relationNames {
relations[strings.ToLower(name)] = true
}
for i := range preloads {
segments := strings.Split(preloads[i].Relation, ".")
if len(segments) > h.config.MaxPreloadDepth {
return NewClientError(CodeLimitExceeded, fmt.Sprintf("preload depth exceeds the maximum of %d", h.config.MaxPreloadDepth))
}
for _, seg := range segments {
if !preloadSegmentRe.MatchString(seg) {
return invalidArg("invalid preload relation")
}
}
if !relations[strings.ToLower(segments[0])] {
return invalidArg("unknown relation %q", segments[0])
}
}
return nil
}
+80
View File
@@ -0,0 +1,80 @@
package resolvemcp
import (
"fmt"
"net/http"
"net/http/httptest"
"sync"
"testing"
"github.com/DATA-DOG/go-sqlmock"
)
func TestHookRegistryConcurrentUse(t *testing.T) {
r := NewHookRegistry()
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(3)
go func() { defer wg.Done(); r.Register(BeforeRead, func(*HookContext) error { return nil }) }()
go func() { defer wg.Done(); _ = r.Execute(BeforeRead, &HookContext{}) }()
go func() { defer wg.Done(); _ = r.HasHooks(BeforeRead); r.Clear(AfterRead) }()
}
wg.Wait()
}
// Update and delete report a missing row with the same error, so ids cannot be probed.
func TestNotFoundErrorsAreUniform(t *testing.T) {
h, mock, ctx := newTxHarness(t)
empty := sqlmock.NewRows([]string{"id", "name"})
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(empty)
mock.ExpectRollback()
_, errU := h.executeUpdate(ctx, "public", "items", "9", map[string]interface{}{"name": "x"})
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}))
mock.ExpectRollback()
_, errD := h.executeDelete(ctx, "public", "items", "9")
if errU == nil || errD == nil || errU.Error() != errD.Error() {
t.Fatalf("update %v / delete %v must be the same error", errU, errD)
}
}
func TestSSEHostAllowlistAndPoolCap(t *testing.T) {
h, _, _ := newTxHarness(t)
h.config.AllowedHosts = []string{"mcp.example.com"}
d := &dynamicSSEHandler{h: h}
r := httptest.NewRequest(http.MethodPost, "/mcp/message", nil)
r.Host = "evil.example.net"
w := httptest.NewRecorder()
d.ServeHTTP(w, r)
if w.Code != http.StatusBadRequest {
t.Errorf("foreign host: status %d, want 400", w.Code)
}
r = httptest.NewRequest(http.MethodPost, "/mcp/message", nil)
r.Host = "mcp.example.com"
r.Header.Set("X-Forwarded-Proto", "javascript")
w = httptest.NewRecorder()
d.ServeHTTP(w, r)
if w.Code != http.StatusBadRequest {
t.Errorf("bad proto: status %d, want 400", w.Code)
}
h.config.AllowedHosts = nil
for i := 0; i < maxSSEPool+5; i++ {
r = httptest.NewRequest(http.MethodPost, "/mcp/message?sessionId=x", nil)
r.Host = fmt.Sprintf("h%d.example.com", i)
d.ServeHTTP(httptest.NewRecorder(), r)
}
if len(d.pool) > maxSSEPool {
t.Errorf("pool grew to %d, cap is %d", len(d.pool), maxSSEPool)
}
r = httptest.NewRequest(http.MethodPost, "/mcp/message", nil)
r.Host = "one-more.example.com"
w = httptest.NewRecorder()
d.ServeHTTP(w, r)
if w.Code != http.StatusServiceUnavailable {
t.Errorf("full pool: status %d, want 503", w.Code)
}
}
+44 -7
View File
@@ -3,6 +3,8 @@ package resolvemcp
import ( import (
"context" "context"
"fmt" "fmt"
"runtime/debug"
"sync"
"github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/logger" "github.com/bitechdev/ResolveSpec/pkg/logger"
@@ -21,12 +23,23 @@ const (
BeforeCreate HookType = "before_create" BeforeCreate HookType = "before_create"
AfterCreate HookType = "after_create" AfterCreate HookType = "after_create"
// BeforeScan fires on update and delete, with hookCtx.Query set to the select that loads the
// target row. Hooks that narrow the query (row security) run here; a row the query does
// not return is reported as not found and never written.
BeforeScan HookType = "before_scan"
BeforeUpdate HookType = "before_update" BeforeUpdate HookType = "before_update"
AfterUpdate HookType = "after_update" AfterUpdate HookType = "after_update"
BeforeDelete HookType = "before_delete" BeforeDelete HookType = "before_delete"
AfterDelete HookType = "after_delete" AfterDelete HookType = "after_delete"
// BeforeCall and AfterCall fire inside the transaction of a call_function call.
// hookCtx.Entity is the function name, Data the validated arguments (BeforeCall may
// replace them) and Result the function's result (AfterCall).
BeforeCall HookType = "before_call"
AfterCall HookType = "after_call"
// OnTxBegin fires once, first, inside every transaction the handler opens // OnTxBegin fires once, first, inside every transaction the handler opens
// (including the second short transaction for post-commit work). hookCtx.Tx is // (including the second short transaction for post-commit work). hookCtx.Tx is
// the transaction; use it to stamp transaction-local state such as RLS // the transaction; use it to stamp transaction-local state such as RLS
@@ -62,6 +75,7 @@ type HookFunc func(*HookContext) error
// HookRegistry manages all registered hooks // HookRegistry manages all registered hooks
type HookRegistry struct { type HookRegistry struct {
mu sync.RWMutex
hooks map[HookType][]HookFunc hooks map[HookType][]HookFunc
} }
@@ -72,11 +86,14 @@ func NewHookRegistry() *HookRegistry {
} }
func (r *HookRegistry) Register(hookType HookType, hook HookFunc) { func (r *HookRegistry) Register(hookType HookType, hook HookFunc) {
r.mu.Lock()
if r.hooks == nil { if r.hooks == nil {
r.hooks = make(map[HookType][]HookFunc) r.hooks = make(map[HookType][]HookFunc)
} }
r.hooks[hookType] = append(r.hooks[hookType], hook) r.hooks[hookType] = append(r.hooks[hookType], hook)
logger.Info("Registered resolvemcp hook for %s (total: %d)", hookType, len(r.hooks[hookType])) total := len(r.hooks[hookType])
r.mu.Unlock()
logger.Info("Registered resolvemcp hook for %s (total: %d)", hookType, total)
} }
func (r *HookRegistry) RegisterMultiple(hookTypes []HookType, hook HookFunc) { func (r *HookRegistry) RegisterMultiple(hookTypes []HookType, hook HookFunc) {
@@ -86,37 +103,57 @@ func (r *HookRegistry) RegisterMultiple(hookTypes []HookType, hook HookFunc) {
} }
func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error { func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
hooks, exists := r.hooks[hookType] // Append-only slices: a snapshot of the slice header is safe to iterate without the lock.
if !exists || len(hooks) == 0 { r.mu.RLock()
hooks := r.hooks[hookType]
r.mu.RUnlock()
if len(hooks) == 0 {
return nil return nil
} }
logger.Debug("Executing %d resolvemcp hook(s) for %s", len(hooks), hookType) logger.Debug("Executing %d resolvemcp hook(s) for %s", len(hooks), hookType)
for i, hook := range hooks { for i, hook := range hooks {
if err := hook(ctx); err != nil { if err := runHook(hook, ctx); err != nil {
logger.Error("resolvemcp hook %d for %s failed: %v", i+1, hookType, err) logger.Error("resolvemcp hook %d for %s failed: %v", i+1, hookType, err)
return fmt.Errorf("hook execution failed: %w", err) return fmt.Errorf("hook execution failed: %w", err)
} }
if ctx.Abort { if ctx.Abort {
logger.Warn("resolvemcp hook %d for %s requested abort: %s", i+1, hookType, ctx.AbortMessage) logger.Warn("resolvemcp hook %d for %s requested abort: %s", i+1, hookType, ctx.AbortMessage)
return fmt.Errorf("operation aborted by hook: %s", ctx.AbortMessage) return fmt.Errorf("operation aborted by hook: %w", NewClientError(CodeForbidden, ctx.AbortMessage))
} }
} }
return nil return nil
} }
// runHook calls hook and turns a panic into an error, so a faulty hook fails the request
// instead of unwinding through the transaction machinery. The stack is logged, not returned.
func runHook(hook HookFunc, ctx *HookContext) (err error) {
defer func() {
if r := recover(); r != nil {
logger.Error("resolvemcp hook panic: %v\n%s", r, debug.Stack())
err = errInternal
}
}()
return hook(ctx)
}
func (r *HookRegistry) Clear(hookType HookType) { func (r *HookRegistry) Clear(hookType HookType) {
r.mu.Lock()
defer r.mu.Unlock()
delete(r.hooks, hookType) delete(r.hooks, hookType)
} }
func (r *HookRegistry) ClearAll() { func (r *HookRegistry) ClearAll() {
r.mu.Lock()
defer r.mu.Unlock()
r.hooks = make(map[HookType][]HookFunc) r.hooks = make(map[HookType][]HookFunc)
} }
func (r *HookRegistry) HasHooks(hookType HookType) bool { func (r *HookRegistry) HasHooks(hookType HookType) bool {
hooks, exists := r.hooks[hookType] r.mu.RLock()
return exists && len(hooks) > 0 defer r.mu.RUnlock()
return len(r.hooks[hookType]) > 0
} }
+151
View File
@@ -0,0 +1,151 @@
package resolvemcp
import (
"context"
"errors"
"strings"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"github.com/bitechdev/ResolveSpec/pkg/common"
)
func ptr[T any](v T) *T { return &v }
func TestConfigDefaults(t *testing.T) {
c := Config{}.withDefaults()
if c.DefaultLimit != 50 || c.MaxLimit != 1000 || c.MaxOffset != 100000 || c.MaxBatch != 100 ||
c.MaxPreloadDepth != 2 || c.MaxWriteRows != 100 || c.QueryTimeout != 30*time.Second || c.ConfirmTTL != 5*time.Minute {
t.Errorf("defaults: %+v", c)
}
if c := (Config{DefaultLimit: 5000, MaxLimit: 200}).withDefaults(); c.DefaultLimit != 200 {
t.Errorf("default limit must not exceed max: %d", c.DefaultLimit)
}
}
func TestCheckReadLimits(t *testing.T) {
h, _, _ := newTxHarness(t)
cases := []struct {
name string
in common.RequestOptions
want int
wantErr string
}{
{"no limit takes default", common.RequestOptions{}, 50, ""},
{"zero takes default", common.RequestOptions{Limit: ptr(0)}, 50, ""},
{"negative takes default", common.RequestOptions{Limit: ptr(-3)}, 50, ""},
{"explicit kept", common.RequestOptions{Limit: ptr(10)}, 10, ""},
{"clamped", common.RequestOptions{Limit: ptr(1 << 30)}, 1000, ""},
{"offset too big", common.RequestOptions{Offset: ptr(100001)}, 0, CodeLimitExceeded},
{"offset negative", common.RequestOptions{Offset: ptr(-1)}, 0, CodeInvalidArgument},
{"offset at max ok", common.RequestOptions{Offset: ptr(100000)}, 50, ""},
}
for _, c := range cases {
opts := c.in
err := h.checkReadLimits(&opts)
if c.wantErr != "" {
var ce *ClientError
if !errors.As(err, &ce) || ce.Code != c.wantErr {
t.Errorf("%s: got %v, want code %s", c.name, err, c.wantErr)
}
continue
}
if err != nil || *opts.Limit != c.want {
t.Errorf("%s: limit %v err %v, want %d", c.name, opts.Limit, err, c.want)
}
}
}
func TestValidatePreloads(t *testing.T) {
type child struct {
ID int `json:"id" bun:"id,pk"`
}
type parent struct {
ID int `json:"id" bun:"id,pk"`
Children []*child `json:"children" bun:"rel:has-many"`
}
h, _, _ := newTxHarness(t)
ok := func(rel string) bool {
return h.validatePreloads(&parent{}, []common.PreloadOption{{Relation: rel}}) == nil
}
if !ok("children") || !ok("Children") || !ok("children.sub") {
t.Error("known relations (and depth 2) must pass")
}
for _, bad := range []string{"nope", "children.a.b", "children; DROP TABLE x", "", "a..b", "id"} {
if ok(bad) {
t.Errorf("preload %q must be rejected", bad)
}
}
}
func TestBatchCap(t *testing.T) {
h, _, ctx := newTxHarness(t)
h.config.MaxBatch = 2
items := []interface{}{map[string]interface{}{}, map[string]interface{}{}, map[string]interface{}{}}
_, err := h.executeCreate(ctx, "public", "items", items)
var ce *ClientError
if !errors.As(err, &ce) || ce.Code != CodeLimitExceeded {
t.Fatalf("want limit_exceeded, got %v", err)
}
}
func TestReadSkipsCountUnlessRequested(t *testing.T) {
h, mock, ctx := newTxHarness(t)
mock.ExpectBegin()
mock.ExpectQuery(`SELECT .* LIMIT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a"))
mock.ExpectCommit()
if _, _, err := h.executeReadCounted(ctx, "public", "items", "", common.RequestOptions{}, false); err != nil {
t.Fatal(err)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestClientFacingHidesInternals(t *testing.T) {
code, msg := clientFacing("t", errors.New(`pq: relation "secret_table" does not exist`))
if code != CodeInternal || strings.Contains(msg, "secret_table") || !strings.Contains(msg, "ref ") {
t.Errorf("raw error leaked: %s %q", code, msg)
}
if code, msg := clientFacing("t", invalidArg("bad %s", "x")); code != CodeInvalidArgument || msg != "bad x" {
t.Errorf("client error: %s %q", code, msg)
}
if code, _ := clientFacing("t", errRecordNotFound); code != CodeNotFound {
t.Errorf("not found: %s", code)
}
wrapped := errors.Join(errors.New("ctx"), NewClientError(CodeForbidden, "update not allowed for x"))
if code, msg := clientFacing("t", wrapped); code != CodeForbidden || msg != "update not allowed for x" {
t.Errorf("wrapped: %s %q", code, msg)
}
}
func TestHookPanicIsRecovered(t *testing.T) {
h, mock, ctx := newTxHarness(t)
h.Hooks().Register(BeforeDelete, func(*HookContext) error { panic("boom: secret") })
mock.ExpectBegin()
mock.ExpectRollback()
_, err := h.executeDelete(ctx, "public", "items", "7")
if err == nil {
t.Fatal("expected error")
}
if _, msg := clientFacing("t", err); strings.Contains(msg, "boom") || strings.Contains(msg, "secret") {
t.Errorf("panic value leaked: %q", msg)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestToolTimeoutBoundsCall(t *testing.T) {
h, _, _ := newTxHarness(t)
h.config.QueryTimeout = 20 * time.Millisecond
ctx, cancel := h.callContext(context.Background())
defer cancel()
select {
case <-ctx.Done():
case <-time.After(time.Second):
t.Fatal("call context did not time out")
}
}
+465
View File
@@ -0,0 +1,465 @@
package resolvemcp
import (
"context"
"fmt"
"reflect"
"sort"
"strings"
"github.com/mark3labs/mcp-go/mcp"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
)
// Operation names used by list_tables and the rule checks.
const (
opSelect = "select"
opInsert = "insert"
opUpdate = "update"
opDelete = "delete"
)
// filterOperators is the operator list shown by describe_table.
var filterOperators = []string{"=", "!=", ">", ">=", "<", "<=", "like", "ilike", "in", "is_null", "is_not_null"}
// registerMetaTools adds the fixed tool set. Their number does not grow with the models.
func registerMetaTools(h *Handler) {
readOnly := mcp.WithReadOnlyHintAnnotation(true)
tableArg := mcp.WithString("table", mcp.Required(), mcp.Description("Table as 'schema.entity' (see list_tables)."))
filtersArg := mcp.WithArray("filters", mcp.Description(`Filter objects, e.g. [{"column":"status","operator":"=","value":"active"}]. Combine with "logic_operator": "AND" (default) or "OR". Operators: `+strings.Join(filterOperators, " ")+"."))
idArg := mcp.WithString("id", mcp.Description("Primary key of one row. Use either id or filters."))
dryRunArg := mcp.WithBoolean("dry_run", mcp.Description("Report how many rows match and a preview of their ids; change nothing."))
confirmArg := mcp.WithString("confirm_token", mcp.Description("Token from the preview a filter-based call returns first; the write only happens with it."))
h.mcpServer.AddTool(mcp.NewTool("list_tables", readOnly,
mcp.WithDescription("List the tables you can use and the operations (select, insert, update, delete) allowed on each.")),
h.handleListTables)
h.mcpServer.AddTool(mcp.NewTool("describe_table", readOnly,
mcp.WithDescription("Describe a table: columns and types, primary key, relations (preloadable), writable fields, allowed operations and server limits."),
tableArg), h.handleDescribeTable)
h.mcpServer.AddTool(mcp.NewTool("select_table", readOnly,
mcp.WithDescription("Read rows from a table. Results are paged: 'limit' defaults to the server default and is capped; use cursor_forward/cursor_backward for deep paging. The total row count is only computed with include_count."),
tableArg, idArg, filtersArg,
mcp.WithArray("sort", mcp.Description(`Sort objects, e.g. [{"column":"created_at","direction":"desc"}].`)),
mcp.WithArray("columns", mcp.Description("Columns to return. Omit for all.")),
mcp.WithArray("omit_columns", mcp.Description("Columns to leave out.")),
mcp.WithArray("preloads", mcp.Description(`Relations to load, e.g. [{"relation":"orders"}]. See describe_table for the names and the maximum depth.`)),
mcp.WithNumber("limit", mcp.Description("Maximum rows to return.")),
mcp.WithNumber("offset", mcp.Description("Rows to skip.")),
mcp.WithString("cursor_forward", mcp.Description("Primary key of the last row of the current page; requires sort.")),
mcp.WithString("cursor_backward", mcp.Description("Primary key of the first row of the current page; requires sort.")),
mcp.WithBoolean("include_count", mcp.Description("Also return the total number of matching rows (slower on large tables).")),
), h.handleSelect)
if !h.config.readOnly {
registerWriteTools(h, tableArg, idArg, filtersArg, dryRunArg, confirmArg)
}
if !h.config.readOnly || h.config.AllowFunctionCalls {
registerFunctionTools(h, readOnly)
}
}
// registerWriteTools adds the tools that change table rows.
func registerWriteTools(h *Handler, tableArg, idArg, filtersArg, dryRunArg, confirmArg mcp.ToolOption) {
h.mcpServer.AddTool(mcp.NewTool("insert_into_table",
mcp.WithDescription("Insert one row (object) or several rows (array, one transaction, capped). Unknown or read-only fields are rejected."),
tableArg, mcp.WithObject("data", mcp.Required(), mcp.Description("A row object or an array of row objects.")),
), h.handleInsert)
h.mcpServer.AddTool(mcp.NewTool("update_table",
mcp.WithDescription("Update rows. Give an id (one row, applied at once) or filters (several rows: the first call returns a preview and a confirm_token, repeat the call with the token to apply; the number of rows is capped). Only the fields in data are changed; null sets NULL."),
tableArg, idArg, filtersArg,
mcp.WithObject("data", mcp.Required(), mcp.Description("Fields to change.")),
dryRunArg, confirmArg,
), h.handleUpdate)
h.mcpServer.AddTool(mcp.NewTool("delete_from_table",
mcp.WithDescription("Delete rows. Give an id (one row, applied at once) or filters (several rows: the first call returns a preview and a confirm_token, repeat the call with the token to delete; the number of rows is capped)."),
mcp.WithDestructiveHintAnnotation(true),
tableArg, idArg, filtersArg, dryRunArg, confirmArg,
), h.handleDelete)
}
// registerFunctionTools adds list_functions and call_function.
func registerFunctionTools(h *Handler, readOnly mcp.ToolOption) {
h.mcpServer.AddTool(mcp.NewTool("list_functions", readOnly,
mcp.WithDescription("List the functions you can call with call_function, with their parameters.")),
h.handleListFunctions)
h.mcpServer.AddTool(mcp.NewTool("call_function",
mcp.WithDescription("Call a registered function by name with validated arguments (see list_functions). Runs in a transaction."),
mcp.WithString("name", mcp.Required(), mcp.Description("Function name.")),
mcp.WithObject("arguments", mcp.Description("Arguments by parameter name.")),
), h.handleCallFunction)
}
// --------------------------------------------------------------------------
// Tables and rules
// --------------------------------------------------------------------------
// splitTable parses "schema.entity" (or a bare entity).
func splitTable(table string) (schema, entity string, err error) {
table = strings.TrimSpace(table)
if table == "" {
return "", "", invalidArg("missing required argument: table")
}
if i := strings.LastIndex(table, "."); i >= 0 {
return table[:i], table[i+1:], nil
}
return "", table, nil
}
// modelRules returns the rules of a registered model (defaults when the registry keeps none).
func (h *Handler) modelRules(schema, entity string) modelregistry.ModelRules {
if reg, ok := h.registry.(*modelregistry.DefaultModelRegistry); ok {
if r, err := reg.GetModelRules(buildModelName(schema, entity)); err == nil {
return r
}
}
return modelregistry.DefaultModelRules()
}
// opsFor lists the operations the rules allow. A read-only server allows select only.
func (h *Handler) opsFor(r modelregistry.ModelRules) []string {
var ops []string
if r.CanRead {
ops = append(ops, opSelect)
}
if h.config.readOnly {
return ops
}
if r.CanCreate {
ops = append(ops, opInsert)
}
if r.CanUpdate {
ops = append(ops, opUpdate)
}
if r.CanDelete {
ops = append(ops, opDelete)
}
return ops
}
// resolveTable finds a registered model and checks the rule for op. A model the rules forbid
// for op is reported like a missing one for reads of the table list, but with a plain
// "not allowed" here so the agent learns the operation is off, not that it misspelled the name.
func (h *Handler) resolveTable(args map[string]any, op string) (schema, entity string, err error) {
table, _ := args["table"].(string)
schema, entity, err = splitTable(table)
if err != nil {
return "", "", err
}
if _, err := h.registry.GetModelByEntity(schema, entity); err != nil {
return "", "", invalidArg("unknown table %q; see list_tables", truncate(table))
}
if op != "" && op != opSelect && h.config.readOnly {
return "", "", NewClientError(CodeForbidden, "this server is read-only: writes are disabled")
}
if op != "" {
allowed := false
for _, o := range h.opsFor(h.modelRules(schema, entity)) {
if o == op {
allowed = true
}
}
if !allowed {
return "", "", NewClientError(CodeForbidden, fmt.Sprintf("%s is not allowed on %s", op, buildModelName(schema, entity)))
}
}
return schema, entity, nil
}
func (h *Handler) handleListTables(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) {
type table struct {
Table string `json:"table"`
Description string `json:"description,omitempty"`
Operations []string `json:"operations"`
}
var tables []table
for name := range h.registry.GetAllModels() {
schema, entity, _ := splitTable(name)
if ops := h.opsFor(h.modelRules(schema, entity)); len(ops) > 0 {
tables = append(tables, table{Table: name, Description: h.modelDocs(schema, entity).Description, Operations: ops})
}
}
sort.Slice(tables, func(i, j int) bool { return tables[i].Table < tables[j].Table })
return marshalResult(map[string]any{"success": true, "tables": tables})
}
func (h *Handler) handleDescribeTable(_ context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
schema, entity, err := h.resolveTable(req.GetArguments(), "")
if err != nil {
return toolError("describe_table", err), nil
}
model, err := h.registry.GetModelByEntity(schema, entity)
if err != nil {
return toolError("describe_table", invalidArg("unknown table")), nil
}
rules := h.modelRules(schema, entity)
if len(h.opsFor(rules)) == 0 {
return toolError("describe_table", invalidArg("unknown table %q; see list_tables", buildModelName(schema, entity))), nil
}
info := buildModelInfo(schema, entity, model)
docs := h.modelDocs(schema, entity)
modelType := reflect.TypeOf(model)
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice) {
modelType = modelType.Elem()
}
writable := map[string]bool{}
if !h.config.readOnly && modelType != nil && modelType.Kind() == reflect.Struct {
for jsonKey := range reflection.BuildJSONToDBColumnMap(modelType) {
writable[jsonKey] = true
}
}
type column struct {
Name string `json:"name"`
Type string `json:"type,omitempty"`
Nullable bool `json:"nullable"`
PrimaryKey bool `json:"primary_key,omitempty"`
Unique bool `json:"unique,omitempty"`
Writable bool `json:"writable"`
Comment string `json:"description,omitempty"`
}
cols := make([]column, 0, len(info.columns))
var writableNames []string
for _, c := range info.columns {
typ := c.sqlType
if typ == "" {
typ = c.goType
}
w := writable[c.jsonName]
cols = append(cols, column{Name: c.jsonName, Type: typ, Nullable: c.nullable, PrimaryKey: c.isPrimary, Unique: c.isUnique, Writable: w, Comment: columnDescription(docs, c)})
if w && !c.isPrimary {
writableNames = append(writableNames, c.jsonName)
}
}
return marshalResult(map[string]any{
"success": true,
"table": info.fullName,
"description": docs.Description,
"purpose": docs.Purpose,
"tags": docs.Tags,
"primary_key": info.pkName,
"columns": cols,
"relations": info.relationNames,
"writable_columns": writableNames,
"operations": h.opsFor(rules),
"read_only": h.config.readOnly,
"filter_operators": filterOperators,
"limits": map[string]any{
"default_limit": h.config.DefaultLimit,
"max_limit": h.config.MaxLimit,
"max_offset": h.config.MaxOffset,
"max_batch": h.config.MaxBatch,
"max_preload_depth": h.config.MaxPreloadDepth,
"max_write_rows": h.config.MaxWriteRows,
},
})
}
// --------------------------------------------------------------------------
// Row tools
// --------------------------------------------------------------------------
// argID reads the id argument, which clients send as a string or a number.
func argID(args map[string]any) string {
switch v := args["id"].(type) {
case string:
return v
case float64:
if v == float64(int64(v)) {
return fmt.Sprintf("%d", int64(v))
}
return fmt.Sprint(v)
}
return ""
}
// parseFiltersStrict is parseFilters for writes: a malformed filter is an error, never dropped.
func parseFiltersStrict(raw any) ([]common.FilterOption, error) {
if raw == nil {
return nil, nil
}
items, ok := raw.([]any)
if !ok {
return nil, invalidArg("filters must be an array")
}
parsed := parseFilters(raw)
if len(parsed) != len(items) {
return nil, invalidArg("every filter needs a column and an operator")
}
return parsed, nil
}
func (h *Handler) handleSelect(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := req.GetArguments()
schema, entity, err := h.resolveTable(args, opSelect)
if err != nil {
return toolError("select_table", err), nil
}
count, _ := args["include_count"].(bool)
data, meta, err := h.executeReadCounted(ctx, schema, entity, argID(args), parseRequestOptions(args), count)
if err != nil {
return toolError("select_table", err), nil
}
if !count && meta != nil {
meta.Total, meta.Filtered = 0, 0
}
return marshalResult(map[string]any{"success": true, "data": data, "metadata": meta})
}
func (h *Handler) handleInsert(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := req.GetArguments()
schema, entity, err := h.resolveTable(args, opInsert)
if err != nil {
return toolError("insert_into_table", err), nil
}
data, ok := args["data"]
if !ok {
return toolError("insert_into_table", invalidArg("missing required argument: data")), nil
}
result, err := h.executeCreate(ctx, schema, entity, data)
if err != nil {
return toolError("insert_into_table", err), nil
}
return marshalResult(map[string]any{"success": true, "data": result})
}
// writeTarget parses id/filters/dry_run/confirm_token shared by update and delete.
func writeTarget(args map[string]any) (id string, filters []common.FilterOption, dryRun bool, token string, err error) {
id = argID(args)
if filters, err = parseFiltersStrict(args["filters"]); err != nil {
return
}
if id != "" && len(filters) > 0 {
err = invalidArg("use either id or filters, not both")
return
}
if id == "" && len(filters) == 0 {
err = invalidArg("provide an id or at least one filter")
return
}
dryRun, _ = args["dry_run"].(bool)
token, _ = args["confirm_token"].(string)
return
}
func (h *Handler) handleUpdate(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := req.GetArguments()
schema, entity, err := h.resolveTable(args, opUpdate)
if err != nil {
return toolError("update_table", err), nil
}
data, ok := args["data"].(map[string]any)
if !ok {
return toolError("update_table", invalidArg("data must be an object")), nil
}
id, filters, dryRun, token, err := writeTarget(args)
if err != nil {
return toolError("update_table", err), nil
}
if id != "" && !dryRun {
result, err := h.executeUpdate(ctx, schema, entity, id, data)
if err != nil {
return toolError("update_table", err), nil
}
return marshalResult(map[string]any{"success": true, "data": result})
}
return h.runWhere(ctx, "update_table", whereRequest{schema: schema, entity: entity, op: "update", filters: h.idFilters(schema, entity, id, filters), data: data, dryRun: dryRun, confirmToken: token})
}
func (h *Handler) handleDelete(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := req.GetArguments()
schema, entity, err := h.resolveTable(args, opDelete)
if err != nil {
return toolError("delete_from_table", err), nil
}
id, filters, dryRun, token, err := writeTarget(args)
if err != nil {
return toolError("delete_from_table", err), nil
}
if id != "" && !dryRun {
result, err := h.executeDelete(ctx, schema, entity, id)
if err != nil {
return toolError("delete_from_table", err), nil
}
return marshalResult(map[string]any{"success": true, "data": result})
}
return h.runWhere(ctx, "delete_from_table", whereRequest{schema: schema, entity: entity, op: "delete", filters: h.idFilters(schema, entity, id, filters), dryRun: dryRun, confirmToken: token})
}
// idFilters turns an id into a primary-key filter so a dry run of an id write goes through the
// same matching as a filter write.
func (h *Handler) idFilters(schema, entity, id string, filters []common.FilterOption) []common.FilterOption {
if id == "" {
return filters
}
model, err := h.registry.GetModelByEntity(schema, entity)
if err != nil {
return filters
}
return []common.FilterOption{{Column: reflection.GetPrimaryKeyName(model), Operator: "eq", Value: id, LogicOperator: "AND"}}
}
func (h *Handler) runWhere(ctx context.Context, op string, req whereRequest) (*mcp.CallToolResult, error) {
res, err := h.executeWhere(ctx, req)
if err != nil {
return toolError(op, err), nil
}
return marshalResult(map[string]any{"success": true, "result": res})
}
// --------------------------------------------------------------------------
// Functions
// --------------------------------------------------------------------------
func (h *Handler) handleListFunctions(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) {
type param struct {
Name string `json:"name"`
Type string `json:"type"`
Required bool `json:"required"`
Description string `json:"description,omitempty"`
}
type fn struct {
Name string `json:"name"`
Description string `json:"description,omitempty"`
Parameters []param `json:"parameters"`
}
out := []fn{}
for _, f := range h.visibleFunctions(ctx) {
item := fn{Name: f.Name, Description: f.Description, Parameters: []param{}}
for _, p := range f.Params {
item.Parameters = append(item.Parameters, param{p.Name, p.Type, p.Required, p.Description})
}
out = append(out, item)
}
return marshalResult(map[string]any{"success": true, "functions": out})
}
func (h *Handler) handleCallFunction(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := req.GetArguments()
name, _ := args["name"].(string)
if name == "" {
return toolError("call_function", invalidArg("missing required argument: name")), nil
}
fnArgs := map[string]any{}
if raw, ok := args["arguments"]; ok && raw != nil {
m, ok := raw.(map[string]any)
if !ok {
return toolError("call_function", invalidArg("arguments must be an object")), nil
}
fnArgs = m
}
result, err := h.executeCall(ctx, name, fnArgs)
if err != nil {
return toolError("call_function", err), nil
}
return marshalResult(map[string]any{"success": true, "result": result})
}
+451
View File
@@ -0,0 +1,451 @@
package resolvemcp
import (
"context"
"encoding/json"
"errors"
"sort"
"strings"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"github.com/mark3labs/mcp-go/mcp"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/security"
)
func callReq(args map[string]any) mcp.CallToolRequest {
var r mcp.CallToolRequest
r.Params.Arguments = args
return r
}
// payload decodes a tool result's text.
func payload(t *testing.T, res *mcp.CallToolResult) map[string]any {
t.Helper()
if len(res.Content) == 0 {
t.Fatal("empty result")
}
tc, ok := res.Content[0].(mcp.TextContent)
if !ok {
t.Fatalf("content is %T", res.Content[0])
}
var m map[string]any
if err := json.Unmarshal([]byte(tc.Text), &m); err != nil {
t.Fatalf("not JSON: %q", tc.Text)
}
return m
}
func errCode(t *testing.T, res *mcp.CallToolResult) string {
t.Helper()
if !res.IsError {
t.Fatalf("expected an error result, got %v", payload(t, res))
}
e, _ := payload(t, res)["error"].(map[string]any)
code, _ := e["code"].(string)
return code
}
func TestMetaToolSetIsFixed(t *testing.T) {
h, _, _ := newTxHarness(t)
for _, name := range []string{"x1", "x2", "x3"} {
if err := h.RegisterModel("public", name, &txItem{}); err != nil {
t.Fatal(err)
}
}
var got []string
for name := range h.mcpServer.ListTools() {
got = append(got, name)
}
sort.Strings(got)
want := "call_function delete_from_table describe_table insert_into_table list_functions list_tables select_table update_table"
if strings.Join(got, " ") != want {
t.Fatalf("tools = %v\nwant %s", got, want)
}
}
func TestListTablesShowsOnlyAllowedOperations(t *testing.T) {
h, _, ctx := newTxHarness(t)
_ = h.RegisterModelWithRules("public", "ro", &txItem{}, modelregistry.ModelRules{CanRead: true})
_ = h.RegisterModelWithRules("public", "hidden", &txItem{}, modelregistry.ModelRules{})
res, _ := h.handleListTables(ctx, callReq(nil))
tables, _ := payload(t, res)["tables"].([]any)
seen := map[string][]any{}
for _, tb := range tables {
m := tb.(map[string]any)
seen[m["table"].(string)] = m["operations"].([]any)
}
if _, ok := seen["public.hidden"]; ok {
t.Error("a table with no allowed operation must not be listed")
}
if ops := seen["public.ro"]; len(ops) != 1 || ops[0] != "select" {
t.Errorf("ro ops = %v", ops)
}
if ops := seen["public.items"]; len(ops) != 4 {
t.Errorf("default rules allow all four, got %v", ops)
}
// describe_table on a table with no allowed operation reads as unknown.
res, _ = h.handleDescribeTable(ctx, callReq(map[string]any{"table": "public.hidden"}))
if errCode(t, res) != CodeInvalidArgument {
t.Error("describe of a hidden table must look like an unknown table")
}
}
func TestDescribeTable(t *testing.T) {
h, _, ctx := newTxHarness(t)
if err := h.RegisterModel("public", "witems", &wItem{}); err != nil {
t.Fatal(err)
}
res, _ := h.handleDescribeTable(ctx, callReq(map[string]any{"table": "public.witems"}))
p := payload(t, res)
if p["primary_key"] != "id" {
t.Errorf("pk = %v", p["primary_key"])
}
w, _ := p["writable_columns"].([]any)
got := map[string]bool{}
for _, c := range w {
got[c.(string)] = true
}
if !got["name"] || !got["fullName"] || got["id"] || got["owner"] {
t.Errorf("writable columns = %v", w)
}
if lim, _ := p["limits"].(map[string]any); lim["max_limit"] != float64(1000) {
t.Errorf("limits = %v", p["limits"])
}
}
func TestSelectRespectsOperationRule(t *testing.T) {
h, _, ctx := newTxHarness(t)
_ = h.RegisterModelWithRules("public", "nowrite", &txItem{}, modelregistry.ModelRules{CanRead: true})
for tool, fn := range map[string]func(context.Context, mcp.CallToolRequest) (*mcp.CallToolResult, error){
"insert": h.handleInsert, "update": h.handleUpdate, "delete": h.handleDelete,
} {
res, _ := fn(ctx, callReq(map[string]any{"table": "public.nowrite", "data": map[string]any{"name": "a"}, "id": "1"}))
if errCode(t, res) != CodeForbidden {
t.Errorf("%s on a read-only table must be forbidden", tool)
}
}
res, _ := h.handleSelect(ctx, callReq(map[string]any{"table": "public.missing"}))
if errCode(t, res) != CodeInvalidArgument {
t.Error("unknown table must be invalid_argument")
}
}
func TestSelectCountOnlyWhenRequested(t *testing.T) {
h, mock, ctx := newTxHarness(t)
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a"))
mock.ExpectCommit()
res, _ := h.handleSelect(ctx, callReq(map[string]any{"table": "public.items"}))
if res.IsError {
t.Fatal(payload(t, res))
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
mock.ExpectBegin()
mock.ExpectQuery(`SELECT COUNT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(41))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a"))
mock.ExpectCommit()
res, _ = h.handleSelect(ctx, callReq(map[string]any{"table": "public.items", "include_count": true}))
meta, _ := payload(t, res)["metadata"].(map[string]any)
if res.IsError || meta["total"] != float64(41) {
t.Fatalf("metadata = %v", meta)
}
}
// --- filter writes ---
func matchRowsQuery(mock sqlmock.Sqlmock, ids ...int) {
rows := sqlmock.NewRows([]string{"id"})
for _, id := range ids {
rows.AddRow(id)
}
mock.ExpectQuery(`SELECT`).WillReturnRows(rows)
}
func updReq(filters any, extra map[string]any) mcp.CallToolRequest {
a := map[string]any{"table": "public.items", "data": map[string]any{"name": "z"}, "filters": filters}
for k, v := range extra {
a[k] = v
}
return callReq(a)
}
var statusFilter = []any{map[string]any{"column": "name", "operator": "=", "value": "a"}}
func TestFilterUpdateNeedsPreviewThenToken(t *testing.T) {
h, mock, ctx := newTxHarness(t)
// 1. preview: counts, lists ids, issues a token, writes nothing.
mock.ExpectBegin()
matchRowsQuery(mock, 1, 2)
mock.ExpectCommit()
res, _ := h.handleUpdate(ctx, updReq(statusFilter, nil))
r, _ := payload(t, res)["result"].(map[string]any)
tok, _ := r["confirm_token"].(string)
if res.IsError || tok == "" || r["requires_confirmation"] != true || r["matched"] != float64(2) {
t.Fatalf("preview = %v", payload(t, res))
}
// 2. with the token: re-matches inside the tx, then writes.
mock.ExpectBegin()
matchRowsQuery(mock, 1, 2)
mock.ExpectExec(`UPDATE .* WHERE "id" IN \(\$2, \$3\)`).WillReturnResult(sqlmock.NewResult(0, 2))
mock.ExpectCommit()
res, _ = h.handleUpdate(ctx, updReq(statusFilter, map[string]any{"confirm_token": tok}))
r, _ = payload(t, res)["result"].(map[string]any)
if res.IsError || r["affected"] != float64(2) {
t.Fatalf("confirmed = %v", payload(t, res))
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
// 3. a token is single use.
mock.ExpectBegin()
matchRowsQuery(mock, 1, 2)
mock.ExpectRollback()
res, _ = h.handleUpdate(ctx, updReq(statusFilter, map[string]any{"confirm_token": tok}))
if errCode(t, res) != CodeInvalidArgument {
t.Error("a spent token must be rejected")
}
}
func TestConfirmTokenBinding(t *testing.T) {
issue := func(t *testing.T) (*Handler, sqlmock.Sqlmock, context.Context, string) {
h, mock, ctx := newTxHarness(t)
mock.ExpectBegin()
matchRowsQuery(mock, 1)
mock.ExpectCommit()
res, _ := h.handleUpdate(ctx, updReq(statusFilter, nil))
r, _ := payload(t, res)["result"].(map[string]any)
return h, mock, ctx, r["confirm_token"].(string)
}
reject := func(t *testing.T, h *Handler, mock sqlmock.Sqlmock, ctx context.Context, req mcp.CallToolRequest, rows ...int) {
t.Helper()
mock.ExpectBegin()
matchRowsQuery(mock, rows...)
mock.ExpectRollback()
res, _ := h.handleUpdate(ctx, req)
if errCode(t, res) != CodeInvalidArgument {
t.Fatal("token must be rejected")
}
}
t.Run("changed data", func(t *testing.T) {
h, mock, ctx, tok := issue(t)
req := updReq(statusFilter, map[string]any{"confirm_token": tok, "data": map[string]any{"name": "different"}})
reject(t, h, mock, ctx, req, 1)
})
t.Run("changed filters", func(t *testing.T) {
h, mock, ctx, tok := issue(t)
other := []any{map[string]any{"column": "name", "operator": "=", "value": "b"}}
reject(t, h, mock, ctx, updReq(other, map[string]any{"confirm_token": tok}), 1)
})
t.Run("rows changed since preview", func(t *testing.T) {
h, mock, ctx, tok := issue(t)
reject(t, h, mock, ctx, updReq(statusFilter, map[string]any{"confirm_token": tok}), 1, 2)
})
t.Run("other user", func(t *testing.T) {
h, mock, ctx, tok := issue(t)
other := context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 99, UserName: "mallory"})
reject(t, h, mock, other, updReq(statusFilter, map[string]any{"confirm_token": tok}), 1)
})
t.Run("other table", func(t *testing.T) {
h, mock, ctx, tok := issue(t)
if err := h.RegisterModel("public", "other", &txItem{}); err != nil {
t.Fatal(err)
}
req := updReq(statusFilter, map[string]any{"confirm_token": tok, "table": "public.other"})
reject(t, h, mock, ctx, req, 1)
})
t.Run("expired", func(t *testing.T) {
h, mock, ctx, tok := issue(t)
h.confirms.now = func() time.Time { return time.Now().Add(time.Hour) }
reject(t, h, mock, ctx, updReq(statusFilter, map[string]any{"confirm_token": tok}), 1)
})
}
func TestFilterWriteGuardrails(t *testing.T) {
h, mock, ctx := newTxHarness(t)
for name, req := range map[string]mcp.CallToolRequest{
"neither id nor filters": callReq(map[string]any{"table": "public.items", "data": map[string]any{"name": "z"}}),
"both id and filters": updReq(statusFilter, map[string]any{"id": "1"}),
"malformed filter": updReq([]any{map[string]any{"column": "name"}}, nil),
"unknown column": updReq([]any{map[string]any{"column": "secret", "operator": "=", "value": 1}}, nil),
"injection in column": updReq([]any{map[string]any{"column": "name) OR (1=1", "operator": "=", "value": 1}}, nil),
"unknown operator": updReq([]any{map[string]any{"column": "name", "operator": "ɸ", "value": 1}}, nil),
"missing value": updReq([]any{map[string]any{"column": "name", "operator": "="}}, nil),
"unknown data field": callReq(map[string]any{"table": "public.items", "data": map[string]any{"role": "x"}, "filters": statusFilter}),
} {
res, _ := h.handleUpdate(ctx, req)
if errCode(t, res) != CodeInvalidArgument {
t.Errorf("%s: want invalid_argument", name)
}
}
// Rejected before any SQL ran.
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestFilterWriteRowCap(t *testing.T) {
h, mock, ctx := newTxHarness(t)
h.config.MaxWriteRows = 2
mock.ExpectBegin()
matchRowsQuery(mock, 1, 2, 3)
mock.ExpectRollback()
res, _ := h.handleDelete(ctx, callReq(map[string]any{"table": "public.items", "filters": statusFilter}))
if errCode(t, res) != CodeLimitExceeded {
t.Fatal("want limit_exceeded")
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestDryRunWritesNothing(t *testing.T) {
h, mock, ctx := newTxHarness(t)
mock.ExpectBegin()
matchRowsQuery(mock, 5)
mock.ExpectCommit()
res, _ := h.handleDelete(ctx, callReq(map[string]any{"table": "public.items", "filters": statusFilter, "dry_run": true}))
r, _ := payload(t, res)["result"].(map[string]any)
if res.IsError || r["dry_run"] != true || r["matched"] != float64(1) || r["confirm_token"] != nil {
t.Fatalf("dry run = %v", payload(t, res))
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestIDWriteNeedsNoToken(t *testing.T) {
h, mock, ctx := newTxHarness(t)
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a"))
mock.ExpectExec(`DELETE`).WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectCommit()
res, _ := h.handleDelete(ctx, callReq(map[string]any{"table": "public.items", "id": float64(7)}))
if res.IsError {
t.Fatal(payload(t, res))
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
// --- functions ---
func TestRegisterFunctionValidation(t *testing.T) {
h, _, _ := newTxHarness(t)
noop := func(context.Context, common.Database, map[string]any) (any, error) { return nil, nil }
bad := map[string]Function{
"bad name": {Name: "1x", Handler: noop},
"neither": {Name: "f"},
"both": {Name: "f", Handler: noop, Procedure: "p"},
"bad procedure": {Name: "f", Procedure: "p(); drop table x"},
"bad param type": {Name: "f", Handler: noop, Params: []FunctionParam{{Name: "a", Type: "blob"}}},
"duplicate param": {Name: "f", Handler: noop, Params: []FunctionParam{{Name: "a", Type: "string"}, {Name: "a", Type: "string"}}},
"bad param name": {Name: "f", Handler: noop, Params: []FunctionParam{{Name: "a b", Type: "string"}}},
}
for name, f := range bad {
if err := h.RegisterFunction(f); err == nil {
t.Errorf("%s: expected error", name)
}
}
if err := h.RegisterFunction(Function{Name: "ok", Handler: noop}); err != nil {
t.Fatal(err)
}
if err := h.RegisterFunction(Function{Name: "ok", Handler: noop}); err == nil {
t.Error("duplicate name must fail")
}
}
func TestCallFunctionValidatesAndRunsInTx(t *testing.T) {
h, mock, ctx := newTxHarness(t)
var gotArgs map[string]any
var gotTx common.Database
if err := h.RegisterFunction(Function{
Name: "greet", Description: "says hi",
Params: []FunctionParam{{Name: "who", Type: ParamString, Required: true}, {Name: "n", Type: ParamInteger}},
Handler: func(_ context.Context, tx common.Database, args map[string]any) (any, error) {
gotArgs, gotTx = args, tx
return map[string]any{"hello": args["who"]}, nil
},
}); err != nil {
t.Fatal(err)
}
tr := traceHooks(h, OnTxBegin, BeforeCall, AfterCall)
for name, args := range map[string]map[string]any{
"missing required": {},
"wrong type": {"who": 5},
"fractional int": {"who": "x", "n": 1.5},
"unknown arg": {"who": "x", "extra": 1},
} {
res, _ := h.handleCallFunction(ctx, callReq(map[string]any{"name": "greet", "arguments": args}))
if errCode(t, res) != CodeInvalidArgument {
t.Errorf("%s: want invalid_argument", name)
}
}
res, _ := h.handleCallFunction(ctx, callReq(map[string]any{"name": "nope"}))
if errCode(t, res) != CodeInvalidArgument {
t.Error("unknown function must be invalid_argument")
}
mock.ExpectBegin()
mock.ExpectCommit()
res, _ = h.handleCallFunction(ctx, callReq(map[string]any{"name": "greet", "arguments": map[string]any{"who": "kim", "n": float64(2)}}))
if res.IsError || gotArgs["who"] != "kim" || gotTx == nil {
t.Fatalf("call = %v", payload(t, res))
}
tr.assertOrder(t, "on_tx_begin", "before_call", "after_call")
if tr.txs["before_call"][0] != tr.txs["on_tx_begin"][0] {
t.Error("the call must run in the OnTxBegin transaction")
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestFunctionAuthorizeHidesAndBlocks(t *testing.T) {
h, _, ctx := newTxHarness(t)
noop := func(context.Context, common.Database, map[string]any) (any, error) { return "ran", nil }
_ = h.RegisterFunction(Function{Name: "open", Handler: noop})
_ = h.RegisterFunction(Function{Name: "admin_only", Handler: noop, Authorize: func(context.Context) error { return errors.New("no") }})
res, _ := h.handleListFunctions(ctx, callReq(nil))
fns, _ := payload(t, res)["functions"].([]any)
if len(fns) != 1 || fns[0].(map[string]any)["name"] != "open" {
t.Fatalf("visible functions = %v", fns)
}
res, _ = h.handleCallFunction(ctx, callReq(map[string]any{"name": "admin_only"}))
if errCode(t, res) != CodeInvalidArgument {
t.Error("an unauthorized function must look unknown")
}
}
func TestProcedureFunctionCallShape(t *testing.T) {
h, mock, ctx := newTxHarness(t)
if err := h.RegisterFunction(Function{
Name: "recalc", Procedure: "app.recalc_totals",
Params: []FunctionParam{{Name: "account", Type: ParamInteger, Required: true}, {Name: "opts", Type: ParamObject}},
}); err != nil {
t.Fatal(err)
}
mock.ExpectBegin()
mock.ExpectQuery(`SELECT \* FROM app\.recalc_totals\(\$1, \$2::jsonb\)`).WithArgs(float64(3), nil).
WillReturnRows(sqlmock.NewRows([]string{"total"}).AddRow(10))
mock.ExpectCommit()
res, _ := h.handleCallFunction(ctx, callReq(map[string]any{"name": "recalc", "arguments": map[string]any{"account": float64(3)}}))
if res.IsError {
t.Fatal(payload(t, res))
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
+2 -34
View File
@@ -114,25 +114,12 @@ func (h *Handler) mountOAuth2Routes(mux *http.ServeMux) {
// context into the request context, making it available to BeforeHandle security hooks. // context into the request context, making it available to BeforeHandle security hooks.
// Unauthenticated requests receive 401 before reaching any MCP tool. // Unauthenticated requests receive 401 before reaching any MCP tool.
func (h *Handler) AuthedSSEServer(securityList *security.SecurityList) http.Handler { func (h *Handler) AuthedSSEServer(securityList *security.SecurityList) http.Handler {
return security.NewAuthMiddleware(securityList)(h.SSEServer()) return Guard(securityList)(h.SSEServer())
}
// OptionalAuthSSEServer wraps SSEServer with optional authentication middleware.
// Unauthenticated requests continue as guest rather than returning 401.
// Use together with RegisterSecurityHooks and per-model CanPublicRead/Write rules
// to allow mixed public/private access.
func (h *Handler) OptionalAuthSSEServer(securityList *security.SecurityList) http.Handler {
return security.NewOptionalAuthMiddleware(securityList)(h.SSEServer())
} }
// AuthedStreamableHTTPServer wraps StreamableHTTPServer with required authentication middleware. // AuthedStreamableHTTPServer wraps StreamableHTTPServer with required authentication middleware.
func (h *Handler) AuthedStreamableHTTPServer(securityList *security.SecurityList) http.Handler { func (h *Handler) AuthedStreamableHTTPServer(securityList *security.SecurityList) http.Handler {
return security.NewAuthMiddleware(securityList)(h.StreamableHTTPServer()) return Guard(securityList)(h.StreamableHTTPServer())
}
// OptionalAuthStreamableHTTPServer wraps StreamableHTTPServer with optional authentication middleware.
func (h *Handler) OptionalAuthStreamableHTTPServer(securityList *security.SecurityList) http.Handler {
return security.NewOptionalAuthMiddleware(securityList)(h.StreamableHTTPServer())
} }
// -------------------------------------------------------------------------- // --------------------------------------------------------------------------
@@ -243,22 +230,3 @@ func SetupMuxOAuth2Routes(muxRouter *mux.Router, auth *security.DatabaseAuthenti
OAuth2CallbackHandler(auth, cfg.ProviderName, cfg.AfterLoginRedirect, cookieOpts...), OAuth2CallbackHandler(auth, cfg.ProviderName, cfg.AfterLoginRedirect, cookieOpts...),
).Methods(http.MethodGet) ).Methods(http.MethodGet)
} }
// SetupMuxRoutesWithAuth mounts the MCP SSE endpoints on a Gorilla Mux router
// with required authentication middleware applied.
func SetupMuxRoutesWithAuth(muxRouter *mux.Router, handler *Handler, securityList *security.SecurityList) {
basePath := handler.config.BasePath
h := handler.AuthedSSEServer(securityList)
muxRouter.Handle(basePath+"/sse", h).Methods(http.MethodGet, http.MethodOptions)
muxRouter.Handle(basePath+"/message", h).Methods(http.MethodPost, http.MethodOptions)
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, h))
}
// SetupMuxStreamableHTTPRoutesWithAuth mounts the MCP streamable HTTP endpoint on a
// Gorilla Mux router with required authentication middleware applied.
func SetupMuxStreamableHTTPRoutesWithAuth(muxRouter *mux.Router, handler *Handler, securityList *security.SecurityList) {
basePath := handler.config.BasePath
h := handler.AuthedStreamableHTTPServer(securityList)
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, h))
}
+168
View File
@@ -0,0 +1,168 @@
package resolvemcp
import (
"context"
"strings"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
)
func newReadOnlyHandler(t *testing.T) *Handler {
t.Helper()
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(),
Config{EnableAnnotations: true})
if err := h.RegisterModel("public", "items", &docItem{}); err != nil {
t.Fatal(err)
}
return h
}
func TestReadOnlyToolSet(t *testing.T) {
h := newReadOnlyHandler(t)
tools := h.mcpServer.ListTools()
for _, name := range []string{"list_tables", "describe_table", "select_table"} {
if tools[name] == nil {
t.Errorf("read tool %s missing", name)
}
}
for _, name := range []string{"insert_into_table", "update_table", "delete_from_table", "call_function", "list_functions", annotationToolName} {
if tools[name] != nil {
t.Errorf("tool %s must not be registered on a read-only server", name)
}
}
}
func TestReadOnlyRefusesWritesAndReportsIt(t *testing.T) {
h := newReadOnlyHandler(t)
ctx := context.Background()
args := map[string]any{"table": "public.items", "data": map[string]any{"name": "x"}, "id": 1}
for name, fn := range map[string]func() map[string]any{
"insert": func() map[string]any { r, _ := h.handleInsert(ctx, callReq(args)); return payload(t, r) },
"update": func() map[string]any { r, _ := h.handleUpdate(ctx, callReq(args)); return payload(t, r) },
"delete": func() map[string]any { r, _ := h.handleDelete(ctx, callReq(args)); return payload(t, r) },
} {
e, _ := fn()["error"].(map[string]any)
if e["code"] != CodeForbidden || !strings.Contains(e["message"].(string), "read-only") {
t.Errorf("%s: error = %v", name, e)
}
}
res, _ := h.handleListTables(ctx, callReq(nil))
tb := payload(t, res)["tables"].([]any)[0].(map[string]any)
if ops := tb["operations"].([]any); len(ops) != 1 || ops[0] != opSelect {
t.Errorf("list_tables operations = %v", ops)
}
res, _ = h.handleDescribeTable(ctx, callReq(map[string]any{"table": "public.items"}))
p := payload(t, res)
if p["read_only"] != true {
t.Errorf("describe_table read_only = %v", p["read_only"])
}
if w, _ := p["writable_columns"].([]any); len(w) != 0 {
t.Errorf("writable_columns = %v", w)
}
cat := h.BuildCatalog()
if !cat.ReadOnly || !strings.Contains(cat.Guide, "READ-ONLY") || !strings.Contains(cat.Markdown(), "read-only") {
t.Error("catalogue must say the server is read-only")
}
for _, c := range cat.Tables[0].Columns {
if c.Writable {
t.Errorf("column %s marked writable", c.Name)
}
}
if !strings.Contains(guideFor(true, false), "READ-ONLY") || strings.Contains(guideFor(false, false), "READ-ONLY") {
t.Error("guideFor")
}
}
func newFnHandler(t *testing.T, cfg Config) *Handler {
t.Helper()
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(), cfg)
for _, name := range []string{"alpha", "beta"} {
name := name
err := h.RegisterFunction(Function{Name: name, Handler: func(context.Context, common.Database, map[string]any) (any, error) {
return name, nil
}})
if err != nil {
t.Fatal(err)
}
}
return h
}
func TestReadOnlyAllowFunctionCalls(t *testing.T) {
h := newFnHandler(t, Config{AllowFunctionCalls: true})
tools := h.mcpServer.ListTools()
if tools["list_functions"] == nil || tools["call_function"] == nil {
t.Error("function tools must be registered")
}
if tools["insert_into_table"] != nil || tools["update_table"] != nil {
t.Error("write tools must stay off")
}
if g := guideFor(true, true); !strings.Contains(g, "READ-ONLY") || !strings.Contains(g, "call_function") {
t.Error("guide must mention functions")
}
if h.mcpServer.ListTools()["call_function"] == nil {
t.Error("call_function missing")
}
}
func TestAllowedFunctions(t *testing.T) {
ctx := context.Background()
for name, tc := range map[string]struct {
allowed []string
visible []string
}{
"empty allows all": {nil, []string{"alpha", "beta"}},
"only listed": {[]string{"beta"}, []string{"beta"}},
"unknown name": {[]string{"zzz"}, nil},
} {
h := newFnHandler(t, Config{AllowedFunctions: tc.allowed})
var got []string
for _, f := range h.visibleFunctions(ctx) {
got = append(got, f.Name)
}
if strings.Join(got, ",") != strings.Join(tc.visible, ",") {
t.Errorf("%s: visible = %v, want %v", name, got, tc.visible)
}
for _, fn := range []string{"alpha", "beta"} {
listed := false
for _, v := range tc.visible {
listed = listed || v == fn
}
if h.functionAllowed(fn) != listed {
t.Errorf("%s: functionAllowed(%s) = %v, want %v", name, fn, !listed, listed)
}
if !listed {
// refused before any database work, and indistinguishable from a missing function
if _, err := h.executeCall(ctx, fn, nil); err == nil || !strings.Contains(err.Error(), "unknown function") {
t.Errorf("%s: %s must be reported unknown, err=%v", name, fn, err)
}
}
}
}
}
func TestReadOnlyDefaultsOnAndCanBeDisabled(t *testing.T) {
if !(Config{}).withDefaults().readOnly {
t.Error("ReadOnly must default to on")
}
if !(Config{ReadOnly: Bool(true)}).withDefaults().readOnly {
t.Error("explicit true")
}
if (Config{ReadOnly: Bool(false)}).withDefaults().readOnly {
t.Error("Bool(false) must enable writes")
}
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(), Config{ReadOnly: Bool(false)})
if h.mcpServer.ListTools()["insert_into_table"] == nil {
t.Error("write tools must register when ReadOnly is Bool(false)")
}
if strings.Contains(h.BuildCatalog().Guide, "READ-ONLY") {
t.Error("guide must not claim read-only")
}
}
+163 -50
View File
@@ -1,23 +1,9 @@
// Package resolvemcp exposes registered database models as Model Context Protocol (MCP) tools
// and resources over HTTP/SSE transport.
//
// It mirrors the resolvespec package patterns:
// - Same model registration API
// - Same filter, sort, cursor pagination, preload options
// - Same lifecycle hook system
//
// Usage:
//
// handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{BaseURL: "http://localhost:8080"})
// handler.RegisterModel("public", "users", &User{})
//
// r := mux.NewRouter()
// resolvemcp.SetupMuxRoutes(r, handler)
package resolvemcp package resolvemcp
import ( import (
"net/http" "net/http"
"runtime/debug" "runtime/debug"
"time"
"github.com/gorilla/mux" "github.com/gorilla/mux"
"github.com/uptrace/bun" "github.com/uptrace/bun"
@@ -28,6 +14,7 @@ import (
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/logger" "github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry" "github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/security"
) )
// Config holds configuration for the resolvemcp handler. // Config holds configuration for the resolvemcp handler.
@@ -39,6 +26,87 @@ type Config struct {
// BasePath is the URL path prefix where the MCP endpoints are mounted (e.g. "/mcp"). // BasePath is the URL path prefix where the MCP endpoints are mounted (e.g. "/mcp").
// If empty, the path is detected from each incoming request automatically. // If empty, the path is detected from each incoming request automatically.
BasePath string BasePath string
// Limits. Zero values take the defaults shown.
// DefaultLimit is the page size when a read gives no limit (50).
DefaultLimit int
// MaxLimit caps a read's limit; larger values are clamped (1000).
MaxLimit int
// MaxOffset rejects a read whose offset is larger (100000).
MaxOffset int
// MaxBatch caps the items in one batch create (100).
MaxBatch int
// MaxPreloadDepth caps the depth of a preload path such as "a.b.c" (2).
MaxPreloadDepth int
// MaxWriteRows caps the rows a filter-based update or delete may touch (100).
MaxWriteRows int
// QueryTimeout bounds one tool call, hooks and queries included (30s).
QueryTimeout time.Duration
// ConfirmTTL is how long a confirmation token for a filter write stays valid (5m).
ConfirmTTL time.Duration
// AllowedHosts restricts the Host header accepted by the SSE transport when BaseURL is
// empty (the message endpoint URL sent to clients is built from it). Empty accepts any
// host, with at most 32 distinct base URLs cached; prefer setting BaseURL.
AllowedHosts []string
// ReadOnly disables every write and is ON when left nil: set it to Bool(false) to allow
// writes. When on, the insert, update, delete and annotation tools are not registered,
// list_tables and describe_table report only the select operation (no writable columns),
// a write attempted anyway is refused with a "forbidden" error, and the server
// instructions tell the agent it cannot write. list_functions/call_function are also
// off, because a registered function may change data, unless AllowFunctionCalls is set.
ReadOnly *bool
// readOnly is ReadOnly after defaults (nil means true).
readOnly bool
// AllowFunctionCalls keeps list_functions and call_function available on a ReadOnly
// server. Only set it for functions that do not change data; pair it with
// AllowedFunctions to name them. It has no effect when writes are enabled (ReadOnly set to Bool(false)) (functions are
// always available then).
AllowFunctionCalls bool
// AllowedFunctions restricts list_functions and call_function to the named functions.
// Empty allows every registered function. A function outside the list is reported as
// unknown, so its existence is not revealed.
AllowedFunctions []string
// EnableAnnotations registers the resolvespec_annotate tool. Off by default: annotations
// are free text that agents read back, so enabling the tool opens a write channel into
// agent-visible text. When on, every call runs the BeforeHandle hooks (operation
// "annotate_set" / "annotate_get") and the writes run in a transaction with OnTxBegin.
EnableAnnotations bool
}
// Bool returns a pointer to v, for the optional boolean fields of Config.
func Bool(v bool) *bool { return &v }
// withDefaults fills the zero limit fields.
func (c Config) withDefaults() Config {
def := func(v *int, d int) {
if *v <= 0 {
*v = d
}
}
def(&c.DefaultLimit, 50)
def(&c.MaxLimit, 1000)
def(&c.MaxOffset, 100000)
def(&c.MaxBatch, 100)
def(&c.MaxPreloadDepth, 2)
def(&c.MaxWriteRows, 100)
if c.DefaultLimit > c.MaxLimit {
c.DefaultLimit = c.MaxLimit
}
c.readOnly = c.ReadOnly == nil || *c.ReadOnly
if c.QueryTimeout <= 0 {
c.QueryTimeout = 30 * time.Second
}
if c.ConfirmTTL <= 0 {
c.ConfirmTTL = 5 * time.Minute
}
return c
} }
// NewHandlerWithGORM creates a Handler backed by a GORM database connection. // NewHandlerWithGORM creates a Handler backed by a GORM database connection.
@@ -57,18 +125,29 @@ func NewHandlerWithDB(db common.Database, cfg Config) *Handler {
} }
// SetupMuxRoutes mounts the MCP HTTP/SSE endpoints on the given Gorilla Mux router // SetupMuxRoutes mounts the MCP HTTP/SSE endpoints on the given Gorilla Mux router
// using the base path from Config.BasePath (falls back to "/mcp" if empty). // using the base path from Config.BasePath, behind Guard(securityList).
// //
// Two routes are registered: // Routes registered:
// - GET {basePath}/sse — SSE connection endpoint (client subscribes here) // - GET {basePath}/sse — SSE connection endpoint (client subscribes here)
// - POST {basePath}/message — JSON-RPC message endpoint (client sends requests here) // - POST {basePath}/message — JSON-RPC message endpoint (client sends requests here)
// //
// To protect these routes with authentication, wrap the mux router or apply middleware // Nothing is mounted (and an error is logged) when securityList has no provider.
// before calling SetupMuxRoutes. func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler, securityList *security.SecurityList) {
func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler) { if !requireGuard("SetupMuxRoutes", securityList) {
basePath := handler.config.BasePath return
h := handler.SSEServer() }
mountMuxSSE(muxRouter, handler, handler.AuthedSSEServer(securityList))
}
// SetupMuxRoutesUnauthenticated is SetupMuxRoutes without the guard. Every caller reaches every
// registered model, so use it only behind another trusted layer. A warning is logged.
func SetupMuxRoutesUnauthenticated(muxRouter *mux.Router, handler *Handler) {
warnUnauthenticated("SetupMuxRoutesUnauthenticated")
mountMuxSSE(muxRouter, handler, handler.SSEServer())
}
func mountMuxSSE(muxRouter *mux.Router, handler *Handler, h http.Handler) {
basePath := handler.config.BasePath
muxRouter.Handle(basePath+"/sse", h).Methods("GET", "OPTIONS") muxRouter.Handle(basePath+"/sse", h).Methods("GET", "OPTIONS")
muxRouter.Handle(basePath+"/message", h).Methods("POST", "OPTIONS") muxRouter.Handle(basePath+"/message", h).Methods("POST", "OPTIONS")
@@ -78,21 +157,32 @@ func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler) {
} }
// SetupBunRouterRoutes mounts the MCP HTTP/SSE endpoints on a bunrouter router // SetupBunRouterRoutes mounts the MCP HTTP/SSE endpoints on a bunrouter router
// using the base path from Config.BasePath. // using the base path from Config.BasePath, behind Guard(securityList).
// //
// Two routes are registered: // Routes registered:
// - GET {basePath}/sse — SSE connection endpoint // - GET {basePath}/sse — SSE connection endpoint
// - POST {basePath}/message — JSON-RPC message endpoint // - POST {basePath}/message — JSON-RPC message endpoint
func SetupBunRouterRoutes(router *bunrouter.Router, handler *Handler) { func SetupBunRouterRoutes(router *bunrouter.Router, handler *Handler, securityList *security.SecurityList) {
if !requireGuard("SetupBunRouterRoutes", securityList) {
return
}
mountBunSSE(router, handler, handler.AuthedSSEServer(securityList))
}
// SetupBunRouterRoutesUnauthenticated is SetupBunRouterRoutes without the guard. A warning is logged.
func SetupBunRouterRoutesUnauthenticated(router *bunrouter.Router, handler *Handler) {
warnUnauthenticated("SetupBunRouterRoutesUnauthenticated")
mountBunSSE(router, handler, handler.SSEServer())
}
func mountBunSSE(router *bunrouter.Router, handler *Handler, h http.Handler) {
defer func() { defer func() {
if rec := recover(); rec != nil { if rec := recover(); rec != nil {
logger.Error("panic in resolvemcp.SetupBunRouterRoutes: %v\n%s", rec, debug.Stack()) logger.Error("panic mounting resolvemcp bunrouter routes: %v\n%s", rec, debug.Stack())
} }
}() }()
basePath := handler.config.BasePath basePath := handler.config.BasePath
h := handler.SSEServer()
router.GET(basePath+"/sse", bunrouter.HTTPHandler(h)) router.GET(basePath+"/sse", bunrouter.HTTPHandler(h))
logger.Info("Registered resolvemcp bunrouter route GET %s/sse", basePath) logger.Info("Registered resolvemcp bunrouter route GET %s/sse", basePath)
@@ -100,45 +190,68 @@ func SetupBunRouterRoutes(router *bunrouter.Router, handler *Handler) {
logger.Info("Registered resolvemcp bunrouter route POST %s/message", basePath) logger.Info("Registered resolvemcp bunrouter route POST %s/message", basePath)
} }
// NewSSEServer returns an http.Handler that serves MCP over SSE. // NewSSEServer returns an http.Handler that serves MCP over SSE behind Guard(securityList).
// If Config.BasePath is set it is used directly; otherwise the base path is // If Config.BasePath is set it is used directly; otherwise the base path is
// detected from each incoming request (by stripping the "/sse" or "/message" suffix). // detected from each incoming request (by stripping the "/sse" or "/message" suffix).
// //
// h := resolvemcp.NewSSEServer(handler) // h := resolvemcp.NewSSEServer(handler, securityList)
// http.Handle("/api/mcp/", h) // http.Handle("/api/mcp/", h)
func NewSSEServer(handler *Handler) http.Handler { func NewSSEServer(handler *Handler, securityList *security.SecurityList) http.Handler {
return handler.SSEServer() return handler.AuthedSSEServer(securityList)
} }
// SetupMuxStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on the given Gorilla Mux router. // SetupMuxStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on the given Gorilla Mux
// The streamable HTTP transport uses a single endpoint (Config.BasePath) for all communication: // router, behind Guard(securityList). The streamable HTTP transport uses a single endpoint
// POST for client→server messages, GET for server→client streaming. // (Config.BasePath) for all communication: POST for client→server messages, GET for
// server→client streaming.
// //
// Example: // Nothing is mounted (and an error is logged) when securityList has no provider.
// func SetupMuxStreamableHTTPRoutes(muxRouter *mux.Router, handler *Handler, securityList *security.SecurityList) {
// resolvemcp.SetupMuxStreamableHTTPRoutes(r, handler) // mounts at Config.BasePath if !requireGuard("SetupMuxStreamableHTTPRoutes", securityList) {
func SetupMuxStreamableHTTPRoutes(muxRouter *mux.Router, handler *Handler) { return
}
basePath := handler.config.BasePath basePath := handler.config.BasePath
h := handler.StreamableHTTPServer() muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, handler.AuthedStreamableHTTPServer(securityList)))
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, h))
} }
// SetupBunRouterStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on a bunrouter router. // SetupMuxStreamableHTTPRoutesUnauthenticated is SetupMuxStreamableHTTPRoutes without the guard.
// The streamable HTTP transport uses a single endpoint (Config.BasePath). // A warning is logged.
func SetupBunRouterStreamableHTTPRoutes(router *bunrouter.Router, handler *Handler) { func SetupMuxStreamableHTTPRoutesUnauthenticated(muxRouter *mux.Router, handler *Handler) {
warnUnauthenticated("SetupMuxStreamableHTTPRoutesUnauthenticated")
basePath := handler.config.BasePath
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, handler.StreamableHTTPServer()))
}
// SetupBunRouterStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on a bunrouter
// router, behind Guard(securityList). The transport uses a single endpoint (Config.BasePath).
func SetupBunRouterStreamableHTTPRoutes(router *bunrouter.Router, handler *Handler, securityList *security.SecurityList) {
if !requireGuard("SetupBunRouterStreamableHTTPRoutes", securityList) {
return
}
mountBunStreamable(router, handler, handler.AuthedStreamableHTTPServer(securityList))
}
// SetupBunRouterStreamableHTTPRoutesUnauthenticated is SetupBunRouterStreamableHTTPRoutes
// without the guard. A warning is logged.
func SetupBunRouterStreamableHTTPRoutesUnauthenticated(router *bunrouter.Router, handler *Handler) {
warnUnauthenticated("SetupBunRouterStreamableHTTPRoutesUnauthenticated")
mountBunStreamable(router, handler, handler.StreamableHTTPServer())
}
func mountBunStreamable(router *bunrouter.Router, handler *Handler, h http.Handler) {
basePath := handler.config.BasePath basePath := handler.config.BasePath
h := handler.StreamableHTTPServer()
router.GET(basePath, bunrouter.HTTPHandler(h)) router.GET(basePath, bunrouter.HTTPHandler(h))
router.POST(basePath, bunrouter.HTTPHandler(h)) router.POST(basePath, bunrouter.HTTPHandler(h))
router.DELETE(basePath, bunrouter.HTTPHandler(h)) router.DELETE(basePath, bunrouter.HTTPHandler(h))
} }
// NewStreamableHTTPHandler returns an http.Handler that serves MCP over the streamable HTTP transport. // NewStreamableHTTPHandler returns an http.Handler that serves MCP over the streamable HTTP
// Mount it at the desired path; that path becomes the MCP endpoint. // transport behind Guard(securityList). Mount it at the desired path; that path becomes the
// MCP endpoint.
// //
// h := resolvemcp.NewStreamableHTTPHandler(handler) // h := resolvemcp.NewStreamableHTTPHandler(handler, securityList)
// http.Handle("/mcp", h) // http.Handle("/mcp", h)
// engine.Any("/mcp", gin.WrapH(h)) // engine.Any("/mcp", gin.WrapH(h))
func NewStreamableHTTPHandler(handler *Handler) http.Handler { func NewStreamableHTTPHandler(handler *Handler, securityList *security.SecurityList) http.Handler {
return handler.StreamableHTTPServer() return handler.AuthedStreamableHTTPServer(securityList)
} }
+25 -3
View File
@@ -31,7 +31,7 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
hookCtx.Abort = true hookCtx.Abort = true
hookCtx.AbortMessage = err.Error() hookCtx.AbortMessage = err.Error()
hookCtx.AbortCode = http.StatusUnauthorized hookCtx.AbortCode = http.StatusUnauthorized
return err return NewClientError(CodeForbidden, err.Error())
} }
return nil return nil
}) })
@@ -62,6 +62,15 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
return security.ApplyRowSecurity(newSecurityContext(hookCtx), securityList) return security.ApplyRowSecurity(newSecurityContext(hookCtx), securityList)
}) })
// BeforeScan: row-level security on the row an update or delete targets. A row the user
// cannot see is "not found" and is never written.
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
if err := security.LoadSecurityRules(newSecurityContext(hookCtx), securityList); err != nil {
return err
}
return security.ApplyRowSecurity(newSecurityContext(hookCtx), securityList)
})
// AfterRead (1st): apply column-level security — mask/hide columns in the result. // AfterRead (1st): apply column-level security — mask/hide columns in the result.
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error { handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
return security.ApplyColumnSecurity(newSecurityContext(hookCtx), securityList) return security.ApplyColumnSecurity(newSecurityContext(hookCtx), securityList)
@@ -72,14 +81,19 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
return security.LogDataAccess(newSecurityContext(hookCtx)) return security.LogDataAccess(newSecurityContext(hookCtx))
}) })
// BeforeCreate: enforce CanCreate rule.
handler.Hooks().Register(BeforeCreate, func(hookCtx *HookContext) error {
return forbidden(security.CheckModelCreateAllowed(newSecurityContext(hookCtx)))
})
// BeforeUpdate: enforce CanUpdate rule. // BeforeUpdate: enforce CanUpdate rule.
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error { handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
return security.CheckModelUpdateAllowed(newSecurityContext(hookCtx)) return forbidden(security.CheckModelUpdateAllowed(newSecurityContext(hookCtx)))
}) })
// BeforeDelete: enforce CanDelete rule. // BeforeDelete: enforce CanDelete rule.
handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error { handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error {
return security.CheckModelDeleteAllowed(newSecurityContext(hookCtx)) return forbidden(security.CheckModelDeleteAllowed(newSecurityContext(hookCtx)))
}) })
logger.Info("Security hooks registered for resolvemcp handler") logger.Info("Security hooks registered for resolvemcp handler")
@@ -153,3 +167,11 @@ func (s *securityContext) GetResult() interface{} {
func (s *securityContext) SetResult(result interface{}) { func (s *securityContext) SetResult(result interface{}) {
s.ctx.Result = result s.ctx.Result = result
} }
// forbidden marks a rule denial as safe to show the client; nil passes through.
func forbidden(err error) error {
if err == nil {
return nil
}
return NewClientError(CodeForbidden, err.Error())
}
+198
View File
@@ -0,0 +1,198 @@
package resolvemcp
import (
"context"
"strings"
"testing"
"github.com/DATA-DOG/go-sqlmock"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/security"
)
type wItem struct {
ID int `json:"id" bun:"id,pk"`
Name string `json:"name" bun:"name"`
FullName string `json:"fullName" bun:"full_name"`
Note *string `json:"note" bun:"note"`
Owner int `json:"-" bun:"-"`
}
func TestWriteColumns(t *testing.T) {
m := &wItem{}
got, err := writeColumns(m, map[string]interface{}{"fullName": "a", "NAME": "b", "note": nil})
if err != nil {
t.Fatal(err)
}
if got["full_name"] != "a" || got["name"] != "b" {
t.Errorf("json/column names must resolve to columns: %v", got)
}
if v, ok := got["note"]; !ok || v != nil {
t.Errorf("explicit null must be kept: %v", got)
}
for name, data := range map[string]map[string]interface{}{
"unknown": {"nope": 1},
"injection": {"name = 'x', id": 1},
"unmapped": {"owner": 1},
"both forms": {"fullName": 1, "full_name": 2},
"empty key": {"": 1},
"quoted char": {`"name"`: 1},
} {
if _, err := writeColumns(m, data); err == nil {
t.Errorf("%s: expected error", name)
}
}
}
func TestCreateRejectsUnknownKeys(t *testing.T) {
h, mock, ctx := newTxHarness(t)
mock.ExpectBegin()
mock.ExpectRollback()
_, err := h.executeCreate(ctx, "public", "items", map[string]interface{}{"name": "a", "is_admin": true})
if err == nil || !strings.Contains(err.Error(), "is_admin") {
t.Fatalf("want unknown field error, got %v", err)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestUpdateSetsOnlyGivenKeysAndAllowsNull(t *testing.T) {
db := wHarness(t)
h, mock, ctx := db.h, db.mock, db.ctx
cols := []string{"id", "name", "full_name", "note"}
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a", "A", "n"))
// Only "note" is set (to NULL); the id in the payload addresses the row and is not rewritten.
mock.ExpectExec(`UPDATE .* SET "?note"? = \$1 WHERE`).WithArgs(nil, "7").WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a", "A", nil))
mock.ExpectCommit()
if _, err := h.executeUpdate(ctx, "public", "witems", "7", map[string]interface{}{"id": 7, "note": nil}); err != nil {
t.Fatal(err)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestUpdateRejectsUnknownKeys(t *testing.T) {
db := wHarness(t)
db.mock.ExpectBegin()
db.mock.ExpectRollback()
if _, err := db.h.executeUpdate(db.ctx, "public", "witems", "7", map[string]interface{}{"role": "admin"}); err == nil {
t.Fatal("expected error")
}
if err := db.mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
type wh struct {
h *Handler
mock sqlmock.Sqlmock
ctx context.Context
}
// wHarness is newTxHarness with a model that has more than id/name.
func wHarness(t *testing.T) wh {
t.Helper()
h, mock, ctx := newTxHarness(t)
if err := h.RegisterModel("public", "witems", &wItem{}); err != nil {
t.Fatal(err)
}
return wh{h, mock, ctx}
}
// rlsProvider returns a fixed row security template.
type rlsProvider struct{ stubProvider }
func (rlsProvider) GetRowSecurity(_ context.Context, userRef any, schema, table string) (security.RowSecurity, error) {
return security.RowSecurity{Schema: schema, Tablename: table, Template: "owner_id = {UserID}", UserID: userRef}, nil
}
func securedHandler(t *testing.T, prov security.SecurityProvider) (*Handler, sqlmock.Sqlmock, context.Context) {
t.Helper()
h, mock, ctx := newTxHarness(t)
list, err := security.NewSecurityList(prov)
if err != nil {
t.Fatal(err)
}
RegisterSecurityHooks(h, list)
ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"})
ctx = context.WithValue(ctx, security.UserIDKey, 7)
return h, mock, ctx
}
// A row hidden by row security is "not found" for update and delete, and nothing is written.
func TestWritesHonourRowSecurity(t *testing.T) {
cols := []string{"id", "name"}
t.Run("update", func(t *testing.T) {
h, mock, ctx := securedHandler(t, rlsProvider{})
mock.ExpectBegin()
mock.ExpectQuery(`SELECT .*owner_id`).WillReturnRows(sqlmock.NewRows(cols))
mock.ExpectRollback()
if _, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "x"}); err == nil {
t.Fatal("expected not found")
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
})
t.Run("delete", func(t *testing.T) {
h, mock, ctx := securedHandler(t, rlsProvider{})
mock.ExpectBegin()
mock.ExpectQuery(`SELECT .*owner_id`).WillReturnRows(sqlmock.NewRows(cols))
mock.ExpectRollback()
if _, err := h.executeDelete(ctx, "public", "items", "7"); err == nil {
t.Fatal("expected not found")
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
})
}
// Rules set with RegisterModelWithRules must reach the security hooks.
func TestModelRulesReachHooks(t *testing.T) {
h, mock, ctx := securedHandler(t, stubProvider{})
if err := h.RegisterModelWithRules("public", "locked", &txItem{}, modelregistry.ModelRules{CanRead: true}); err != nil {
t.Fatal(err)
}
for name, op := range map[string]func() error{
"create": func() error {
_, err := h.executeCreate(ctx, "public", "locked", map[string]interface{}{"name": "a"})
return err
},
"update": func() error {
_, err := h.executeUpdate(ctx, "public", "locked", "7", map[string]interface{}{"name": "a"})
return err
},
"delete": func() error {
_, err := h.executeDelete(ctx, "public", "locked", "7")
return err
},
} {
// Each denies inside its transaction, before any statement.
mock.ExpectBegin()
mock.ExpectRollback()
if err := op(); err == nil || !strings.Contains(err.Error(), "not allowed") {
t.Errorf("%s: want 'not allowed', got %v", name, err)
}
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestAnnotationToolIsOptIn(t *testing.T) {
h, _, _ := newTxHarness(t)
if h.mcpServer.GetTool(annotationToolName) != nil {
t.Fatal("annotation tool must be off by default")
}
on := NewHandler(h.db, modelregistry.NewModelRegistry(), Config{EnableAnnotations: true, ReadOnly: Bool(false)})
if on.mcpServer.GetTool(annotationToolName) == nil {
t.Fatal("annotation tool missing when enabled")
}
}
+9 -370
View File
@@ -1,7 +1,6 @@
package resolvemcp package resolvemcp
import ( import (
"context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"reflect" "reflect"
@@ -10,34 +9,10 @@ import (
"github.com/mark3labs/mcp-go/mcp" "github.com/mark3labs/mcp-go/mcp"
"github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/logger" "github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/reflection" "github.com/bitechdev/ResolveSpec/pkg/reflection"
) )
// toolName builds the MCP tool name for a given operation and model.
func toolName(operation, schema, entity string) string {
if schema == "" {
return fmt.Sprintf("%s_%s", operation, entity)
}
return fmt.Sprintf("%s_%s_%s", operation, schema, entity)
}
// registerModelTools registers the four CRUD tools and resource for a model.
func registerModelTools(h *Handler, schema, entity string, model interface{}) {
info := buildModelInfo(schema, entity, model)
registerReadTool(h, schema, entity, info)
registerCreateTool(h, schema, entity, info)
registerUpdateTool(h, schema, entity, info)
registerDeleteTool(h, schema, entity, info)
registerModelResource(h, schema, entity, info)
logger.Info("[resolvemcp] Registered MCP tools for %s", info.fullName)
}
// --------------------------------------------------------------------------
// Model introspection
// --------------------------------------------------------------------------
// modelInfo holds pre-computed metadata for a model used in tool descriptions. // modelInfo holds pre-computed metadata for a model used in tool descriptions.
type modelInfo struct { type modelInfo struct {
fullName string // e.g. "public.users" fullName string // e.g. "public.users"
@@ -56,6 +31,7 @@ type columnInfo struct {
isUnique bool isUnique bool
isFK bool isFK bool
nullable bool nullable bool
comment string // from struct tags (gorm/bun comment:, comment/note/desc tags)
} }
// buildModelInfo extracts column metadata and pre-builds the schema documentation string. // buildModelInfo extracts column metadata and pre-builds the schema documentation string.
@@ -126,7 +102,12 @@ func buildModelInfo(schema, entity string, model interface{}) modelInfo {
isPrimary := d.SQLKey == "primary_key" || isPrimary := d.SQLKey == "primary_key" ||
(info.pkName != "" && (sqlName == info.pkName || jsonName == info.pkName)) (info.pkName != "" && (sqlName == info.pkName || jsonName == info.pkName))
comment := ""
if found {
comment = modelregistry.FieldComment(fieldType)
}
ci := columnInfo{ ci := columnInfo{
comment: comment,
jsonName: jsonName, jsonName: jsonName,
sqlName: sqlName, sqlName: sqlName,
goType: goType, goType: goType,
@@ -227,350 +208,8 @@ func buildSchemaDoc(info modelInfo) string {
return sb.String() return sb.String()
} }
// columnNameList returns a comma-separated list of JSON column names (for descriptions). // parseRequestOptions reads the paging, filter, sort, column and preload arguments shared
func columnNameList(cols []columnInfo) string { // by the read tools.
names := make([]string, len(cols))
for i, c := range cols {
names[i] = c.jsonName
}
return strings.Join(names, ", ")
}
// writableColumnNames returns JSON names for all non-primary-key columns.
func writableColumnNames(cols []columnInfo) []string {
var names []string
for _, c := range cols {
if !c.isPrimary {
names = append(names, c.jsonName)
}
}
return names
}
// --------------------------------------------------------------------------
// Read tool
// --------------------------------------------------------------------------
func registerReadTool(h *Handler, schema, entity string, info modelInfo) {
name := toolName("read", schema, entity)
var descParts []string
descParts = append(descParts, fmt.Sprintf("Read records from the '%s' database table.", info.fullName))
if info.pkName != "" {
descParts = append(descParts, fmt.Sprintf("Primary key: '%s'. Pass it via 'id' to fetch a single record.", info.pkName))
}
if info.schemaDoc != "" {
descParts = append(descParts, info.schemaDoc)
}
descParts = append(descParts,
"Pagination: use 'limit'/'offset' for offset-based paging, or 'cursor_forward'/'cursor_backward' (pass the primary key value of the last/first record on the current page) for cursor-based paging.",
"Filtering: each filter object requires 'column' (JSON field name) and 'operator'. Supported operators: = != > < >= <= like ilike in is_null is_not_null. Combine with 'logic_operator': AND (default) or OR.",
"Sorting: each sort object requires 'column' and 'direction' (asc or desc).",
)
if len(info.relationNames) > 0 {
descParts = append(descParts, fmt.Sprintf("Preloadable relations: %s. Pass relation name in 'preloads'.", strings.Join(info.relationNames, ", ")))
}
description := strings.Join(descParts, "\n\n")
filterDesc := `Array of filter objects. Example: [{"column":"status","operator":"=","value":"active"},{"column":"age","operator":">","value":18,"logic_operator":"AND"}]`
if len(info.columns) > 0 {
filterDesc += fmt.Sprintf(" Available columns: %s.", columnNameList(info.columns))
}
sortDesc := `Array of sort objects. Example: [{"column":"created_at","direction":"desc"}]`
if len(info.columns) > 0 {
sortDesc += fmt.Sprintf(" Available columns: %s.", columnNameList(info.columns))
}
tool := mcp.NewTool(name,
mcp.WithDescription(description),
mcp.WithString("id",
mcp.Description(fmt.Sprintf("Primary key (%s) of a single record to fetch. Omit to return multiple records.", info.pkName)),
),
mcp.WithNumber("limit",
mcp.Description("Maximum number of records to return per page. Recommended: 10–100."),
),
mcp.WithNumber("offset",
mcp.Description("Number of records to skip (for offset-based pagination). Use with 'limit'."),
),
mcp.WithString("cursor_forward",
mcp.Description(fmt.Sprintf("Cursor for the next page: pass the '%s' value of the last record on the current page. Requires 'sort' to be set.", info.pkName)),
),
mcp.WithString("cursor_backward",
mcp.Description(fmt.Sprintf("Cursor for the previous page: pass the '%s' value of the first record on the current page. Requires 'sort' to be set.", info.pkName)),
),
mcp.WithArray("columns",
mcp.Description(fmt.Sprintf("Columns to include in the result. Omit to return all columns. Available: %s.", columnNameList(info.columns))),
),
mcp.WithArray("omit_columns",
mcp.Description(fmt.Sprintf("Columns to exclude from the result. Available: %s.", columnNameList(info.columns))),
),
mcp.WithArray("filters",
mcp.Description(filterDesc),
),
mcp.WithArray("sort",
mcp.Description(sortDesc),
),
mcp.WithArray("preloads",
mcp.Description(buildPreloadDesc(info)),
),
)
h.mcpServer.AddTool(tool, func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := req.GetArguments()
id, _ := args["id"].(string)
options := parseRequestOptions(args)
data, metadata, err := h.executeRead(ctx, schema, entity, id, options)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
return marshalResult(map[string]interface{}{
"success": true,
"data": data,
"metadata": metadata,
})
})
}
func buildPreloadDesc(info modelInfo) string {
if len(info.relationNames) == 0 {
return `Array of relation preload objects. Each object: {"relation":"RelationName"}. No relations defined on this model.`
}
return fmt.Sprintf(
`Array of relation preload objects. Each object: {"relation":"RelationName","columns":["col1","col2"]}. Available relations: %s.`,
strings.Join(info.relationNames, ", "),
)
}
// --------------------------------------------------------------------------
// Create tool
// --------------------------------------------------------------------------
func registerCreateTool(h *Handler, schema, entity string, info modelInfo) {
name := toolName("create", schema, entity)
writable := writableColumnNames(info.columns)
var descParts []string
descParts = append(descParts, fmt.Sprintf("Create one or more new records in the '%s' table.", info.fullName))
if len(writable) > 0 {
descParts = append(descParts, fmt.Sprintf("Writable fields: %s.", strings.Join(writable, ", ")))
}
if info.pkName != "" {
descParts = append(descParts, fmt.Sprintf("The primary key ('%s') is typically auto-generated — omit it unless you need to supply it explicitly.", info.pkName))
}
descParts = append(descParts,
"Pass a single JSON object to 'data' to create one record. Pass an array of objects to create multiple records in a single transaction (all succeed or all fail).",
)
if info.schemaDoc != "" {
descParts = append(descParts, info.schemaDoc)
}
description := strings.Join(descParts, "\n\n")
dataDesc := "Record fields to create."
if len(writable) > 0 {
dataDesc += fmt.Sprintf(" Writable fields: %s.", strings.Join(writable, ", "))
}
dataDesc += " Pass a single object or an array of objects."
tool := mcp.NewTool(name,
mcp.WithDescription(description),
mcp.WithObject("data",
mcp.Description(dataDesc),
mcp.Required(),
),
)
h.mcpServer.AddTool(tool, func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := req.GetArguments()
data, ok := args["data"]
if !ok {
return mcp.NewToolResultError("missing required argument: data"), nil
}
result, err := h.executeCreate(ctx, schema, entity, data)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
return marshalResult(map[string]interface{}{
"success": true,
"data": result,
})
})
}
// --------------------------------------------------------------------------
// Update tool
// --------------------------------------------------------------------------
func registerUpdateTool(h *Handler, schema, entity string, info modelInfo) {
name := toolName("update", schema, entity)
writable := writableColumnNames(info.columns)
var descParts []string
descParts = append(descParts, fmt.Sprintf("Update an existing record in the '%s' table.", info.fullName))
if info.pkName != "" {
descParts = append(descParts, fmt.Sprintf("Identify the record by its primary key ('%s') via the 'id' argument or by including '%s' inside 'data'.", info.pkName, info.pkName))
}
if len(writable) > 0 {
descParts = append(descParts, fmt.Sprintf("Updatable fields: %s.", strings.Join(writable, ", ")))
}
descParts = append(descParts,
"Only non-null, non-empty fields in 'data' are applied — existing values are preserved for fields you omit. Returns the merged record as stored.",
)
if info.schemaDoc != "" {
descParts = append(descParts, info.schemaDoc)
}
description := strings.Join(descParts, "\n\n")
idDesc := fmt.Sprintf("Primary key ('%s') of the record to update. Can also be included inside 'data'.", info.pkName)
dataDesc := "Fields to update (non-null, non-empty values are merged into the existing record)."
if len(writable) > 0 {
dataDesc += fmt.Sprintf(" Updatable fields: %s.", strings.Join(writable, ", "))
}
tool := mcp.NewTool(name,
mcp.WithDescription(description),
mcp.WithString("id",
mcp.Description(idDesc),
),
mcp.WithObject("data",
mcp.Description(dataDesc),
mcp.Required(),
),
)
h.mcpServer.AddTool(tool, func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := req.GetArguments()
id, _ := args["id"].(string)
data, ok := args["data"]
if !ok {
return mcp.NewToolResultError("missing required argument: data"), nil
}
dataMap, ok := data.(map[string]interface{})
if !ok {
return mcp.NewToolResultError("data must be an object"), nil
}
result, err := h.executeUpdate(ctx, schema, entity, id, dataMap)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
return marshalResult(map[string]interface{}{
"success": true,
"data": result,
})
})
}
// --------------------------------------------------------------------------
// Delete tool
// --------------------------------------------------------------------------
func registerDeleteTool(h *Handler, schema, entity string, info modelInfo) {
name := toolName("delete", schema, entity)
descParts := []string{
fmt.Sprintf("Delete a record from the '%s' table by its primary key.", info.fullName),
}
if info.pkName != "" {
descParts = append(descParts, fmt.Sprintf("Pass the '%s' value of the record to delete via the 'id' argument.", info.pkName))
}
descParts = append(descParts, "Returns the deleted record. This operation is irreversible.")
description := strings.Join(descParts, " ")
tool := mcp.NewTool(name,
mcp.WithDescription(description),
mcp.WithString("id",
mcp.Description(fmt.Sprintf("Primary key ('%s') of the record to delete.", info.pkName)),
mcp.Required(),
),
)
h.mcpServer.AddTool(tool, func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := req.GetArguments()
id, _ := args["id"].(string)
result, err := h.executeDelete(ctx, schema, entity, id)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
return marshalResult(map[string]interface{}{
"success": true,
"data": result,
})
})
}
// --------------------------------------------------------------------------
// Resource registration
// --------------------------------------------------------------------------
func registerModelResource(h *Handler, schema, entity string, info modelInfo) {
resourceURI := info.fullName
var resourceDesc strings.Builder
fmt.Fprintf(&resourceDesc, "Database table: %s", info.fullName)
if info.pkName != "" {
fmt.Fprintf(&resourceDesc, " (primary key: %s)", info.pkName)
}
if info.schemaDoc != "" {
resourceDesc.WriteString("\n\n")
resourceDesc.WriteString(info.schemaDoc)
}
resource := mcp.NewResource(
resourceURI,
entity,
mcp.WithResourceDescription(resourceDesc.String()),
mcp.WithMIMEType("application/json"),
)
h.mcpServer.AddResource(resource, func(ctx context.Context, req mcp.ReadResourceRequest) ([]mcp.ResourceContents, error) {
limit := 100
options := common.RequestOptions{Limit: &limit}
data, metadata, err := h.executeRead(ctx, schema, entity, "", options)
if err != nil {
return nil, err
}
payload := map[string]interface{}{
"data": data,
"metadata": metadata,
}
jsonBytes, err := json.Marshal(payload)
if err != nil {
return nil, fmt.Errorf("error marshaling resource: %w", err)
}
return []mcp.ResourceContents{
mcp.TextResourceContents{
URI: req.Params.URI,
MIMEType: "application/json",
Text: string(jsonBytes),
},
}, nil
})
}
// --------------------------------------------------------------------------
// Argument parsing helpers
// --------------------------------------------------------------------------
// parseRequestOptions converts raw MCP tool arguments into common.RequestOptions.
func parseRequestOptions(args map[string]interface{}) common.RequestOptions { func parseRequestOptions(args map[string]interface{}) common.RequestOptions {
options := common.RequestOptions{} options := common.RequestOptions{}
+36 -29
View File
@@ -29,7 +29,7 @@ func newTxHarness(t *testing.T) (*Handler, sqlmock.Sqlmock, context.Context) {
// connection and fails on the context timeout. // connection and fails on the context timeout.
db.SetMaxOpenConns(1) db.SetMaxOpenConns(1)
t.Cleanup(func() { _ = db.Close() }) t.Cleanup(func() { _ = db.Close() })
h := NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry(), Config{}) h := NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry(), Config{ReadOnly: Bool(false)})
if err := h.RegisterModel("public", "items", &txItem{}); err != nil { if err := h.RegisterModel("public", "items", &txItem{}); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -128,14 +128,12 @@ func TestReadRunsInOneTransaction(t *testing.T) {
} }
} }
func TestCreateSingleUsesTwoTransactions(t *testing.T) { func TestCreateSingleRunsInOneTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(t) h, mock, ctx := newTxHarness(t)
tr := traceHooks(h, OnTxBegin, BeforeCreate, AfterCreate) tr := traceHooks(h, OnTxBegin, BeforeCreate, AfterCreate)
mock.ExpectBegin() mock.ExpectBegin()
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7)) mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a"))
mock.ExpectCommit() mock.ExpectCommit()
@@ -145,24 +143,19 @@ func TestCreateSingleUsesTwoTransactions(t *testing.T) {
if err := mock.ExpectationsWereMet(); err != nil { if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
tr.assertOrder(t, "on_tx_begin", "before_create", "on_tx_begin", "after_create") tr.assertOrder(t, "on_tx_begin", "before_create", "after_create")
if tr.txs["before_create"][0] != tr.txs["on_tx_begin"][0] { if tr.txs["before_create"][0] != tr.txs["on_tx_begin"][0] || tr.txs["after_create"][0] != tr.txs["on_tx_begin"][0] {
t.Fatal("BeforeCreate must run on the first transaction") t.Fatal("BeforeCreate, the re-fetch and AfterCreate must share one transaction")
}
if tr.txs["after_create"][0] != tr.txs["on_tx_begin"][1] || tr.txs["on_tx_begin"][0] == tr.txs["on_tx_begin"][1] {
t.Fatal("AfterCreate must run on a second, distinct transaction")
} }
} }
func TestCreateBatchRefetchOnSecondTransaction(t *testing.T) { func TestCreateBatchRefetchInSameTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(t) h, mock, ctx := newTxHarness(t)
tr := traceHooks(h, OnTxBegin, AfterCreate) tr := traceHooks(h, OnTxBegin, AfterCreate)
mock.ExpectBegin() mock.ExpectBegin()
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1)) mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2)) mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a")) mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a"))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(2, "b")) mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(2, "b"))
mock.ExpectCommit() mock.ExpectCommit()
@@ -174,10 +167,10 @@ func TestCreateBatchRefetchOnSecondTransaction(t *testing.T) {
if err := mock.ExpectationsWereMet(); err != nil { if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
tr.assertOrder(t, "on_tx_begin", "on_tx_begin", "after_create") tr.assertOrder(t, "on_tx_begin", "after_create")
} }
func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) { func TestUpdateRefetchRunsInSameTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(t) h, mock, ctx := newTxHarness(t)
tr := traceHooks(h, OnTxBegin, BeforeUpdate, AfterUpdate) tr := traceHooks(h, OnTxBegin, BeforeUpdate, AfterUpdate)
@@ -185,25 +178,38 @@ func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) {
mock.ExpectBegin() mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a"))
mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b")) mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b"))
mock.ExpectCommit() mock.ExpectCommit()
if _, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "b"}); err != nil { res, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "b"})
if err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := mock.ExpectationsWereMet(); err != nil { if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
tr.assertOrder(t, "on_tx_begin", "before_update", "after_update", "on_tx_begin") if m, _ := res.(map[string]interface{}); m["name"] != "b" {
if tr.txs["on_tx_begin"][0] == tr.txs["on_tx_begin"][1] { t.Fatalf("result must be the re-fetched row, got %v", res)
t.Fatal("re-fetch must run on a second transaction")
} }
for _, ht := range []string{"before_update", "after_update"} { tr.assertOrder(t, "on_tx_begin", "before_update", "after_update")
if tr.txs[ht][0] != tr.txs["on_tx_begin"][0] { }
t.Fatalf("%s must run on the first transaction", ht)
} // A failing AfterCreate must roll the insert back: the client sees an error, so nothing may
// have been committed that a retry would duplicate.
func TestAfterCreateErrorRollsBackInsert(t *testing.T) {
h, mock, ctx := newTxHarness(t)
h.Hooks().Register(AfterCreate, func(*HookContext) error { return sql.ErrConnDone })
mock.ExpectBegin()
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a"))
mock.ExpectRollback()
if _, err := h.executeCreate(ctx, "public", "items", map[string]interface{}{"name": "a"}); err == nil {
t.Fatal("expected error")
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
} }
} }
@@ -297,6 +303,10 @@ func (stubProvider) GetColumnSecurity(context.Context, int, string, string) ([]s
return nil, nil return nil, nil
} }
func (stubProvider) GetRowSecurity(context.Context, any, string, string) (security.RowSecurity, error) {
return security.RowSecurity{}, nil
}
func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) { func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(t) h, mock, ctx := newTxHarness(t)
list, err := security.NewSecurityList(stubProvider{}) list, err := security.NewSecurityList(stubProvider{})
@@ -310,15 +320,12 @@ func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) {
ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"})
ctx = context.WithValue(ctx, security.UserIDKey, 7) ctx = context.WithValue(ctx, security.UserIDKey, 7)
// Update opens two transactions; each must be stamped before any other SQL. // The transaction is stamped before any other SQL.
cols := []string{"id", "name"} cols := []string{"id", "name"}
mock.ExpectBegin() mock.ExpectBegin()
mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0)) mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a"))
mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b")) mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b"))
mock.ExpectCommit() mock.ExpectCommit()
+55
View File
@@ -0,0 +1,55 @@
package resolvemcp
import (
"reflect"
"sort"
"strings"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
)
// maxKeyEcho caps how much of a rejected client key is echoed back in an error.
const maxKeyEcho = 64
// writeColumns validates the keys of a create/update payload against the model and returns the
// values keyed by database column name. A key may be the json name or the column name of a
// writable field (case-insensitive); relations, scan-only and unexported fields are not
// writable. Unknown keys are rejected rather than dropped, so a client learns that a write
// did not take effect, and no client-chosen identifier reaches SQL.
func writeColumns(model interface{}, data map[string]interface{}) (map[string]interface{}, error) {
modelType := reflect.TypeOf(model)
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice) {
modelType = modelType.Elem()
}
if modelType == nil || modelType.Kind() != reflect.Struct {
return nil, errInternal
}
accepted := make(map[string]string)
for jsonKey, col := range reflection.BuildJSONToDBColumnMap(modelType) {
accepted[strings.ToLower(jsonKey)] = col
accepted[strings.ToLower(col)] = col
}
out := make(map[string]interface{}, len(data))
var unknown []string
for key, value := range data {
col, ok := accepted[strings.ToLower(key)]
if !ok {
if len(key) > maxKeyEcho {
key = key[:maxKeyEcho] + "..."
}
unknown = append(unknown, key)
continue
}
if _, dup := out[col]; dup {
return nil, invalidArg("column %q given more than once", col)
}
out[col] = value
}
if len(unknown) > 0 {
sort.Strings(unknown)
return nil, invalidArg("unknown or read-only fields: %s", strings.Join(unknown, ", "))
}
return out, nil
}
+260
View File
@@ -0,0 +1,260 @@
package resolvemcp
import (
"context"
"fmt"
"reflect"
"regexp"
"sort"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
"github.com/bitechdev/ResolveSpec/pkg/security"
)
// previewRows is how many matched primary keys a preview lists.
const previewRows = 10
var filterColumnRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
// writeOperators are the filter operators a filter write accepts.
var writeOperators = map[string]bool{
"eq": true, "=": true, "neq": true, "!=": true, "<>": true, "gt": true, ">": true, "gte": true, ">=": true,
"lt": true, "<": true, "lte": true, "<=": true, "like": true, "ilike": true, "in": true,
"is_null": true, "is_not_null": true,
}
// validateWriteFilters requires at least one filter and that every filter is usable. Reads
// silently drop filters they cannot apply; a write must not, because a dropped filter widens
// the set of rows the write touches.
func validateWriteFilters(model interface{}, filters []common.FilterOption) error {
if len(filters) == 0 {
return invalidArg("provide an id or at least one filter")
}
v := common.NewColumnValidator(model)
for _, f := range filters {
if !filterColumnRe.MatchString(f.Column) || !v.IsValidColumn(f.Column) {
return invalidArg("unknown filter column %q", truncate(f.Column))
}
if !writeOperators[strings.ToLower(f.Operator)] {
return invalidArg("unsupported filter operator %q", truncate(f.Operator))
}
op := strings.ToLower(f.Operator)
if op != "is_null" && op != "is_not_null" && f.Value == nil {
return invalidArg("filter on %q needs a value", f.Column)
}
}
return nil
}
// whereRequest describes a filter-based update or delete.
type whereRequest struct {
schema, entity string
op string // "update" or "delete"
filters []common.FilterOption
data map[string]interface{} // update only
dryRun bool
confirmToken string
}
// whereResult is what a filter write returns to the client.
type whereResult struct {
DryRun bool `json:"dry_run,omitempty"`
RequiresConfirm bool `json:"requires_confirmation,omitempty"`
Matched int `json:"matched"`
Preview []interface{} `json:"preview,omitempty"`
ConfirmToken string `json:"confirm_token,omitempty"`
ExpiresInSec int `json:"expires_in_seconds,omitempty"`
Affected int `json:"affected,omitempty"`
IDs []interface{} `json:"ids,omitempty"`
}
// executeWhere runs a filter-based write behind the guardrails: filters are validated, the
// matching rows are counted inside the transaction (aborting above Config.MaxWriteRows), and
// without dry_run the write only happens with a confirm token from a preview of the same
// request that matched the same rows.
func (h *Handler) executeWhere(ctx context.Context, req whereRequest) (_ *whereResult, retErr error) {
defer recoverPanic(&retErr)
ctx, cancel := h.callContext(ctx)
defer cancel()
model, err := h.registry.GetModelByEntity(req.schema, req.entity)
if err != nil {
return nil, invalidArg("model not found: %s", buildModelName(req.schema, req.entity))
}
unwrapped, err := common.ValidateAndUnwrapModel(model)
if err != nil {
return nil, errInternal
}
model = unwrapped.Model
if err := validateWriteFilters(model, req.filters); err != nil {
return nil, err
}
pkName := reflection.GetPrimaryKeyName(model)
if pkName == "" {
return nil, invalidArg("table has no primary key; filter writes are not available")
}
tableName := h.getTableName(req.schema, req.entity, model)
ctx = withRequestData(h.withModelRules(ctx, req.schema, req.entity), req.schema, req.entity, tableName, model, unwrapped.ModelPtr)
var setCols map[string]interface{}
if req.op == "update" {
if setCols, err = writeColumns(model, req.data); err != nil {
return nil, err
}
for col := range setCols {
if strings.EqualFold(col, pkName) {
delete(setCols, col)
}
}
if len(setCols) == 0 {
return nil, invalidArg("no updatable fields in data")
}
}
hookCtx := &HookContext{
Context: ctx, Handler: h, Schema: req.schema, Entity: req.entity, Model: model,
Operation: req.op, Tx: h.db,
}
if req.op == "update" {
hookCtx.Data = req.data
}
if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil {
return nil, err
}
user := callerKey(ctx)
table := buildModelName(req.schema, req.entity)
res := &whereResult{}
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
before, after := BeforeUpdate, AfterUpdate
if req.op == "delete" {
before, after = BeforeDelete, AfterDelete
}
if err := h.hooks.Execute(before, hookCtx); err != nil {
return err
}
if req.op == "update" {
// A hook (column security) may have narrowed the payload.
if m, ok := hookCtx.Data.(map[string]interface{}); ok {
if setCols, err = writeColumns(model, m); err != nil {
return err
}
for col := range setCols {
if strings.EqualFold(col, pkName) {
delete(setCols, col)
}
}
if len(setCols) == 0 {
return invalidArg("no updatable fields in data")
}
}
}
ids, err := h.matchRows(ctx, tx, hookCtx, model, unwrapped.ModelType, tableName, pkName, req.filters)
if err != nil {
return err
}
res.Matched = len(ids)
for i := 0; i < len(ids) && i < previewRows; i++ {
res.Preview = append(res.Preview, ids[i])
}
if req.dryRun {
res.DryRun = true
return nil
}
binding, err := bindingHash(req.op, req.filters, req.data, ids)
if err != nil {
return err
}
if req.confirmToken == "" {
tok, err := h.confirms.issue(user, table, req.op, binding, h.config.ConfirmTTL)
if err != nil {
return err
}
res.RequiresConfirm = true
res.ConfirmToken = tok
res.ExpiresInSec = int(h.config.ConfirmTTL / time.Second)
return nil
}
if err := h.confirms.consume(req.confirmToken, user, table, req.op, binding); err != nil {
return err
}
if len(ids) == 0 {
return nil
}
inList := make([]string, len(ids))
for i := range ids {
inList[i] = "?"
}
cond := fmt.Sprintf("%s IN (%s)", common.QuoteIdent(pkName), strings.Join(inList, ", "))
var affected int64
if req.op == "update" {
reflection.RemoveNonWritableColumns(model, setCols)
r, err := tx.NewUpdate().Table(tableName).SetMap(setCols).Where(cond, ids...).Exec(ctx)
if err != nil {
return fmt.Errorf("error updating records: %w", err)
}
affected = r.RowsAffected()
} else {
r, err := tx.NewDelete().Table(tableName).Where(cond, ids...).Exec(ctx)
if err != nil {
return fmt.Errorf("delete error: %w", err)
}
affected = r.RowsAffected()
}
if int(affected) != len(ids) {
// Rows changed under us: roll back rather than report a partial write.
return invalidArg("matched rows changed during the write; repeat the preview")
}
res.Affected = int(affected)
res.IDs = ids
hookCtx.Result = map[string]interface{}{"ids": ids, "count": len(ids)}
return h.hooks.Execute(after, hookCtx)
})
if err != nil {
return nil, err
}
return res, nil
}
// matchRows returns the primary keys of the rows the filters select, narrowed by the BeforeScan
// hooks (row security). It fails when more than Config.MaxWriteRows match.
func (h *Handler) matchRows(ctx context.Context, tx common.Database, hookCtx *HookContext, model interface{}, modelType reflect.Type, tableName, pkName string, filters []common.FilterOption) ([]interface{}, error) {
sliceType := reflect.SliceOf(reflect.PointerTo(modelType))
dest := reflect.New(sliceType)
q := tx.NewSelect().Model(dest.Interface())
if provider, ok := reflect.New(modelType).Interface().(common.TableNameProvider); !ok || provider.TableName() == "" {
q = q.Table(tableName)
}
q = h.applyFilters(q.Column(pkName), filters, model)
hookCtx.Query = q
if err := h.hooks.Execute(BeforeScan, hookCtx); err != nil {
return nil, err
}
if err := hookCtx.Query.Limit(h.config.MaxWriteRows + 1).ScanModel(ctx); err != nil {
return nil, fmt.Errorf("error matching records: %w", err)
}
rows := dest.Elem()
if rows.Len() > h.config.MaxWriteRows {
return nil, NewClientError(CodeLimitExceeded, fmt.Sprintf("the filters match more than %d rows; narrow them", h.config.MaxWriteRows))
}
ids := make([]interface{}, 0, rows.Len())
for i := 0; i < rows.Len(); i++ {
ids = append(ids, reflection.GetPrimaryKeyValue(rows.Index(i).Interface()))
}
sort.Slice(ids, func(i, j int) bool { return fmt.Sprint(ids[i]) < fmt.Sprint(ids[j]) })
return ids, nil
}
// callerKey identifies the caller for confirm-token binding.
func callerKey(ctx context.Context) string {
if uc, ok := security.GetUserContext(ctx); ok && uc != nil {
return fmt.Sprintf("%d/%s", uc.UserID, uc.UserName)
}
return "anonymous"
}
+27 -1
View File
@@ -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))
@@ -1244,6 +1247,16 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
return return
} }
// A primary key change is only honoured when the ID was given in the URL and
// the body carries a different, non-null primary key value.
var newPK interface{}
pkChanged := false
if urlID != "" {
if v, ok := updates[pkName]; ok && v != nil && !reflection.IsEmptyValue(v) && fmt.Sprintf("%v", v) != urlID {
newPK, pkChanged = v, true
}
}
// Wrap in transaction to ensure BeforeUpdate hook is inside transaction // Wrap in transaction to ensure BeforeUpdate hook is inside transaction
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error { err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
// Execute BeforeUpdate hooks inside transaction, before any queries run. // Execute BeforeUpdate hooks inside transaction, before any queries run.
@@ -1312,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)
@@ -1337,6 +1351,14 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
return fmt.Errorf("no records found to update") return fmt.Errorf("no records found to update")
} }
// SetMap skips primary key columns, so apply a PK change explicitly.
if pkChanged {
if _, err := tx.NewUpdate().Table(tableName).Set(pkName, newPK).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID).Exec(ctx); err != nil {
return fmt.Errorf("error updating primary key: %w", err)
}
}
// Execute AfterUpdate hooks inside transaction // Execute AfterUpdate hooks inside transaction
hookCtx.Result = updates hookCtx.Result = updates
hookCtx.Error = nil hookCtx.Error = nil
@@ -1362,7 +1384,9 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
updatedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() updatedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error { if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
fetchQuery := tx.NewSelect().Model(updatedRecord).Column(reflection.GetSQLModelColumns(model)...) fetchQuery := tx.NewSelect().Model(updatedRecord).Column(reflection.GetSQLModelColumns(model)...)
if urlID != "" { if pkChanged {
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), newPK)
} else if urlID != "" {
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), urlID) fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), urlID)
} else if reqID != nil { } else if reqID != nil {
switch id := reqID.(type) { switch id := reqID.(type) {
@@ -1487,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 {
@@ -1642,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 {
+61
View File
@@ -0,0 +1,61 @@
package resolvespec
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/uptrace/bunrouter"
)
type wrapCtxKey struct{}
// The auth wrapper must hand the handler the middleware-enriched request
// without dropping the bunrouter route params.
func TestWrapBunRouterHandler_PreservesRouteParams(t *testing.T) {
var gotSchema, gotEntity, gotID string
var gotCtxVal any
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotSchema = req.Param("schema")
gotEntity = req.Param("entity")
gotID = req.Param("id")
gotCtxVal = req.Context().Value(wrapCtxKey{})
return nil
}
auth := func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), wrapCtxKey{}, "enriched")))
})
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, auth))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/public/users/42", nil))
if gotSchema != "public" || gotEntity != "users" || gotID != "42" {
t.Errorf("route params lost: schema=%q entity=%q id=%q", gotSchema, gotEntity, gotID)
}
if gotCtxVal != "enriched" {
t.Errorf("handler did not see middleware-enriched context, got %v", gotCtxVal)
}
}
func TestWrapBunRouterHandler_NilAuthPassesThrough(t *testing.T) {
var gotID string
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotID = req.Param("id")
return nil
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, nil))
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/public/users/7", nil))
if gotID != "7" {
t.Errorf("id = %q, want 7", gotID)
}
}
+149 -64
View File
@@ -241,9 +241,11 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
h.sendError(w, http.StatusBadRequest, "invalid_request", "Invalid request body", err) h.sendError(w, http.StatusBadRequest, "invalid_request", "Invalid request body", err)
return return
} }
validId, _ := strconv.ParseInt(id, 10, 64) // A URL id is valid when it is a positive integer or any non-numeric
// string (string primary keys); "", "0" and negatives mean no id.
validId, parseErr := strconv.ParseInt(id, 10, 64)
updateID := id updateID := id
isUpdate := validId > 0 isUpdate := id != "" && (parseErr != nil || validId > 0)
if !isUpdate { if !isUpdate {
// No valid /:id in the URL - check whether the body itself carries // No valid /:id in the URL - check whether the body itself carries
// a valid primary key value and treat this as an update if so. // a valid primary key value and treat this as an update if so.
@@ -415,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
@@ -644,77 +654,92 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
// This may need to be handled differently per database adapter // This may need to be handled differently per database adapter
} }
// Apply filters - validate and adjust for column types first // Client-controlled conditions (filters + x-custom-sql-w) are built by a closure so they can
// Group consecutive OR filters together to prevent OR logic from escaping // be wrapped in a single group together with x-custom-sql-or: the OR then only widens the
for i := 0; i < len(options.Filters); { // client's own conditions and can never escape the server-side filters ANDed around it.
filter := &options.Filters[i] applyUserConds := func(query common.SelectQuery) common.SelectQuery {
// Apply filters - validate and adjust for column types first
// Group consecutive OR filters together to prevent OR logic from escaping
for i := 0; i < len(options.Filters); {
filter := &options.Filters[i]
// Validate and adjust filter based on column type // Validate and adjust filter based on column type
castInfo := h.ValidateAndAdjustFilterForColumnType(filter, model) castInfo := h.ValidateAndAdjustFilterForColumnType(filter, model)
// Default to AND if LogicOperator is not set // Default to AND if LogicOperator is not set
logicOp := filter.LogicOperator logicOp := filter.LogicOperator
if logicOp == "" { if logicOp == "" {
logicOp = "AND" logicOp = "AND"
}
// Check if this is the start of an OR group
if logicOp == "OR" {
// Collect all consecutive OR filters
orFilters := []*common.FilterOption{filter}
orCastInfo := []ColumnCastInfo{castInfo}
j := i + 1
for j < len(options.Filters) {
nextFilter := &options.Filters[j]
nextLogicOp := nextFilter.LogicOperator
if nextLogicOp == "" {
nextLogicOp = "AND"
}
if nextLogicOp == "OR" {
nextCastInfo := h.ValidateAndAdjustFilterForColumnType(nextFilter, model)
orFilters = append(orFilters, nextFilter)
orCastInfo = append(orCastInfo, nextCastInfo)
j++
} else {
break
}
} }
// Apply the OR group as a single grouped condition // Check if this is the start of an OR group
logger.Debug("Applying OR filter group with %d conditions", len(orFilters)) if logicOp == "OR" {
query = h.applyOrFilterGroup(query, orFilters, orCastInfo, tableName, model) // Collect all consecutive OR filters
i = j orFilters := []*common.FilterOption{filter}
} else { orCastInfo := []ColumnCastInfo{castInfo}
// Single AND filter - apply normally
logger.Debug("Applying filter: %s %s %v (needsCast=%v, logic=%s)", filter.Column, filter.Operator, filter.Value, castInfo.NeedsCast, logicOp) j := i + 1
query = h.applyFilter(query, *filter, tableName, castInfo.NeedsCast, logicOp, model) for j < len(options.Filters) {
i++ nextFilter := &options.Filters[j]
nextLogicOp := nextFilter.LogicOperator
if nextLogicOp == "" {
nextLogicOp = "AND"
}
if nextLogicOp == "OR" {
nextCastInfo := h.ValidateAndAdjustFilterForColumnType(nextFilter, model)
orFilters = append(orFilters, nextFilter)
orCastInfo = append(orCastInfo, nextCastInfo)
j++
} else {
break
}
}
// Apply the OR group as a single grouped condition
logger.Debug("Applying OR filter group with %d conditions", len(orFilters))
query = h.applyOrFilterGroup(query, orFilters, orCastInfo, tableName, model)
i = j
} else {
// Single AND filter - apply normally
logger.Debug("Applying filter: %s %s %v (needsCast=%v, logic=%s)", filter.Column, filter.Operator, filter.Value, castInfo.NeedsCast, logicOp)
query = h.applyFilter(query, *filter, tableName, castInfo.NeedsCast, logicOp, model)
i++
}
} }
// Apply custom SQL WHERE clause (AND condition)
if options.CustomSQLWhere != "" {
logger.Debug("Applying custom SQL WHERE: %s", options.CustomSQLWhere)
// First add table prefixes to unqualified columns (but skip columns inside function calls)
prefixedWhere := common.AddTablePrefixToColumns(options.CustomSQLWhere, reflection.ExtractTableNameOnly(tableName))
// Then sanitize and allow preload table prefixes since custom SQL may reference multiple tables
sanitizedWhere := common.SanitizeWhereClause(prefixedWhere, reflection.ExtractTableNameOnly(tableName), &options.RequestOptions)
// Ensure outer parentheses to prevent OR logic from escaping
sanitizedWhere = common.EnsureOuterParentheses(sanitizedWhere)
if sanitizedWhere != "" {
query = query.Where(sanitizedWhere)
}
}
return query
} }
// Apply custom SQL WHERE clause (AND condition) sanitizedOr := ""
if options.CustomSQLWhere != "" {
logger.Debug("Applying custom SQL WHERE: %s", options.CustomSQLWhere)
// First add table prefixes to unqualified columns (but skip columns inside function calls)
prefixedWhere := common.AddTablePrefixToColumns(options.CustomSQLWhere, reflection.ExtractTableNameOnly(tableName))
// Then sanitize and allow preload table prefixes since custom SQL may reference multiple tables
sanitizedWhere := common.SanitizeWhereClause(prefixedWhere, reflection.ExtractTableNameOnly(tableName), &options.RequestOptions)
// Ensure outer parentheses to prevent OR logic from escaping
sanitizedWhere = common.EnsureOuterParentheses(sanitizedWhere)
if sanitizedWhere != "" {
query = query.Where(sanitizedWhere)
}
}
// Apply custom SQL WHERE clause (OR condition)
if options.CustomSQLOr != "" { if options.CustomSQLOr != "" {
logger.Debug("Applying custom SQL OR: %s", options.CustomSQLOr) logger.Debug("Applying custom SQL OR: %s", options.CustomSQLOr)
customOr := common.AddTablePrefixToColumns(options.CustomSQLOr, reflection.ExtractTableNameOnly(tableName)) customOr := common.AddTablePrefixToColumns(options.CustomSQLOr, reflection.ExtractTableNameOnly(tableName))
// Sanitize and allow preload table prefixes since custom SQL may reference multiple tables // Sanitize and allow preload table prefixes since custom SQL may reference multiple tables
sanitizedOr := common.SanitizeWhereClause(customOr, reflection.ExtractTableNameOnly(tableName), &options.RequestOptions) sanitizedOr = common.SanitizeWhereClause(customOr, reflection.ExtractTableNameOnly(tableName), &options.RequestOptions)
// Ensure outer parentheses to prevent OR logic from escaping // Ensure outer parentheses to prevent OR logic from escaping
sanitizedOr = common.EnsureOuterParentheses(sanitizedOr) sanitizedOr = common.EnsureOuterParentheses(sanitizedOr)
}
if grouper, ok := query.(common.WhereGrouper); ok && sanitizedOr != "" {
query = grouper.WhereGroup(func(q common.SelectQuery) common.SelectQuery {
return applyUserConds(q).WhereOr(sanitizedOr)
})
} else {
query = applyUserConds(query)
if sanitizedOr != "" { if sanitizedOr != "" {
query = query.WhereOr(sanitizedOr) query = query.WhereOr(sanitizedOr)
} }
@@ -1393,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...)
@@ -1539,6 +1567,10 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
// Variable to store the updated record // Variable to store the updated record
var updatedRecord interface{} var updatedRecord interface{}
// ID used to re-fetch the record after the update; differs from targetID
// when the request changes the primary key.
finalID := targetID
// Hook context used inside and outside transaction // Hook context used inside and outside transaction
hookCtx := &HookContext{ hookCtx := &HookContext{
Context: ctx, Context: ctx,
@@ -1604,11 +1636,25 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
nestedRelations = relations nestedRelations = relations
} }
// Capture a changed primary key from the request before merging. The row is
// located by the original targetID (WHERE), while the new value is written via SET.
// Only honoured when an ID was given in the URL (id != "").
var newPK interface{}
var pkChanged bool
if id != "" {
newPK, pkChanged = h.requestedPrimaryKey(model, pkName, dataMap, targetID)
}
// 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, dataMap, h.disallowNulls) common.MergeUpdateValues(existingMap, dataMap, h.disallowNulls)
// Ensure ID is in the data map for the update // Ensure ID is in the data map for the update (new value if the PK is being changed)
existingMap[pkName] = targetID if pkChanged {
existingMap[pkName] = newPK
finalID = newPK
} else {
existingMap[pkName] = targetID
}
dataMap = existingMap dataMap = existingMap
// Populate model instance from dataMap to preserve custom types (like SqlJSONB) // Populate model instance from dataMap to preserve custom types (like SqlJSONB)
@@ -1622,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
@@ -1651,6 +1700,20 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
} }
_ = result _ = result
// Primary key changes are not part of the struct SET, so apply them explicitly.
if pkChanged {
pkResult, err := tx.NewUpdate().Table(tableName).
Set(pkName, newPK).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID).
Exec(ctx)
if err != nil {
return fmt.Errorf("failed to update primary key: %w", err)
}
if pkResult.RowsAffected() == 0 {
return fmt.Errorf("primary key update affected no rows for ID: %v", targetID)
}
}
return nil return nil
}) })
@@ -1667,7 +1730,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
var errCode, errMsg string var errCode, errMsg string
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface() fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface()
selectQuery := tx.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID) selectQuery := tx.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), finalID)
// Execute BeforeScan hooks so row security is re-applied to the post-update // Execute BeforeScan hooks so row security is re-applied to the post-update
// re-fetch, same as it is for the initial read and the update query itself. // re-fetch, same as it is for the initial read and the update query itself.
@@ -1711,7 +1774,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
return return
} }
logger.Info("Successfully updated record with ID: %v", targetID) logger.Info("Successfully updated record with ID: %v", finalID)
// Invalidate cache for this table // Invalidate cache for this table
cacheTags := buildCacheTags(schema, tableName) cacheTags := buildCacheTags(schema, tableName)
if err := invalidateCacheForTags(ctx, cacheTags); err != nil { if err := invalidateCacheForTags(ctx, cacheTags); err != nil {
@@ -1720,6 +1783,28 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
h.sendResponseWithOptions(w, mergedData, nil, &options) h.sendResponseWithOptions(w, mergedData, nil, &options)
} }
// requestedPrimaryKey returns the primary key value carried in the request body
// (by column name or JSON key) when it differs from the current target ID.
func (h *Handler) requestedPrimaryKey(model interface{}, pkName string, dataMap map[string]interface{}, targetID interface{}) (interface{}, bool) {
val, exists := dataMap[pkName]
if !exists {
modelType := reflection.GetPointerElement(reflect.TypeOf(model))
for jsonKey, col := range reflection.BuildJSONToDBColumnMap(modelType) {
if col == pkName {
val, exists = dataMap[jsonKey]
break
}
}
}
if !exists || val == nil || reflection.IsEmptyValue(val) {
return nil, false
}
if fmt.Sprintf("%v", val) == fmt.Sprintf("%v", targetID) {
return nil, false
}
return val, true
}
func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id string, data interface{}) { func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id string, data interface{}) {
// Capture panics and return error response // Capture panics and return error response
defer func() { defer func() {
@@ -0,0 +1,153 @@
//go:build integration
package restheadspec
import (
"context"
"database/sql"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
_ "github.com/jackc/pgx/v5/stdlib"
"github.com/stretchr/testify/require"
"github.com/testcontainers/testcontainers-go"
"github.com/testcontainers/testcontainers-go/wait"
"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"
)
type pkAsset struct {
bun.BaseModel `bun:"table:public.t_pkasset,alias:t_pkasset"`
Category string `json:"category" bun:"category,type:citext"`
Description string `json:"description" bun:"description,type:citext,pk"`
}
func setupPKTestDB(t *testing.T) *sql.DB {
t.Helper()
ctx := context.Background()
pg, err := testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{
ContainerRequest: testcontainers.ContainerRequest{
Image: "postgres:15-alpine",
ExposedPorts: []string{"5432/tcp"},
Env: map[string]string{
"POSTGRES_USER": "testuser", "POSTGRES_PASSWORD": "testpass", "POSTGRES_DB": "testdb",
},
WaitingFor: wait.ForLog("database system is ready to accept connections").
WithOccurrence(2).WithStartupTimeout(60 * time.Second),
},
Started: true,
})
require.NoError(t, err)
t.Cleanup(func() { _ = pg.Terminate(ctx) })
host, err := pg.Host(ctx)
require.NoError(t, err)
port, err := pg.MappedPort(ctx, "5432")
require.NoError(t, err)
db, err := sql.Open("pgx", fmt.Sprintf("postgres://testuser:testpass@%s:%s/testdb?sslmode=disable", host, port.Port()))
require.NoError(t, err)
t.Cleanup(func() { _ = db.Close() })
require.NoError(t, db.Ping())
_, err = db.Exec(`
CREATE EXTENSION IF NOT EXISTS citext;
CREATE TABLE public.t_pkasset (
category citext,
description citext PRIMARY KEY
);
INSERT INTO public.t_pkasset VALUES ('cat', 'old-pk');`)
require.NoError(t, err)
return db
}
func pkUpdateCtx(base context.Context) context.Context {
ctx := WithSchema(base, "public")
ctx = WithEntity(ctx, "t_pkasset")
ctx = WithTableName(ctx, "t_pkasset")
return WithModel(ctx, pkAsset{})
}
func countPK(t *testing.T, db *sql.DB, pk string) int {
t.Helper()
var n int
require.NoError(t, db.QueryRow(`SELECT count(*) FROM public.t_pkasset WHERE description = $1`, pk).Scan(&n))
return n
}
func TestUpdateChangesPrimaryKeyWhenURLIDGiven(t *testing.T) {
db := setupPKTestDB(t)
h := NewHandler(database.NewBunAdapter(bun.NewDB(db, pgdialect.New())), modelregistry.NewModelRegistry())
update := func(id string, body map[string]interface{}) *httptest.ResponseRecorder {
rec := httptest.NewRecorder()
w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPut, "/", nil))
base, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
h.handleUpdate(pkUpdateCtx(base), w, id, nil, body, ExtendedRequestOptions{})
return rec
}
// URL id = old PK, body carries a different PK: the PK is changed.
rec := update("old-pk", map[string]interface{}{"description": "new-pk", "category": "cat2"})
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
require.Equal(t, 0, countPK(t, db, "old-pk"))
require.Equal(t, 1, countPK(t, db, "new-pk"))
var cat string
require.NoError(t, db.QueryRow(`SELECT category FROM public.t_pkasset WHERE description = 'new-pk'`).Scan(&cat))
require.Equal(t, "cat2", cat)
// Body PK equal to the URL id: ordinary update, PK untouched.
rec = update("new-pk", map[string]interface{}{"description": "new-pk", "category": "cat3"})
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
require.Equal(t, 1, countPK(t, db, "new-pk"))
// Body without a PK: ordinary update.
rec = update("new-pk", map[string]interface{}{"category": "cat4"})
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
// Null PK in the body is ignored, not applied.
rec = update("new-pk", map[string]interface{}{"description": nil, "category": "cat5"})
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
require.Equal(t, 1, countPK(t, db, "new-pk"))
var total int
require.NoError(t, db.QueryRow(`SELECT count(*) FROM public.t_pkasset`).Scan(&total))
require.Equal(t, 1, total)
require.NoError(t, db.QueryRow(`SELECT category FROM public.t_pkasset WHERE description = 'new-pk'`).Scan(&cat))
require.Equal(t, "cat5", cat)
}
// Exercises the real dispatch: a POST with a string id in the URL and a different
// PK in the body must update the row identified by the URL id.
func TestHandlePostWithStringURLIDChangesPrimaryKey(t *testing.T) {
db := setupPKTestDB(t)
reg := modelregistry.NewModelRegistry()
require.NoError(t, reg.RegisterModel("public.t_pkasset", pkAsset{}))
h := NewHandler(database.NewBunAdapter(bun.NewDB(db, pgdialect.New())), reg)
for _, method := range []string{http.MethodPost, http.MethodPut} {
t.Run(method, func(t *testing.T) {
from, to := "old-pk", "new-pk-"+method
if method == http.MethodPut {
from = "new-pk-" + http.MethodPost
}
req := httptest.NewRequest(method, "/public/t_pkasset/"+from,
strings.NewReader(`{"description":"`+to+`","category":"c"}`))
rec := httptest.NewRecorder()
w, r := common.WrapHTTPRequest(rec, req)
h.Handle(w, r, map[string]string{"schema": "public", "entity": "t_pkasset", "id": from})
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
require.Equal(t, 0, countPK(t, db, from))
require.Equal(t, 1, countPK(t, db, to))
})
}
}
+157
View File
@@ -0,0 +1,157 @@
package restheadspec
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"testing"
"github.com/DATA-DOG/go-sqlmock"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/pgdialect"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
)
// readCapturingSQL runs handleRead and returns every SELECT it issued.
func readCapturingSQL(t *testing.T, id string, options ExtendedRequestOptions) []string {
queries, _ := readCapturingSQLAndBody(t, id, options)
return queries
}
// readCapturingSQLAndBody is readCapturingSQL that also returns the response body.
// The mocked row carries the requested id so the body can be checked against it.
func readCapturingSQLAndBody(t *testing.T, id string, options ExtendedRequestOptions) ([]string, string) {
t.Helper()
resetTotalCache(t)
var queries []string
matcher := sqlmock.QueryMatcherFunc(func(_, actual string) error {
queries = append(queries, actual)
return nil
})
sqlDB, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(matcher))
if err != nil {
t.Fatal(err)
}
sqlDB.SetMaxOpenConns(1)
t.Cleanup(func() { _ = sqlDB.Close() })
h := NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry())
rowID, err := strconv.Atoi(id)
if err != nil {
rowID = 7
}
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(rowID, "a"))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectCommit()
rec := httptest.NewRecorder()
w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil))
h.handleRead(itemCtx(t), w, id, options)
if rec.Code != http.StatusOK {
t.Fatalf("status %d body %s", rec.Code, rec.Body)
}
return queries, rec.Body.String()
}
func TestReadByIDIgnoresLimitOffsetAndCursor(t *testing.T) {
limit, offset := 50, 10
queries := readCapturingSQL(t, "7", ExtendedRequestOptions{
RequestOptions: common.RequestOptions{
Limit: &limit,
Offset: &offset,
},
})
last := queries[len(queries)-1]
if !strings.Contains(last, "LIMIT 1") || strings.Contains(last, "OFFSET") {
t.Fatalf("read by id must be LIMIT 1 with no OFFSET: %s", last)
}
}
func TestReadWithoutIDKeepsRequestedLimit(t *testing.T) {
limit := 50
queries := readCapturingSQL(t, "", ExtendedRequestOptions{
RequestOptions: common.RequestOptions{Limit: &limit},
})
if last := queries[len(queries)-1]; !strings.Contains(last, "LIMIT 50") {
t.Fatalf("list read must keep its limit: %s", last)
}
}
// topLevelOr reports whether the WHERE clause has an OR outside any parentheses,
// i.e. one that would let rows bypass the AND-ed primary key condition.
func topLevelOr(sql string) bool {
where := sql[strings.Index(sql, "WHERE")+len("WHERE"):]
depth, inStr := 0, false
for i := 0; i < len(where); i++ {
switch c := where[i]; {
case c == '\'':
inStr = !inStr
case inStr:
case c == '(':
depth++
case c == ')':
depth--
case depth == 0 && strings.HasPrefix(where[i:], " OR "):
return true
}
}
return false
}
func TestReadByIDCustomSQLOrCannotEscapePrimaryKey(t *testing.T) {
queries := readCapturingSQL(t, "7", ExtendedRequestOptions{
RequestOptions: common.RequestOptions{
Filters: []common.FilterOption{{Column: "name", Operator: "eq", Value: "a"}},
},
CustomSQLOr: "name = 'x'",
})
last := queries[len(queries)-1]
if !strings.Contains(last, `"id" = '7'`) && !strings.Contains(last, `"id" = 7`) {
t.Fatalf("primary key filter missing: %s", last)
}
if topLevelOr(last) {
t.Fatalf("OR escapes the primary key filter: %s", last)
}
}
func TestReadByIDFiltersAndReturnsRequestedRecord(t *testing.T) {
queries, body := readCapturingSQLAndBody(t, "42", ExtendedRequestOptions{})
last := queries[len(queries)-1]
if !strings.Contains(last, `"items"."id" = '42'`) && !strings.Contains(last, `"items"."id" = 42`) {
t.Fatalf("query must filter the primary key to 42: %s", last)
}
if strings.Contains(last, "= 7") || strings.Contains(last, "= '7'") {
t.Fatalf("query filters a different id: %s", last)
}
// every query that touches rows (count and select) must carry the id filter
for _, q := range queries {
if strings.Contains(q, "FROM") && !strings.Contains(q, "42") {
t.Fatalf("query without the id filter: %s", q)
}
}
var rows []struct {
ID int `json:"id"`
}
data := body
if i := strings.Index(body, `"data"`); i >= 0 {
data = body[i+len(`"data"`):]
}
if i := strings.Index(data, "["); i >= 0 {
data = data[i:]
}
dec := json.NewDecoder(strings.NewReader(data))
if err := dec.Decode(&rows); err != nil {
t.Fatalf("decode %q: %v", body, err)
}
if len(rows) != 1 || rows[0].ID != 42 {
t.Fatalf("response must contain exactly the record with id 42: %s", body)
}
}
+61
View File
@@ -0,0 +1,61 @@
package restheadspec
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/uptrace/bunrouter"
)
type wrapCtxKey struct{}
// The auth wrapper must hand the handler the middleware-enriched request
// without dropping the bunrouter route params.
func TestWrapBunRouterHandler_PreservesRouteParams(t *testing.T) {
var gotSchema, gotEntity, gotID string
var gotCtxVal any
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotSchema = req.Param("schema")
gotEntity = req.Param("entity")
gotID = req.Param("id")
gotCtxVal = req.Context().Value(wrapCtxKey{})
return nil
}
auth := func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), wrapCtxKey{}, "enriched")))
})
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, auth))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/public/users/42", nil))
if gotSchema != "public" || gotEntity != "users" || gotID != "42" {
t.Errorf("route params lost: schema=%q entity=%q id=%q", gotSchema, gotEntity, gotID)
}
if gotCtxVal != "enriched" {
t.Errorf("handler did not see middleware-enriched context, got %v", gotCtxVal)
}
}
func TestWrapBunRouterHandler_NilAuthPassesThrough(t *testing.T) {
var gotID string
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotID = req.Param("id")
return nil
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, nil))
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/public/users/7", nil))
if gotID != "7" {
t.Errorf("id = %q, want 7", gotID)
}
}
+18 -23
View File
@@ -19,7 +19,7 @@ In-memory store seeded from a static list. Suitable for a small, fixed set of se
```go ```go
// Pre-load keys from config (KeyHash = SHA-256 hex of the raw key) // Pre-load keys from config (KeyHash = SHA-256 hex of the raw key)
store := security.NewConfigKeyStore([]security.UserKey{ store := providers.NewConfigKeyStore([]security.UserKey{
{ {
UserID: 1, UserID: 1,
KeyType: security.KeyTypeGenericAPI, KeyType: security.KeyTypeGenericAPI,
@@ -33,7 +33,7 @@ store := security.NewConfigKeyStore([]security.UserKey{
### DatabaseKeyStore ### DatabaseKeyStore
Backed by PostgreSQL stored procedures. Supports optional caching (default 2-minute TTL). Apply `keystore_schema.sql` before use. Backed by PostgreSQL stored procedures by default, or by the `user_keys` table in direct mode (any supported dialect). Supports optional caching (default 2-minute TTL). Apply `lookup/keystore_schema.sql` before use.
```go ```go
db, _ := sql.Open("postgres", dsn) db, _ := sql.Open("postgres", dsn)
@@ -43,8 +43,8 @@ store := security.NewDatabaseKeyStore(db)
// With options // With options
store = security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{ store = security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{
CacheTTL: 5 * time.Minute, CacheTTL: 5 * time.Minute,
SQLNames: &security.KeyStoreSQLNames{ Lookup: lookup.Config{
ValidateKey: "myapp_keystore_validate", // override one procedure name Procs: lookup.ProcNames{KeystoreValidateKey: "myapp_keystore_validate"}, // override one procedure name
}, },
}) })
``` ```
@@ -85,9 +85,9 @@ Keys are extracted from the request in this order:
3. `X-API-Key: <key>` 3. `X-API-Key: <key>`
```go ```go
auth := security.NewKeyStoreAuthenticator(store, "") // "" = accept any key type auth := providers.NewKeyStoreAuthenticator(store, "") // "" = accept any key type
// Restrict to a specific type: // Restrict to a specific type:
auth = security.NewKeyStoreAuthenticator(store, security.KeyTypeGenericAPI) auth = providers.NewKeyStoreAuthenticator(store, security.KeyTypeGenericAPI)
``` ```
Plug it into a handler: Plug it into a handler:
@@ -109,10 +109,10 @@ On successful validation the request context receives a `UserContext` where:
## Database setup ## Database setup
Apply `keystore_schema.sql` to your PostgreSQL database. It requires the `users` table from the main `database_schema.sql`. Apply `lookup/keystore_schema.sql` to your PostgreSQL database. It requires the `users` table from the main `lookup/database_schema.sql`.
```sql ```sql
\i pkg/security/keystore_schema.sql \i pkg/security/lookup/keystore_schema.sql
``` ```
This creates: This creates:
@@ -123,28 +123,23 @@ This creates:
- `resolvespec_keystore_delete_key(p_user_id, p_key_id)` - `resolvespec_keystore_delete_key(p_user_id, p_key_id)`
- `resolvespec_keystore_validate_key(p_key_hash, p_key_type)` - `resolvespec_keystore_validate_key(p_key_hash, p_key_type)`
### Custom procedure names ### Custom names and modes
```go ```go
store := security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{ store := security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{
SQLNames: &security.KeyStoreSQLNames{ Lookup: lookup.Config{
GetUserKeys: "myschema_get_keys", Procs: lookup.ProcNames{
CreateKey: "myschema_create_key", KeystoreGetUserKeys: "myschema_get_keys",
DeleteKey: "myschema_delete_key", KeystoreCreateKey: "myschema_create_key",
ValidateKey: "myschema_validate_key", KeystoreDeleteKey: "myschema_delete_key",
KeystoreValidateKey: "myschema_validate_key",
},
}, },
}) })
// Validate names at startup
names := &security.KeyStoreSQLNames{
GetUserKeys: "myschema_get_keys",
// ...
}
if err := security.ValidateKeyStoreSQLNames(names); err != nil {
log.Fatal(err)
}
``` ```
Names are validated when the store is first used. On Postgres the key store calls the procedures by default; on SQLite, MySQL and SQL Server (or with `Mode: lookup.ModeDirect`) it reads and writes the `user_keys` table directly (see `lookup/ddl`). Table and column names are configurable through `lookup.Config.Schema`.
## Security notes ## Security notes
- Raw keys are never stored. Only the SHA-256 hex digest is persisted. - Raw keys are never stored. Only the SHA-256 hex digest is persisted.
+6 -3
View File
@@ -4,6 +4,8 @@
The security package provides OAuth2 authentication support for any OAuth2-compliant provider including Google, GitHub, Microsoft, Facebook, and custom providers. The security package provides OAuth2 authentication support for any OAuth2-compliant provider including Google, GitHub, Microsoft, Facebook, and custom providers.
> **Full OAuth 2.1 / OpenID Connect**: this guide covers the plain OAuth2 client login. For the OIDC relying party (discovery, PKCE, nonce, id_token validation, logout) and the complete authorization server (consent, refresh rotation, DPoP, PAR, device grant, token exchange) see [OAUTH2_SERVER.md](OAUTH2_SERVER.md).
## Features ## Features
- **Universal OAuth2 Support**: Works with any OAuth2 provider - **Universal OAuth2 Support**: Works with any OAuth2 provider
@@ -14,6 +16,7 @@ The security package provides OAuth2 authentication support for any OAuth2-compl
- **Token Refresh**: Automatic token refresh support - **Token Refresh**: Automatic token refresh support
- **State Validation**: Built-in CSRF protection - **State Validation**: Built-in CSRF protection
- **User Auto-Creation**: Automatically creates users on first login - **User Auto-Creation**: Automatically creates users on first login
- **OpenID Connect** (opt-in): `WithOIDC` discovery, PKCE, nonce and id_token validation, RP-initiated logout
- **Unified Authentication**: OAuth2 and traditional auth share same session storage - **Unified Authentication**: OAuth2 and traditional auth share same session storage
## Quick Start ## Quick Start
@@ -21,7 +24,7 @@ The security package provides OAuth2 authentication support for any OAuth2-compl
### 1. Database Setup ### 1. Database Setup
```sql ```sql
-- Run the schema from database_schema.sql -- Run the schema from lookup/database_schema.sql
CREATE TABLE IF NOT EXISTS users ( CREATE TABLE IF NOT EXISTS users (
id SERIAL PRIMARY KEY, id SERIAL PRIMARY KEY,
username VARCHAR(255) NOT NULL UNIQUE, username VARCHAR(255) NOT NULL UNIQUE,
@@ -53,7 +56,7 @@ CREATE TABLE IF NOT EXISTS user_sessions (
); );
-- OAuth2 stored procedures (7 functions) -- OAuth2 stored procedures (7 functions)
-- See database_schema.sql for full implementation -- See lookup/database_schema.sql for full implementation
``` ```
### 2. Google OAuth2 ### 2. Google OAuth2
@@ -397,7 +400,7 @@ UserInfoParser: func(userInfo map[string]any) (*security.UserContext, error) {
## Implementation Details ## Implementation Details
All database operations use stored procedures for consistency and security: On PostgreSQL, database operations use stored procedures by default (other dialects use direct SQL through `pkg/security/lookup`):
- `resolvespec_oauth_getorcreateuser` - Find or create OAuth2 user - `resolvespec_oauth_getorcreateuser` - Find or create OAuth2 user
- `resolvespec_oauth_createsession` - Create OAuth2 session - `resolvespec_oauth_createsession` - Create OAuth2 session
- `resolvespec_oauth_getsession` - Validate and retrieve session - `resolvespec_oauth_getsession` - Validate and retrieve session
@@ -1,5 +1,7 @@
# OAuth2 Refresh Token - Quick Reference # OAuth2 Refresh Token - Quick Reference
> This covers refreshing tokens of an upstream provider with `OAuth2RefreshToken`. For refresh tokens issued by `OAuthServer` (rotation, reuse detection, downscoping) see [OAUTH2_SERVER.md](OAUTH2_SERVER.md#refresh-token-rotation).
## Quick Setup (3 Steps) ## Quick Setup (3 Steps)
### 1. Initialize Authenticator ### 1. Initialize Authenticator
@@ -276,6 +278,6 @@ authURL += "&access_type=offline&prompt=consent"
## Complete Example ## Complete Example
See `/pkg/security/oauth2_examples.go` line 250 for full working example. See `/pkg/security/oauth2_examples.go` (`ExampleOAuth2TokenRefresh`) for full working example.
For detailed documentation see `/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md`. For detailed documentation see `/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md`.
@@ -43,16 +43,16 @@ CREATE TABLE IF NOT EXISTS user_sessions (
**`resolvespec_oauth_getrefreshtoken(p_refresh_token)`** **`resolvespec_oauth_getrefreshtoken(p_refresh_token)`**
- Gets OAuth2 session data by refresh token - Gets OAuth2 session data by refresh token
- Returns: `{user_id, access_token, token_type, expiry}` - Returns: `{user_id, access_token, token_type, expiry}`
- Location: `database_schema.sql:714` - Location: `lookup/database_schema.sql` (section 15); direct mode: `lookup/direct` `OAuthUserStore.GetByRefreshToken`
**`resolvespec_oauth_updaterefreshtoken(p_update_data)`** **`resolvespec_oauth_updaterefreshtoken(p_update_data)`**
- Updates session with new tokens after refresh - Updates session with new tokens after refresh
- Input: `{user_id, old_refresh_token, new_session_token, new_access_token, new_refresh_token, expires_at}` - Input: `{user_id, old_refresh_token, new_session_token, new_access_token, new_refresh_token, expires_at}`
- Location: `database_schema.sql:752` - Location: `lookup/database_schema.sql` (section 16); direct mode: `lookup/direct` `OAuthUserStore.UpdateRefreshToken`
**`resolvespec_oauth_getuser(p_user_id)`** **`resolvespec_oauth_getuser(p_user_id)`**
- Gets user data by ID for building UserContext - Gets user data by ID for building UserContext
- Location: `database_schema.sql:791` - Location: `lookup/database_schema.sql` (section 17); direct mode: `lookup/direct` `OAuthUserStore.GetUser`
--- ---
@@ -68,7 +68,7 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(
) (*LoginResponse, error) ) (*LoginResponse, error)
``` ```
**Location:** `pkg/security/oauth2_methods.go:375` **Location:** `pkg/security/oauth2_methods.go` (`OAuth2RefreshToken`)
### Implementation Flow ### Implementation Flow
@@ -476,7 +476,7 @@ auth.OAuth2RefreshToken(ctx, token, "google") // Must match ProviderName
## 8. Complete Working Example ## 8. Complete Working Example
See `pkg/security/oauth2_examples.go:250` for full working example with token refresh. See `pkg/security/oauth2_examples.go` (`ExampleOAuth2TokenRefresh`) for full working example with token refresh.
--- ---
+248
View File
@@ -0,0 +1,248 @@
# OAuth 2.1 / OpenID Connect Authorization Server
`OAuthServer` turns a `DatabaseAuthenticator` into a standards-based identity provider, and `OIDCConfig` / `WithOIDC` make the same package a relying party for any OpenID Connect provider. This guide covers both. For the older "log in with Google/GitHub" client flow see [OAUTH2.md](OAUTH2.md); a runnable end-to-end wiring is in [`oauth2_full_example.go`](oauth2_full_example.go) (`ExampleOAuth2FullServer`, `ExampleOAuth2FullClient`).
Every feature beyond the original authorization-code flow is **opt-in**: the zero value of each `OAuthServerConfig` option keeps the previous behaviour.
## Contents
1. [Architecture](#architecture)
2. [Quick start](#quick-start)
3. [Endpoints](#endpoints)
4. [Configuration reference](#configuration-reference)
5. [Flows](#flows)
6. [Resource servers](#resource-servers)
7. [Client side: logging in with an OpenID Connect provider](#client-side-relying-party)
8. [Database setup](#database-setup)
9. [Security checklist](#security-checklist)
10. [Not supported](#not-supported)
## Architecture
```
browser / app ──► OAuthServer.HTTPHandler() ──► DatabaseAuthenticator (users, sessions, login)
│
└────────────────────► lookup.Provider ──► procedure | direct backend
oauth_clients, oauth_codes, oauth_consents,
oauth_refresh_tokens, oauth_device_codes,
oauth_par_requests, oauth_jti
```
- **State is in the database** (clients, codes, consents, rotating refresh tokens, device codes, pushed requests, the replay cache), so any number of instances can serve one issuer.
- **Stateless pieces** (the SSO cookie, the login/consent form state, the device and logout state) are HMAC-sealed with `CookieSecret`. Instances must share it, or share the first signing key it is derived from.
- **Signing keys**: RSA or ECDSA (P-256/P-384); `kid` is the RFC 7638 thumbprint. All keys are published in the JWKS.
- Access tokens are either the opaque session token (default) or RFC 9068 JWTs. Either way an *access-grant record* stores scope, client, audience and DPoP binding, so introspection, revocation and logout work for both.
## Quick start
```go
auth := security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{})
srv := security.NewOAuthServer(security.OAuthServerConfig{
Issuer: "https://auth.example.com",
PersistClients: true,
PersistCodes: true,
RequireConsent: true,
ManagedRefreshTokens: true,
}, auth)
defer srv.Close()
mux := http.NewServeMux()
mux.Handle("/", srv.HTTPHandler())
```
Register a first-party client from code (no consent screen), or let clients register themselves with `POST /oauth/register`:
```go
client, secret, err := srv.RegisterTrustedClient(ctx, security.OAuthServerClient{
ClientName: "Admin console",
RedirectURIs: []string{"https://console.example.com/callback"},
})
```
## Endpoints
| Method | Path | Spec | Notes |
|---|---|---|---|
| GET | `/.well-known/oauth-authorization-server` | RFC 8414 | Also `/{path}` variants for issuers with a path |
| GET | `/.well-known/openid-configuration` | OIDC Discovery | Same document, plus OIDC fields |
| GET | `/.well-known/oauth-protected-resource` | RFC 9728 | |
| POST | `/oauth/register` | RFC 7591 | Needs `InitialAccessToken` when configured |
| GET PUT DELETE | `/oauth/register/{client_id}` | RFC 7592 | `registration_access_token` as Bearer |
| POST | `/oauth/register/{client_id}/rotate-secret` | RFC 7592 | |
| GET POST | `/oauth/authorize` | RFC 6749, OIDC Core | PKCE S256 required; `response_mode` `query` or `form_post` |
| POST | `/oauth/token` | RFC 6749 | `authorization_code`, `refresh_token`, `client_credentials`, device code, token exchange |
| POST | `/oauth/par` | RFC 9126 | `EnablePAR` / `RequirePAR` |
| POST | `/oauth/device_authorization` | RFC 8628 | `EnableDeviceFlow` |
| GET POST | `/oauth/device` | RFC 8628 | Verification page (login, consent) |
| POST | `/oauth/revoke` | RFC 7009 | Client authentication required |
| POST | `/oauth/introspect` | RFC 7662 | Client authentication required |
| GET POST | `/oauth/userinfo` | OIDC Core | Scope-filtered; signed response if the client asks |
| GET | `/oauth/jwks.json` | RFC 7517 | |
| GET POST | `/oauth/logout` | OIDC RP-Initiated + Back-Channel Logout | |
| GET | `ProviderCallbackPath` | | Callback of a federated upstream provider |
Client authentication methods: `client_secret_basic`, `client_secret_post`, `private_key_jwt` (`jwks` or `jwks_uri`), `none` (public clients).
## Configuration reference
| Option | Default | Purpose |
|---|---|---|
| `Issuer` | required | Public base URL; `iss` of every token. May contain a path |
| `SigningKeys` / `SigningKey` | generated RSA-2048 | Persistent keys for multi-instance and restarts; first signs |
| `CookieSecret` | derived from first key | HMAC key of cookie and form state |
| `SSOCookie` | `resolvespec_sso`, 8h, Lax, Secure when issuer is https | Enables `prompt=none`, `max_age`, single sign-on, logout |
| `PersistClients`, `PersistCodes` | false | Store clients/codes in the DB (needed for several instances) |
| `RequireConsent`, `ConsentTTL` | false, 90 days | Consent screen for non-first-party clients; remembered per user and client |
| `ScopeDescriptions` | built-in for the OIDC scopes | Text on the consent screen |
| `ManagedRefreshTokens`, `RefreshTokenTTL` | false, 30 days | Server-issued rotating refresh tokens with reuse detection |
| `JWTAccessTokens`, `AccessTokenAudience` | false, `ResourceIdentifier` | RFC 9068 access tokens |
| `AccessTokenTTL`, `AuthCodeTTL` | 24h, 2 min | |
| `EnableDPoP` | false | RFC 9449 sender-constrained tokens |
| `EnablePAR`, `RequirePAR`, `PARTTL` | false, false, 90s | |
| `EnableDeviceFlow`, `DeviceCodeTTL`, `DevicePollSeconds` | false, 10 min, 5 | |
| `EnableTokenExchange` | false | RFC 8693 |
| `ClaimsProvider` | `sub`, `preferred_username`, `email` | Source of profile/email/address/phone/custom claims |
| `SupportedACR` | none | Advertised and accepted `acr_values` |
| `DisableLogout` | false | Do not serve `/oauth/logout` |
| `InitialAccessToken` | none | Bearer secret required to register clients |
| `AllowAnonymousIntrospection` | false | Skip client authentication at introspect/revoke |
| `RateLimiter` | none | `func(r, endpoint) bool`; false answers 429 |
| `AllowPrivateNetworkFetch` | false | Allow `jwks_uri` fetches to private addresses (SSRF guard) |
| `LoginTemplate`, `ConsentTemplate` | built-in | `html/template` overrides (`OAuthLoginPage`, `OAuthConsentPage`) |
## Flows
### Authorization code with PKCE
```
GET /oauth/authorize?response_type=code&client_id=ID&redirect_uri=https://app/cb
&scope=openid%20profile&state=S&nonce=N
&code_challenge=BASE64URL(SHA256(V))&code_challenge_method=S256
```
The user signs in (a cookie keeps the session), approves the consent screen if required, and is redirected to `redirect_uri?code=…&state=S&iss=<issuer>` (RFC 9207; compare `iss`). Exchange it:
```
curl -X POST https://auth.example.com/oauth/token \
-d grant_type=authorization_code -d code=CODE -d redirect_uri=https://app/cb \
-d client_id=ID -d code_verifier=V
```
Request parameters: `prompt` (`none`, `login`, `consent`, `select_account`), `max_age`, `id_token_hint`, `login_hint`, `acr_values`, `claims`, `resource` (RFC 8707), `response_mode=form_post`. Once the `redirect_uri` is validated, errors are returned to the client as redirects (`error`, `state`, `iss`); before that they are shown to the user. `prompt=none` without a session answers `login_required`.
The id_token carries `iss`, `sub`, `aud`, `exp`, `iat`, `nonce`, `auth_time`, `acr`, `amr`, `sid`, `at_hash`, `azp` and the claims the granted scopes entitle the client to.
### Consent and scopes
Requested scopes are intersected with the client's `AllowedScopes`; an empty result is `invalid_scope`. With `RequireConsent` (or per client `require_consent`), a third-party client sees a consent screen unless a stored consent already covers the scopes. A client marked `first_party` (`RegisterTrustedClient`) never does. Approval can be remembered for `ConsentTTL`; a denial redirects with `access_denied`.
### Refresh token rotation
With `ManagedRefreshTokens`, every refresh returns a **new** refresh token and invalidates the old one:
```
curl -X POST …/oauth/token -d grant_type=refresh_token -d refresh_token=R1 -d client_id=ID
```
Presenting a rotated token again (a stolen copy) answers `invalid_grant` and **revokes the whole family**, so the legitimate holder has to sign in again. A refresh may downscope (`scope=`) but never widen. Confidential clients must authenticate. Request `offline_access` or allow the `refresh_token` grant to receive one. Without `ManagedRefreshTokens` the refresh token of the underlying authenticator is passed through as in earlier versions.
### JWT access tokens (RFC 9068)
`JWTAccessTokens` issues `at+jwt` tokens with `iss sub aud exp iat jti client_id scope` (and `cnf` for DPoP). See [Resource servers](#resource-servers).
### DPoP (RFC 9449)
With `EnableDPoP` a client sends a `DPoP` proof header (typ `dpop+jwt`, `htm`, `htu`, `iat`, `jti`, public `jwk`) to the token endpoint. The access and refresh tokens are then bound to the proof key; the response has `token_type: DPoP`. Use them as `Authorization: DPoP <token>` plus a proof carrying `ath = base64url(SHA256(token))`. Proof `jti`s are single use (replay cache in `oauth_jti`). A DPoP-bound token is refused as a Bearer token. Set the client's `dpop_bound_access_tokens` to require proofs.
### Pushed authorization requests (RFC 9126)
```
curl -X POST …/oauth/par -d client_id=ID -d response_type=code … -d code_challenge=… # → {"request_uri": "urn:ietf:params:oauth:request_uri:…", "expires_in": 90}
GET /oauth/authorize?client_id=ID&request_uri=urn:ietf:params:oauth:request_uri:…
```
Confidential clients authenticate at the PAR endpoint. A `request_uri` is single use. `RequirePAR` rejects plain authorization requests.
### Device grant (RFC 8628)
```
curl -X POST …/oauth/device_authorization -d client_id=ID -d scope=openid
# → device_code, user_code, verification_uri, verification_uri_complete, interval
curl -X POST …/oauth/token -d grant_type=urn:ietf:params:oauth:grant-type:device_code -d device_code=… -d client_id=ID
```
The user opens `verification_uri` on another device, signs in, enters the code and approves. The device polls; answers are `authorization_pending`, `slow_down` (polled faster than `interval`), `access_denied`, `expired_token`. A code is consumed by the first successful poll.
### Token exchange (RFC 8693)
A confidential client holding the `urn:ietf:params:oauth:grant-type:token-exchange` grant swaps a user's access token for a narrower one (`scope`, `audience`/`resource`). The scope can only shrink; DPoP-bound subject tokens and `actor_token` are not accepted.
### Client credentials
Unchanged: `grant_type=client_credentials` with client authentication returns a token for the client itself.
### Dynamic registration (RFC 7591/7592)
`POST /oauth/register` accepts `redirect_uris` (https, loopback http, or a custom scheme; no fragments), `grant_types`, `response_types`, `token_endpoint_auth_method`, `scope`, `jwks` / `jwks_uri`, `client_name`, `client_uri`, `logo_uri`, `contacts`, `post_logout_redirect_uris`, `backchannel_logout_uri`, `id_token_signed_response_alg`, `userinfo_signed_response_alg`, `dpop_bound_access_tokens`. The response includes `registration_access_token` and `registration_client_uri`; use them to read, update or delete the client and to rotate its secret. Secrets are stored hashed and shown once. Loopback redirect URIs ignore the port (RFC 8252).
### Logout
`/oauth/logout?id_token_hint=…&post_logout_redirect_uri=…&state=…` ends the SSO session, revokes the session's tokens and refresh families, and redirects to a `post_logout_redirect_uri` the client registered. Without a hint the user is asked to confirm. Clients with `backchannel_logout_uri` receive a signed `logout_token` (best effort, in the background).
### Federation
`srv.RegisterExternalProvider(auth, "google")` lets users sign in through an upstream provider; the server remains the issuer for your clients and your users get the same tokens, consent and logout behaviour. `login_hint` pre-fills the built-in login form.
## Resource servers
```go
claims, err := srv.VerifyAccessToken(ctx, token, security.VerifyAccessTokenOptions{
Audience: "https://api.example.com",
Scopes: []string{"orders:read"},
})
```
JWT tokens are verified locally against the key set; opaque tokens are looked up. `claims` holds `Subject`, `UserID`, `ClientID`, `Scopes`, `Audience`, `JTI` and the DPoP key thumbprint. For DPoP-bound tokens also check the proof (`verifyDPoP` runs on the server's own endpoints; an external API validates the `DPoP` header and compares `DPoPKey`). A ready-made middleware is in `ExampleOAuth2FullServer`. Services in another process can call `/oauth/introspect` with their client credentials, or verify the JWT with `GET /oauth/jwks.json`.
## Client side (relying party)
`WithOIDC` discovers the endpoints from `{issuer}/.well-known/openid-configuration` and registers a provider with these protections switched on:
- PKCE (S256) and a `nonce`, both kept with the `state` and used once.
- The `id_token` is verified: signature against the provider's JWKS (refetched once when a `kid` is unknown), algorithm allow-list (`RS256 PS256 ES256 ES384` by default), `iss`, `aud`/`azp`, `exp` (1 minute skew), `nonce`, `at_hash`.
- The userinfo `sub` must equal the id_token `sub`; the RFC 9207 `iss` parameter must match.
```go
auth.WithOIDC(ctx, security.OIDCConfig{Issuer: "https://auth.example.com", ClientID: id, ClientSecret: secret,
RedirectURL: "https://app/auth/callback", ProviderName: "company"})
url, _ := auth.OAuth2GetAuthURLWithOptions("company", state, security.OAuth2AuthOptions{Prompt: "login"})
login, err := auth.OAuth2HandleCallbackRequest(ctx, "company", r) // r is the callback request
logout, _ := auth.OAuth2LogoutURL(ctx, "company", login.Meta["id_token"].(string), "https://app/", "")
```
`OAuth2HandleCallback(ctx, provider, code, state)` still works. The raw `id_token` is returned in `LoginResponse.Meta["id_token"]` (keep it for the logout hint); `OAuth2RefreshToken` re-validates a new id_token when the provider returns one. The Google preset validates id_tokens as well; for other providers set `Issuer` and `JWKSURL` on `OAuth2Config`, or use `WithOIDC`.
## Database setup
Apply the schema for your backend (see the lookup section of the [README](README.md)):
- Postgres procedures: `lookup/database_schema.sql`
- Table-only (any dialect, direct mode): `lookup/ddl/{postgres,sqlite,mysql,mssql}.sql`
The OAuth additions are the `oauth_clients.metadata` and `oauth_codes.extra` JSON columns plus the tables `oauth_consents`, `oauth_refresh_tokens`, `oauth_device_codes`, `oauth_par_requests` and `oauth_jti`. Existing installs: see [breaking_changes.md](breaking_changes.md#step-8-full-oauth2--openid-connect) for the ALTER statements. Purge expired rows of the new tables periodically (`expires_at < now`).
## Security checklist
- Serve the issuer over HTTPS only; the SSO cookie is `Secure` automatically then.
- Persist `SigningKeys` and share `CookieSecret` across instances.
- Set `InitialAccessToken` unless open dynamic registration is intended.
- Use `ManagedRefreshTokens` for public clients; keep `RequireConsent` on for third-party clients.
- Put a `RateLimiter` in front of `token`, `authorize`, `device` and `register`.
- Only PKCE S256 is accepted; redirect URIs match exactly (except the loopback port).
- Keep `AllowPrivateNetworkFetch` off; `jwks_uri` fetches are SSRF-guarded.
- Authenticate callers of `introspect` and `revoke` (the default).
## Not supported
`client_secret_jwt`, signed request objects (`request` / `request_uri` to a remote JWT), the DPoP server nonce, `c_hash`, encrypted id_tokens, and actor tokens in token exchange. Discovery does not advertise them.
+2 -2
View File
@@ -6,9 +6,9 @@ Passkey authentication (WebAuthn/FIDO2) is now integrated into the DatabaseAuthe
## Setup ## Setup
### Database Schema ### Database Schema
Run the passkey SQL schema (in database_schema.sql): Run the passkey SQL schema (in lookup/database_schema.sql):
- Creates `user_passkey_credentials` table - Creates `user_passkey_credentials` table
- Adds stored procedures for passkey operations - Adds stored procedures for passkey operations (Postgres procedure backend; other dialects use `lookup/ddl` tables with direct SQL)
### Go Code ### Go Code
```go ```go
+20 -18
View File
@@ -6,7 +6,7 @@
// Step 1: Create security providers // Step 1: Create security providers
auth := security.NewDatabaseAuthenticator(db) // Session-based (recommended) auth := security.NewDatabaseAuthenticator(db) // Session-based (recommended)
// OR: auth := security.NewJWTAuthenticator("secret-key", db) // OR: auth := security.NewJWTAuthenticator("secret-key", db)
// OR: auth := security.NewHeaderAuthenticator() // OR: auth := providers.NewHeaderAuthenticator()
// OR: auth := security.NewGoogleAuthenticator(clientID, secret, redirectURL, db) // OAuth2 // OR: auth := security.NewGoogleAuthenticator(clientID, secret, redirectURL, db) // OAuth2
colSec := security.NewDatabaseColumnSecurityProvider(db) colSec := security.NewDatabaseColumnSecurityProvider(db)
@@ -16,7 +16,8 @@ rowSec := security.NewDatabaseRowSecurityProvider(db)
provider, _ := security.NewCompositeSecurityProvider(auth, colSec, rowSec) provider, _ := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
// Step 3: Setup and apply middleware // Step 3: Setup and apply middleware
securityList, _ := security.SetupSecurityProvider(handler, provider) securityList, _ := security.NewSecurityList(provider)
restheadspec.RegisterSecurityHooks(handler, securityList)
router.Use(security.NewAuthMiddleware(securityList)) router.Use(security.NewAuthMiddleware(securityList))
router.Use(security.SetSecurityMiddleware(securityList)) router.Use(security.SetSecurityMiddleware(securityList))
``` ```
@@ -25,7 +26,7 @@ router.Use(security.SetSecurityMiddleware(securityList))
## Stored Procedures ## Stored Procedures
**All database operations use PostgreSQL stored procedures** with `resolvespec_*` naming: **On PostgreSQL, database operations use stored procedures by default** with `resolvespec_*` naming (other dialects use direct SQL; see `lookup.Config` in README.md):
### Database Authenticators ### Database Authenticators
```go ```go
@@ -55,7 +56,7 @@ All stored procedures return structured results:
- Session/Login: `(p_success bool, p_error text, p_data jsonb)` - Session/Login: `(p_success bool, p_error text, p_data jsonb)`
- Security: `(p_success bool, p_error text, p_rules jsonb)` - Security: `(p_success bool, p_error text, p_rules jsonb)`
See `database_schema.sql` for complete definitions. See `lookup/database_schema.sql` for complete definitions.
--- ---
@@ -182,7 +183,7 @@ auth := security.NewDatabaseAuthenticator(db)
// Requires these tables: // Requires these tables:
// - users (id, username, email, password, user_level, roles, is_active) // - users (id, username, email, password, user_level, roles, is_active)
// - user_sessions (session_token, user_id, expires_at, created_at, last_activity_at) // - user_sessions (session_token, user_id, expires_at, created_at, last_activity_at)
// See database_schema.sql for full schema // See lookup/database_schema.sql for full schema
// Features: // Features:
// - Login with username/password // - Login with username/password
@@ -313,16 +314,15 @@ func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context,
} }
query := ` query := `
SELECT control, accesstype, jsonvalue SELECT schema_name || '.' || table_name || '.' || column_path AS control,
FROM core.secaccess access_type AS accesstype, COALESCE(extra_filters, '') AS jsonvalue
WHERE rid_hub IN ( FROM sec_column_rules
SELECT rid_hub_parent FROM core.hub_link WHERE is_active = true
WHERE rid_hub_child = ? AND parent_hubtype = 'secgroup' AND lower(schema_name) = lower(?) AND lower(table_name) = lower(?)
) AND (user_id = ? OR group_id IN (SELECT group_id FROM sec_group_members WHERE user_id = ?))
AND control ILIKE ?
` `
err := p.db.WithContext(ctx).Raw(query, userID, fmt.Sprintf("%s.%s%%", schema, table)).Scan(&records).Error err := p.db.WithContext(ctx).Raw(query, schema, table, userID, userID).Scan(&records).Error
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -378,19 +378,19 @@ func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userID i
```go ```go
// Test Authenticator // Test Authenticator
auth := security.NewHeaderAuthenticator() auth := providers.NewHeaderAuthenticator()
req := httptest.NewRequest("GET", "/", nil) req := httptest.NewRequest("GET", "/", nil)
req.Header.Set("X-User-ID", "123") req.Header.Set("X-User-ID", "123")
userCtx, err := auth.Authenticate(req) userCtx, err := auth.Authenticate(req)
assert.Equal(t, 123, userCtx.UserID) assert.Equal(t, 123, userCtx.UserID)
// Test ColumnSecurityProvider // Test ColumnSecurityProvider
colSec := security.NewConfigColumnSecurityProvider(rules) colSec := providers.NewConfigColumnSecurityProvider(rules)
cols, err := colSec.GetColumnSecurity(context.Background(), 123, "public", "employees") cols, err := colSec.GetColumnSecurity(context.Background(), 123, "public", "employees")
assert.Equal(t, "mask", cols[0].Accesstype) assert.Equal(t, "mask", cols[0].Accesstype)
// Test RowSecurityProvider // Test RowSecurityProvider
rowSec := security.NewConfigRowSecurityProvider(templates, blocked) rowSec := providers.NewConfigRowSecurityProvider(templates, blocked)
row, err := rowSec.GetRowSecurity(context.Background(), 123, "public", "orders") row, err := rowSec.GetRowSecurity(context.Background(), 123, "public", "orders")
assert.Equal(t, "user_id = {UserID}", row.Template) assert.Equal(t, "user_id = {UserID}", row.Template)
``` ```
@@ -633,7 +633,8 @@ func main() {
// Setup security // Setup security
provider := &SimpleProvider{} provider := &SimpleProvider{}
securityList := security.SetupSecurityProvider(handler, provider) securityList, _ := security.NewSecurityList(provider)
restheadspec.RegisterSecurityHooks(handler, securityList)
// Apply middleware // Apply middleware
router := mux.NewRouter() router := mux.NewRouter()
@@ -762,7 +763,8 @@ auth := security.NewJWTAuthenticator("secret", db)
colSec := security.NewDatabaseColumnSecurityProvider(db) colSec := security.NewDatabaseColumnSecurityProvider(db)
rowSec := security.NewDatabaseRowSecurityProvider(db) rowSec := security.NewDatabaseRowSecurityProvider(db)
provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec) provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
securityList := security.SetupSecurityProvider(handler, provider) securityList, _ := security.NewSecurityList(provider)
restheadspec.RegisterSecurityHooks(handler, securityList)
// ===== INTERFACE METHODS ===== // ===== INTERFACE METHODS =====
Authenticate(r *http.Request) (*UserContext, error) Authenticate(r *http.Request) (*UserContext, error)
+113 -102
View File
@@ -13,12 +13,12 @@ Type-safe, composable security system for ResolveSpec with support for authentic
- ✅ **Extensible** - Implement custom providers for your needs - ✅ **Extensible** - Implement custom providers for your needs
- ✅ **Stored Procedures** - Database operations use PostgreSQL stored procedures where available, for security and maintainability - ✅ **Stored Procedures** - Database operations use PostgreSQL stored procedures where available, for security and maintainability
- ✅ **Direct Mode** - Portable Go/SQL fallback for SQLite, MySQL, or Postgres without the stored procedures installed — no code changes required - ✅ **Direct Mode** - Portable Go/SQL fallback for SQLite, MySQL, or Postgres without the stored procedures installed — no code changes required
- ✅ **OAuth2 Authorization Server** - Built-in OAuth 2.1 + PKCE server (RFC 8414, 7591, 7009, 7662) with login form and external provider federation - ✅ **OAuth2 / OpenID Connect** - Built-in OAuth 2.1 + PKCE authorization server and OIDC provider (RFC 8414, 7591/7592, 7009, 7662, 9068, 9126, 9207, 9449, 8628, 8693): consent, rotating refresh tokens, JWT access tokens, logout, federation; plus an OIDC relying-party client. See [OAUTH2_SERVER.md](OAUTH2_SERVER.md)
- ✅ **Password Reset** - Self-service password reset with secure token generation and session invalidation - ✅ **Password Reset** - Self-service password reset with secure token generation and session invalidation
## Stored Procedure Architecture ## Stored Procedure Architecture
**All database-backed security providers use PostgreSQL stored procedures exclusively.** No raw SQL queries are executed from Go code. **On PostgreSQL, database-backed security providers use stored procedures by default.** `pkg/security` itself contains no SQL; all database access lives in [`pkg/security/lookup`](lookup), which can also run the same operations as direct SQL on tables (see [Database access (lookup)](#database-access-lookup)).
### Benefits ### Benefits
@@ -35,6 +35,7 @@ Type-safe, composable security system for ResolveSpec with support for authentic
| `resolvespec_login` | Session-based login | DatabaseAuthenticator | | `resolvespec_login` | Session-based login | DatabaseAuthenticator |
| `resolvespec_logout` | Session invalidation | DatabaseAuthenticator | | `resolvespec_logout` | Session invalidation | DatabaseAuthenticator |
| `resolvespec_session` | Session validation | DatabaseAuthenticator | | `resolvespec_session` | Session validation | DatabaseAuthenticator |
| `resolvespec_login_api_key` | Exchange a raw header/generic API key for a session (defined in `lookup/keystore_schema.sql`; direct mode reads `user_keys`) | DatabaseAuthenticator.LoginWithAPIKey |
| `resolvespec_session_update` | Update session activity | DatabaseAuthenticator | | `resolvespec_session_update` | Update session activity | DatabaseAuthenticator |
| `resolvespec_refresh_token` | Token refresh | DatabaseAuthenticator | | `resolvespec_refresh_token` | Token refresh | DatabaseAuthenticator |
| `resolvespec_jwt_login` | JWT user validation | JWTAuthenticator | | `resolvespec_jwt_login` | JWT user validation | JWTAuthenticator |
@@ -50,95 +51,96 @@ Type-safe, composable security system for ResolveSpec with support for authentic
| `resolvespec_password_reset_request` | Create password reset token | DatabaseAuthenticator | | `resolvespec_password_reset_request` | Create password reset token | DatabaseAuthenticator |
| `resolvespec_password_reset` | Validate token and set new password | DatabaseAuthenticator | | `resolvespec_password_reset` | Validate token and set new password | DatabaseAuthenticator |
See `database_schema.sql` for complete stored procedure definitions and examples. See `lookup/database_schema.sql` for complete stored procedure definitions and examples.
**Not on Postgres, or don't have the procedures installed?** See [Direct Mode](#direct-mode-portable-sql-without-stored-procedures) below — every provider that calls a `resolvespec_*` procedure also has a portable Go/SQL implementation that works on SQLite, MySQL, or plain Postgres. **Not on Postgres, or don't have the procedures installed?** See [Database access (lookup)](#database-access-lookup) below: every provider can also work directly on tables, on SQLite, MySQL, SQL Server or plain Postgres.
## Direct Mode (portable SQL without stored procedures) ## Database access (lookup)
Every database-backed provider (`DatabaseAuthenticator`, `JWTAuthenticator`, `DatabaseTwoFactorProvider`, `DatabasePasskeyProvider`, the OAuth2 methods/server, `DatabaseKeyStore`) has two code paths: `pkg/security` itself contains no SQL. Every database-backed provider (`DatabaseAuthenticator`, `JWTAuthenticator`, column/row security, `DatabaseTwoFactorProvider`, `DatabasePasskeyProvider`, the OAuth2 methods/server, `DatabaseKeyStore`) calls a store interface from `pkg/security/lookup`. Two backends implement each store:
- **Procedure mode** — calls the configured `resolvespec_*` stored procedure (original behavior, Postgres-only). - **procedure** (`lookup/procedure`): calls the `resolvespec_*` stored procedures (`p_success` / `p_error` / `p_data` contract). Postgres only.
- **Direct mode** — reimplements the same logic in Go using plain parameterized SQL against configurable table names. Works on SQLite, MySQL, or a Postgres database where the procedures were never deployed. - **direct** (`lookup/direct`): plain parameterized SQL on tables, rendered by a per-database dialect (`lookup/dialect`: postgres, sqlite, mysql, mssql). Table and column names are configurable.
### QueryMode `lookup/backends.New(db, cfg, opts)` builds a `lookup.Provider` (all stores) and routes each operation to a backend. The security constructors do this for you from `lookup.Config`; pass a ready `*lookup.Provider` with `LookupProvider` / `WithLookupProvider` to share one between components.
Selection is controlled per-provider by a `QueryMode`: ### Choosing the mode
```go ```go
type QueryMode int type Config struct {
Dialect string // "postgres", "sqlite", "mysql", "mssql", or one you registered; empty = detect from the driver
const ( Mode lookup.Mode // default for every operation
ModeAuto QueryMode = iota // default Overrides map[lookup.Op]lookup.Mode // per-operation mode, e.g. lookup.OpSession: lookup.ModeDirect
ModeProcedure Procs lookup.ProcNames // procedure names, empty fields keep the default
ModeDirect Schema lookup.Schema // table/column names, missing entries keep the default
)
```
- **`ModeAuto`** (default, zero value) — auto-detects per connection:
- SQLite/MySQL drivers → Direct mode, no probing.
- Postgres drivers (`lib/pq`, `pgx`) → probes `pg_proc` for the configured procedure name and uses it **only if it actually exists**; otherwise falls back to Direct mode. The result is cached per procedure name and reset on reconnect.
- Any other/unrecognized driver (including `sqlmock` test doubles) → defaults to Procedure mode, preserving existing behavior for callers that don't expose an identifiable driver type.
- **`ModeProcedure`** — always calls the stored procedure, regardless of dialect.
- **`ModeDirect`** — always uses the portable Go/SQL path, never the stored procedure.
Set it via the provider's `Options` struct or `With...` chain method:
```go
auth := security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{
QueryMode: security.ModeDirect, // force Direct mode, e.g. for SQLite
})
tfaProvider := security.NewDatabaseTwoFactorProvider(sqliteDB, nil).
WithQueryMode(security.ModeDirect)
```
On a real SQLite/MySQL connection you can usually leave `QueryMode` unset — `ModeAuto` detects the dialect and uses Direct mode automatically.
### TableNames / KeyStoreTableNames
Direct mode reads/writes plain tables instead of calling procedures, so table names are configurable the same way procedure names are (`SQLNames`):
```go
type TableNames struct {
Users string // default: "users"
UserSessions string // default: "user_sessions"
TokenBlacklist string // default: "token_blacklist"
UserTOTPBackupCodes string // default: "user_totp_backup_codes"
UserPasskeyCredentials string // default: "user_passkey_credentials"
UserPasswordResets string // default: "user_password_resets"
OAuthClients string // default: "oauth_clients"
OAuthCodes string // default: "oauth_codes"
}
type KeyStoreTableNames struct {
UserKeys string // default: "user_keys" — used by DatabaseKeyStore
} }
``` ```
`DefaultTableNames()` / `MergeTableNames()` / `ValidateTableNames()` mirror `DefaultSQLNames()` / `MergeSQLNames()` / `ValidateSQLNames()`. Set custom names via the same `Options`/`With...` surface as `QueryMode`: | Mode | Behaviour |
|---|---|
| `ModeDefault` (zero value) | stored procedure on Postgres, direct SQL on every other dialect |
| `ModeProcedure` | always the procedure; an error on a non-Postgres dialect |
| `ModeDirect` | always direct SQL |
| `ModeAuto` | Postgres: probe `pg_proc` once per procedure (cached, reset on reconnect), use it if present, else direct. Other dialects: direct |
An impossible combination (procedure on SQLite) fails when the provider is built, not on the first request. If the dialect is not set and cannot be detected from the driver, Postgres is assumed. A bad configuration makes every call return the error (fail closed).
```go ```go
auth := security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{ // SQLite or MySQL: nothing to configure, direct SQL is the default.
TableNames: &security.TableNames{Users: "app_users"}, // only override what differs auth := security.NewDatabaseAuthenticator(sqliteDB)
// Postgres without the procedures installed: use tables only.
auth = security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{
Lookup: lookup.Config{Mode: lookup.ModeDirect},
}) })
// Postgres, procedures for everything except session lookups.
auth = security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{
Lookup: lookup.Config{Overrides: map[lookup.Op]lookup.Mode{lookup.OpSession: lookup.ModeDirect}},
})
tfa := security.NewDatabaseTwoFactorProvider(db, nil).WithLookup(lookup.Config{Mode: lookup.ModeDirect})
``` ```
`oauth2_methods.go` and `oauth_server_db.go` are methods on `*DatabaseAuthenticator` and reuse its `TableNames`/`QueryMode`; there's no separate config for them. Other components take the same `Lookup` / `LookupProvider` options (or `WithLookup` / `WithLookupProvider` on the chain-style types).
### Schema ### Custom names
`database_schema_sqlite.sql` is the portable companion to `database_schema.sql` — plain `CREATE TABLE` statements only (no functions, no triggers, no `jsonb`/`bytea`/array types), covering every table Direct mode reads or writes. Use it to stand up a SQLite (or adapt for MySQL) database for Direct mode. ```go
cfg := lookup.Config{
Procs: lookup.ProcNames{Login: "myapp_login"}, // only override what differs
Schema: lookup.Schema{
lookup.EntityUsers: {Name: "app_users", Columns: map[string]string{"username": "login_name"}},
},
}
```
### What's NOT covered Procedure names, table names and column names are validated as identifiers at construction. `Schema` entries may also set `Schema` to qualify a table (`schema.table`).
`ColumnSecurityProvider`/`RowSecurityProvider` (`resolvespec_column_security` / `resolvespec_row_security`) query an external `core.secaccess`/`core.hub_link` schema this package doesn't own. Direct mode has no portable equivalent to fabricate for these and returns `security.ErrDirectModeUnsupported` — use `ConfigColumnSecurityProvider`/`ConfigRowSecurityProvider` instead when not running against Postgres with those procedures installed. ### Schemas
| File | Purpose |
|---|---|
| `lookup/database_schema.sql` | Postgres: tables **and** stored procedures (procedure backend) |
| `lookup/keystore_schema.sql` | Postgres: `user_keys` table and key store procedures |
| `lookup/ddl/{postgres,sqlite,mysql,mssql}.sql` | tables only, for the direct backend, with the default names |
Read them from Go with `ddl.SQL("sqlite")` or, for drivers that reject multi-statement execution (MySQL, SQL Server), `ddl.Statements("mysql")`. Do not mix `ddl/postgres.sql` with `database_schema.sql`: the procedure schema stores passkey credential ids as `bytea` and OAuth lists as `text[]`, the direct backend stores base64 / JSON text. The `ddl` files are starting points: adjust types and collations to your deployment, and set `lookup.Config.Schema` if you rename anything.
### Column and row security
`ColumnSecurityProvider` / `RowSecurityProvider` read `sec_group_members` (optional), `sec_column_rules` and `sec_row_rules`, in both backends. A rule belongs to one user or one group; rules apply to the exact schema and table (case-insensitive, never a prefix); a blocking row rule wins, otherwise row templates are combined with `AND`. A non-numeric user reference is an error, and no rule means no restriction from this provider. `WithNoGroupTables()` skips the membership table.
### Behavioral notes ### Behavioral notes
- Direct mode matches Procedure mode's current behavior exactly, including its TODOs — e.g. passwords are compared as-is (the stored procedures don't verify bcrypt hashes yet either; see the TODO in `resolvespec_login`/`resolvespec_password_reset`). - Direct login, register, refresh, API-key login, password reset and passkey login run in one transaction.
- Session tokens generated by Direct mode use the same `sess_<hex>_<unix-timestamp>` shape as the plpgsql procedures. - Passwords are stored as bcrypt; legacy cleartext values are accepted at login and only rewritten when `UpgradePasswordHash` is enabled.
- `bytea`/array/`jsonb` Postgres-only columns (passkey credentials, OAuth2 client scopes, keystore `meta`) are stored as base64/JSON-encoded `TEXT` in Direct mode — transparent to callers, since the Go-level API already deals in those same encodings. - Session tokens use the shape `sess_<hex>_<unix-timestamp>` in both backends.
- Direct mode stores `bytea` / array / `jsonb` values (passkey credentials, OAuth client lists, key meta) as base64 / JSON text; the Go API is unchanged.
- OAuth authorization codes are consumed atomically.
- Adding a database: implement `dialect.Dialect`, register it with `dialect.Register`, then set `Config.Dialect`.
- Backend conformance: `lookup/conformance` is one behavioural suite run against every backend (`go test ./pkg/security/lookup/backends -run TestConformance`). SQLite runs always; Postgres (procedure and direct), MySQL and SQL Server run when `RESOLVESPEC_TEST_PG_DSN`, `RESOLVESPEC_TEST_PG_DIRECT_DSN`, `RESOLVESPEC_TEST_MYSQL_DSN` or `RESOLVESPEC_TEST_MSSQL_DSN` is set (see the comment in `backends/conformance_test.go`). Rows are prefixed and removed afterwards. With `RESOLVESPEC_TEST_CONTAINERS=1` (and not `-short`) the container tests start a throwaway database with podman or docker (podman first), run the suite and remove the container, so no DSN is needed.
- Migration from the old `SQLNames` / `TableNames` / `QueryMode` API: see `breaking_changes.md`.
## Quick Start ## Quick Start
@@ -297,7 +299,7 @@ type UserContext struct {
**HeaderAuthenticator** - Simple header-based authentication: **HeaderAuthenticator** - Simple header-based authentication:
```go ```go
auth := security.NewHeaderAuthenticator() auth := providers.NewHeaderAuthenticator()
// Expects: X-User-ID, X-User-Name, X-User-Level, etc. // Expects: X-User-ID, X-User-Name, X-User-Level, etc.
``` ```
@@ -307,7 +309,7 @@ auth := security.NewDatabaseAuthenticator(db)
// Supports: Login, Logout, Session management, Token refresh // Supports: Login, Logout, Session management, Token refresh
// All operations use stored procedures: resolvespec_login, resolvespec_logout, // All operations use stored procedures: resolvespec_login, resolvespec_logout,
// resolvespec_session, resolvespec_session_update, resolvespec_refresh_token // resolvespec_session, resolvespec_session_update, resolvespec_refresh_token
// Requires: users and user_sessions tables + stored procedures (see database_schema.sql) // Requires: users and user_sessions tables + stored procedures (see lookup/database_schema.sql)
``` ```
**JWTAuthenticator** - JWT token authentication with login/logout: **JWTAuthenticator** - JWT token authentication with login/logout:
@@ -318,19 +320,19 @@ auth := security.NewJWTAuthenticator("secret-key", db)
// Note: Requires JWT library installation for token signing/verification // Note: Requires JWT library installation for token signing/verification
``` ```
**TwoFactorAuthenticator** - Wraps any authenticator with TOTP 2FA: **totp.Authenticator** - Wraps any authenticator with TOTP 2FA:
```go ```go
baseAuth := security.NewDatabaseAuthenticator(db) baseAuth := security.NewDatabaseAuthenticator(db)
// Use in-memory provider (for testing) // Use in-memory provider (for testing)
tfaProvider := security.NewMemoryTwoFactorProvider(nil) tfaProvider := totp.NewMemoryProvider(nil)
// Or use database provider (for production) // Or use database provider (for production)
tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil) tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil)
// Requires: users table with totp fields, user_totp_backup_codes table // Requires: users table with totp fields, user_totp_backup_codes table
// Requires: resolvespec_totp_* stored procedures (see totp_database_schema.sql) // Requires: resolvespec_totp_* stored procedures (see lookup/database_schema.sql)
auth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil) auth := totp.NewAuthenticator(baseAuth, tfaProvider, nil)
// Supports: TOTP codes, backup codes, QR code generation // Supports: TOTP codes, backup codes, QR code generation
// Compatible with Google Authenticator, Microsoft Authenticator, Authy, etc. // Compatible with Google Authenticator, Microsoft Authenticator, Authy, etc.
``` ```
@@ -341,7 +343,7 @@ auth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil)
```go ```go
colSec := security.NewDatabaseColumnSecurityProvider(db) colSec := security.NewDatabaseColumnSecurityProvider(db)
// Uses stored procedure: resolvespec_column_security // Uses stored procedure: resolvespec_column_security
// Queries core.secaccess and core.hub_link tables // Reads sec_column_rules (user rules + rules of the user's sec_group_members groups)
``` ```
**ConfigColumnSecurityProvider** - Static configuration: **ConfigColumnSecurityProvider** - Static configuration:
@@ -351,7 +353,7 @@ rules := map[string][]security.ColumnSecurity{
{Path: []string{"ssn"}, Accesstype: "mask", MaskStart: 5}, {Path: []string{"ssn"}, Accesstype: "mask", MaskStart: 5},
}, },
} }
colSec := security.NewConfigColumnSecurityProvider(rules) colSec := providers.NewConfigColumnSecurityProvider(rules)
``` ```
### Row Security Providers ### Row Security Providers
@@ -370,7 +372,7 @@ templates := map[string]string{
blocked := map[string]bool{ blocked := map[string]bool{
"public.admin_logs": true, "public.admin_logs": true,
} }
rowSec := security.NewConfigRowSecurityProvider(templates, blocked) rowSec := providers.NewConfigRowSecurityProvider(templates, blocked)
``` ```
## Usage Examples ## Usage Examples
@@ -381,7 +383,7 @@ rowSec := security.NewConfigRowSecurityProvider(templates, blocked)
func main() { func main() {
db := setupDatabase() db := setupDatabase()
// Run migrations (see database_schema.sql) // Run migrations (see lookup/database_schema.sql)
// db.Exec("CREATE TABLE users ...") // db.Exec("CREATE TABLE users ...")
// db.Exec("CREATE TABLE user_sessions ...") // db.Exec("CREATE TABLE user_sessions ...")
@@ -475,8 +477,8 @@ func handleRefresh(securityList *security.SecurityList) http.HandlerFunc {
```go ```go
// 1. Wrap existing authenticator with 2FA support // 1. Wrap existing authenticator with 2FA support
baseAuth := security.NewDatabaseAuthenticator(db) baseAuth := security.NewDatabaseAuthenticator(db)
tfaProvider := security.NewMemoryTwoFactorProvider(nil) // Use custom DB implementation in production tfaProvider := totp.NewMemoryProvider(nil) // Use custom DB implementation in production
tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil) tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, nil)
// 2. Use as normal authenticator // 2. Use as normal authenticator
provider := security.NewCompositeSecurityProvider(tfaAuth, colSec, rowSec) provider := security.NewCompositeSecurityProvider(tfaAuth, colSec, rowSec)
@@ -548,18 +550,18 @@ has2FA, err := tfaProvider.Get2FAStatus(userID)
// Uses PostgreSQL stored procedures for all operations // Uses PostgreSQL stored procedures for all operations
db := setupDatabase() db := setupDatabase()
// Run migrations from totp_database_schema.sql // Run migrations from lookup/database_schema.sql
// - Add totp_secret, totp_enabled, totp_enabled_at to users table // - Add totp_secret, totp_enabled, totp_enabled_at to users table
// - Create user_totp_backup_codes table // - Create user_totp_backup_codes table
// - Create resolvespec_totp_* stored procedures // - Create resolvespec_totp_* stored procedures
tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil) tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil)
tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil) tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, nil)
``` ```
**Option 2: Implement Custom Provider** **Option 2: Implement Custom Provider**
Implement `TwoFactorAuthProvider` for custom storage: Implement `totp.AuthProvider` for custom storage:
```go ```go
type DBTwoFactorProvider struct { type DBTwoFactorProvider struct {
@@ -585,15 +587,15 @@ func (p *DBTwoFactorProvider) Get2FASecret(userID int) (string, error) {
### Configuration ### Configuration
```go ```go
config := &security.TwoFactorConfig{ config := &totp.Config{
Algorithm: "SHA256", // SHA1, SHA256, SHA512 Algorithm: "SHA256", // SHA1, SHA256, SHA512
Digits: 8, // 6 or 8 Digits: 8, // 6 or 8
Period: 30, // Seconds per code Period: 30, // Seconds per code
SkewWindow: 2, // Accept codes ±2 periods SkewWindow: 2, // Accept codes ±2 periods
} }
totp := security.NewTOTPGenerator(config) totp := totp.NewGenerator(config)
tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, config) tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, config)
``` ```
### API Response Structure ### API Response Structure
@@ -662,9 +664,9 @@ func main() {
} }
// Create providers // Create providers
auth := security.NewHeaderAuthenticator() auth := providers.NewHeaderAuthenticator()
colSec := security.NewConfigColumnSecurityProvider(columnRules) colSec := providers.NewConfigColumnSecurityProvider(columnRules)
rowSec := security.NewConfigRowSecurityProvider(rowTemplates, nil) rowSec := providers.NewConfigRowSecurityProvider(rowTemplates, nil)
// Combine providers and register hooks // Combine providers and register hooks
provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec) provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
@@ -934,7 +936,8 @@ func TestMyHandler(t *testing.T) {
&MockRowSecurity{}, &MockRowSecurity{},
) )
securityList := security.SetupSecurityProvider(handler, provider) securityList, _ := security.NewSecurityList(provider)
restheadspec.RegisterSecurityHooks(handler, securityList)
// ... test your handler // ... test your handler
} }
``` ```
@@ -1013,7 +1016,7 @@ restheadspec.RegisterSecurityHooks(handler, securityList) // or funcspec/resolve
### DB Requirements ### DB Requirements
Run the migrations in `database_schema.sql`: Run the migrations in `lookup/database_schema.sql`:
- `user_password_resets` table (`user_id`, `token_hash` SHA-256, `expires_at`, `used`, `used_at`) - `user_password_resets` table (`user_id`, `token_hash` SHA-256, `expires_at`, `used`, `used_at`)
- `resolvespec_password_reset_request` stored procedure - `resolvespec_password_reset_request` stored procedure
- `resolvespec_password_reset` stored procedure - `resolvespec_password_reset` stored procedure
@@ -1050,13 +1053,14 @@ err = auth.CompletePasswordReset(ctx, security.PasswordResetCompleteRequest{
- `RequestPasswordReset` always returns success even when the email/username is not found, preventing user enumeration - `RequestPasswordReset` always returns success even when the email/username is not found, preventing user enumeration
- Hash the new password with bcrypt before storing (pgcrypto `crypt`/`gen_salt`) — see the TODO comment in `resolvespec_password_reset` - Hash the new password with bcrypt before storing (pgcrypto `crypt`/`gen_salt`) — see the TODO comment in `resolvespec_password_reset`
### SQLNames ### Procedure names
Set through `lookup.Config.Procs` (`lookup.ProcNames`):
```go ```go
type SQLNames struct { lookup.ProcNames{
// ... PasswordResetRequest: "resolvespec_password_reset_request", // default
PasswordResetRequest string // default: "resolvespec_password_reset_request" PasswordResetComplete: "resolvespec_password_reset", // default
PasswordResetComplete string // default: "resolvespec_password_reset"
} }
``` ```
@@ -1064,6 +1068,8 @@ type SQLNames struct {
## OAuth2 Authorization Server ## OAuth2 Authorization Server
> The complete guide (consent, OIDC, refresh rotation, DPoP, PAR, device grant, token exchange, logout, relying-party client) is in [OAUTH2_SERVER.md](OAUTH2_SERVER.md). The table below lists the original endpoints.
`OAuthServer` is a generic OAuth 2.1 + PKCE authorization server. It is not tied to any spec — `pkg/resolvemcp` uses it, but it can be used standalone with any `http.ServeMux`. `OAuthServer` is a generic OAuth 2.1 + PKCE authorization server. It is not tied to any spec — `pkg/resolvemcp` uses it, but it can be used standalone with any `http.ServeMux`.
### Endpoints ### Endpoints
@@ -1151,7 +1157,7 @@ http.ListenAndServe(":8080", mux)
When `PersistClients: true` or `PersistCodes: true`, the server calls the corresponding `DatabaseAuthenticator` methods. Both flags default to `false` (in-memory maps). Enable both for multi-instance deployments. When `PersistClients: true` or `PersistCodes: true`, the server calls the corresponding `DatabaseAuthenticator` methods. Both flags default to `false` (in-memory maps). Enable both for multi-instance deployments.
Requires `oauth_clients` and `oauth_codes` tables + 6 stored procedures from `database_schema.sql`. Requires `oauth_clients` and `oauth_codes` tables + 6 stored procedures from `lookup/database_schema.sql`.
#### New DB Types #### New DB Types
@@ -1198,10 +1204,10 @@ auth.OAuthIntrospectToken(ctx, token) // RFC 7662 — returns OAuthTokenInfo
auth.OAuthRevokeToken(ctx, token) // RFC 7009 — revoke session auth.OAuthRevokeToken(ctx, token) // RFC 7009 — revoke session
``` ```
#### SQLNames Fields #### Procedure names
```go ```go
type SQLNames struct { type ProcNames struct {
// ... existing fields ... // ... existing fields ...
OAuthRegisterClient string // default: "resolvespec_oauth_register_client" OAuthRegisterClient string // default: "resolvespec_oauth_register_client"
OAuthGetClient string // default: "resolvespec_oauth_get_client" OAuthGetClient string // default: "resolvespec_oauth_get_client"
@@ -1224,9 +1230,14 @@ The main changes:
| File | Description | | File | Description |
|------|-------------| |------|-------------|
| **QUICK_REFERENCE.md** | Quick reference guide with examples | | **QUICK_REFERENCE.md** | Quick reference guide with examples |
| **INTERFACE_GUIDE.md** | Complete implementation guide | | **KEYSTORE.md** | Per-user auth keys and key stores |
| **examples.go** | Working provider implementations | | **OAUTH2.md** | OAuth2 client login |
| **setup_example.go** | 6 complete integration examples | | **OAUTH2_SERVER.md** | OAuth 2.1 / OpenID Connect server and relying-party client (full guide) |
| **OAUTH2_REFRESH_QUICK_REFERENCE.md** / **OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md** | OAuth2 refresh tokens |
| **PASSKEY_QUICK_REFERENCE.md** | WebAuthn passkeys |
| **SECURITY_FEATURES.md** | Security feature overview |
| **breaking_changes.md** | Migration notes (`lookup` refactor, full OAuth2/OIDC schema changes) |
| **examples.go**, **examples_funcspec.go**, **oauth2_examples.go**, **oauth2_full_example.go**, **passkey_examples.go** | Working provider implementations |
## API Reference ## API Reference
+159
View File
@@ -0,0 +1,159 @@
# Breaking changes
Appended as each step of `audit/sec_query_builder.plan.md` lands.
## Step 0: shared types moved to `sectypes` (no action needed)
The plain data types now live in `pkg/security/sectypes` and are aliased in `pkg/security`
(`types.go`), so `security.UserContext` and `sectypes.UserContext` are the same type.
Moved: `UserContext`, `LoginRequest/Response`, `RegisterRequest`, `LogoutRequest`,
`PasswordReset*`, `KeyType` (+ constants), `UserKey`, `CreateKey*`, `OAuthServerClient`,
`OAuthCode`, `OAuthTokenInfo`, `Passkey*` data structs, `TwoFactorSecret`, `ColumnSecurity`,
`RowSecurity` (incl. `GetTemplate`).
## Step 0b (totp): moved to `pkg/security/totp`
Import `github.com/bitechdev/ResolveSpec/pkg/security/totp`. No aliases (import cycle).
| Old | New |
|---|---|
| `security.TwoFactorAuthProvider` | `totp.AuthProvider` |
| `security.TwoFactorConfig` / `DefaultTwoFactorConfig` | `totp.Config` / `totp.DefaultConfig` |
| `security.TOTPGenerator` / `NewTOTPGenerator` | `totp.Generator` / `totp.NewGenerator` |
| `security.GenerateBackupCodes` | `totp.GenerateBackupCodes` |
| `security.MemoryTwoFactorProvider` / `NewMemoryTwoFactorProvider` | `totp.MemoryProvider` / `totp.NewMemoryProvider` |
| `security.TwoFactorAuthenticator` / `NewTwoFactorAuthenticator` | `totp.Authenticator` / `totp.NewAuthenticator` |
`totp.NewAuthenticator` takes a `totp.BaseAuthenticator` (Login, Logout, Authenticate) instead of
`security.Authenticator`; any `security.Authenticator` satisfies it.
`DatabaseTwoFactorProvider` stays in `security` (it now calls the lookup `TOTPStore`). Core imports
`totp`, so `totp` must not import `security`.
## Step 0b (providers, first part): moved to `pkg/security/providers`
Import `github.com/bitechdev/ResolveSpec/pkg/security/providers`. Names unchanged, no aliases (import cycle).
| Old | New |
|---|---|
| `security.HeaderAuthenticator` / `NewHeaderAuthenticator` | `providers.HeaderAuthenticator` / `providers.NewHeaderAuthenticator` |
| `security.ConfigKeyStore` / `NewConfigKeyStore` | `providers.ConfigKeyStore` / `providers.NewConfigKeyStore` |
| `security.KeyStoreAuthenticator` / `NewKeyStoreAuthenticator` | `providers.KeyStoreAuthenticator` / `providers.NewKeyStoreAuthenticator` |
| `security.ConfigColumnSecurityProvider` / `NewConfigColumnSecurityProvider` | `providers.ConfigColumnSecurityProvider` / `providers.NewConfigColumnSecurityProvider` |
| `security.ConfigRowSecurityProvider` / `NewConfigRowSecurityProvider` | `providers.ConfigRowSecurityProvider` / `providers.NewConfigRowSecurityProvider` |
The SHA-256 key hash helper is now `sectypes.HashKey`. The database-backed providers
(`DatabaseAuthenticator`, `JWTAuthenticator`, `DatabaseKeyStore`, `DatabaseColumn/RowSecurityProvider`)
stay in `security`; they call the lookup stores (see step 5).
## Additions (no action needed)
- `common.SQLDBProvider` (`SQLDB() *sql.DB`) is implemented by the bun, gorm and pgsql adapters
(not their transaction adapters). `lookup.FromDatabase(common.Database)` uses it, plus the
adapter's `DriverName()`, to get the `*sql.DB` and dialect name.
## Steps 2–3: lookup dialects and procedure backend
- New `lookup/dialect` (postgres, sqlite, mysql, mssql) and `lookup/procedure` packages. No
existing exported `security` API changed in these steps.
- Procedure-mode code paths in `DatabaseAuthenticator`, `JWTAuthenticator`, the policy providers,
`DatabaseKeyStore`, `DatabaseTwoFactorProvider` and `DatabasePasskeyProvider` now delegate to
`lookup/procedure`. Error texts are unchanged.
- Behaviour change (improvement): these procedure paths now reconnect once on a closed `*sql.DB`
(JWT logout, key create, TOTP, passkey, OAuth previously used the handle directly).
## Step 4: lookup/direct backend
- New `lookup/direct` package: table-backed stores for auth, keys, OAuth (client + user), passkey,
TOTP and policy, built from `lookup.Schema` and the dialect. Nothing in `pkg/security` calls it
yet (wiring happens in step 5), so no existing API changes here.
- Direct `LoginAPIKey` is new: `header_api` / `api` keys only; unknown, expired, inactive and
wrong-type keys (and inactive users) all return `lookup.ErrInvalidAPIKey`.
- Policy tables (`sec_group_members`, `sec_column_rules`, `sec_row_rules`) are required for the
direct policy store; `PolicyOptions.NoGroups` skips the membership table.
- Direct behaviour that changes when step 5 switches over: login, register, refresh, API-key login,
password reset and passkey login now write in one transaction; `Keys.Create` stores NULL (not the
text `null`) for empty scopes/meta; OAuth code exchange consumes the code atomically; a
non-numeric row-security user reference is an error instead of loading no rules.
## Step 5: pkg/security uses lookup
`pkg/security` no longer contains SQL (guarded by `TestCoreContainsNoSQL`). Every database call goes
through a `lookup.Provider` built by `lookup/backends.New`.
Removed (replaced by `lookup.Config`: `Dialect`, `Mode`, `Overrides`, `Procs`, `Schema`):
- Types and functions `SQLNames`, `DefaultSQLNames`, `MergeSQLNames`, `ValidateSQLNames`,
`TableNames` (+ Default/Merge/Validate), `KeyStoreSQLNames`, `KeyStoreTableNames` (+ same),
`QueryMode`, `ModeAuto`/`ModeProcedure`/`ModeDirect`, `ErrDirectModeUnsupported`.
- Options fields `SQLNames`, `TableNames`, `QueryMode` on `DatabaseAuthenticatorOptions`,
`DatabaseKeyStoreOptions`, `DatabasePasskeyProviderOptions`; replaced by `Lookup lookup.Config`
and `LookupProvider *lookup.Provider`.
- Builders `WithQueryMode`, `WithTableNames` on `JWTAuthenticator`, the column/row providers and
`DatabaseTwoFactorProvider`; replaced by `WithLookup(cfg)` and `WithLookupProvider(p)`.
- The variadic `names ...*SQLNames` argument of `NewJWTAuthenticator`,
`NewDatabaseColumnSecurityProvider`, `NewDatabaseRowSecurityProvider` and
`NewDatabaseTwoFactorProvider`.
Behaviour changes:
- Default mode is per dialect: stored procedures on Postgres, direct SQL elsewhere. `ModeAuto`
(probe `pg_proc` once per procedure) is now opt-in via `lookup.ModeAuto`; it used to be the
default everywhere. Procedure mode on a non-Postgres dialect is a configuration error.
- A bad lookup configuration no longer panics or is silently ignored: the component logs it and
every call returns the error.
- Dialect is detected from the driver; if detection fails the postgres dialect is assumed.
- Column and row security now work in direct mode (tables `sec_group_members`, `sec_column_rules`,
`sec_row_rules`); they used to return `ErrDirectModeUnsupported`. `WithNoGroupTables()` skips the
group membership table.
- `LoginWithAPIKey` works in direct mode; `DatabaseAuthenticator.Logout` now clears the session
cache in every mode (direct mode used to skip it).
- `Authenticate` now holds the session lookup in the configured backend only; the activity update no
longer silently falls back to a direct write when the procedure is missing.
- Direct-mode behaviour changes listed under step 4 take effect here.
## Step 6: schemas and docs
- SQL files moved from `pkg/security/` to `pkg/security/lookup/` (`database_schema.sql`,
`keystore_schema.sql`). `database_schema_sqlite.sql` is superseded by `lookup/ddl/sqlite.sql`
(now also includes `sec_group_members`, `sec_column_rules`, `sec_row_rules`).
- New `lookup/ddl` package: embedded reference table schemas for `postgres`, `sqlite`, `mysql`,
`mssql` (`ddl.SQL(dialect)`, `ddl.Statements(dialect)`). `ddl/postgres.sql` is tables only and uses
base64 / JSON text columns, so it cannot be combined with the procedure schema
(`database_schema.sql`, `bytea` / `text[]` columns) on the same tables.
- `security.ApplyTxSettings` is unchanged; its SQL moved to `lookup.ApplyTxSettings(ctx, tx, settings)`.
- Removed the unexported `password.go` from `pkg/security` (bcrypt helpers live in `lookup/direct`).
- `README.md`, `KEYSTORE.md` and the root README describe `lookup.Config` instead of `QueryMode`,
`SQLNames` and `TableNames`.
## Step 8: full OAuth2 / OpenID Connect
Full guide: [OAUTH2_SERVER.md](OAUTH2_SERVER.md). New features are opt-in; the items below are what existing installs must do or notice.
### Schema (existing installs)
Fresh installs use `lookup/database_schema.sql` or `lookup/ddl/<dialect>.sql`. Existing databases need:
- `ALTER TABLE oauth_clients ADD COLUMN metadata <json>` (client metadata: logout URIs, jwks, require_consent, first_party, dpop_bound, signing algs, ...)
- `ALTER TABLE oauth_codes ADD COLUMN extra <json>` (nonce, auth_time, acr, amr, claims, user_id, dpop_jkt, resource)
- New tables `oauth_consents`, `oauth_refresh_tokens`, `oauth_device_codes`, `oauth_par_requests`, `oauth_jti` (copy them from the schema files). Access-grant records are stored in `oauth_refresh_tokens`.
- Postgres procedure mode: reapply `lookup/database_schema.sql` (new `resolvespec_oauth_*` functions, listed in `lookup/procs.go`).
`<json>` is `jsonb` on Postgres, `TEXT` on SQLite, `JSON` on MySQL and `NVARCHAR(MAX)` on SQL Server. New lookup operations and `lookup.OAuthGrantStore` (`Provider.OAuthGrant`) are added to the procedure and direct backends and to the conformance suite; custom `lookup.Config.Procs` overrides gain the new names.
### New API (no action needed)
`OAuthServerConfig` options (see the guide), `OAuthSigningKey`, `OAuthServer.RegisterTrustedClient`, `VerifyAccessToken`, `OAuthClaimsProvider`; `OIDCConfig`, `DatabaseAuthenticator.WithOIDC`, `OAuth2GetAuthURLWithOptions`, `OAuth2HandleCallbackRequest`, `OAuth2LogoutURL`; `OAuth2Config` gains `Issuer`, `JWKSURL`, `EndSessionURL`, `UsePKCE`, `AllowedAlgs`, `AuthStyle`, `HTTPClient`, `ClockSkew`. `DatabaseAuthenticator` gains `OAuthUpdateClient`, `OAuthDeleteClient`, `OAuthGetUser`, `OAuthGrants`.
### Behaviour changes
- `/oauth/introspect` and `/oauth/revoke` require client authentication. Set `AllowAnonymousIntrospection` for the old behaviour.
- Once the `redirect_uri` is validated, authorization errors are redirected to the client (`error`, `state`, `iss`) instead of being returned as JSON. Authorization responses carry `iss` (RFC 9207).
- Only PKCE `S256` is accepted.
- The login form is an `html/template` page with a signed state field; direct form POSTs of earlier versions are still accepted.
- Default grant types of a dynamically registered client include `refresh_token`.
- Authorization-code grants mint a fresh session for the grant. Tokens saved directly with `OAuthSaveCode(SessionToken: ...)` keep working.
- `OAuth2Provider` keeps its PKCE verifier and nonce with the `state`; `Google` preset now validates id_tokens and uses the OpenID Connect endpoints.
- Unauthenticated `userinfo` and discovery routes are unchanged; `userinfo` also answers POST and releases only the claims the granted scopes allow.
### Not supported
`client_secret_jwt`, signed request objects, the DPoP server nonce, `c_hash`, encrypted id_tokens and `actor_token`.
+15
View File
@@ -55,3 +55,18 @@ func (c *ChainAuthenticator) Logout(ctx context.Context, req LogoutRequest) erro
func (c *ChainAuthenticator) LogoutWithCookie(ctx context.Context, req LogoutRequest, w http.ResponseWriter) error { func (c *ChainAuthenticator) LogoutWithCookie(ctx context.Context, req LogoutRequest, w http.ResponseWriter) error {
return c.authenticators[0].LogoutWithCookie(ctx, req, w) return c.authenticators[0].LogoutWithCookie(ctx, req, w)
} }
// LoginWithAPIKey tries each authenticator that supports API key login and
// returns the first success. Failures collapse to one generic error.
func (c *ChainAuthenticator) LoginWithAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*LoginResponse, error) {
for _, a := range c.authenticators {
l, ok := a.(APIKeyLoginable)
if !ok {
continue
}
if resp, err := l.LoginWithAPIKey(ctx, rawKey, claims); err == nil {
return resp, nil
}
}
return nil, errInvalidAPIKey
}
+8
View File
@@ -88,6 +88,14 @@ func (c *CompositeSecurityProvider) RefreshToken(ctx context.Context, refreshTok
return nil, fmt.Errorf("authenticator does not support token refresh") return nil, fmt.Errorf("authenticator does not support token refresh")
} }
// LoginWithAPIKey implements APIKeyLoginable if the authenticator supports it
func (c *CompositeSecurityProvider) LoginWithAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*LoginResponse, error) {
if l, ok := c.auth.(APIKeyLoginable); ok {
return l.LoginWithAPIKey(ctx, rawKey, claims)
}
return nil, fmt.Errorf("authenticator does not support API key login")
}
// ValidateToken implements Validatable if the authenticator supports it // ValidateToken implements Validatable if the authenticator supports it
func (c *CompositeSecurityProvider) ValidateToken(ctx context.Context, token string) (bool, error) { func (c *CompositeSecurityProvider) ValidateToken(ctx context.Context, token string) (bool, error) {
if validatable, ok := c.auth.(Validatable); ok { if validatable, ok := c.auth.(Validatable); ok {
-143
View File
@@ -1,143 +0,0 @@
-- Portable schema for Direct-mode (non-stored-procedure) operation.
-- Plain CREATE TABLE statements only, no functions/triggers, using types
-- understood by SQLite (and portable to MySQL). Used by Direct-mode tests
-- and as a reference for deployments without Postgres.
CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username VARCHAR(255) NOT NULL UNIQUE,
email VARCHAR(255) NOT NULL UNIQUE,
password VARCHAR(255), -- bcrypt hash (nullable for OAuth2 users); legacy cleartext is accepted at login (upgrade to bcrypt is opt-in)
user_level INTEGER DEFAULT 0,
roles VARCHAR(500),
is_active BOOLEAN DEFAULT 1,
created_at TIMESTAMP,
updated_at TIMESTAMP,
last_login_at TIMESTAMP,
program_user_id INTEGER DEFAULT 0,
program_user_table VARCHAR(255) DEFAULT '',
remote_id VARCHAR(255),
auth_provider VARCHAR(50),
totp_secret VARCHAR(255),
totp_enabled BOOLEAN DEFAULT 0,
totp_enabled_at TIMESTAMP
);
CREATE TABLE IF NOT EXISTS user_sessions (
id INTEGER PRIMARY KEY AUTOINCREMENT,
session_token VARCHAR(500) NOT NULL UNIQUE,
user_id INTEGER NOT NULL,
expires_at TIMESTAMP NOT NULL,
created_at TIMESTAMP,
last_activity_at TIMESTAMP,
ip_address VARCHAR(45),
user_agent TEXT,
access_token TEXT,
refresh_token TEXT,
token_type VARCHAR(50) DEFAULT 'Bearer',
auth_provider VARCHAR(50)
);
CREATE INDEX IF NOT EXISTS idx_session_token ON user_sessions(session_token);
CREATE INDEX IF NOT EXISTS idx_user_id ON user_sessions(user_id);
CREATE INDEX IF NOT EXISTS idx_expires_at ON user_sessions(expires_at);
CREATE INDEX IF NOT EXISTS idx_refresh_token ON user_sessions(refresh_token);
CREATE TABLE IF NOT EXISTS token_blacklist (
id INTEGER PRIMARY KEY AUTOINCREMENT,
token VARCHAR(500) NOT NULL,
user_id INTEGER,
expires_at TIMESTAMP NOT NULL,
created_at TIMESTAMP
);
CREATE TABLE IF NOT EXISTS user_totp_backup_codes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
code_hash VARCHAR(64) NOT NULL,
used BOOLEAN DEFAULT 0,
used_at TIMESTAMP,
created_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_totp_user_id ON user_totp_backup_codes(user_id);
CREATE INDEX IF NOT EXISTS idx_totp_code_hash ON user_totp_backup_codes(code_hash);
CREATE TABLE IF NOT EXISTS user_passkey_credentials (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
credential_id TEXT NOT NULL UNIQUE, -- base64 text (Direct mode), not native bytea
public_key TEXT NOT NULL, -- base64 text
attestation_type VARCHAR(50) DEFAULT 'none',
aaguid TEXT, -- base64 text
sign_count INTEGER DEFAULT 0,
clone_warning BOOLEAN DEFAULT 0,
transports TEXT, -- JSON-encoded []string
backup_eligible BOOLEAN DEFAULT 0,
backup_state BOOLEAN DEFAULT 0,
name VARCHAR(255),
created_at TIMESTAMP,
last_used_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_passkey_user_id ON user_passkey_credentials(user_id);
CREATE INDEX IF NOT EXISTS idx_passkey_credential_id ON user_passkey_credentials(credential_id);
CREATE TABLE IF NOT EXISTS user_password_resets (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
token_hash VARCHAR(64) NOT NULL UNIQUE,
expires_at TIMESTAMP NOT NULL,
created_at TIMESTAMP,
used BOOLEAN DEFAULT 0,
used_at TIMESTAMP
);
CREATE TABLE IF NOT EXISTS oauth_clients (
id INTEGER PRIMARY KEY AUTOINCREMENT,
client_id VARCHAR(255) NOT NULL UNIQUE,
redirect_uris TEXT NOT NULL, -- JSON-encoded []string
client_name VARCHAR(255),
grant_types TEXT, -- JSON-encoded []string
allowed_scopes TEXT, -- JSON-encoded []string
client_secret_hash TEXT, -- sha256 hex of the confidential-client secret; NULL for public clients
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
is_active BOOLEAN DEFAULT 1,
created_at TIMESTAMP
);
CREATE TABLE IF NOT EXISTS oauth_codes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
code VARCHAR(255) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
redirect_uri TEXT NOT NULL,
client_state TEXT,
code_challenge VARCHAR(255) NOT NULL,
code_challenge_method VARCHAR(10) DEFAULT 'S256',
session_token TEXT NOT NULL,
refresh_token TEXT,
scopes TEXT, -- JSON-encoded []string
expires_at TIMESTAMP NOT NULL,
created_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_codes_code ON oauth_codes(code);
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
CREATE TABLE IF NOT EXISTS user_keys (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
key_type VARCHAR(50) NOT NULL,
key_hash VARCHAR(64) NOT NULL UNIQUE,
name VARCHAR(255) NOT NULL DEFAULT '',
scopes TEXT, -- JSON-encoded []string
meta TEXT, -- JSON-encoded map
expires_at TIMESTAMP,
created_at TIMESTAMP,
last_used_at TIMESTAMP,
is_active BOOLEAN DEFAULT 1
);
CREATE INDEX IF NOT EXISTS idx_user_keys_user_id ON user_keys(user_id);
CREATE INDEX IF NOT EXISTS idx_user_keys_key_hash ON user_keys(key_hash);
CREATE INDEX IF NOT EXISTS idx_user_keys_key_type ON user_keys(key_type);
+28 -36
View File
@@ -5,37 +5,41 @@ import (
"database/sql" "database/sql"
"encoding/base64" "encoding/base64"
"net/http" "net/http"
"os" "strings"
"path/filepath"
"testing" "testing"
"time" "time"
_ "github.com/mattn/go-sqlite3" _ "github.com/glebarez/go-sqlite"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl"
) )
// directConfig forces the direct (table) backend for every operation.
var directConfig = lookup.Config{Mode: lookup.ModeDirect}
func futureTime() time.Time { func futureTime() time.Time {
return time.Now().Add(1 * time.Hour) return time.Now().Add(1 * time.Hour)
} }
// newDirectTestDB opens a fresh in-memory SQLite database and applies the // newDirectTestDB opens a fresh in-memory SQLite database and applies the
// portable Direct-mode schema (database_schema_sqlite.sql), giving every // portable Direct-mode schema (lookup/ddl/sqlite.sql), giving every
// Direct-mode test a real, isolated database to exercise end-to-end. // Direct-mode test a real, isolated database to exercise end-to-end.
func newDirectTestDB(t *testing.T) *sql.DB { func newDirectTestDB(t *testing.T) *sql.DB {
t.Helper() t.Helper()
db, err := sql.Open("sqlite3", "file::memory:?cache=shared") db, err := sql.Open("sqlite", ":memory:")
if err != nil { if err != nil {
t.Fatalf("failed to open sqlite db: %v", err) t.Fatalf("failed to open sqlite db: %v", err)
} }
db.SetMaxOpenConns(1) // keep the shared in-memory db single-connection so state isn't lost db.SetMaxOpenConns(1) // keep the shared in-memory db single-connection so state isn't lost
t.Cleanup(func() { _ = db.Close() }) t.Cleanup(func() { _ = db.Close() })
schemaPath := filepath.Join("database_schema_sqlite.sql") schema, err := ddl.SQL("sqlite")
schema, err := os.ReadFile(schemaPath)
if err != nil { if err != nil {
t.Fatalf("failed to read schema: %v", err) t.Fatalf("failed to read schema: %v", err)
} }
if _, err := db.Exec(string(schema)); err != nil { if _, err := db.Exec(schema); err != nil {
t.Fatalf("failed to apply schema: %v", err) t.Fatalf("failed to apply schema: %v", err)
} }
return db return db
@@ -49,7 +53,7 @@ func authenticatedRequest(token string) *http.Request {
func TestDirectMode_RegisterThenLogin(t *testing.T) { func TestDirectMode_RegisterThenLogin(t *testing.T) {
db := newDirectTestDB(t) db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
ctx := context.Background() ctx := context.Background()
regResp, err := auth.Register(ctx, RegisterRequest{ regResp, err := auth.Register(ctx, RegisterRequest{
@@ -105,7 +109,7 @@ func TestDirectMode_RegisterThenLogin(t *testing.T) {
func TestDirectMode_SessionLifecycle(t *testing.T) { func TestDirectMode_SessionLifecycle(t *testing.T) {
db := newDirectTestDB(t) db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
ctx := context.Background() ctx := context.Background()
loginResp, err := auth.Register(ctx, RegisterRequest{Username: "bob", Password: "p", Email: "bob@example.com"}) loginResp, err := auth.Register(ctx, RegisterRequest{Username: "bob", Password: "p", Email: "bob@example.com"})
@@ -140,7 +144,7 @@ func TestDirectMode_SessionLifecycle(t *testing.T) {
func TestDirectMode_PasswordReset(t *testing.T) { func TestDirectMode_PasswordReset(t *testing.T) {
db := newDirectTestDB(t) db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
ctx := context.Background() ctx := context.Background()
if _, err := auth.Register(ctx, RegisterRequest{Username: "carol", Password: "old", Email: "carol@example.com"}); err != nil { if _, err := auth.Register(ctx, RegisterRequest{Username: "carol", Password: "old", Email: "carol@example.com"}); err != nil {
@@ -167,8 +171,8 @@ func TestDirectMode_PasswordReset(t *testing.T) {
func TestDirectMode_JWTLoginAndLogout(t *testing.T) { func TestDirectMode_JWTLoginAndLogout(t *testing.T) {
db := newDirectTestDB(t) db := newDirectTestDB(t)
jwtAuth := NewJWTAuthenticator("secret", db).WithQueryMode(ModeDirect) jwtAuth := NewJWTAuthenticator("secret", db).WithLookup(directConfig)
directAuth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) directAuth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
ctx := context.Background() ctx := context.Background()
if _, err := directAuth.Register(ctx, RegisterRequest{Username: "dave", Password: "p", Email: "dave@example.com"}); err != nil { if _, err := directAuth.Register(ctx, RegisterRequest{Username: "dave", Password: "p", Email: "dave@example.com"}); err != nil {
@@ -190,7 +194,7 @@ func TestDirectMode_JWTLoginAndLogout(t *testing.T) {
func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) { func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) {
db := newDirectTestDB(t) db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
ctx := context.Background() ctx := context.Background()
regResp, err := auth.Register(ctx, RegisterRequest{Username: "erin", Password: "p", Email: "erin@example.com"}) regResp, err := auth.Register(ctx, RegisterRequest{Username: "erin", Password: "p", Email: "erin@example.com"})
@@ -199,7 +203,7 @@ func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) {
} }
userID := regResp.User.UserID userID := regResp.User.UserID
totp := NewDatabaseTwoFactorProvider(db, nil).WithQueryMode(ModeDirect) totp := NewDatabaseTwoFactorProvider(db, nil).WithLookup(directConfig)
if err := totp.Enable2FA(userID, "SECRET123", []string{"code1", "code2"}); err != nil { if err := totp.Enable2FA(userID, "SECRET123", []string{"code1", "code2"}); err != nil {
t.Fatalf("Enable2FA() error = %v", err) t.Fatalf("Enable2FA() error = %v", err)
@@ -248,7 +252,7 @@ func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) {
func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) { func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) {
db := newDirectTestDB(t) db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
ctx := context.Background() ctx := context.Background()
regResp, err := auth.Register(ctx, RegisterRequest{Username: "frank", Password: "p", Email: "frank@example.com"}) regResp, err := auth.Register(ctx, RegisterRequest{Username: "frank", Password: "p", Email: "frank@example.com"})
@@ -258,7 +262,7 @@ func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) {
userID := regResp.User.UserID userID := regResp.User.UserID
passkeys := NewDatabasePasskeyProvider(db, DatabasePasskeyProviderOptions{ passkeys := NewDatabasePasskeyProvider(db, DatabasePasskeyProviderOptions{
RPID: "example.com", RPName: "Example", RPOrigin: "https://example.com", QueryMode: ModeDirect, RPID: "example.com", RPName: "Example", RPOrigin: "https://example.com", Lookup: directConfig,
}) })
cred, err := passkeys.CompleteRegistration(ctx, userID, PasskeyRegistrationResponse{ cred, err := passkeys.CompleteRegistration(ctx, userID, PasskeyRegistrationResponse{
@@ -317,7 +321,7 @@ func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) {
func TestDirectMode_OAuthGetOrCreateUserAndSession(t *testing.T) { func TestDirectMode_OAuthGetOrCreateUserAndSession(t *testing.T) {
db := newDirectTestDB(t) db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
ctx := context.Background() ctx := context.Background()
userCtx := &UserContext{UserName: "gina", Email: "gina@example.com", Roles: []string{"user"}} userCtx := &UserContext{UserName: "gina", Email: "gina@example.com", Roles: []string{"user"}}
@@ -341,7 +345,7 @@ func TestDirectMode_OAuthGetOrCreateUserAndSession(t *testing.T) {
func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) { func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) {
db := newDirectTestDB(t) db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
ctx := context.Background() ctx := context.Background()
regResp, err := auth.Register(ctx, RegisterRequest{Username: "henry", Password: "p", Email: "henry@example.com"}) regResp, err := auth.Register(ctx, RegisterRequest{Username: "henry", Password: "p", Email: "henry@example.com"})
@@ -349,7 +353,7 @@ func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) {
t.Fatalf("Register() error = %v", err) t.Fatalf("Register() error = %v", err)
} }
ks := NewDatabaseKeyStore(db, DatabaseKeyStoreOptions{QueryMode: ModeDirect}) ks := NewDatabaseKeyStore(db, DatabaseKeyStoreOptions{Lookup: directConfig})
createResp, err := ks.CreateKey(ctx, CreateKeyRequest{ createResp, err := ks.CreateKey(ctx, CreateKeyRequest{
UserID: regResp.User.UserID, UserID: regResp.User.UserID,
@@ -399,7 +403,7 @@ func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) {
func TestDirectMode_OAuthServerClientAndCode(t *testing.T) { func TestDirectMode_OAuthServerClientAndCode(t *testing.T) {
db := newDirectTestDB(t) db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect}) auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
ctx := context.Background() ctx := context.Background()
client := &OAuthServerClient{ client := &OAuthServerClient{
@@ -489,7 +493,7 @@ func TestDirectMode_OAuthServerClientAndCode(t *testing.T) {
func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) { func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) {
for _, enabled := range []bool{false, true} { for _, enabled := range []bool{false, true} {
db := newDirectTestDB(t) db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect, UpgradePasswordHash: enabled}) auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig, UpgradePasswordHash: enabled})
ctx := context.Background() ctx := context.Background()
if _, err := db.Exec(`DELETE FROM users`); err != nil { if _, err := db.Exec(`DELETE FROM users`); err != nil {
@@ -518,18 +522,6 @@ func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) {
} }
} }
func TestVerifyPasswordEdgeCases(t *testing.T) { func isBcryptHash(s string) bool {
h, _ := hashPassword("pw") return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$")
if ok, _ := verifyPassword(h, "pw"); !ok {
t.Error("bcrypt match failed")
}
if ok, _ := verifyPassword("", "pw"); ok {
t.Error("empty stored must not match")
}
if ok, _ := verifyPassword("pw", ""); ok {
t.Error("empty supplied must not match")
}
if _, err := hashPassword(string(make([]byte, 73))); err == nil {
t.Error("73-byte password must be rejected")
}
} }
+16
View File
@@ -147,6 +147,9 @@ func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error
secCtx.SetQuery(q.Where(whereClause, whereArgs...)) secCtx.SetQuery(q.Where(whereClause, whereArgs...))
case common.DeleteQuery: case common.DeleteQuery:
secCtx.SetQuery(q.Where(whereClause, whereArgs...)) secCtx.SetQuery(q.Where(whereClause, whereArgs...))
case common.InsertQuery:
// Inserts read no existing rows, so there is nothing to filter.
logger.Debug("Row security filter not applicable to insert on %s.%s", schema, tablename)
default: default:
return fmt.Errorf("row security: query type %T on %s.%s does not support Where", secCtx.GetQuery(), schema, tablename) return fmt.Errorf("row security: query type %T on %s.%s does not support Where", secCtx.GetQuery(), schema, tablename)
} }
@@ -438,6 +441,19 @@ func resolveModelRules(secCtx SecurityContext) (modelregistry.ModelRules, bool)
return rules, true return rules, true
} }
// CheckModelCreateAllowed returns an error if CanCreate is false for the model. Rules are read
// from context with a fallback to the model registry; an unregistered model is allowed.
func CheckModelCreateAllowed(secCtx SecurityContext) error {
rules, ok := resolveModelRules(secCtx)
if !ok {
return nil // model not registered, allow by default
}
if !rules.CanCreate {
return fmt.Errorf("create not allowed for %s", secCtx.GetEntity())
}
return nil
}
// CheckModelUpdateAllowed is the public wrapper for checkModelUpdateAllowed. // CheckModelUpdateAllowed is the public wrapper for checkModelUpdateAllowed.
func CheckModelUpdateAllowed(secCtx SecurityContext) error { func CheckModelUpdateAllowed(secCtx SecurityContext) error {
return checkModelUpdateAllowed(secCtx) return checkModelUpdateAllowed(secCtx)
+7 -75
View File
@@ -5,81 +5,6 @@ import (
"net/http" "net/http"
) )
// UserContext holds authenticated user information
type UserContext struct {
UserID int `json:"user_id"`
UserName string `json:"user_name"`
UserLevel int `json:"user_level"`
SessionID string `json:"session_id"`
SessionRID int64 `json:"session_rid"`
RemoteID string `json:"remote_id"`
Roles []string `json:"roles"`
Email string `json:"email"`
Claims map[string]any `json:"claims"`
Meta map[string]any `json:"meta"` // Additional metadata that can hold any JSON-serializable values
TwoFactorEnabled bool `json:"two_factor_enabled"` // Indicates if 2FA is enabled for this user
ProgramUserID int `json:"program_user_id"`
ProgramUserTable string `json:"program_user_table"`
}
// LoginRequest contains credentials for login
type LoginRequest struct {
Username string `json:"username"`
Password string `json:"password"`
TwoFactorCode string `json:"two_factor_code,omitempty"` // TOTP or backup code
Claims map[string]any `json:"claims"` // Additional login data
Meta map[string]any `json:"meta"` // Additional metadata to be set on user context
}
// RegisterRequest contains information for new user registration
type RegisterRequest struct {
Username string `json:"username"`
Password string `json:"password"`
Email string `json:"email"`
UserLevel int `json:"user_level"`
Roles []string `json:"roles"`
Claims map[string]any `json:"claims"` // Additional registration data
Meta map[string]any `json:"meta"` // Additional metadata
}
// LoginResponse contains the result of a login attempt
type LoginResponse struct {
Token string `json:"token"`
RefreshToken string `json:"refresh_token"`
User *UserContext `json:"user"`
ExpiresIn int64 `json:"expires_in"` // Token expiration in seconds
Requires2FA bool `json:"requires_2fa"` // True if 2FA code is required
TwoFactorSetupData *TwoFactorSecret `json:"two_factor_setup,omitempty"` // Present when setting up 2FA
Meta map[string]any `json:"meta"` // Additional metadata to be set on user context
}
// LogoutRequest contains information for logout
type LogoutRequest struct {
Token string `json:"token"`
UserID int `json:"user_id"`
}
// PasswordResetRequest initiates a password reset for a user
type PasswordResetRequest struct {
Email string `json:"email,omitempty"`
Username string `json:"username,omitempty"`
}
// PasswordResetResponse is returned when a reset is initiated
type PasswordResetResponse struct {
// Token is the reset token to be delivered out-of-band (e.g. email).
// The stored procedure may return it for delivery or leave it empty
// if the delivery is handled entirely in the database.
Token string `json:"token"`
ExpiresIn int64 `json:"expires_in"` // seconds
}
// PasswordResetCompleteRequest completes a password reset using the token
type PasswordResetCompleteRequest struct {
Token string `json:"token"`
NewPassword string `json:"new_password"`
}
// Authenticator handles user authentication operations // Authenticator handles user authentication operations
type Authenticator interface { type Authenticator interface {
// Login authenticates credentials and returns a token // Login authenticates credentials and returns a token
@@ -144,6 +69,13 @@ type Refreshable interface {
RefreshToken(ctx context.Context, refreshToken string) (*LoginResponse, error) RefreshToken(ctx context.Context, refreshToken string) (*LoginResponse, error)
} }
// APIKeyLoginable allows providers to exchange a raw API key for a session.
type APIKeyLoginable interface {
// LoginWithAPIKey validates the raw API key and creates a session for its user.
// Unknown, expired and inactive keys all yield the same generic error.
LoginWithAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*LoginResponse, error)
}
// Validatable allows providers to validate tokens without full authentication // Validatable allows providers to validate tokens without full authentication
type Validatable interface { type Validatable interface {
// ValidateToken checks if a token is valid without extracting full user context // ValidateToken checks if a token is valid without extracting full user context
+4 -56
View File
@@ -2,64 +2,12 @@ package security
import ( import (
"context" "context"
"crypto/sha256"
"encoding/hex" "github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
"time"
) )
// hashSHA256Hex returns the lowercase hex SHA-256 digest of the given string. // hashSHA256Hex is kept as a short alias for sectypes.HashKey inside this package.
// Used by all keystore implementations to hash raw keys before storage or lookup. func hashSHA256Hex(raw string) string { return sectypes.HashKey(raw) }
func hashSHA256Hex(raw string) string {
sum := sha256.Sum256([]byte(raw))
return hex.EncodeToString(sum[:])
}
// KeyType identifies the category of an auth key.
type KeyType string
const (
// KeyTypeJWTSecret is a per-user JWT signing secret for token generation.
KeyTypeJWTSecret KeyType = "jwt_secret"
// KeyTypeHeaderAPI is a static API key sent via a request header.
KeyTypeHeaderAPI KeyType = "header_api"
// KeyTypeOAuth2 holds OAuth2 client credentials (client_id / client_secret).
KeyTypeOAuth2 KeyType = "oauth2"
// KeyTypeGenericAPI is a generic application API key.
KeyTypeGenericAPI KeyType = "api"
)
// UserKey represents a single named auth key belonging to a user.
// KeyHash stores the SHA-256 hex digest of the raw key; the raw key is never persisted.
type UserKey struct {
ID int64 `json:"id"`
UserID int `json:"user_id"`
KeyType KeyType `json:"key_type"`
KeyHash string `json:"key_hash"` // SHA-256 hex; never the raw key
Name string `json:"name"`
Scopes []string `json:"scopes,omitempty"`
Meta map[string]any `json:"meta,omitempty"`
ExpiresAt *time.Time `json:"expires_at,omitempty"`
CreatedAt time.Time `json:"created_at"`
LastUsedAt *time.Time `json:"last_used_at,omitempty"`
IsActive bool `json:"is_active"`
}
// CreateKeyRequest specifies the parameters for a new key.
type CreateKeyRequest struct {
UserID int
KeyType KeyType
Name string
Scopes []string
Meta map[string]any
ExpiresAt *time.Time
}
// CreateKeyResponse is returned exactly once when a key is created.
// The caller is responsible for persisting RawKey; it is not stored anywhere.
type CreateKeyResponse struct {
Key UserKey
RawKey string // crypto/rand 32 bytes, base64url-encoded
}
// KeyStore manages per-user auth keys with pluggable storage backends. // KeyStore manages per-user auth keys with pluggable storage backends.
// Implementations: ConfigKeyStore (static list) and DatabaseKeyStore (stored procedures). // Implementations: ConfigKeyStore (static list) and DatabaseKeyStore (stored procedures).
+33 -193
View File
@@ -5,16 +5,16 @@ import (
"crypto/rand" "crypto/rand"
"database/sql" "database/sql"
"encoding/base64" "encoding/base64"
"encoding/json"
"errors" "errors"
"fmt" "fmt"
"sync"
"time" "time"
"golang.org/x/sync/singleflight" "golang.org/x/sync/singleflight"
"github.com/bitechdev/ResolveSpec/pkg/cache" "github.com/bitechdev/ResolveSpec/pkg/cache"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace" "github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/backends"
) )
// DatabaseKeyStoreOptions configures DatabaseKeyStore. // DatabaseKeyStoreOptions configures DatabaseKeyStore.
@@ -24,36 +24,28 @@ type DatabaseKeyStoreOptions struct {
// CacheTTL is the duration to cache ValidateKey results. // CacheTTL is the duration to cache ValidateKey results.
// Default: 2 minutes. // Default: 2 minutes.
CacheTTL time.Duration CacheTTL time.Duration
// SQLNames provides custom procedure names. If nil, uses DefaultKeyStoreSQLNames(). // Lookup selects dialect, query mode and procedure/table/column names.
SQLNames *KeyStoreSQLNames // The zero value uses stored procedures on Postgres and direct SQL elsewhere.
// TableNames provides custom table names for Direct mode. If nil, uses DefaultKeyStoreTableNames(). Lookup lookup.Config
TableNames *KeyStoreTableNames // LookupProvider, when set, is used instead of building one from Lookup and the db.
// QueryMode selects stored-procedure vs Direct-mode SQL. Default (zero value) is ModeAuto. LookupProvider *lookup.Provider
QueryMode QueryMode
// DBFactory is called to obtain a fresh *sql.DB when the existing connection is closed. // DBFactory is called to obtain a fresh *sql.DB when the existing connection is closed.
// If nil, reconnection is disabled. // If nil, reconnection is disabled.
DBFactory func() (*sql.DB, error) DBFactory func() (*sql.DB, error)
} }
// DatabaseKeyStore is a KeyStore backed by PostgreSQL stored procedures. // DatabaseKeyStore is a KeyStore backed by the lookup package (stored procedures on
// All DB operations go through configurable procedure names; the raw key is // Postgres by default, direct SQL elsewhere). The raw key is never passed to the database.
// never passed to the database.
// //
// See keystore_schema.sql for the required table and procedure definitions. // See lookup/keystore_schema.sql for the required table and procedure definitions.
// //
// Note: DeleteKey invalidates the cache entry for the deleted key. Due to the // Note: DeleteKey invalidates the cache entry for the deleted key. Due to the
// cache TTL, a deleted key may continue to authenticate for up to CacheTTL // cache TTL, a deleted key may continue to authenticate for up to CacheTTL
// (default 2 minutes) if the cache entry cannot be invalidated. // (default 2 minutes) if the cache entry cannot be invalidated.
type DatabaseKeyStore struct { type DatabaseKeyStore struct {
db *sql.DB src *lookupSource
dbMu sync.RWMutex cache *cache.Cache
dbFactory func() (*sql.DB, error) cacheTTL time.Duration
sqlNames *KeyStoreSQLNames
tableNames *KeyStoreTableNames
queryMode QueryMode
capability *dbCapability
cache *cache.Cache
cacheTTL time.Duration
// validateLoads collapses concurrent key lookups for the same key // validateLoads collapses concurrent key lookups for the same key
validateLoads singleflight.Group validateLoads singleflight.Group
@@ -72,42 +64,14 @@ func NewDatabaseKeyStore(db *sql.DB, opts ...DatabaseKeyStoreOptions) *DatabaseK
if c == nil { if c == nil {
c = cache.GetDefaultCache() c = cache.GetDefaultCache()
} }
names := MergeKeyStoreSQLNames(DefaultKeyStoreSQLNames(), o.SQLNames) src := newLookupSource(db)
tableNames := resolveKeyStoreTableNames(o.TableNames) src.cfg = o.Lookup
return &DatabaseKeyStore{ src.provider = o.LookupProvider
db: db, src.opts = backends.Options{DBFactory: o.DBFactory}
dbFactory: o.DBFactory, return &DatabaseKeyStore{src: src, cache: c, cacheTTL: o.CacheTTL}
sqlNames: names,
tableNames: tableNames,
queryMode: o.QueryMode,
capability: newDBCapability(),
cache: c,
cacheTTL: o.CacheTTL,
}
} }
func (ks *DatabaseKeyStore) getDB() *sql.DB { func (ks *DatabaseKeyStore) keys() lookup.KeyStore { return ks.src.get().Keys }
ks.dbMu.RLock()
defer ks.dbMu.RUnlock()
return ks.db
}
func (ks *DatabaseKeyStore) reconnectDB() error {
if ks.dbFactory == nil {
return fmt.Errorf("no db factory configured for reconnect")
}
newDB, err := ks.dbFactory()
if err != nil {
return err
}
ks.dbMu.Lock()
ks.db = newDB
ks.dbMu.Unlock()
if ks.capability != nil {
ks.capability.reset()
}
return nil
}
// CreateKey generates a raw key, stores its SHA-256 hash via the create procedure, // CreateKey generates a raw key, stores its SHA-256 hash via the create procedure,
// and returns the raw key once. // and returns the raw key once.
@@ -119,110 +83,29 @@ func (ks *DatabaseKeyStore) CreateKey(ctx context.Context, req CreateKeyRequest)
rawKey := base64.RawURLEncoding.EncodeToString(rawBytes) rawKey := base64.RawURLEncoding.EncodeToString(rawBytes)
hash := hashSHA256Hex(rawKey) hash := hashSHA256Hex(rawKey)
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.CreateKey) { key, err := ks.keys().Create(ctx, req, hash)
key, err := ks.createKeyDirect(ctx, req, hash)
if err != nil {
return nil, err
}
return &CreateKeyResponse{Key: *key, RawKey: rawKey}, nil
}
type createRequest struct {
UserID int `json:"user_id"`
KeyType KeyType `json:"key_type"`
KeyHash string `json:"key_hash"`
Name string `json:"name"`
Scopes []string `json:"scopes,omitempty"`
Meta map[string]any `json:"meta,omitempty"`
ExpiresAt *time.Time `json:"expires_at,omitempty"`
}
reqJSON, err := json.Marshal(createRequest{
UserID: req.UserID,
KeyType: req.KeyType,
KeyHash: hash,
Name: req.Name,
Scopes: req.Scopes,
Meta: req.Meta,
ExpiresAt: req.ExpiresAt,
})
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to marshal create key request: %w", err) return nil, err
} }
return &CreateKeyResponse{Key: *key, RawKey: rawKey}, nil
var success bool
var errorMsg sql.NullString
var keyJSON sql.NullString
query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1::jsonb)`, ks.sqlNames.CreateKey)
if err = ks.getDB().QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &keyJSON); err != nil {
return nil, fmt.Errorf("create key procedure failed: %w", err)
}
if !success {
return nil, errors.New(nullStringOr(errorMsg, "create key failed"))
}
var key UserKey
if err = json.Unmarshal([]byte(keyJSON.String), &key); err != nil {
return nil, fmt.Errorf("failed to parse created key: %w", err)
}
return &CreateKeyResponse{Key: key, RawKey: rawKey}, nil
} }
// GetUserKeys returns all active, non-expired keys for the given user. // GetUserKeys returns all active, non-expired keys for the given user.
// Pass an empty KeyType to return all types. // Pass an empty KeyType to return all types.
func (ks *DatabaseKeyStore) GetUserKeys(ctx context.Context, userID int, keyType KeyType) ([]UserKey, error) { func (ks *DatabaseKeyStore) GetUserKeys(ctx context.Context, userID int, keyType KeyType) ([]UserKey, error) {
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.GetUserKeys) { return ks.keys().List(ctx, userID, keyType)
return ks.getUserKeysDirect(ctx, userID, keyType)
}
var success bool
var errorMsg sql.NullString
var keysJSON sql.NullString
query := fmt.Sprintf(`SELECT p_success, p_error, p_keys::text FROM %s($1, $2)`, ks.sqlNames.GetUserKeys)
if err := ks.getDB().QueryRowContext(ctx, query, userID, string(keyType)).Scan(&success, &errorMsg, &keysJSON); err != nil {
return nil, fmt.Errorf("get user keys procedure failed: %w", err)
}
if !success {
return nil, errors.New(nullStringOr(errorMsg, "get user keys failed"))
}
var keys []UserKey
if keysJSON.Valid && keysJSON.String != "" && keysJSON.String != "[]" {
if err := json.Unmarshal([]byte(keysJSON.String), &keys); err != nil {
return nil, fmt.Errorf("failed to parse user keys: %w", err)
}
}
if keys == nil {
keys = []UserKey{}
}
return keys, nil
} }
// DeleteKey soft-deletes a key after verifying ownership and invalidates its cache entry. // DeleteKey soft-deletes a key after verifying ownership and invalidates its cache entry.
// The delete procedure returns the key_hash so no separate lookup is needed. // The delete procedure returns the key_hash so no separate lookup is needed.
// Note: cache invalidation is best-effort; a cached entry may persist for up to CacheTTL. // Note: cache invalidation is best-effort; a cached entry may persist for up to CacheTTL.
func (ks *DatabaseKeyStore) DeleteKey(ctx context.Context, userID int, keyID int64) error { func (ks *DatabaseKeyStore) DeleteKey(ctx context.Context, userID int, keyID int64) error {
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.DeleteKey) { keyHash, err := ks.keys().Delete(ctx, userID, keyID)
return ks.deleteKeyDirect(ctx, userID, keyID) if err != nil {
return err
} }
if keyHash != "" && ks.cache != nil {
var success bool _ = ks.cache.Delete(ctx, keystoreCacheKey(keyHash))
var errorMsg sql.NullString
var keyHash sql.NullString
query := fmt.Sprintf(`SELECT p_success, p_error, p_key_hash FROM %s($1, $2)`, ks.sqlNames.DeleteKey)
if err := ks.getDB().QueryRowContext(ctx, query, userID, keyID).Scan(&success, &errorMsg, &keyHash); err != nil {
return fmt.Errorf("delete key procedure failed: %w", err)
}
if !success {
return errors.New(nullStringOr(errorMsg, "delete key failed"))
}
if keyHash.Valid && keyHash.String != "" && ks.cache != nil {
_ = ks.cache.Delete(ctx, keystoreCacheKey(keyHash.String))
} }
return nil return nil
} }
@@ -261,61 +144,18 @@ func (ks *DatabaseKeyStore) ValidateKey(ctx context.Context, rawKey string, keyT
// validateKeyLoad validates against the database and fills the cache. // validateKeyLoad validates against the database and fills the cache.
func (ks *DatabaseKeyStore) validateKeyLoad(ctx context.Context, hash, cacheKey string, keyType KeyType) (*UserKey, error) { func (ks *DatabaseKeyStore) validateKeyLoad(ctx context.Context, hash, cacheKey string, keyType KeyType) (*UserKey, error) {
dbtrace.Raw(ctx, "keystore.validate") dbtrace.Raw(ctx, "keystore.validate")
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.ValidateKey) { key, err := ks.keys().Validate(ctx, hash, keyType)
key, err := ks.validateKeyDirect(ctx, hash, keyType) if err != nil {
if err != nil { return nil, err
return nil, err
}
if ks.cache != nil {
_ = ks.cache.Set(ctx, cacheKey, *key, ks.cacheTTL)
}
return key, nil
}
var success bool
var errorMsg sql.NullString
var keyJSON sql.NullString
runQuery := func() error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1, $2)`, ks.sqlNames.ValidateKey)
return ks.getDB().QueryRowContext(ctx, query, hash, string(keyType)).Scan(&success, &errorMsg, &keyJSON)
}
if err := runQuery(); err != nil {
if isDBClosed(err) {
if reconnErr := ks.reconnectDB(); reconnErr == nil {
err = runQuery()
}
if err != nil {
return nil, fmt.Errorf("validate key procedure failed: %w", err)
}
} else {
return nil, fmt.Errorf("validate key procedure failed: %w", err)
}
}
if !success {
return nil, errors.New(nullStringOr(errorMsg, "invalid or expired key"))
}
var key UserKey
if err := json.Unmarshal([]byte(keyJSON.String), &key); err != nil {
return nil, fmt.Errorf("failed to parse validated key: %w", err)
} }
if ks.cache != nil { if ks.cache != nil {
_ = ks.cache.Set(ctx, cacheKey, key, ks.cacheTTL) _ = ks.cache.Set(ctx, cacheKey, *key, ks.cacheTTL)
} }
return &key, nil return key, nil
} }
func keystoreCacheKey(hash string) string { func keystoreCacheKey(hash string) string {
return "keystore:validate:" + hash return "keystore:validate:" + hash
} }
// nullStringOr returns s.String if valid, otherwise the fallback.
func nullStringOr(s sql.NullString, fallback string) string {
if s.Valid && s.String != "" {
return s.String
}
return fallback
}

Some files were not shown because too many files have changed in this diff Show More