Compare commits

...
23 Commits
Author SHA1 Message Date
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
201 changed files with 25109 additions and 7881 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
+21
View File
@@ -275,6 +275,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 +537,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
+25
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
@@ -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)
}
})
}
+44 -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)
@@ -762,6 +795,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
@@ -892,23 +928,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")
}
}
+18
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
@@ -309,3 +319,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
}
+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
+122 -134
View File
@@ -1,46 +1,53 @@
# 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",
}) })
// 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 +70,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 +112,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 +123,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 +145,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 +153,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 +161,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 +179,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 +217,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 +316,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 +327,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:
@@ -356,121 +394,58 @@ handler.SetModelRules("public", "users", modelregistry.ModelRules{
## MCP Tools ## MCP Tools
### Tool Naming 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 |
{operation}_{schema}_{entity} // e.g. read_public_users |---|---|
{operation}_{entity} // e.g. read_users (when schema is empty) | `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` |
Operations: `read`, `create`, `update`, `delete`. ### `select_table`
### Read Tool — `read_{schema}_{entity}`
Fetch one or many records.
| Argument | Type | Description | | Argument | Type | Description |
|---|---|---| |---|---|---|
| `id` | string | Primary key value. Omit to return multiple records. | | `table` | string (required) | `schema.entity` |
| `limit` | number | Max records per page (recommended: 10–100). | | `id` | string | Primary key of one row |
| `offset` | number | Records to skip (offset-based pagination). | | `filters`, `sort` | array | See [Filtering](#filtering), [Sorting](#sorting) |
| `cursor_forward` | string | PK of the **last** record on the current page (next-page cursor). | | `columns`, `omit_columns` | array | Column selection |
| `cursor_backward` | string | PK of the **first** record on the current page (prev-page cursor). | | `preloads` | array | Relations (validated against the model, max depth `MaxPreloadDepth`) |
| `columns` | array | Column names to include. Omit for all columns. | | `limit`, `offset` | number | Clamped to `MaxLimit` / rejected above `MaxOffset` |
| `omit_columns` | array | Column names to exclude. | | `cursor_forward`, `cursor_backward` | string | PK cursor, requires `sort` |
| `filters` | array | Filter objects (see [Filtering](#filtering)). | | `include_count` | boolean | Also compute totals (slower); otherwise `total`/`filtered` are 0 |
| `sort` | array | Sort objects (see [Sorting](#sorting)). |
| `preloads` | array | Relation preload objects (see [Preloading](#preloading)). |
**Response:** Response: `{"success":true,"data":[...],"metadata":{"total","filtered","count","limit","offset"}}`
```json
{
"success": true,
"data": [...],
"metadata": {
"total": 100,
"filtered": 100,
"count": 10,
"limit": 10,
"offset": 0
}
}
```
### Create Tool — `create_{schema}_{entity}` ### `insert_into_table`
Insert one or more records. `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.
| Argument | Type | Description | ### `update_table` / `delete_from_table`
|---|---|---|
| `data` | object \| array | Single object or array of objects to insert. |
Array input runs inside a single transaction — all succeed or all fail. Either `id` or `filters` is required.
**Response:** | Mode | Behaviour |
```json |---|---|
{ "success": true, "data": { ... } } | `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. |
### Update Tool — `update_{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.
Partially update an existing record. Only non-null, non-empty fields in `data` are applied; existing values are preserved for omitted fields. ### `call_function`
| Argument | Type | Description | `name` and `arguments` (object). See [Functions](#functions).
|---|---|---|
| `id` | string | Primary key of the record. Can also be included inside `data`. |
| `data` | object (required) | Fields to update. |
**Response:** ### `resolvespec_annotate`
```json
{ "success": true, "data": { ...merged record... } }
```
### Delete Tool — `delete_{schema}_{entity}` 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.
Delete a record by primary key. **Irreversible.**
| Argument | Type | Description |
|---|---|---|
| `id` | string (required) | Primary key of the record to delete. |
**Response:**
```json
{ "success": true, "data": { ...deleted record... } }
```
### Annotation Tool — `resolvespec_annotate`
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 +545,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 +577,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 +611,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 +631,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 +647,13 @@ 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
- 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
} }
+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)
}
+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))
}
+280
View File
@@ -0,0 +1,280 @@
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
}
// 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 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 {
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")
}
}
+222 -95
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"
@@ -30,6 +33,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.
@@ -39,11 +44,15 @@ func NewHandler(db common.Database, registry common.ModelRegistry, cfg Config) *
registry: registry, registry: registry,
hooks: NewHookRegistry(), hooks: NewHookRegistry(),
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"), mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"),
config: cfg, config: cfg.withDefaults(),
confirms: newConfirmStore(),
name: "resolvemcp", name: "resolvemcp",
version: "1.0.0", version: "1.0.0",
} }
registerAnnotationTool(h) registerMetaTools(h)
if cfg.EnableAnnotations {
registerAnnotationTool(h)
}
return h return h
} }
@@ -97,7 +106,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 +128,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 +142,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 +169,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 +191,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 +241,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 +281,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 +308,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 +318,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 +369,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 +382,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 +401,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 +425,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 +434,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 +483,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 +500,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 +522,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 +542,25 @@ 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")
}
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 +576,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 +607,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 +624,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 +638,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 +651,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 +673,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 +722,57 @@ 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 q := tx.NewUpdate().Table(tableName).SetMap(setCols).
for key, newValue := range updates {
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 +782,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 +812,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 +831,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 +976,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")
}
}
+438
View File
@@ -0,0 +1,438 @@
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)
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)
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()
}
func opsFor(r modelregistry.ModelRules) []string {
var ops []string
if r.CanRead {
ops = append(ops, opSelect)
}
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 != "" {
allowed := false
for _, o := range 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"`
Operations []string `json:"operations"`
}
var tables []table
for name := range h.registry.GetAllModels() {
schema, entity, _ := splitTable(name)
if ops := opsFor(h.modelRules(schema, entity)); len(ops) > 0 {
tables = append(tables, table{Table: name, 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(opsFor(rules)) == 0 {
return toolError("describe_table", invalidArg("unknown table %q; see list_tables", buildModelName(schema, entity))), nil
}
info := buildModelInfo(schema, entity, model)
modelType := reflect.TypeOf(model)
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice) {
modelType = modelType.Elem()
}
writable := map[string]bool{}
if 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"`
}
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})
if w && !c.isPrimary {
writableNames = append(writableNames, c.jsonName)
}
}
return marshalResult(map[string]any{
"success": true,
"table": info.fullName,
"primary_key": info.pkName,
"columns": cols,
"relations": info.relationNames,
"writable_columns": writableNames,
"operations": opsFor(rules),
"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))
}
+138 -36
View File
@@ -12,12 +12,13 @@
// handler.RegisterModel("public", "users", &User{}) // handler.RegisterModel("public", "users", &User{})
// //
// r := mux.NewRouter() // r := mux.NewRouter()
// resolvemcp.SetupMuxRoutes(r, handler) // resolvemcp.SetupMuxRoutes(r, handler, securityList) // requires an authenticated caller
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 +29,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 +41,61 @@ 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
// 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
}
// 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
}
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 +114,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 +146,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 +179,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})
if on.mcpServer.GetTool(annotationToolName) == nil {
t.Fatal("annotation tool missing when enabled")
}
}
+2 -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,9 @@ 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/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"
@@ -227,350 +201,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{}
+35 -28
View File
@@ -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
}
+259
View File
@@ -0,0 +1,259 @@
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" {
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"
}
+21 -1
View File
@@ -1244,6 +1244,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.
@@ -1337,6 +1347,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 +1380,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) {
+135 -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.
@@ -644,77 +646,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 != "" && common.Hardening().SQLStrict {
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)
} }
@@ -1539,6 +1556,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 +1625,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)
@@ -1651,6 +1686,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 +1716,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 +1760,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 +1769,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))
})
}
}
+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
}
-216
View File
@@ -1,216 +0,0 @@
package security
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"time"
)
// Direct-mode implementations mirroring the resolvespec_keystore_* stored
// procedures in keystore_schema.sql using plain SQL against
// TableNames.UserKeys. meta/scopes are stored as JSON-encoded TEXT instead
// of Postgres JSONB.
func (ks *DatabaseKeyStore) createKeyDirect(ctx context.Context, req CreateKeyRequest, keyHash string) (*UserKey, error) {
scopesJSON, err := json.Marshal(req.Scopes)
if err != nil {
return nil, fmt.Errorf("failed to marshal scopes: %w", err)
}
var metaJSON []byte
if req.Meta != nil {
metaJSON, err = json.Marshal(req.Meta)
if err != nil {
return nil, fmt.Errorf("failed to marshal meta: %w", err)
}
}
now := time.Now()
var id int64
err = ks.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (user_id, key_type, key_hash, name, scopes, meta, expires_at, created_at, is_active) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
ks.tableNames.UserKeys))
res, err := db.ExecContext(ctx, query, req.UserID, string(req.KeyType), keyHash, req.Name, string(scopesJSON), nullableString(metaJSON), req.ExpiresAt, now, true)
if err != nil {
return err
}
id, err = res.LastInsertId()
return err
})
if err != nil {
return nil, fmt.Errorf("create key query failed: %w", err)
}
return &UserKey{
ID: id,
UserID: req.UserID,
KeyType: req.KeyType,
KeyHash: keyHash,
Name: req.Name,
Scopes: req.Scopes,
Meta: req.Meta,
ExpiresAt: req.ExpiresAt,
CreatedAt: now,
IsActive: true,
}, nil
}
func (ks *DatabaseKeyStore) runDBOpWithReconnect(run func(*sql.DB) error) error {
db := ks.getDB()
if db == nil {
return fmt.Errorf("database connection is nil")
}
err := run(db)
if isDBClosed(err) {
if reconnErr := ks.reconnectDB(); reconnErr == nil {
err = run(ks.getDB())
}
}
return err
}
func (ks *DatabaseKeyStore) getUserKeysDirect(ctx context.Context, userID int, keyType KeyType) ([]UserKey, error) {
keys := []UserKey{}
err := ks.runDBOpWithReconnect(func(db *sql.DB) error {
var query string
var args []any
if keyType == "" {
query = rewritePlaceholders(db, fmt.Sprintf(
`SELECT id, user_id, key_type, name, scopes, meta, expires_at, created_at, last_used_at, is_active
FROM %s WHERE user_id = ? AND is_active = ? AND (expires_at IS NULL OR expires_at > ?)`, ks.tableNames.UserKeys))
args = []any{userID, true, time.Now()}
} else {
query = rewritePlaceholders(db, fmt.Sprintf(
`SELECT id, user_id, key_type, name, scopes, meta, expires_at, created_at, last_used_at, is_active
FROM %s WHERE user_id = ? AND is_active = ? AND (expires_at IS NULL OR expires_at > ?) AND key_type = ?`, ks.tableNames.UserKeys))
args = []any{userID, true, time.Now(), string(keyType)}
}
rows, err := db.QueryContext(ctx, query, args...)
if err != nil {
return err
}
defer rows.Close()
for rows.Next() {
var k UserKey
var kt string
var scopesJSON, metaJSON sql.NullString
var expiresAt, lastUsedAt sql.NullTime
if err := rows.Scan(&k.ID, &k.UserID, &kt, &k.Name, &scopesJSON, &metaJSON, &expiresAt, &k.CreatedAt, &lastUsedAt, &k.IsActive); err != nil {
return err
}
k.KeyType = KeyType(kt)
if scopesJSON.Valid && scopesJSON.String != "" {
_ = json.Unmarshal([]byte(scopesJSON.String), &k.Scopes)
}
if metaJSON.Valid && metaJSON.String != "" {
_ = json.Unmarshal([]byte(metaJSON.String), &k.Meta)
}
if expiresAt.Valid {
t := expiresAt.Time
k.ExpiresAt = &t
}
if lastUsedAt.Valid {
t := lastUsedAt.Time
k.LastUsedAt = &t
}
keys = append(keys, k)
}
return rows.Err()
})
if err != nil {
return nil, fmt.Errorf("get user keys query failed: %w", err)
}
return keys, nil
}
func (ks *DatabaseKeyStore) deleteKeyDirect(ctx context.Context, userID int, keyID int64) error {
var keyHash string
err := ks.runDBOpWithReconnect(func(db *sql.DB) error {
selQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT key_hash FROM %s WHERE id = ? AND user_id = ? AND is_active = ?`, ks.tableNames.UserKeys))
if err := db.QueryRowContext(ctx, selQuery, keyID, userID, true).Scan(&keyHash); err != nil {
return err
}
updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET is_active = ? WHERE id = ? AND user_id = ? AND is_active = ?`, ks.tableNames.UserKeys))
_, err := db.ExecContext(ctx, updQuery, false, keyID, userID, true)
return err
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return errors.New("key not found or already deleted")
}
return fmt.Errorf("delete key query failed: %w", err)
}
if keyHash != "" && ks.cache != nil {
_ = ks.cache.Delete(ctx, keystoreCacheKey(keyHash))
}
return nil
}
func (ks *DatabaseKeyStore) validateKeyDirect(ctx context.Context, keyHash string, keyType KeyType) (*UserKey, error) {
var k UserKey
var kt string
var scopesJSON, metaJSON sql.NullString
var expiresAt, lastUsedAt sql.NullTime
err := ks.runDBOpWithReconnect(func(db *sql.DB) error {
var query string
var args []any
if keyType == "" {
query = rewritePlaceholders(db, fmt.Sprintf(
`SELECT id, user_id, key_type, name, scopes, meta, expires_at, created_at, is_active
FROM %s WHERE key_hash = ? AND is_active = ? AND (expires_at IS NULL OR expires_at > ?)`, ks.tableNames.UserKeys))
args = []any{keyHash, true, time.Now()}
} else {
query = rewritePlaceholders(db, fmt.Sprintf(
`SELECT id, user_id, key_type, name, scopes, meta, expires_at, created_at, is_active
FROM %s WHERE key_hash = ? AND is_active = ? AND (expires_at IS NULL OR expires_at > ?) AND key_type = ?`, ks.tableNames.UserKeys))
args = []any{keyHash, true, time.Now(), string(keyType)}
}
if err := db.QueryRowContext(ctx, query, args...).Scan(&k.ID, &k.UserID, &kt, &k.Name, &scopesJSON, &metaJSON, &expiresAt, &k.CreatedAt, &k.IsActive); err != nil {
return err
}
now := time.Now()
updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET last_used_at = ? WHERE id = ?`, ks.tableNames.UserKeys))
_, err := db.ExecContext(ctx, updQuery, now, k.ID)
return err
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, errors.New("invalid or expired key")
}
return nil, fmt.Errorf("validate key query failed: %w", err)
}
k.KeyType = KeyType(kt)
k.KeyHash = keyHash
if scopesJSON.Valid && scopesJSON.String != "" {
_ = json.Unmarshal([]byte(scopesJSON.String), &k.Scopes)
}
if metaJSON.Valid && metaJSON.String != "" {
_ = json.Unmarshal([]byte(metaJSON.String), &k.Meta)
}
if expiresAt.Valid {
t := expiresAt.Time
k.ExpiresAt = &t
}
_ = lastUsedAt
now := time.Now()
k.LastUsedAt = &now
return &k, nil
}
func nullableString(b []byte) any {
if b == nil {
return nil
}
return string(b)
}
-61
View File
@@ -1,61 +0,0 @@
package security
import "fmt"
// KeyStoreSQLNames holds the configurable stored procedure names used by DatabaseKeyStore.
// Use DefaultKeyStoreSQLNames() for defaults and MergeKeyStoreSQLNames() for partial overrides.
type KeyStoreSQLNames struct {
GetUserKeys string // default: "resolvespec_keystore_get_user_keys"
CreateKey string // default: "resolvespec_keystore_create_key"
DeleteKey string // default: "resolvespec_keystore_delete_key"
ValidateKey string // default: "resolvespec_keystore_validate_key"
}
// DefaultKeyStoreSQLNames returns a KeyStoreSQLNames with all default resolvespec_keystore_* values.
func DefaultKeyStoreSQLNames() *KeyStoreSQLNames {
return &KeyStoreSQLNames{
GetUserKeys: "resolvespec_keystore_get_user_keys",
CreateKey: "resolvespec_keystore_create_key",
DeleteKey: "resolvespec_keystore_delete_key",
ValidateKey: "resolvespec_keystore_validate_key",
}
}
// MergeKeyStoreSQLNames returns a copy of base with any non-empty fields from override applied.
// If override is nil, a copy of base is returned.
func MergeKeyStoreSQLNames(base, override *KeyStoreSQLNames) *KeyStoreSQLNames {
if override == nil {
copied := *base
return &copied
}
merged := *base
if override.GetUserKeys != "" {
merged.GetUserKeys = override.GetUserKeys
}
if override.CreateKey != "" {
merged.CreateKey = override.CreateKey
}
if override.DeleteKey != "" {
merged.DeleteKey = override.DeleteKey
}
if override.ValidateKey != "" {
merged.ValidateKey = override.ValidateKey
}
return &merged
}
// ValidateKeyStoreSQLNames checks that all non-empty procedure names are valid SQL identifiers.
func ValidateKeyStoreSQLNames(names *KeyStoreSQLNames) error {
fields := map[string]string{
"GetUserKeys": names.GetUserKeys,
"CreateKey": names.CreateKey,
"DeleteKey": names.DeleteKey,
"ValidateKey": names.ValidateKey,
}
for field, val := range fields {
if val != "" && !validSQLIdentifier.MatchString(val) {
return fmt.Errorf("KeyStoreSQLNames.%s contains invalid characters: %q", field, val)
}
}
return nil
}
-44
View File
@@ -1,44 +0,0 @@
package security
import "fmt"
// KeyStoreTableNames holds the configurable table name used by DatabaseKeyStore
// in Direct mode. Use DefaultKeyStoreTableNames() for defaults and
// MergeKeyStoreTableNames() for partial overrides.
type KeyStoreTableNames struct {
UserKeys string // default: "user_keys"
}
// DefaultKeyStoreTableNames returns a KeyStoreTableNames with default table names.
func DefaultKeyStoreTableNames() *KeyStoreTableNames {
return &KeyStoreTableNames{
UserKeys: "user_keys",
}
}
// MergeKeyStoreTableNames returns a copy of base with any non-empty fields from override applied.
// If override is nil, a copy of base is returned.
func MergeKeyStoreTableNames(base, override *KeyStoreTableNames) *KeyStoreTableNames {
if override == nil {
copied := *base
return &copied
}
merged := *base
if override.UserKeys != "" {
merged.UserKeys = override.UserKeys
}
return &merged
}
// ValidateKeyStoreTableNames checks that all non-empty table names are valid SQL identifiers.
func ValidateKeyStoreTableNames(names *KeyStoreTableNames) error {
if names.UserKeys != "" && !validSQLIdentifier.MatchString(names.UserKeys) {
return fmt.Errorf("KeyStoreTableNames.UserKeys contains invalid characters: %q", names.UserKeys)
}
return nil
}
// resolveKeyStoreTableNames merges an optional override with defaults.
func resolveKeyStoreTableNames(override *KeyStoreTableNames) *KeyStoreTableNames {
return MergeKeyStoreTableNames(DefaultKeyStoreTableNames(), override)
}
+39
View File
@@ -0,0 +1,39 @@
package security
import (
"context"
"errors"
"testing"
)
type fakeAPIKeyAuth struct {
Authenticator
key string
}
func (f *fakeAPIKeyAuth) LoginWithAPIKey(_ context.Context, rawKey string, _ map[string]any) (*LoginResponse, error) {
if rawKey != f.key {
return nil, errInvalidAPIKey
}
return &LoginResponse{Token: "tok"}, nil
}
func TestChainLoginWithAPIKey(t *testing.T) {
ctx := context.Background()
chain := NewChainAuthenticator(&fakeAPIKeyAuth{key: "a"}, &fakeAPIKeyAuth{key: "b"})
for _, k := range []string{"a", "b"} {
if resp, err := chain.LoginWithAPIKey(ctx, k, nil); err != nil || resp.Token != "tok" {
t.Errorf("chain LoginWithAPIKey(%q) = %v, %v", k, resp, err)
}
}
if _, err := chain.LoginWithAPIKey(ctx, "bad", nil); !errors.Is(err, errInvalidAPIKey) {
t.Errorf("chain bad key error = %v, want errInvalidAPIKey", err)
}
}
func TestDatabaseAuthenticatorLoginWithAPIKey_EmptyKey(t *testing.T) {
auth := NewDatabaseAuthenticatorWithOptions(newDirectTestDB(t), DatabaseAuthenticatorOptions{})
if _, err := auth.LoginWithAPIKey(context.Background(), "", nil); !errors.Is(err, errInvalidAPIKey) {
t.Errorf("empty key error = %v, want errInvalidAPIKey", err)
}
}
+166
View File
@@ -0,0 +1,166 @@
// Package backends assembles a lookup.Provider: it builds the procedure and direct stores for
// one database and routes every operation to one of them according to lookup.Config.
// It lives apart from package lookup because both backends import lookup.
package backends
import (
"context"
"database/sql"
"fmt"
"sync"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/direct"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/procedure"
)
// Options are the settings that are not naming or mode.
type Options struct {
// DBFactory is called to obtain a fresh *sql.DB when the current one has been closed.
// Nil disables reconnecting.
DBFactory func() (*sql.DB, error)
// UpgradePasswordHash rewrites a legacy cleartext password as bcrypt after a successful
// direct-mode login. Off by default.
UpgradePasswordHash bool
// NoGroupTables skips the group membership table when loading direct-mode policy rules.
NoGroupTables bool
}
// New builds a Provider for db. cfg is merged with the defaults and validated; the dialect
// is cfg.Dialect or detected from the driver. Every operation's mode is resolved up front so
// an impossible combination (procedure mode on a non-Postgres dialect) fails here, not on the
// first request.
func New(db *sql.DB, cfg lookup.Config, opts Options) (*lookup.Provider, error) {
if db == nil {
return nil, fmt.Errorf("backends: nil database")
}
res, err := cfg.Resolve()
if err != nil {
return nil, err
}
d, err := cfg.ResolveDialect(db)
if err != nil {
return nil, err
}
for _, op := range lookup.AllOps() {
if _, err := res.EffectiveMode(op, d.Name()); err != nil {
return nil, err
}
}
c := &chooser{cfg: res, dialect: d.Name(), procs: res.Procs}
run := procedure.NewDB(db, opts.DBFactory, c.resetProbes)
c.db = run
base, err := direct.NewBase(run, d, res.Schema)
if err != nil {
return nil, err
}
p := procedure.NewPasskey(run, res.Procs)
return &lookup.Provider{
Auth: &authRouter{c: c,
proc: procedure.NewAuth(run, res.Procs),
direct: direct.NewAuth(base, direct.AuthOptions{UpgradePasswordHash: opts.UpgradePasswordHash})},
Keys: &keysRouter{c: c,
proc: procedure.NewKeys(run, res.Procs),
direct: direct.NewKeys(base)},
OAuthClient: &oauthClientRouter{c: c,
proc: procedure.NewOAuthClients(run, res.Procs),
direct: direct.NewOAuthClients(base)},
OAuthUser: &oauthUserRouter{c: c,
proc: procedure.NewOAuthUsers(run, res.Procs),
direct: direct.NewOAuthUsers(base)},
OAuthGrant: &oauthGrantRouter{c: c,
proc: procedure.NewOAuthGrants(run, res.Procs),
direct: direct.NewOAuthGrants(base)},
Passkey: &passkeyRouter{c: c, proc: p, direct: direct.NewPasskey(base)},
TOTP: &totpRouter{c: c,
proc: procedure.NewTOTP(run, res.Procs),
direct: direct.NewTOTP(base)},
Policy: &policyRouter{c: c,
proc: procedure.NewPolicy(run, res.Procs),
direct: direct.NewPolicy(base, direct.PolicyOptions{NoGroups: opts.NoGroupTables})},
}, nil
}
// Failed returns a Provider whose every operation returns err. Constructors that cannot
// return an error use it so a bad configuration fails closed on first use.
func Failed(err error) *lookup.Provider {
c := &chooser{fail: err}
return &lookup.Provider{
Auth: &authRouter{c: c},
Keys: &keysRouter{c: c},
OAuthClient: &oauthClientRouter{c: c},
OAuthUser: &oauthUserRouter{c: c},
OAuthGrant: &oauthGrantRouter{c: c},
Passkey: &passkeyRouter{c: c},
TOTP: &totpRouter{c: c},
Policy: &policyRouter{c: c},
}
}
// chooser decides per operation whether the procedure or the direct store runs.
type chooser struct {
cfg *lookup.Resolved
dialect string
procs lookup.ProcNames
db *procedure.DB
probes sync.Map // proc name -> bool
fail error // set by Failed: every operation returns it
}
func (c *chooser) resetProbes() {
c.probes.Range(func(k, _ any) bool { c.probes.Delete(k); return true })
}
// useProc reports whether op should call the stored procedure proc. In auto mode on Postgres
// the catalog is probed once per procedure (cached until a reconnect).
func (c *chooser) useProc(ctx context.Context, op lookup.Op, proc string) (bool, error) {
if c.fail != nil {
return false, c.fail
}
m, err := c.cfg.EffectiveMode(op, c.dialect)
if err != nil {
return false, err
}
switch m {
case lookup.ModeProcedure:
return true, nil
case lookup.ModeAuto:
if v, ok := c.probes.Load(proc); ok {
return v.(bool), nil
}
exists := probeProcedure(ctx, c.db.Get(), proc)
c.probes.Store(proc, exists)
return exists, nil
}
return false, nil
}
// probeProcedure asks the Postgres catalog whether a function exists. Any failure counts as
// "does not exist" so the probe can never block an operation.
func probeProcedure(ctx context.Context, db *sql.DB, proc string) (exists bool) {
if db == nil {
return false
}
defer func() { _ = recover() }()
dbtrace.Raw(ctx, "probe.pg_proc")
if err := db.QueryRowContext(ctx, `SELECT EXISTS (SELECT 1 FROM pg_proc WHERE proname = $1 LIMIT 1)`, proc).Scan(&exists); err != nil {
return false
}
return exists
}
// pick returns the store that should serve op.
func pick[T any](c *chooser, ctx context.Context, op lookup.Op, proc string, p, d T) (T, error) {
use, err := c.useProc(ctx, op, proc)
if err != nil {
var zero T
return zero, err
}
if use {
return p, nil
}
return d, nil
}
@@ -0,0 +1,119 @@
package backends
import (
"context"
"database/sql"
"errors"
"testing"
"github.com/DATA-DOG/go-sqlmock"
_ "github.com/glebarez/go-sqlite"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
func sqliteDB(t *testing.T) *sql.DB {
t.Helper()
db, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatal(err)
}
db.SetMaxOpenConns(1)
t.Cleanup(func() { _ = db.Close() })
ddl, err := ddl.SQL("sqlite")
if err != nil {
t.Fatal(err)
}
if _, err := db.Exec(ddl); err != nil {
t.Fatal(err)
}
return db
}
func TestSQLiteDefaultsToDirect(t *testing.T) {
ctx := context.Background()
p, err := New(sqliteDB(t), lookup.Config{}, Options{})
if err != nil {
t.Fatal(err)
}
reg, err := p.Auth.Register(ctx, sectypes.RegisterRequest{Username: "a", Email: "a@x.io", Password: "pw"})
if err != nil {
t.Fatal(err)
}
if _, err := p.Auth.Session(ctx, reg.Token, "authenticate"); err != nil {
t.Fatal(err)
}
if on, err := p.TOTP.Status(ctx, reg.User.UserID); err != nil || on {
t.Fatalf("%v %v", on, err)
}
}
func TestProcedureModeRejectedOnSQLite(t *testing.T) {
_, err := New(sqliteDB(t), lookup.Config{Overrides: map[lookup.Op]lookup.Mode{lookup.OpLogin: lookup.ModeProcedure}}, Options{})
if err == nil {
t.Fatal("expected error")
}
}
func TestCustomSchemaAndUnknownDialect(t *testing.T) {
if _, err := New(sqliteDB(t), lookup.Config{Dialect: "nosuch"}, Options{}); err == nil {
t.Fatal("unknown dialect accepted")
}
bad := lookup.Config{Schema: lookup.Schema{lookup.EntityUsers: {Name: "x; drop"}}}
if _, err := New(sqliteDB(t), bad, Options{}); err == nil {
t.Fatal("unsafe schema accepted")
}
if _, err := New(nil, lookup.Config{}, Options{}); err == nil {
t.Fatal("nil db accepted")
}
}
func TestPostgresDefaultsToProcedure(t *testing.T) {
db, mock, _ := sqlmock.New()
defer db.Close()
p, err := New(db, lookup.Config{Dialect: lookup.DialectPostgres}, Options{})
if err != nil {
t.Fatal(err)
}
mock.ExpectQuery("resolvespec_totp_get_status").WithArgs(1).
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_enabled"}).AddRow(true, nil, true))
if on, err := p.TOTP.Status(context.Background(), 1); err != nil || !on {
t.Fatalf("%v %v", on, err)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestPostgresAutoProbesOnce(t *testing.T) {
db, mock, _ := sqlmock.New()
defer db.Close()
p, err := New(db, lookup.Config{Dialect: lookup.DialectPostgres, Mode: lookup.ModeAuto}, Options{})
if err != nil {
t.Fatal(err)
}
mock.ExpectQuery("pg_proc").WithArgs("resolvespec_totp_get_status").
WillReturnRows(sqlmock.NewRows([]string{"e"}).AddRow(true))
for i := 0; i < 2; i++ { // second call must reuse the cached probe
mock.ExpectQuery("resolvespec_totp_get_status").WithArgs(1).
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_enabled"}).AddRow(true, nil, true))
if on, err := p.TOTP.Status(context.Background(), 1); err != nil || !on {
t.Fatalf("%v %v", on, err)
}
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestFailed(t *testing.T) {
p := Failed(errors.New("boom"))
if _, err := p.Auth.Login(context.Background(), sectypes.LoginRequest{}); err == nil || err.Error() != "boom" {
t.Fatalf("got %v", err)
}
if _, err := p.Policy.RowSecurity(context.Background(), 1, "s", "t"); err == nil {
t.Fatal("want error")
}
}
@@ -0,0 +1,131 @@
package backends
import (
"crypto/rand"
"database/sql"
"encoding/hex"
"fmt"
"os"
"slices"
"testing"
_ "github.com/jackc/pgx/v5/stdlib"
_ "github.com/microsoft/go-mssqldb"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/conformance"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
)
// Every backend/dialect runs the same suite. SQLite runs always. The others run only when a
// DSN is set, and only against a database you are happy to add rows to (all names the suite
// creates carry a unique "cf<hex>_" prefix and are removed afterwards):
//
// RESOLVESPEC_TEST_PG_DSN Postgres with the procedure schema installed
// (lookup/database_schema.sql + keystore_schema.sql): procedure mode
// RESOLVESPEC_TEST_PG_DIRECT_DSN Postgres for direct mode; ddl/postgres.sql is applied if the
// tables are missing. Do not point it at the procedure schema:
// the column types differ.
// RESOLVESPEC_TEST_MYSQL_DSN MySQL (needs a "mysql" database/sql driver linked into the test binary)
// RESOLVESPEC_TEST_MSSQL_DSN SQL Server (driver "sqlserver")
func TestConformance(t *testing.T) {
t.Run("sqlite/direct", func(t *testing.T) {
runConformance(t, sqliteDB(t), "sqlite", lookup.Config{Mode: lookup.ModeDirect}, false)
})
t.Run("sqlite/default", func(t *testing.T) {
runConformance(t, sqliteDB(t), "sqlite", lookup.Config{}, false)
})
real := []struct {
name, env, driver, dialect string
cfg lookup.Config
applyDDL bool
}{
{"postgres/procedure", "RESOLVESPEC_TEST_PG_DSN", "pgx", "postgres", lookup.Config{Mode: lookup.ModeProcedure}, false},
{"postgres/direct", "RESOLVESPEC_TEST_PG_DIRECT_DSN", "pgx", "postgres", lookup.Config{Mode: lookup.ModeDirect}, true},
{"mysql/direct", "RESOLVESPEC_TEST_MYSQL_DSN", "mysql", "mysql", lookup.Config{}, true},
{"mssql/direct", "RESOLVESPEC_TEST_MSSQL_DSN", "sqlserver", "mssql", lookup.Config{}, true},
}
for _, r := range real {
t.Run(r.name, func(t *testing.T) {
dsn := os.Getenv(r.env)
if dsn == "" {
t.Skipf("%s not set", r.env)
}
runOnServer(t, r.driver, dsn, r.dialect, r.cfg, r.applyDDL)
})
}
}
// runOnServer opens dsn, optionally applies the reference DDL, and runs the suite.
func runOnServer(t *testing.T, driver, dsn, dialectName string, cfg lookup.Config, applyDDL bool) {
t.Helper()
if !slices.Contains(sql.Drivers(), driver) {
t.Skipf("database/sql driver %q is not linked into this test binary", driver)
}
db, err := sql.Open(driver, dsn)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
if err := db.Ping(); err != nil {
t.Fatalf("ping: %v", err)
}
if applyDDL {
stmts, err := ddl.Statements(dialectName)
if err != nil {
t.Fatal(err)
}
for _, s := range stmts {
if _, err := db.Exec(s); err != nil {
t.Fatalf("apply ddl: %v\n%s", err, s)
}
}
}
runConformance(t, db, dialectName, cfg, true)
}
func runConformance(t *testing.T, db *sql.DB, dialectName string, cfg lookup.Config, shared bool) {
t.Helper()
d, err := dialect.Get(dialectName)
if err != nil {
t.Fatal(err)
}
cfg.Dialect = dialectName
p, err := New(db, cfg, Options{})
if err != nil {
t.Fatal(err)
}
var b [3]byte
_, _ = rand.Read(b[:])
env := conformance.Env{Provider: p, DB: db, Dialect: d, Prefix: "cf" + hex.EncodeToString(b[:]) + "_"}
if shared {
env.Cleanup = func(t *testing.T) { cleanup(t, db, d, env.Prefix) }
}
conformance.Run(t, env)
}
// cleanup deletes the rows a conformance run created, by prefix. Child rows go with their user
// through the foreign keys; tables without one are cleaned explicitly.
func cleanup(t *testing.T, db *sql.DB, d dialect.Dialect, prefix string) {
t.Helper()
like := prefix + "%"
for _, q := range []struct{ table, col string }{
{"oauth_codes", "code"},
{"oauth_consents", "client_id"},
{"oauth_refresh_tokens", "client_id"},
{"oauth_device_codes", "client_id"},
{"oauth_par_requests", "client_id"},
{"oauth_jti", "jti_key"},
{"oauth_clients", "client_id"},
{"token_blacklist", "token"},
{"sec_column_rules", "schema_name"},
{"sec_row_rules", "schema_name"},
{"users", "username"},
} {
if _, err := db.Exec(fmt.Sprintf("DELETE FROM %s WHERE %s LIKE %s", q.table, q.col, d.Placeholder(1)), like); err != nil {
t.Logf("cleanup %s: %v", q.table, err)
}
}
}
@@ -0,0 +1,180 @@
package backends
import (
"bytes"
"context"
"database/sql"
"fmt"
"net"
"os"
"os/exec"
"strings"
"testing"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// Container tests start a throwaway database server with podman or docker (whichever is
// installed, podman first) and run the conformance suite against it. 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 containerPassword = "Resolve_Spec_1"
func containerRuntime(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 run(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())
}
// startContainer runs image publishing containerPort on a random localhost port and returns
// the host port. The container is force-removed on cleanup.
func startContainer(t *testing.T, rt, image, containerPort string, env map[string]string) string {
t.Helper()
args := []string{"run", "-d", "--rm", "-p", "127.0.0.1::" + containerPort}
for k, v := range env {
args = append(args, "-e", k+"="+v)
}
args = append(args, image)
id := run(t, 10*time.Minute, rt, args...) // first run may pull the image
t.Cleanup(func() { _ = exec.Command(rt, "rm", "-f", id).Run() })
// "127.0.0.1:49153" (docker may print one line per address family)
out := run(t, 30*time.Second, rt, "port", id, containerPort)
line := strings.Fields(out)[len(strings.Fields(out))-1]
for _, l := range strings.Split(out, "\n") {
if strings.HasPrefix(strings.TrimSpace(l), "127.0.0.1") || 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)
}
return port
}
// waitReady retries until the server accepts queries or the deadline passes.
func waitReady(t *testing.T, driver, dsn string, d time.Duration) *sql.DB {
t.Helper()
deadline := time.Now().Add(d)
var last error
for time.Now().Before(deadline) {
db, err := sql.Open(driver, dsn)
if err == nil {
if last = db.Ping(); last == nil {
t.Cleanup(func() { _ = db.Close() })
return db
}
_ = db.Close()
} else {
last = err
}
time.Sleep(time.Second)
}
t.Fatalf("database did not become ready within %s: %v", d, last)
return nil
}
func TestConformancePostgresContainer(t *testing.T) {
rt := containerRuntime(t)
port := startContainer(t, rt, "docker.io/library/postgres:16-alpine", "5432", map[string]string{"POSTGRES_PASSWORD": containerPassword})
dsn := func(db string) string {
return fmt.Sprintf("postgres://postgres:%s@127.0.0.1:%s/%s?sslmode=disable", containerPassword, port, db)
}
admin := waitReady(t, "pgx", dsn("postgres"), 90*time.Second)
// The official image restarts once during init: make sure the second start is the one we use.
time.Sleep(2 * time.Second)
admin = waitReady(t, "pgx", dsn("postgres"), 60*time.Second)
for _, name := range []string{"cf_proc", "cf_direct"} {
if _, err := admin.Exec("CREATE DATABASE " + name); err != nil {
t.Fatal(err)
}
}
t.Run("procedure", func(t *testing.T) {
db, err := sql.Open("pgx", dsn("cf_proc"))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
for _, f := range []string{"../database_schema.sql", "../keystore_schema.sql"} {
b, err := os.ReadFile(f)
if err != nil {
t.Fatal(err)
}
if _, err := db.Exec(string(b)); err != nil {
t.Fatalf("apply %s: %v", f, err)
}
}
runConformance(t, db, "postgres", lookup.Config{Mode: lookup.ModeProcedure}, true)
})
t.Run("direct", func(t *testing.T) {
runOnServer(t, "pgx", dsn("cf_direct"), "postgres", lookup.Config{Mode: lookup.ModeDirect}, true)
})
}
func TestConformanceMSSQLContainer(t *testing.T) {
rt := containerRuntime(t)
port := startContainer(t, rt, "mcr.microsoft.com/mssql/server:2022-latest", "1433", map[string]string{
"ACCEPT_EULA": "Y", "MSSQL_SA_PASSWORD": containerPassword,
})
dsn := func(db string) string {
return fmt.Sprintf("sqlserver://sa:%s@127.0.0.1:%s?database=%s&encrypt=disable", containerPassword, port, db)
}
admin := waitReady(t, "sqlserver", dsn("master"), 3*time.Minute)
if _, err := admin.Exec("CREATE DATABASE cf_direct"); err != nil {
t.Fatal(err)
}
runOnServer(t, "sqlserver", dsn("cf_direct"), "mssql", lookup.Config{}, true)
}
// TestContainerLifecycle checks the start/stop plumbing the container tests rely on: the
// container comes up and accepts connections, and after stop it is gone (it runs with --rm).
func TestContainerLifecycle(t *testing.T) {
rt := containerRuntime(t)
port := startContainer(t, rt, "docker.io/library/postgres:16-alpine", "5432", map[string]string{"POSTGRES_PASSWORD": containerPassword})
waitReady(t, "pgx", fmt.Sprintf("postgres://postgres:%s@127.0.0.1:%s/postgres?sslmode=disable", containerPassword, port), 90*time.Second)
listed := func() string {
return run(t, 30*time.Second, rt, "ps", "-q", "--filter", "ancestor=docker.io/library/postgres:16-alpine")
}
id := listed()
if id == "" {
t.Fatal("container is not running after start")
}
run(t, time.Minute, rt, "stop", "-t", "2", id)
deadline := time.Now().Add(30 * time.Second)
for listed() != "" {
if time.Now().After(deadline) {
t.Fatal("container still present after stop")
}
time.Sleep(500 * time.Millisecond)
}
}
+543
View File
@@ -0,0 +1,543 @@
package backends
import (
"context"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// Routers: each store method asks the chooser which backend serves that operation.
type authRouter struct {
c *chooser
proc, direct lookup.AuthStore
}
var _ lookup.AuthStore = (*authRouter)(nil)
func (r *authRouter) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpLogin, r.c.procs.Login, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.Login(ctx, req)
}
func (r *authRouter) Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpRegister, r.c.procs.Register, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.Register(ctx, req)
}
func (r *authRouter) Logout(ctx context.Context, req sectypes.LogoutRequest) error {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpLogout, r.c.procs.Logout, r.proc, r.direct)
if err != nil {
return err
}
return st.Logout(ctx, req)
}
func (r *authRouter) Session(ctx context.Context, token, reference string) (*sectypes.UserContext, error) {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpSession, r.c.procs.Session, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.Session(ctx, token, reference)
}
func (r *authRouter) TouchSession(ctx context.Context, token string, user *sectypes.UserContext) error {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpTouchSession, r.c.procs.SessionUpdate, r.proc, r.direct)
if err != nil {
return err
}
return st.TouchSession(ctx, token, user)
}
func (r *authRouter) Refresh(ctx context.Context, refreshToken string) (*sectypes.LoginResponse, error) {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpRefresh, r.c.procs.RefreshToken, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.Refresh(ctx, refreshToken)
}
func (r *authRouter) LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error) {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpLoginAPIKey, r.c.procs.LoginAPIKey, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.LoginAPIKey(ctx, rawKey, claims)
}
func (r *authRouter) JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpJWTLogin, r.c.procs.JWTLogin, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.JWTLogin(ctx, req)
}
func (r *authRouter) JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpJWTLogout, r.c.procs.JWTLogout, r.proc, r.direct)
if err != nil {
return err
}
return st.JWTLogout(ctx, req)
}
func (r *authRouter) ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpResetRequest, r.c.procs.PasswordResetRequest, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.ResetRequest(ctx, req)
}
func (r *authRouter) ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error {
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpResetComplete, r.c.procs.PasswordResetComplete, r.proc, r.direct)
if err != nil {
return err
}
return st.ResetComplete(ctx, req)
}
type keysRouter struct {
c *chooser
proc, direct lookup.KeyStore
}
var _ lookup.KeyStore = (*keysRouter)(nil)
func (r *keysRouter) Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error) {
st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyCreate, r.c.procs.KeystoreCreateKey, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.Create(ctx, req, keyHash)
}
func (r *keysRouter) List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) {
st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyList, r.c.procs.KeystoreGetUserKeys, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.List(ctx, userID, keyType)
}
func (r *keysRouter) Delete(ctx context.Context, userID int, keyID int64) (string, error) {
st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyDelete, r.c.procs.KeystoreDeleteKey, r.proc, r.direct)
if err != nil {
return "", err
}
return st.Delete(ctx, userID, keyID)
}
func (r *keysRouter) Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error) {
st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyValidate, r.c.procs.KeystoreValidateKey, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.Validate(ctx, keyHash, keyType)
}
type oauthClientRouter struct {
c *chooser
proc, direct lookup.OAuthClientStore
}
var _ lookup.OAuthClientStore = (*oauthClientRouter)(nil)
func (r *oauthClientRouter) RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error) {
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthRegisterClient, r.c.procs.OAuthRegisterClient, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.RegisterClient(ctx, client)
}
func (r *oauthClientRouter) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthGetClient, r.c.procs.OAuthGetClient, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.GetClient(ctx, clientID)
}
func (r *oauthClientRouter) SaveCode(ctx context.Context, code *sectypes.OAuthCode) error {
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthSaveCode, r.c.procs.OAuthSaveCode, r.proc, r.direct)
if err != nil {
return err
}
return st.SaveCode(ctx, code)
}
func (r *oauthClientRouter) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) {
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthExchangeCode, r.c.procs.OAuthExchangeCode, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.ExchangeCode(ctx, code)
}
func (r *oauthClientRouter) Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error) {
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthIntrospect, r.c.procs.OAuthIntrospect, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.Introspect(ctx, token)
}
func (r *oauthClientRouter) Revoke(ctx context.Context, token string) error {
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthRevoke, r.c.procs.OAuthRevoke, r.proc, r.direct)
if err != nil {
return err
}
return st.Revoke(ctx, token)
}
func (r *oauthClientRouter) UpdateClient(ctx context.Context, client *sectypes.OAuthServerClient) error {
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthUpdateClient, r.c.procs.OAuthUpdateClient, r.proc, r.direct)
if err != nil {
return err
}
return st.UpdateClient(ctx, client)
}
func (r *oauthClientRouter) DeleteClient(ctx context.Context, clientID string) error {
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthDeleteClient, r.c.procs.OAuthDeleteClient, r.proc, r.direct)
if err != nil {
return err
}
return st.DeleteClient(ctx, clientID)
}
type oauthUserRouter struct {
c *chooser
proc, direct lookup.OAuthUserStore
}
var _ lookup.OAuthUserStore = (*oauthUserRouter)(nil)
func (r *oauthUserRouter) GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error) {
st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthGetOrCreateUser, r.c.procs.OAuthGetOrCreateUser, r.proc, r.direct)
if err != nil {
return 0, err
}
return st.GetOrCreateUser(ctx, user, provider)
}
func (r *oauthUserRouter) CreateSession(ctx context.Context, session lookup.OAuthSession) error {
st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthCreateSession, r.c.procs.OAuthCreateSession, r.proc, r.direct)
if err != nil {
return err
}
return st.CreateSession(ctx, session)
}
func (r *oauthUserRouter) GetByRefreshToken(ctx context.Context, refreshToken string) (*lookup.OAuthRefreshSession, error) {
st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthGetRefreshToken, r.c.procs.OAuthGetRefreshToken, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.GetByRefreshToken(ctx, refreshToken)
}
func (r *oauthUserRouter) UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error {
st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthUpdateRefreshToken, r.c.procs.OAuthUpdateRefreshToken, r.proc, r.direct)
if err != nil {
return err
}
return st.UpdateRefreshToken(ctx, userID, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken, expiresAt)
}
func (r *oauthUserRouter) GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error) {
st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthGetUser, r.c.procs.OAuthGetUser, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.GetUser(ctx, userID)
}
type passkeyRouter struct {
c *chooser
proc, direct lookup.PasskeyStore
}
var _ lookup.PasskeyStore = (*passkeyRouter)(nil)
func (r *passkeyRouter) Store(ctx context.Context, rec lookup.PasskeyCredentialRecord) (int64, error) {
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyStore, r.c.procs.PasskeyStoreCredential, r.proc, r.direct)
if err != nil {
return 0, err
}
return st.Store(ctx, rec)
}
func (r *passkeyRouter) Get(ctx context.Context, credentialID string) (userID int, signCount uint32, err error) {
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyGet, r.c.procs.PasskeyGetCredential, r.proc, r.direct)
if err != nil {
return 0, 0, err
}
return st.Get(ctx, credentialID)
}
func (r *passkeyRouter) UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (bool, error) {
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyUpdateCounter, r.c.procs.PasskeyUpdateCounter, r.proc, r.direct)
if err != nil {
return false, err
}
return st.UpdateCounter(ctx, credentialID, newCounter)
}
func (r *passkeyRouter) List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error) {
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyList, r.c.procs.PasskeyGetUserCredentials, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.List(ctx, userID)
}
func (r *passkeyRouter) Delete(ctx context.Context, userID int, credentialID string) error {
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyDelete, r.c.procs.PasskeyDeleteCredential, r.proc, r.direct)
if err != nil {
return err
}
return st.Delete(ctx, userID, credentialID)
}
func (r *passkeyRouter) Rename(ctx context.Context, userID int, credentialID, name string) error {
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyRename, r.c.procs.PasskeyUpdateName, r.proc, r.direct)
if err != nil {
return err
}
return st.Rename(ctx, userID, credentialID, name)
}
func (r *passkeyRouter) ByUsername(ctx context.Context, username string) (int, []lookup.PasskeyCredentialRef, error) {
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyByUsername, r.c.procs.PasskeyGetCredsByUsername, r.proc, r.direct)
if err != nil {
return 0, nil, err
}
return st.ByUsername(ctx, username)
}
func (r *passkeyRouter) Login(ctx context.Context, userID int, claims map[string]any) (*sectypes.LoginResponse, error) {
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyLogin, r.c.procs.PasskeyLogin, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.Login(ctx, userID, claims)
}
type totpRouter struct {
c *chooser
proc, direct lookup.TOTPStore
}
var _ lookup.TOTPStore = (*totpRouter)(nil)
func (r *totpRouter) Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error {
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPEnable, r.c.procs.TOTPEnable, r.proc, r.direct)
if err != nil {
return err
}
return st.Enable(ctx, userID, secret, hashedCodes)
}
func (r *totpRouter) Disable(ctx context.Context, userID int) error {
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPDisable, r.c.procs.TOTPDisable, r.proc, r.direct)
if err != nil {
return err
}
return st.Disable(ctx, userID)
}
func (r *totpRouter) Status(ctx context.Context, userID int) (bool, error) {
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPStatus, r.c.procs.TOTPGetStatus, r.proc, r.direct)
if err != nil {
return false, err
}
return st.Status(ctx, userID)
}
func (r *totpRouter) Secret(ctx context.Context, userID int) (string, error) {
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPSecret, r.c.procs.TOTPGetSecret, r.proc, r.direct)
if err != nil {
return "", err
}
return st.Secret(ctx, userID)
}
func (r *totpRouter) RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error {
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPRegenerateBackup, r.c.procs.TOTPRegenerateBackup, r.proc, r.direct)
if err != nil {
return err
}
return st.RegenerateBackupCodes(ctx, userID, hashedCodes)
}
func (r *totpRouter) ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error) {
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPValidateBackupCode, r.c.procs.TOTPValidateBackupCode, r.proc, r.direct)
if err != nil {
return false, err
}
return st.ValidateBackupCode(ctx, userID, codeHash)
}
type policyRouter struct {
c *chooser
proc, direct lookup.PolicyStore
}
var _ lookup.PolicyStore = (*policyRouter)(nil)
func (r *policyRouter) ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) {
st, err := pick[lookup.PolicyStore](r.c, ctx, lookup.OpColumnSecurity, r.c.procs.ColumnSecurity, r.proc, r.direct)
if err != nil {
return nil, err
}
return st.ColumnSecurity(ctx, userID, schema, table)
}
func (r *policyRouter) RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error) {
st, err := pick[lookup.PolicyStore](r.c, ctx, lookup.OpRowSecurity, r.c.procs.RowSecurity, r.proc, r.direct)
if err != nil {
return sectypes.RowSecurity{}, err
}
return st.RowSecurity(ctx, userRef, schema, table)
}
type oauthGrantRouter struct {
c *chooser
proc, direct lookup.OAuthGrantStore
}
var _ lookup.OAuthGrantStore = (*oauthGrantRouter)(nil)
func (r *oauthGrantRouter) pick(ctx context.Context, op lookup.Op, proc string) (lookup.OAuthGrantStore, error) {
return pick[lookup.OAuthGrantStore](r.c, ctx, op, proc, r.proc, r.direct)
}
func (r *oauthGrantRouter) SaveConsent(ctx context.Context, c lookup.Consent) error {
st, err := r.pick(ctx, lookup.OpOAuthSaveConsent, r.c.procs.OAuthSaveConsent)
if err != nil {
return err
}
return st.SaveConsent(ctx, c)
}
func (r *oauthGrantRouter) GetConsent(ctx context.Context, userID int, clientID string) (*lookup.Consent, error) {
st, err := r.pick(ctx, lookup.OpOAuthGetConsent, r.c.procs.OAuthGetConsent)
if err != nil {
return nil, err
}
return st.GetConsent(ctx, userID, clientID)
}
func (r *oauthGrantRouter) RevokeConsent(ctx context.Context, userID int, clientID string) error {
st, err := r.pick(ctx, lookup.OpOAuthRevokeConsent, r.c.procs.OAuthRevokeConsent)
if err != nil {
return err
}
return st.RevokeConsent(ctx, userID, clientID)
}
func (r *oauthGrantRouter) SaveRefresh(ctx context.Context, t lookup.RefreshToken) error {
st, err := r.pick(ctx, lookup.OpOAuthSaveRefresh, r.c.procs.OAuthSaveRefresh)
if err != nil {
return err
}
return st.SaveRefresh(ctx, t)
}
func (r *oauthGrantRouter) RotateRefresh(ctx context.Context, oldHash string, next lookup.RefreshToken) (*lookup.RefreshToken, error) {
st, err := r.pick(ctx, lookup.OpOAuthRotateRefresh, r.c.procs.OAuthRotateRefresh)
if err != nil {
return nil, err
}
return st.RotateRefresh(ctx, oldHash, next)
}
func (r *oauthGrantRouter) PeekRefresh(ctx context.Context, hash string) (*lookup.RefreshToken, error) {
st, err := r.pick(ctx, lookup.OpOAuthPeekRefresh, r.c.procs.OAuthPeekRefresh)
if err != nil {
return nil, err
}
return st.PeekRefresh(ctx, hash)
}
func (r *oauthGrantRouter) RevokeRefreshFamily(ctx context.Context, familyID string) error {
st, err := r.pick(ctx, lookup.OpOAuthRevokeRefreshFamily, r.c.procs.OAuthRevokeRefreshFamily)
if err != nil {
return err
}
return st.RevokeRefreshFamily(ctx, familyID)
}
func (r *oauthGrantRouter) RevokeRefreshBySession(ctx context.Context, sessionToken string) error {
st, err := r.pick(ctx, lookup.OpOAuthRevokeRefreshByUser, r.c.procs.OAuthRevokeRefreshByUser)
if err != nil {
return err
}
return st.RevokeRefreshBySession(ctx, sessionToken)
}
func (r *oauthGrantRouter) CreateDevice(ctx context.Context, d lookup.DeviceCode) error {
st, err := r.pick(ctx, lookup.OpOAuthCreateDevice, r.c.procs.OAuthCreateDevice)
if err != nil {
return err
}
return st.CreateDevice(ctx, d)
}
func (r *oauthGrantRouter) DeviceByUserCode(ctx context.Context, userCode string) (*lookup.DeviceCode, error) {
st, err := r.pick(ctx, lookup.OpOAuthDeviceByUserCode, r.c.procs.OAuthDeviceByUserCode)
if err != nil {
return nil, err
}
return st.DeviceByUserCode(ctx, userCode)
}
func (r *oauthGrantRouter) DeviceDecide(ctx context.Context, userCode string, approve bool, userID int, sessionToken string) error {
st, err := r.pick(ctx, lookup.OpOAuthDeviceDecide, r.c.procs.OAuthDeviceDecide)
if err != nil {
return err
}
return st.DeviceDecide(ctx, userCode, approve, userID, sessionToken)
}
func (r *oauthGrantRouter) DevicePoll(ctx context.Context, deviceHash string) (*lookup.DeviceCode, error) {
st, err := r.pick(ctx, lookup.OpOAuthDevicePoll, r.c.procs.OAuthDevicePoll)
if err != nil {
return nil, err
}
return st.DevicePoll(ctx, deviceHash)
}
func (r *oauthGrantRouter) SavePushedRequest(ctx context.Context, req lookup.PushedRequest) error {
st, err := r.pick(ctx, lookup.OpOAuthSavePAR, r.c.procs.OAuthSavePAR)
if err != nil {
return err
}
return st.SavePushedRequest(ctx, req)
}
func (r *oauthGrantRouter) ConsumePushedRequest(ctx context.Context, requestURI string) (*lookup.PushedRequest, error) {
st, err := r.pick(ctx, lookup.OpOAuthConsumePAR, r.c.procs.OAuthConsumePAR)
if err != nil {
return nil, err
}
return st.ConsumePushedRequest(ctx, requestURI)
}
func (r *oauthGrantRouter) SeenJTI(ctx context.Context, key string, expires time.Time) (bool, error) {
st, err := r.pick(ctx, lookup.OpOAuthSeenJTI, r.c.procs.OAuthSeenJTI)
if err != nil {
return false, err
}
return st.SeenJTI(ctx, key, expires)
}
@@ -0,0 +1,872 @@
// Package conformance is the shared behavioural suite every lookup backend must pass.
// It only uses the store interfaces, so the same cases run against the direct backend on
// every dialect and against the procedure backend on Postgres. Error messages are not
// asserted (backends word them differently), only whether an operation succeeds or fails
// and the values it returns.
//
// The suite names everything it creates with Env.Prefix and never assumes empty tables, so
// it can run against a shared database. Env.Cleanup, when set, removes the prefixed rows.
package conformance
import (
"context"
"database/sql"
"encoding/base64"
"errors"
"fmt"
"strings"
"testing"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// Env is one backend under test.
type Env struct {
Provider *lookup.Provider
// DB and Dialect are used only to seed policy rules, which have no store method.
DB *sql.DB
Dialect dialect.Dialect
// Prefix makes every created name unique to this run.
Prefix string
// Cleanup removes rows whose names start with Prefix. Optional.
Cleanup func(t *testing.T)
}
// Run executes the suite.
func Run(t *testing.T, env Env) {
if env.Cleanup != nil {
t.Cleanup(func() { env.Cleanup(t) })
}
s := &suite{Env: env}
t.Run("AuthSessionLifecycle", s.authSessionLifecycle)
t.Run("AuthRejectsBadCredentials", s.authRejectsBadCredentials)
t.Run("RegisterIgnoresPrivileges", s.registerIgnoresPrivileges)
t.Run("RegisterRejectsDuplicates", s.registerRejectsDuplicates)
t.Run("PasswordReset", s.passwordReset)
t.Run("JWT", s.jwt)
t.Run("Keys", s.keys)
t.Run("LoginAPIKey", s.loginAPIKey)
t.Run("OAuthClientAndCodes", s.oauthClientAndCodes)
t.Run("OAuthIntrospectRevoke", s.oauthIntrospectRevoke)
t.Run("OAuthUsers", s.oauthUsers)
t.Run("OAuthClientMetadata", s.oauthClientMetadata)
t.Run("OAuthGrantConsent", s.oauthGrantConsent)
t.Run("OAuthGrantRefresh", s.oauthGrantRefresh)
t.Run("OAuthGrantDevice", s.oauthGrantDevice)
t.Run("OAuthGrantPAR", s.oauthGrantPAR)
t.Run("OAuthGrantJTI", s.oauthGrantJTI)
t.Run("Passkey", s.passkey)
t.Run("TOTP", s.totp)
t.Run("Policy", s.policy)
}
type suite struct{ Env }
var ctx = context.Background()
func (s *suite) name(n string) string { return s.Prefix + n }
func (s *suite) register(t *testing.T, n string) *sectypes.LoginResponse {
t.Helper()
resp, err := s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{
Username: s.name(n), Email: s.name(n) + "@example.test", Password: "pw-" + n,
})
if err != nil {
t.Fatalf("register %s: %v", n, err)
}
if resp == nil || resp.User == nil || resp.Token == "" || resp.User.UserID == 0 {
t.Fatalf("register %s: incomplete response %+v", n, resp)
}
return resp
}
func rejected(t *testing.T, what string, err error) {
t.Helper()
if err == nil {
t.Fatalf("%s: expected an error", what)
}
}
// notOK asserts an operation did not validate: it either failed or returned false.
func notOK(t *testing.T, what string, ok bool, err error) {
t.Helper()
if err == nil && ok {
t.Fatalf("%s: accepted", what)
}
}
func (s *suite) authSessionLifecycle(t *testing.T) {
a := s.Provider.Auth
reg := s.register(t, "life")
login, err := a.Login(ctx, sectypes.LoginRequest{Username: s.name("life"), Password: "pw-life",
Claims: map[string]any{"ip_address": "10.0.0.1", "user_agent": "conformance"}})
if err != nil || login.Token == "" || login.User.UserName != s.name("life") {
t.Fatalf("login: %+v %v", login, err)
}
if login.Token == reg.Token {
t.Fatal("login reused the registration session")
}
u, err := a.Session(ctx, login.Token, "authenticate")
if err != nil || u.UserName != s.name("life") || u.UserID != reg.User.UserID {
t.Fatalf("session: %+v %v", u, err)
}
if err := a.TouchSession(ctx, login.Token, u); err != nil {
t.Fatalf("touch: %v", err)
}
_, err = a.Session(ctx, s.name("no-such-token"), "authenticate")
rejected(t, "unknown session", err)
ref, err := a.Refresh(ctx, login.Token)
if err != nil || ref.Token == "" || ref.Token == login.Token {
t.Fatalf("refresh: %+v %v", ref, err)
}
_, err = a.Session(ctx, login.Token, "")
rejected(t, "session after refresh", err)
_, err = a.Refresh(ctx, login.Token)
rejected(t, "second refresh of the same token", err)
if err := a.Logout(ctx, sectypes.LogoutRequest{Token: ref.Token, UserID: ref.User.UserID}); err != nil {
t.Fatalf("logout: %v", err)
}
_, err = a.Session(ctx, ref.Token, "")
rejected(t, "session after logout", err)
}
func (s *suite) authRejectsBadCredentials(t *testing.T) {
a := s.Provider.Auth
s.register(t, "creds")
_, err := a.Login(ctx, sectypes.LoginRequest{Username: s.name("creds"), Password: "wrong"})
rejected(t, "wrong password", err)
_, err = a.Login(ctx, sectypes.LoginRequest{Username: s.name("creds")})
rejected(t, "empty password", err)
_, err = a.Login(ctx, sectypes.LoginRequest{Username: s.name("nobody"), Password: "pw"})
rejected(t, "unknown user", err)
}
func (s *suite) registerIgnoresPrivileges(t *testing.T) {
resp, err := s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{
Username: s.name("priv"), Email: s.name("priv") + "@example.test", Password: "x",
UserLevel: 99, Roles: []string{"admin"},
})
if err != nil {
t.Fatal(err)
}
if resp.User.UserLevel != 0 || len(resp.User.Roles) != 0 {
t.Fatalf("client-supplied privileges honoured: %+v", resp.User)
}
}
func (s *suite) registerRejectsDuplicates(t *testing.T) {
s.register(t, "dup")
_, err := s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{Username: s.name("dup"), Email: s.name("dup2") + "@example.test", Password: "x"})
rejected(t, "duplicate username", err)
_, err = s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{Username: s.name("dup2"), Email: s.name("dup") + "@example.test", Password: "x"})
rejected(t, "duplicate email", err)
}
func (s *suite) passwordReset(t *testing.T) {
a := s.Provider.Auth
reg := s.register(t, "reset")
if r, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: s.name("nobody") + "@example.test"}); err != nil || (r != nil && r.Token != "") {
t.Fatalf("unknown email must succeed without a token (user enumeration): %+v %v", r, err)
}
req, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: s.name("reset") + "@example.test"})
if err != nil || req == nil || req.Token == "" {
t.Fatalf("reset request: %+v %v", req, err)
}
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: "bogus", NewPassword: "x"}); err == nil {
t.Fatal("bogus reset token accepted")
}
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: req.Token, NewPassword: "new-pw"}); err != nil {
t.Fatalf("reset complete: %v", err)
}
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: req.Token, NewPassword: "again"}); err == nil {
t.Fatal("reset token reused")
}
_, err = a.Session(ctx, reg.Token, "")
rejected(t, "session surviving a password reset", err)
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: s.name("reset"), Password: "new-pw"}); err != nil {
t.Fatalf("login with new password: %v", err)
}
_, err = a.Login(ctx, sectypes.LoginRequest{Username: s.name("reset"), Password: "pw-reset"})
rejected(t, "old password after reset", err)
}
func (s *suite) jwt(t *testing.T) {
a := s.Provider.Auth
reg := s.register(t, "jwt")
resp, err := a.JWTLogin(ctx, sectypes.LoginRequest{Username: s.name("jwt"), Password: "pw-jwt"})
if err != nil || resp.Token == "" || resp.User.UserID != reg.User.UserID {
t.Fatalf("jwt login: %+v %v", resp, err)
}
_, err = a.JWTLogin(ctx, sectypes.LoginRequest{Username: s.name("jwt"), Password: "bad"})
rejected(t, "jwt login with wrong password", err)
if err := a.JWTLogout(ctx, sectypes.LogoutRequest{Token: s.name("jwt-tok"), UserID: reg.User.UserID}); err != nil {
t.Fatalf("jwt logout: %v", err)
}
}
func (s *suite) createKey(t *testing.T, uid int, typ sectypes.KeyType, raw string, exp *time.Time) *sectypes.UserKey {
t.Helper()
k, err := s.Provider.Keys.Create(ctx, sectypes.CreateKeyRequest{UserID: uid, KeyType: typ, Name: s.name("key"),
Scopes: []string{"read"}, ExpiresAt: exp}, sectypes.HashKey(raw))
if err != nil || k == nil || k.ID == 0 {
t.Fatalf("create key: %+v %v", k, err)
}
return k
}
func (s *suite) keys(t *testing.T) {
k := s.Provider.Keys
uid := s.register(t, "keys").User.UserID
raw := s.name("raw-keys")
created := s.createKey(t, uid, sectypes.KeyTypeHeaderAPI, raw, nil)
s.createKey(t, uid, sectypes.KeyTypeJWTSecret, s.name("raw-keys-jwt"), nil)
s.createKey(t, uid, sectypes.KeyTypeGenericAPI, s.name("raw-keys-old"), ptr(time.Now().Add(-time.Hour)))
all, err := k.List(ctx, uid, "")
if err != nil || len(all) != 2 {
t.Fatalf("list must hide expired keys: %d %v", len(all), err)
}
one, err := k.List(ctx, uid, sectypes.KeyTypeHeaderAPI)
if err != nil || len(one) != 1 || one[0].ID != created.ID || len(one[0].Scopes) != 1 {
t.Fatalf("typed list: %+v %v", one, err)
}
got, err := k.Validate(ctx, sectypes.HashKey(raw), sectypes.KeyTypeHeaderAPI)
if err != nil || got.UserID != uid {
t.Fatalf("validate: %+v %v", got, err)
}
_, err = k.Validate(ctx, sectypes.HashKey(raw), sectypes.KeyTypeGenericAPI)
rejected(t, "wrong key type", err)
_, err = k.Validate(ctx, sectypes.HashKey(s.name("raw-keys-old")), "")
rejected(t, "expired key", err)
_, err = k.Validate(ctx, sectypes.HashKey(s.name("unknown")), "")
rejected(t, "unknown key", err)
_, err = k.Delete(ctx, uid+1_000_000, created.ID)
rejected(t, "deleting another user's key", err)
if _, err := k.Delete(ctx, uid, created.ID); err != nil {
t.Fatalf("delete: %v", err)
}
_, err = k.Delete(ctx, uid, created.ID)
rejected(t, "deleting twice", err)
_, err = k.Validate(ctx, sectypes.HashKey(raw), "")
rejected(t, "deleted key", err)
}
func (s *suite) loginAPIKey(t *testing.T) {
a := s.Provider.Auth
uid := s.register(t, "apikey").User.UserID
good, generic, jwtKey, off, old := s.name("ak-good"), s.name("ak-generic"), s.name("ak-jwt"), s.name("ak-off"), s.name("ak-old")
s.createKey(t, uid, sectypes.KeyTypeHeaderAPI, good, nil)
s.createKey(t, uid, sectypes.KeyTypeGenericAPI, generic, nil)
s.createKey(t, uid, sectypes.KeyTypeJWTSecret, jwtKey, nil)
inactive := s.createKey(t, uid, sectypes.KeyTypeGenericAPI, off, nil)
if _, err := s.Provider.Keys.Delete(ctx, uid, inactive.ID); err != nil {
t.Fatal(err)
}
s.createKey(t, uid, sectypes.KeyTypeGenericAPI, old, ptr(time.Now().Add(-time.Hour)))
for _, raw := range []string{good, generic} {
resp, err := a.LoginAPIKey(ctx, raw, map[string]any{"ip_address": "10.0.0.2"})
if err != nil || resp.User.UserName != s.name("apikey") || resp.Token == "" {
t.Fatalf("api key login: %+v %v", resp, err)
}
if _, err := a.Session(ctx, resp.Token, ""); err != nil {
t.Fatalf("session from api key login: %v", err)
}
}
for _, raw := range []string{"", s.name("ak-missing"), jwtKey, off, old} {
_, err := a.LoginAPIKey(ctx, raw, nil)
if !errors.Is(err, lookup.ErrInvalidAPIKey) {
t.Fatalf("key %q: want ErrInvalidAPIKey, got %v", raw, err)
}
}
}
func (s *suite) oauthClientAndCodes(t *testing.T) {
c := s.Provider.OAuthClient
cid := s.name("client")
reg, err := c.RegisterClient(ctx, &sectypes.OAuthServerClient{ClientID: cid, RedirectURIs: []string{"https://app.example.test/cb"}, ClientName: "App"})
if err != nil || reg.ClientID != cid {
t.Fatalf("register client: %+v %v", reg, err)
}
got, err := c.GetClient(ctx, cid)
if err != nil || got.ClientName != "App" || len(got.RedirectURIs) != 1 || got.RedirectURIs[0] != "https://app.example.test/cb" {
t.Fatalf("get client: %+v %v", got, err)
}
_, err = c.GetClient(ctx, s.name("no-client"))
rejected(t, "unknown client", err)
code := &sectypes.OAuthCode{Code: s.name("code1"), ClientID: cid, RedirectURI: "https://app.example.test/cb",
CodeChallenge: "challenge", SessionToken: s.name("sess"), Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Minute)}
if err := c.SaveCode(ctx, code); err != nil {
t.Fatalf("save code: %v", err)
}
ex, err := c.ExchangeCode(ctx, code.Code)
if err != nil || ex.Code != code.Code || ex.ClientID != cid || ex.SessionToken != code.SessionToken || len(ex.Scopes) != 1 {
t.Fatalf("exchange: %+v %v", ex, err)
}
_, err = c.ExchangeCode(ctx, code.Code)
rejected(t, "code reuse", err)
expired := *code
expired.Code, expired.ExpiresAt = s.name("code2"), time.Now().Add(-time.Minute)
if err := c.SaveCode(ctx, &expired); err != nil {
t.Fatal(err)
}
_, err = c.ExchangeCode(ctx, expired.Code)
rejected(t, "expired code", err)
}
func (s *suite) oauthIntrospectRevoke(t *testing.T) {
c := s.Provider.OAuthClient
reg := s.register(t, "intro")
info, err := c.Introspect(ctx, reg.Token)
if err != nil || !info.Active || info.Username != s.name("intro") {
t.Fatalf("introspect: %+v %v", info, err)
}
if err := c.Revoke(ctx, reg.Token); err != nil {
t.Fatalf("revoke: %v", err)
}
if info, err := c.Introspect(ctx, reg.Token); err != nil || info.Active {
t.Fatalf("revoked token still active: %+v %v", info, err)
}
if err := c.Revoke(ctx, s.name("unknown-token")); err != nil {
t.Fatalf("revoking an unknown token must succeed (RFC 7009): %v", err)
}
if info, err := c.Introspect(ctx, s.name("unknown-token")); err != nil || info.Active {
t.Fatalf("unknown token: %+v %v", info, err)
}
}
func (s *suite) oauthUsers(t *testing.T) {
o := s.Provider.OAuthUser
id, err := o.GetOrCreateUser(ctx, &sectypes.UserContext{UserName: s.name("gh"), Email: s.name("gh") + "@example.test", RemoteID: s.name("remote-1")}, "github")
if err != nil || id == 0 {
t.Fatalf("get or create: %d %v", id, err)
}
again, err := o.GetOrCreateUser(ctx, &sectypes.UserContext{UserName: s.name("gh"), Email: s.name("gh") + "@example.test", RemoteID: s.name("remote-1")}, "github")
if err != nil || again != id {
t.Fatalf("second login must return the same user: %d %v", again, err)
}
exp := time.Now().Add(time.Hour)
sess := lookup.OAuthSession{SessionToken: s.name("os1"), UserID: id, AccessToken: "a1", RefreshToken: s.name("or1"), TokenType: "Bearer", ExpiresAt: exp, Provider: "github"}
if err := o.CreateSession(ctx, sess); err != nil {
t.Fatalf("create session: %v", err)
}
ref, err := o.GetByRefreshToken(ctx, sess.RefreshToken)
if err != nil || ref.UserID != id || ref.AccessToken != "a1" {
t.Fatalf("by refresh token: %+v %v", ref, err)
}
_, err = o.GetByRefreshToken(ctx, s.name("or-missing"))
rejected(t, "unknown refresh token", err)
if err := o.UpdateRefreshToken(ctx, id, sess.RefreshToken, s.name("os2"), "a2", s.name("or2"), exp); err != nil {
t.Fatalf("update refresh token: %v", err)
}
if _, err := o.GetByRefreshToken(ctx, s.name("or2")); err != nil {
t.Fatalf("rotated refresh token not found: %v", err)
}
u, err := o.GetUser(ctx, id)
if err != nil || u.UserName != s.name("gh") {
t.Fatalf("get user: %+v %v", u, err)
}
_, err = o.GetUser(ctx, id+1_000_000)
rejected(t, "unknown user", err)
}
func b64(s string) string { return base64.StdEncoding.EncodeToString([]byte(s)) }
func (s *suite) passkey(t *testing.T) {
p := s.Provider.Passkey
reg := s.register(t, "pk")
uid := reg.User.UserID
c1, c2 := b64(s.name("cred1")), b64(s.name("cred2"))
rec := lookup.PasskeyCredentialRecord{UserID: uid, CredentialID: c1, PublicKey: b64("pubkey"), AttestationType: "none",
Transports: []string{"usb", "nfc"}, Name: "Key 1"}
if id, err := p.Store(ctx, rec); err != nil || id == 0 {
t.Fatalf("store: %d %v", id, err)
}
_, err := p.Store(ctx, rec)
rejected(t, "duplicate credential", err)
rec.CredentialID, rec.Name = c2, "Key 2"
if _, err := p.Store(ctx, rec); err != nil {
t.Fatal(err)
}
owner, count, err := p.Get(ctx, c1)
if err != nil || owner != uid || count != 0 {
t.Fatalf("get: %d %d %v", owner, count, err)
}
_, _, err = p.Get(ctx, b64(s.name("missing")))
rejected(t, "unknown credential", err)
if clone, err := p.UpdateCounter(ctx, c1, 5); err != nil || clone {
t.Fatalf("advance counter: clone=%v %v", clone, err)
}
if clone, err := p.UpdateCounter(ctx, c1, 5); err != nil || !clone {
t.Fatalf("replayed counter must raise a clone warning: clone=%v %v", clone, err)
}
list, err := p.List(ctx, uid)
if err != nil || len(list) != 2 {
t.Fatalf("list: %d %v", len(list), err)
}
if err := p.Rename(ctx, uid, c1, "Renamed"); err != nil {
t.Fatalf("rename: %v", err)
}
rejected(t, "renaming another user's credential", p.Rename(ctx, uid+1_000_000, c1, "x"))
gotID, refs, err := p.ByUsername(ctx, s.name("pk"))
if err != nil || gotID != uid || len(refs) != 2 {
t.Fatalf("by username: %d %+v %v", gotID, refs, err)
}
_, _, err = p.ByUsername(ctx, s.name("ghost"))
rejected(t, "unknown username", err)
resp, err := p.Login(ctx, uid, map[string]any{"ip_address": "10.0.0.3"})
if err != nil || resp.Token == "" || resp.User.UserName != s.name("pk") {
t.Fatalf("passkey login: %+v %v", resp, err)
}
if _, err := s.Provider.Auth.Session(ctx, resp.Token, ""); err != nil {
t.Fatalf("session from passkey login: %v", err)
}
rejected(t, "deleting another user's credential", p.Delete(ctx, uid+1_000_000, c1))
if err := p.Delete(ctx, uid, c1); err != nil {
t.Fatalf("delete: %v", err)
}
rejected(t, "deleting twice", p.Delete(ctx, uid, c1))
}
func (s *suite) totp(t *testing.T) {
st := s.Provider.TOTP
uid := s.register(t, "totp").User.UserID
if on, err := st.Status(ctx, uid); err != nil || on {
t.Fatalf("initial status: %v %v", on, err)
}
_, err := st.Secret(ctx, uid)
rejected(t, "secret without 2FA", err)
if err := st.Enable(ctx, uid, "SECRET", []string{s.name("h1"), s.name("h2")}); err != nil {
t.Fatalf("enable: %v", err)
}
if on, _ := st.Status(ctx, uid); !on {
t.Fatal("not enabled")
}
if sec, err := st.Secret(ctx, uid); err != nil || sec != "SECRET" {
t.Fatalf("secret: %q %v", sec, err)
}
if ok, err := st.ValidateBackupCode(ctx, uid, s.name("h1")); err != nil || !ok {
t.Fatalf("backup code: %v %v", ok, err)
}
ok, err := st.ValidateBackupCode(ctx, uid, s.name("h1"))
notOK(t, "backup code reuse", ok, err)
ok, err = st.ValidateBackupCode(ctx, uid, s.name("nope"))
notOK(t, "unknown backup code", ok, err)
if err := st.RegenerateBackupCodes(ctx, uid, []string{s.name("n1")}); err != nil {
t.Fatalf("regenerate: %v", err)
}
ok, err = st.ValidateBackupCode(ctx, uid, s.name("h2"))
notOK(t, "old backup code after regenerate", ok, err)
if ok, err := st.ValidateBackupCode(ctx, uid, s.name("n1")); err != nil || !ok {
t.Fatalf("new backup code: %v %v", ok, err)
}
if err := st.Disable(ctx, uid); err != nil {
t.Fatalf("disable: %v", err)
}
if on, _ := st.Status(ctx, uid); on {
t.Fatal("still enabled after disable")
}
}
// seed inserts one row with dialect placeholders. Values are bound, booleans converted.
func (s *suite) seed(t *testing.T, table string, cols []string, vals ...any) {
t.Helper()
ph := make([]string, len(vals))
args := make([]any, len(vals))
for i, v := range vals {
ph[i] = s.Dialect.Placeholder(i + 1)
if b, ok := v.(bool); ok {
v = s.Dialect.Bool(b)
}
args[i] = v
}
q := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)", table, strings.Join(cols, ", "), strings.Join(ph, ", ")) //nolint:gosec // test seeding with fixed table names
if _, err := s.DB.ExecContext(ctx, q, args...); err != nil {
t.Fatalf("seed %s: %v", table, err)
}
}
func (s *suite) policy(t *testing.T) {
p := s.Provider.Policy
u1 := s.register(t, "pol1").User.UserID
u2 := s.register(t, "pol2").User.UserID
group := 7_000_000 + u1
schema, users, orders, secret := s.name("pub"), "Users", "orders", "secret"
s.seed(t, "sec_group_members", []string{"group_id", "user_id"}, group, u1)
colCols := []string{"user_id", "group_id", "schema_name", "table_name", "column_path", "access_type", "is_active"}
s.seed(t, "sec_column_rules", colCols, u1, nil, schema, users, "email", "mask", true)
s.seed(t, "sec_column_rules", colCols, nil, group, schema, strings.ToLower(users), "profile.ssn", "hide", true)
s.seed(t, "sec_column_rules", colCols, u2, nil, schema, strings.ToLower(users), "other", "hide", true)
s.seed(t, "sec_column_rules", colCols, u1, nil, schema, strings.ToLower(users), "inactive", "hide", false)
s.seed(t, "sec_column_rules", colCols, u1, nil, schema, orders, "x", "hide", true)
s.seed(t, "sec_column_rules", colCols, u1, nil, schema, "users_archive", "y", "hide", true)
rules, err := p.ColumnSecurity(ctx, u1, schema, "users")
if err != nil || len(rules) != 2 {
t.Fatalf("column rules (user + group, exact table, active only): %d %v %+v", len(rules), err, rules)
}
paths := map[string]bool{}
for i := range rules {
paths[strings.Join(rules[i].Path, ".")] = true
}
if !paths["email"] || !paths["profile.ssn"] {
t.Fatalf("paths: %v", paths)
}
if r, err := p.ColumnSecurity(ctx, u2, schema, "users"); err != nil || len(r) != 1 {
t.Fatalf("other user's rules: %d %v", len(r), err)
}
if r, err := p.ColumnSecurity(ctx, u1+u2+1_000_000, schema, "users"); err != nil || len(r) != 0 {
t.Fatalf("no rules must be empty, not an error: %d %v", len(r), err)
}
rowCols := []string{"user_id", "group_id", "schema_name", "table_name", "template", "has_block", "is_active"}
s.seed(t, "sec_row_rules", rowCols, u1, nil, schema, orders, "owner_id = {UserID}", false, true)
s.seed(t, "sec_row_rules", rowCols, nil, group, schema, orders, "region = 1", false, true)
s.seed(t, "sec_row_rules", rowCols, nil, group, schema, orders, "ignored = 1", false, false)
s.seed(t, "sec_row_rules", rowCols, u2, nil, schema, secret, nil, true, true)
s.seed(t, "sec_row_rules", rowCols, nil, group, schema, secret, "x = 1", false, true)
rs, err := p.RowSecurity(ctx, u1, schema, orders)
if err != nil || rs.HasBlock || !strings.Contains(rs.Template, "owner_id = {UserID}") || !strings.Contains(rs.Template, "region = 1") || strings.Contains(rs.Template, "ignored") {
t.Fatalf("row template: %+v %v", rs, err)
}
if rs, err := p.RowSecurity(ctx, u2, schema, secret); err != nil || !rs.HasBlock {
t.Fatalf("blocking rule must win: %+v %v", rs, err)
}
if rs, err := p.RowSecurity(ctx, u1+u2+1_000_000, schema, orders); err != nil || rs.HasBlock || rs.Template != "" {
t.Fatalf("no rules: %+v %v", rs, err)
}
if _, err := p.RowSecurity(ctx, "not-a-number", schema, orders); err == nil {
t.Fatal("non-numeric user reference accepted (must fail closed)")
}
}
func ptr[T any](v T) *T { return &v }
// --- OAuth server grant state -------------------------------------------------------------
func (s *suite) oauthClientMetadata(t *testing.T) {
st := s.Provider.OAuthClient
id := s.name("meta-client")
reg, err := st.RegisterClient(ctx, &sectypes.OAuthServerClient{
ClientID: id, RedirectURIs: []string{"https://app.example/cb"}, ClientName: "Meta",
PostLogoutRedirectURIs: []string{"https://app.example/bye"}, RequireConsent: true, FirstParty: false,
IDTokenSignedResponseAlg: "RS256", Contacts: []string{"ops@example.test"}, DPoPBoundAccessTokens: true,
})
if err != nil || reg.ClientID != id {
t.Fatalf("register: %+v %v", reg, err)
}
got, err := st.GetClient(ctx, id)
if err != nil {
t.Fatalf("get: %v", err)
}
if !got.RequireConsent || !got.DPoPBoundAccessTokens || got.IDTokenSignedResponseAlg != "RS256" ||
len(got.PostLogoutRedirectURIs) != 1 || got.PostLogoutRedirectURIs[0] != "https://app.example/bye" ||
len(got.Contacts) != 1 {
t.Fatalf("metadata lost: %+v", got)
}
got.ClientName = "Renamed"
got.RequireConsent = false
got.RedirectURIs = []string{"https://app.example/cb", "https://app.example/cb2"}
if err := st.UpdateClient(ctx, got); err != nil {
t.Fatalf("update: %v", err)
}
again, err := st.GetClient(ctx, id)
if err != nil || again.ClientName != "Renamed" || again.RequireConsent || len(again.RedirectURIs) != 2 || !again.DPoPBoundAccessTokens {
t.Fatalf("after update: %+v %v", again, err)
}
if err := st.DeleteClient(ctx, id); err != nil {
t.Fatalf("delete: %v", err)
}
_, err = st.GetClient(ctx, id)
rejected(t, "deleted client", err)
// Code extras round-trip.
code := s.name("meta-code")
err = st.SaveCode(ctx, &sectypes.OAuthCode{
Code: code, ClientID: id, RedirectURI: "https://app.example/cb", CodeChallenge: "chal", CodeChallengeMethod: "S256",
SessionToken: "sess", Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Minute),
Nonce: "n-0S6", AuthTime: 1700000000, ACR: "urn:acr:1", AMR: []string{"pwd"}, UserID: 7,
Claims: map[string]any{"id_token": map[string]any{"email": nil}}, Resource: []string{"https://api.example"}, DPoPJKT: "jkt",
})
if err != nil {
t.Fatalf("save code: %v", err)
}
c, err := st.ExchangeCode(ctx, code)
if err != nil {
t.Fatalf("exchange: %v", err)
}
if c.Nonce != "n-0S6" || c.AuthTime != 1700000000 || c.ACR != "urn:acr:1" || c.UserID != 7 || c.DPoPJKT != "jkt" ||
len(c.AMR) != 1 || len(c.Resource) != 1 || c.Claims["id_token"] == nil {
t.Fatalf("code extra lost: %+v", c)
}
}
func (s *suite) oauthGrantConsent(t *testing.T) {
g := s.Provider.OAuthGrant
client := s.name("consent-client")
if _, err := g.GetConsent(ctx, 1, client); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("missing consent: %v", err)
}
if err := g.SaveConsent(ctx, lookup.Consent{UserID: 1, ClientID: client, Scopes: []string{"openid", "email"}, ExpiresAt: time.Now().Add(time.Hour)}); err != nil {
t.Fatalf("save: %v", err)
}
c, err := g.GetConsent(ctx, 1, client)
if err != nil || len(c.Scopes) != 2 {
t.Fatalf("get: %+v %v", c, err)
}
// Saving again replaces.
if err := g.SaveConsent(ctx, lookup.Consent{UserID: 1, ClientID: client, Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Hour)}); err != nil {
t.Fatalf("resave: %v", err)
}
if c, err = g.GetConsent(ctx, 1, client); err != nil || len(c.Scopes) != 1 {
t.Fatalf("replaced: %+v %v", c, err)
}
// Another user is separate.
if _, err := g.GetConsent(ctx, 2, client); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("other user: %v", err)
}
// Expired consents are not returned.
if err := g.SaveConsent(ctx, lookup.Consent{UserID: 3, ClientID: client, Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(-time.Minute)}); err != nil {
t.Fatalf("save expired: %v", err)
}
if _, err := g.GetConsent(ctx, 3, client); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("expired consent: %v", err)
}
if err := g.RevokeConsent(ctx, 1, client); err != nil {
t.Fatalf("revoke: %v", err)
}
if _, err := g.GetConsent(ctx, 1, client); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("revoked consent: %v", err)
}
}
func (s *suite) oauthGrantRefresh(t *testing.T) {
g := s.Provider.OAuthGrant
client := s.name("refresh-client")
mk := func(n string, session string) lookup.RefreshToken {
return lookup.RefreshToken{TokenHash: s.name(n), FamilyID: s.name("fam-" + n), ClientID: client, UserID: 5,
SessionToken: session, Scopes: []string{"openid", "offline_access"},
Extra: map[string]any{"nonce": "abc"}, ExpiresAt: time.Now().Add(time.Hour)}
}
first := mk("r1", s.name("sess1"))
if err := g.SaveRefresh(ctx, first); err != nil {
t.Fatalf("save: %v", err)
}
peek, err := g.PeekRefresh(ctx, first.TokenHash)
if err != nil || peek.UserID != 5 || peek.ClientID != client || len(peek.Scopes) != 2 || peek.Extra["nonce"] != "abc" {
t.Fatalf("peek: %+v %v", peek, err)
}
if _, err := g.PeekRefresh(ctx, s.name("nope")); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("peek unknown: %v", err)
}
next := lookup.RefreshToken{TokenHash: s.name("r2"), ExpiresAt: time.Now().Add(time.Hour), Scopes: []string{"openid"}}
old, err := g.RotateRefresh(ctx, first.TokenHash, next)
if err != nil || old.FamilyID != first.FamilyID || old.SessionToken != first.SessionToken {
t.Fatalf("rotate: %+v %v", old, err)
}
// The new token belongs to the same family, client and user.
n, err := g.PeekRefresh(ctx, next.TokenHash)
if err != nil || n.FamilyID != first.FamilyID || n.ClientID != client || n.UserID != 5 || n.SessionToken != first.SessionToken {
t.Fatalf("next: %+v %v", n, err)
}
// The consumed token is still visible to Peek, so that presenting it reaches RotateRefresh.
if _, err := g.PeekRefresh(ctx, first.TokenHash); err != nil {
t.Fatalf("peek consumed: %v", err)
}
// Presenting the consumed token again is reuse: the family (including the new token) dies.
third := lookup.RefreshToken{TokenHash: s.name("r3"), ExpiresAt: time.Now().Add(time.Hour)}
reused, err := g.RotateRefresh(ctx, first.TokenHash, third)
if !errors.Is(err, lookup.ErrRefreshReused) || reused == nil || reused.FamilyID != first.FamilyID {
t.Fatalf("reuse: %+v %v", reused, err)
}
if _, err := g.PeekRefresh(ctx, next.TokenHash); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("family survived reuse: %v", err)
}
if _, err := g.RotateRefresh(ctx, next.TokenHash, lookup.RefreshToken{TokenHash: s.name("r4"), ExpiresAt: time.Now().Add(time.Hour)}); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("rotate revoked: %v", err)
}
if _, err := g.PeekRefresh(ctx, third.TokenHash); err == nil {
t.Fatal("a rejected rotation must not store the next token")
}
// Expired tokens cannot rotate.
exp := mk("rexp", s.name("sess2"))
exp.ExpiresAt = time.Now().Add(-time.Minute)
if err := g.SaveRefresh(ctx, exp); err != nil {
t.Fatalf("save expired: %v", err)
}
if _, err := g.RotateRefresh(ctx, exp.TokenHash, lookup.RefreshToken{TokenHash: s.name("rexp2"), ExpiresAt: time.Now().Add(time.Hour)}); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("rotate expired: %v", err)
}
// Revoking by family and by session.
fam := mk("rfam", s.name("sess3"))
if err := g.SaveRefresh(ctx, fam); err != nil {
t.Fatal(err)
}
if err := g.RevokeRefreshFamily(ctx, fam.FamilyID); err != nil {
t.Fatalf("revoke family: %v", err)
}
if _, err := g.PeekRefresh(ctx, fam.TokenHash); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("family revoked: %v", err)
}
bs := mk("rsess", s.name("sess4"))
if err := g.SaveRefresh(ctx, bs); err != nil {
t.Fatal(err)
}
if err := g.RevokeRefreshBySession(ctx, bs.SessionToken); err != nil {
t.Fatalf("revoke session: %v", err)
}
if _, err := g.PeekRefresh(ctx, bs.TokenHash); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("session revoked: %v", err)
}
}
func (s *suite) oauthGrantDevice(t *testing.T) {
g := s.Provider.OAuthGrant
client := s.name("device-client")
mk := func(n string) lookup.DeviceCode {
return lookup.DeviceCode{DeviceHash: s.name("dh-" + n), UserCode: strings.ToUpper(s.name("uc-" + n)), ClientID: client,
Scopes: []string{"openid"}, Interval: 1, ExpiresAt: time.Now().Add(time.Minute)}
}
// pending -> approved
d := mk("a")
if err := g.CreateDevice(ctx, d); err != nil {
t.Fatalf("create: %v", err)
}
got, err := g.DeviceByUserCode(ctx, strings.ToLower(d.UserCode)) // user codes are case-insensitive
if err != nil || got.ClientID != client || got.DeviceHash != d.DeviceHash || len(got.Scopes) != 1 {
t.Fatalf("by user code: %+v %v", got, err)
}
if _, err := g.DevicePoll(ctx, d.DeviceHash); !errors.Is(err, lookup.ErrDevicePending) {
t.Fatalf("first poll: %v", err)
}
if _, err := g.DevicePoll(ctx, d.DeviceHash); !errors.Is(err, lookup.ErrDeviceSlowDown) {
t.Fatalf("immediate re-poll: %v", err)
}
if err := g.DeviceDecide(ctx, d.UserCode, true, 9, s.name("dsess")); err != nil {
t.Fatalf("approve: %v", err)
}
if _, err := g.DeviceByUserCode(ctx, d.UserCode); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("decided code still pending: %v", err)
}
if err := g.DeviceDecide(ctx, d.UserCode, true, 9, "x"); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("second decision: %v", err)
}
time.Sleep(1100 * time.Millisecond)
done, err := g.DevicePoll(ctx, d.DeviceHash)
if err != nil || done.UserID != 9 || done.SessionToken != s.name("dsess") || done.ClientID != client {
t.Fatalf("approved poll: %+v %v", done, err)
}
if _, err := g.DevicePoll(ctx, d.DeviceHash); !errors.Is(err, lookup.ErrDeviceExpired) {
t.Fatalf("consumed code: %v", err)
}
// denied
dd := mk("d")
if err := g.CreateDevice(ctx, dd); err != nil {
t.Fatal(err)
}
if err := g.DeviceDecide(ctx, dd.UserCode, false, 0, ""); err != nil {
t.Fatalf("deny: %v", err)
}
if _, err := g.DevicePoll(ctx, dd.DeviceHash); !errors.Is(err, lookup.ErrDeviceDenied) {
t.Fatalf("denied poll: %v", err)
}
// expired
de := mk("e")
de.ExpiresAt = time.Now().Add(-time.Second)
if err := g.CreateDevice(ctx, de); err != nil {
t.Fatal(err)
}
if _, err := g.DevicePoll(ctx, de.DeviceHash); !errors.Is(err, lookup.ErrDeviceExpired) {
t.Fatalf("expired poll: %v", err)
}
if _, err := g.DeviceByUserCode(ctx, de.UserCode); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("expired by user code: %v", err)
}
if _, err := g.DevicePoll(ctx, s.name("unknown")); !errors.Is(err, lookup.ErrDeviceExpired) {
t.Fatalf("unknown poll: %v", err)
}
}
func (s *suite) oauthGrantPAR(t *testing.T) {
g := s.Provider.OAuthGrant
uri := "urn:ietf:params:oauth:request_uri:" + s.name("par")
if len(uri) > 255 {
t.Fatal("test request_uri too long")
}
req := lookup.PushedRequest{RequestURI: uri, ClientID: s.name("par-client"),
Params: map[string]string{"redirect_uri": "https://app.example/cb", "scope": "openid"}, ExpiresAt: time.Now().Add(time.Minute)}
if err := g.SavePushedRequest(ctx, req); err != nil {
t.Fatalf("save: %v", err)
}
got, err := g.ConsumePushedRequest(ctx, uri)
if err != nil || got.ClientID != req.ClientID || got.Params["scope"] != "openid" {
t.Fatalf("consume: %+v %v", got, err)
}
if _, err := g.ConsumePushedRequest(ctx, uri); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("second consume: %v", err)
}
exp := lookup.PushedRequest{RequestURI: uri + "-exp", ClientID: s.name("par-client"), Params: map[string]string{"a": "b"}, ExpiresAt: time.Now().Add(-time.Minute)}
if err := g.SavePushedRequest(ctx, exp); err != nil {
t.Fatal(err)
}
if _, err := g.ConsumePushedRequest(ctx, exp.RequestURI); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("expired consume: %v", err)
}
}
func (s *suite) oauthGrantJTI(t *testing.T) {
g := s.Provider.OAuthGrant
key := s.name("jti")
seen, err := g.SeenJTI(ctx, key, time.Now().Add(time.Minute))
if err != nil || seen {
t.Fatalf("first: %v %v", seen, err)
}
seen, err = g.SeenJTI(ctx, key, time.Now().Add(time.Minute))
if err != nil || !seen {
t.Fatalf("replay: %v %v", seen, err)
}
// An expired entry is forgotten.
old := s.name("jti-old")
if _, err := g.SeenJTI(ctx, old, time.Now().Add(-time.Minute)); err != nil {
t.Fatal(err)
}
seen, err = g.SeenJTI(ctx, old, time.Now().Add(time.Minute))
if err != nil || seen {
t.Fatalf("after expiry: %v %v", seen, err)
}
}
+44
View File
@@ -0,0 +1,44 @@
package lookup
import (
"database/sql"
"fmt"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
)
// FromDatabase extracts the *sql.DB and the dialect name from an application's
// common.Database (bun, gorm or pgsql adapter), so callers do not have to dig the
// connection out or set Config.Dialect by hand. The returned name is the adapter's
// normalised DriverName ("postgres", "sqlite", "mssql", "mysql") and is empty when
// the adapter reports a driver the dialect registry does not know; set
// Config.Dialect explicitly in that case.
//
// Transaction adapters do not expose a *sql.DB and are rejected.
func FromDatabase(db common.Database) (*sql.DB, string, error) {
if db == nil {
return nil, "", fmt.Errorf("lookup: nil database")
}
p, ok := db.(common.SQLDBProvider)
if !ok {
return nil, "", fmt.Errorf("lookup: %T does not expose a *sql.DB (transaction adapter or unsupported adapter)", db)
}
sqlDB := p.SQLDB()
if sqlDB == nil {
return nil, "", fmt.Errorf("lookup: %T has no *sql.DB", db)
}
name := db.DriverName()
if _, err := dialect.Get(name); err != nil {
name = ""
}
return sqlDB, name, nil
}
// ResolveDialect returns the dialect for db: the configured one, or detected from the driver.
func (c Config) ResolveDialect(db *sql.DB) (dialect.Dialect, error) {
if c.Dialect != "" {
return dialect.Get(c.Dialect)
}
return dialect.Detect(db)
}
@@ -428,30 +428,78 @@ EXCEPTION
END; END;
$$ LANGUAGE plpgsql; $$ LANGUAGE plpgsql;
-- ============================================
-- Column / row security tables
-- ============================================
-- A rule applies either to one user (user_id) or to every member of a group
-- (group_id via sec_group_members); exactly one of the two is set.
CREATE TABLE IF NOT EXISTS sec_group_members (
group_id INTEGER NOT NULL,
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
PRIMARY KEY (group_id, user_id)
);
CREATE TABLE IF NOT EXISTS sec_column_rules (
id SERIAL PRIMARY KEY,
user_id INTEGER REFERENCES users(id) ON DELETE CASCADE,
group_id INTEGER,
schema_name TEXT NOT NULL,
table_name TEXT NOT NULL,
column_path TEXT NOT NULL, -- dot path under the table: col or col.sub.field
access_type TEXT NOT NULL, -- e.g. mask, hide, read
mask_start INTEGER DEFAULT 0,
mask_end INTEGER DEFAULT 0,
mask_invert BOOLEAN DEFAULT false,
mask_char TEXT DEFAULT '*',
extra_filters TEXT, -- JSON object
is_active BOOLEAN NOT NULL DEFAULT true,
CHECK ((user_id IS NULL) <> (group_id IS NULL))
);
CREATE INDEX IF NOT EXISTS idx_sec_column_rules_table ON sec_column_rules(lower(schema_name), lower(table_name));
CREATE TABLE IF NOT EXISTS sec_row_rules (
id SERIAL PRIMARY KEY,
user_id INTEGER REFERENCES users(id) ON DELETE CASCADE,
group_id INTEGER,
schema_name TEXT NOT NULL,
table_name TEXT NOT NULL,
template TEXT, -- SQL fragment, e.g. 'user_id = {UserID}'
has_block BOOLEAN NOT NULL DEFAULT false,
is_active BOOLEAN NOT NULL DEFAULT true,
CHECK ((user_id IS NULL) <> (group_id IS NULL))
);
CREATE INDEX IF NOT EXISTS idx_sec_row_rules_table ON sec_row_rules(lower(schema_name), lower(table_name));
-- 8. resolvespec_column_security - Loads column security rules for user -- 8. resolvespec_column_security - Loads column security rules for user
-- Input: user_id (int), schema (text), table_name (text) -- Input: user_id (int), schema (text), table_name (text)
-- Output: p_success (bool), p_error (text), p_rules (array of security rules as jsonb) -- Output: p_success (bool), p_error (text), p_rules (array of security rules as jsonb)
-- Rules are the active sec_column_rules for the exact schema + table (case-insensitive)
-- that belong to the user or to a group the user is a member of.
-- 'control' is returned as schema.table.column_path.
CREATE OR REPLACE FUNCTION resolvespec_column_security(p_user_id integer, p_schema text, p_table_name text) CREATE OR REPLACE FUNCTION resolvespec_column_security(p_user_id integer, p_schema text, p_table_name text)
RETURNS TABLE(p_success boolean, p_error text, p_rules jsonb) AS $$ RETURNS TABLE(p_success boolean, p_error text, p_rules jsonb) AS $$
DECLARE DECLARE
v_rules jsonb; v_rules jsonb;
BEGIN BEGIN
-- Query column security rules from core.secaccess
SELECT jsonb_agg( SELECT jsonb_agg(
jsonb_build_object( jsonb_build_object(
'control', control, 'control', r.schema_name || '.' || r.table_name || '.' || r.column_path,
'accesstype', accesstype, 'accesstype', r.access_type,
'jsonvalue', jsonvalue 'jsonvalue', COALESCE(r.extra_filters, '')
) )
) )
INTO v_rules INTO v_rules
FROM core.secaccess FROM sec_column_rules r
WHERE rid_hub IN ( WHERE r.is_active = true
SELECT rid_hub_parent AND lower(r.schema_name) = lower(p_schema)
FROM core.hub_link AND lower(r.table_name) = lower(p_table_name)
WHERE rid_hub_child = p_user_id AND parent_hubtype = 'secgroup' AND (
) r.user_id = p_user_id
AND control ILIKE (p_schema || '.' || p_table_name || '%'); OR r.group_id IN (SELECT m.group_id FROM sec_group_members m WHERE m.user_id = p_user_id)
);
IF v_rules IS NULL THEN IF v_rules IS NULL THEN
v_rules := '[]'::jsonb; v_rules := '[]'::jsonb;
@@ -464,20 +512,36 @@ EXCEPTION
END; END;
$$ LANGUAGE plpgsql; $$ LANGUAGE plpgsql;
-- 9. resolvespec_row_security - Loads row security template for user (replaces core.api_sec_rowtemplate) -- 9. resolvespec_row_security - Loads the row security template for user
-- Input: schema (text), table_name (text), user_id (int) -- Input: schema (text), table_name (text), user_id (int)
-- Output: p_template (text), p_block (bool) -- Output: p_template (text), p_block (bool)
-- Applicable rules = active sec_row_rules of the user and of the user's groups for the exact
-- schema + table. Any has_block wins (template empty); otherwise templates are AND-combined,
-- each wrapped in parentheses.
CREATE OR REPLACE FUNCTION resolvespec_row_security(p_schema text, p_table_name text, p_user_id integer) CREATE OR REPLACE FUNCTION resolvespec_row_security(p_schema text, p_table_name text, p_user_id integer)
RETURNS TABLE(p_template text, p_block boolean) AS $$ RETURNS TABLE(p_template text, p_block boolean) AS $$
DECLARE
v_block boolean;
v_template text;
BEGIN BEGIN
-- Call the existing core function if it exists, or implement your own logic SELECT COALESCE(bool_or(r.has_block), false),
-- This is a placeholder that you should customize based on your core.api_sec_rowtemplate logic COALESCE(string_agg('(' || r.template || ')', ' AND ' ORDER BY r.id)
RETURN QUERY SELECT ''::text, false; FILTER (WHERE r.template IS NOT NULL AND r.template <> ''), '')
INTO v_block, v_template
FROM sec_row_rules r
WHERE r.is_active = true
AND lower(r.schema_name) = lower(p_schema)
AND lower(r.table_name) = lower(p_table_name)
AND (
r.user_id = p_user_id
OR r.group_id IN (SELECT m.group_id FROM sec_group_members m WHERE m.user_id = p_user_id)
);
-- Example implementation: IF v_block THEN
-- RETURN QUERY SELECT template, has_block v_template := '';
-- FROM core.row_security_config END IF;
-- WHERE schema_name = p_schema AND table_name = p_table_name AND user_id = p_user_id;
RETURN QUERY SELECT v_template, v_block;
END; END;
$$ LANGUAGE plpgsql; $$ LANGUAGE plpgsql;
@@ -650,7 +714,7 @@ BEGIN
v_auth_provider := COALESCE(p_user_data->>'auth_provider', 'oauth2'); v_auth_provider := COALESCE(p_user_data->>'auth_provider', 'oauth2');
-- Convert roles array to comma-separated string -- Convert roles array to comma-separated string
SELECT array_to_string(ARRAY(SELECT jsonb_array_elements_text(p_user_data->'roles')), ',') SELECT array_to_string(ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_user_data->'roles') = 'array' THEN p_user_data->'roles' ELSE '[]'::jsonb END)), ',')
INTO v_roles; INTO v_roles;
-- Try to find existing user by email -- Try to find existing user by email
@@ -701,7 +765,7 @@ BEGIN
v_access_token := p_session_data->>'access_token'; v_access_token := p_session_data->>'access_token';
v_refresh_token := p_session_data->>'refresh_token'; v_refresh_token := p_session_data->>'refresh_token';
v_token_type := COALESCE(p_session_data->>'token_type', 'Bearer'); v_token_type := COALESCE(p_session_data->>'token_type', 'Bearer');
v_expires_at := (p_session_data->>'expires_at')::timestamp; v_expires_at := (p_session_data->>'expires_at')::timestamptz::timestamp;
v_auth_provider := COALESCE(p_session_data->>'auth_provider', 'oauth2'); v_auth_provider := COALESCE(p_session_data->>'auth_provider', 'oauth2');
-- Insert or update session -- Insert or update session
@@ -857,7 +921,7 @@ BEGIN
v_new_session_token := p_update_data->>'new_session_token'; v_new_session_token := p_update_data->>'new_session_token';
v_new_access_token := p_update_data->>'new_access_token'; v_new_access_token := p_update_data->>'new_access_token';
v_new_refresh_token := p_update_data->>'new_refresh_token'; v_new_refresh_token := p_update_data->>'new_refresh_token';
v_expires_at := (p_update_data->>'expires_at')::timestamp; v_expires_at := (p_update_data->>'expires_at')::timestamptz::timestamp;
-- Update session in user_sessions table -- Update session in user_sessions table
UPDATE user_sessions UPDATE user_sessions
@@ -1214,7 +1278,7 @@ BEGIN
-- Convert transports array -- Convert transports array
IF p_credential->'transports' IS NOT NULL THEN IF p_credential->'transports' IS NOT NULL THEN
SELECT ARRAY(SELECT jsonb_array_elements_text(p_credential->'transports')) SELECT ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_credential->'transports') = 'array' THEN p_credential->'transports' ELSE '[]'::jsonb END))
INTO v_transports; INTO v_transports;
END IF; END IF;
@@ -1304,12 +1368,11 @@ BEGIN
'name', name, 'name', name,
'created_at', created_at, 'created_at', created_at,
'last_used_at', last_used_at 'last_used_at', last_used_at
) ) ORDER BY created_at DESC
), '[]'::jsonb) ), '[]'::jsonb)
INTO v_credentials INTO v_credentials
FROM user_passkey_credentials FROM user_passkey_credentials
WHERE user_id = p_user_id WHERE user_id = p_user_id;
ORDER BY created_at DESC;
RETURN QUERY SELECT true, NULL::text, v_credentials; RETURN QUERY SELECT true, NULL::text, v_credentials;
EXCEPTION EXCEPTION
@@ -1451,6 +1514,64 @@ EXCEPTION
END; END;
$$ LANGUAGE plpgsql; $$ LANGUAGE plpgsql;
-- 8. resolvespec_passkey_login - Creates a session for a user whose passkey assertion was verified
-- Input: p_request (jsonb) {user_id: int, ip_address: string, user_agent: string}
-- Output: p_success (bool), p_error (text), p_data (LoginResponse as jsonb)
CREATE OR REPLACE FUNCTION resolvespec_passkey_login(p_request jsonb)
RETURNS TABLE(p_success boolean, p_error text, p_data jsonb) AS $$
DECLARE
v_user_id INTEGER;
v_username TEXT;
v_email TEXT;
v_user_level INTEGER;
v_roles TEXT;
v_program_user_id INTEGER;
v_program_user_table TEXT;
v_session_token TEXT;
BEGIN
v_user_id := (p_request->>'user_id')::integer;
SELECT username, email, user_level, roles, program_user_id, program_user_table
INTO v_username, v_email, v_user_level, v_roles, v_program_user_id, v_program_user_table
FROM users
WHERE id = v_user_id AND is_active = true;
IF NOT FOUND THEN
RETURN QUERY SELECT false, 'User not found'::text, NULL::jsonb;
RETURN;
END IF;
v_session_token := 'sess_' || encode(gen_random_bytes(32), 'hex') || '_' || extract(epoch from now())::bigint::text;
INSERT INTO user_sessions (session_token, user_id, expires_at, ip_address, user_agent, last_activity_at)
VALUES (v_session_token, v_user_id, now() + interval '24 hours',
p_request->>'ip_address', p_request->>'user_agent', now());
UPDATE users SET last_login_at = now() WHERE id = v_user_id;
RETURN QUERY SELECT
true,
NULL::text,
jsonb_build_object(
'token', v_session_token,
'user', jsonb_build_object(
'user_id', v_user_id,
'user_name', v_username,
'email', v_email,
'user_level', v_user_level,
'roles', string_to_array(COALESCE(v_roles, ''), ','),
'session_id', v_session_token,
'program_user_id', COALESCE(v_program_user_id, 0),
'program_user_table', COALESCE(v_program_user_table, '')
),
'expires_in', 86400
);
EXCEPTION
WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM::text, NULL::jsonb;
END;
$$ LANGUAGE plpgsql;
-- ============================================ -- ============================================
-- Example: Test Passkey stored procedures -- Example: Test Passkey stored procedures
-- ============================================ -- ============================================
@@ -1644,8 +1765,10 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
client_secret_hash TEXT, -- sha256 hex of the confidential-client secret; NULL for public clients client_secret_hash TEXT, -- sha256 hex of the confidential-client secret; NULL for public clients
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none', token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
is_active BOOLEAN DEFAULT true, is_active BOOLEAN DEFAULT true,
metadata jsonb, -- every other RFC 7591 field (see sectypes.OAuthServerClient)
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
); );
ALTER TABLE oauth_clients ADD COLUMN IF NOT EXISTS metadata jsonb;
-- oauth_codes: short-lived authorization codes (for multi-instance deployments) -- oauth_codes: short-lived authorization codes (for multi-instance deployments)
-- Note: client_id is stored without a foreign key so codes can be persisted even -- Note: client_id is stored without a foreign key so codes can be persisted even
@@ -1662,8 +1785,10 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
refresh_token TEXT, refresh_token TEXT,
scopes TEXT[], scopes TEXT[],
expires_at TIMESTAMP NOT NULL, expires_at TIMESTAMP NOT NULL,
extra jsonb, -- nonce, auth_time, acr, claims, user_id, dpop_jkt ... (see sectypes.OAuthCode)
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
); );
ALTER TABLE oauth_codes ADD COLUMN IF NOT EXISTS extra jsonb;
CREATE INDEX IF NOT EXISTS idx_oauth_codes_code ON oauth_codes(code); 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 INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
@@ -1671,26 +1796,27 @@ CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
-- OAuth2 Server Stored Procedures -- OAuth2 Server Stored Procedures
-- ============================================ -- ============================================
CREATE OR REPLACE FUNCTION resolvespec_oauth_register_client(p_data jsonb) CREATE OR REPLACE FUNCTION resolvespec_oauth_register_client(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb) RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$ LANGUAGE plpgsql AS $$
DECLARE DECLARE
v_client_id text; v_client_id text;
v_row jsonb; v_row jsonb;
BEGIN BEGIN
v_client_id := p_data->>'client_id'; v_client_id := p_request->>'client_id';
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method) INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method, metadata)
VALUES ( VALUES (
v_client_id, v_client_id,
ARRAY(SELECT jsonb_array_elements_text(p_data->'redirect_uris')), ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)),
p_data->>'client_name', p_request->>'client_name',
COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'grant_types')), ARRAY['authorization_code']), CASE WHEN jsonb_typeof(p_request->'grant_types') = 'array' AND jsonb_array_length(p_request->'grant_types') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'grant_types')) ELSE ARRAY['authorization_code'] END,
COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'allowed_scopes')), ARRAY['openid','profile','email']), CASE WHEN jsonb_typeof(p_request->'allowed_scopes') = 'array' AND jsonb_array_length(p_request->'allowed_scopes') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'allowed_scopes')) ELSE ARRAY['openid','profile','email'] END,
NULLIF(p_data->>'client_secret_hash', ''), NULLIF(p_request->>'client_secret_hash', ''),
COALESCE(NULLIF(p_data->>'token_endpoint_auth_method', ''), 'none') COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), 'none'),
NULLIF(p_request - ARRAY['client_id','redirect_uris','client_name','grant_types','allowed_scopes','client_secret_hash','token_endpoint_auth_method'], '{}'::jsonb)
) )
RETURNING to_jsonb(oauth_clients.*) INTO v_row; RETURNING (to_jsonb(oauth_clients.*) - 'metadata') || COALESCE(metadata, '{}'::jsonb) INTO v_row;
RETURN QUERY SELECT true, null::text, v_row; RETURN QUERY SELECT true, null::text, v_row;
EXCEPTION WHEN OTHERS THEN EXCEPTION WHEN OTHERS THEN
@@ -1704,7 +1830,7 @@ LANGUAGE plpgsql AS $$
DECLARE DECLARE
v_row jsonb; v_row jsonb;
BEGIN BEGIN
SELECT to_jsonb(oauth_clients.*) SELECT (to_jsonb(oauth_clients.*) - 'metadata') || COALESCE(metadata, '{}'::jsonb)
INTO v_row INTO v_row
FROM oauth_clients FROM oauth_clients
WHERE client_id = p_client_id AND is_active = true; WHERE client_id = p_client_id AND is_active = true;
@@ -1717,22 +1843,23 @@ BEGIN
END; END;
$$; $$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_code(p_data jsonb) CREATE OR REPLACE FUNCTION resolvespec_oauth_save_code(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text) RETURNS TABLE(p_success bool, p_error text)
LANGUAGE plpgsql AS $$ LANGUAGE plpgsql AS $$
BEGIN BEGIN
INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at) INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at, extra)
VALUES ( VALUES (
p_data->>'code', p_request->>'code',
p_data->>'client_id', p_request->>'client_id',
p_data->>'redirect_uri', p_request->>'redirect_uri',
p_data->>'client_state', p_request->>'client_state',
p_data->>'code_challenge', p_request->>'code_challenge',
COALESCE(p_data->>'code_challenge_method', 'S256'), COALESCE(p_request->>'code_challenge_method', 'S256'),
p_data->>'session_token', p_request->>'session_token',
p_data->>'refresh_token', p_request->>'refresh_token',
ARRAY(SELECT jsonb_array_elements_text(p_data->'scopes')), ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'scopes') = 'array' THEN p_request->'scopes' ELSE '[]'::jsonb END)),
(p_data->>'expires_at')::timestamp (p_request->>'expires_at')::timestamptz::timestamp,
NULLIF(p_request - ARRAY['code','client_id','redirect_uri','client_state','code_challenge','code_challenge_method','session_token','refresh_token','scopes','expires_at'], '{}'::jsonb)
); );
RETURN QUERY SELECT true, null::text; RETURN QUERY SELECT true, null::text;
@@ -1758,7 +1885,7 @@ BEGIN
'session_token', session_token, 'session_token', session_token,
'refresh_token', refresh_token, 'refresh_token', refresh_token,
'scopes', to_jsonb(scopes) 'scopes', to_jsonb(scopes)
) INTO v_row; ) || COALESCE(extra, '{}'::jsonb) INTO v_row;
IF v_row IS NULL THEN IF v_row IS NULL THEN
RETURN QUERY SELECT false, 'invalid or expired code'::text, null::jsonb; RETURN QUERY SELECT false, 'invalid or expired code'::text, null::jsonb;
@@ -1809,3 +1936,374 @@ BEGIN
RETURN QUERY SELECT true, null::text; RETURN QUERY SELECT true, null::text;
END; END;
$$; $$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_update_client(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text)
LANGUAGE plpgsql AS $$
DECLARE
v_rows int;
BEGIN
UPDATE oauth_clients SET
redirect_uris = ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)),
client_name = p_request->>'client_name',
grant_types = CASE WHEN jsonb_typeof(p_request->'grant_types') = 'array' AND jsonb_array_length(p_request->'grant_types') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'grant_types')) ELSE grant_types END,
allowed_scopes = CASE WHEN jsonb_typeof(p_request->'allowed_scopes') = 'array' AND jsonb_array_length(p_request->'allowed_scopes') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'allowed_scopes')) ELSE allowed_scopes END,
client_secret_hash = NULLIF(p_request->>'client_secret_hash', ''),
token_endpoint_auth_method = COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), token_endpoint_auth_method),
metadata = NULLIF(p_request - ARRAY['client_id','redirect_uris','client_name','grant_types','allowed_scopes','client_secret_hash','token_endpoint_auth_method'], '{}'::jsonb)
WHERE client_id = p_request->>'client_id' AND is_active = true;
GET DIAGNOSTICS v_rows = ROW_COUNT;
IF v_rows = 0 THEN
RETURN QUERY SELECT false, 'client not found'::text;
ELSE
RETURN QUERY SELECT true, null::text;
END IF;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_delete_client(p_client_id text)
RETURNS TABLE(p_success bool, p_error text)
LANGUAGE plpgsql AS $$
BEGIN
UPDATE oauth_clients SET is_active = false WHERE client_id = p_client_id;
RETURN QUERY SELECT true, null::text;
END;
$$;
-- ============================================
-- OAuth2 Server grant state (consents, refresh tokens, device codes, PAR, replay cache)
-- ============================================
-- Procedure-backend tables use jsonb for scopes/extra/params. Every procedure takes one jsonb
-- request and returns (p_success, p_error, p_data). p_error carries a stable code for the
-- failures the Go side maps to errors: not_found, refresh_invalid, refresh_reused,
-- device_pending, device_slowdown, device_denied, device_expired.
CREATE TABLE IF NOT EXISTS oauth_consents (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL,
client_id VARCHAR(255) NOT NULL,
scopes jsonb,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
id SERIAL PRIMARY KEY,
token_hash VARCHAR(64) NOT NULL UNIQUE, -- sha256 hex of the raw refresh token
family_id VARCHAR(64) NOT NULL,
client_id VARCHAR(255) NOT NULL,
user_id INTEGER NOT NULL,
session_token VARCHAR(255),
scopes jsonb,
extra jsonb,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
used_at TIMESTAMP,
revoked_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
CREATE TABLE IF NOT EXISTS oauth_device_codes (
id SERIAL PRIMARY KEY,
device_hash VARCHAR(64) NOT NULL UNIQUE,
user_code VARCHAR(32) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
scopes jsonb,
status VARCHAR(16) NOT NULL DEFAULT 'pending',
user_id INTEGER,
session_token VARCHAR(255),
poll_interval INTEGER NOT NULL DEFAULT 5,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
last_polled_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_device_expires ON oauth_device_codes(expires_at);
CREATE TABLE IF NOT EXISTS oauth_par_requests (
id SERIAL PRIMARY KEY,
request_uri VARCHAR(255) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
params jsonb,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_par_expires ON oauth_par_requests(expires_at);
CREATE TABLE IF NOT EXISTS oauth_jti (
id SERIAL PRIMARY KEY,
jti_key VARCHAR(255) NOT NULL UNIQUE,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_jti_expires ON oauth_jti(expires_at);
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_consent(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
DELETE FROM oauth_consents
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id';
INSERT INTO oauth_consents (user_id, client_id, scopes, expires_at)
VALUES ((p_request->>'user_id')::int, p_request->>'client_id', COALESCE(p_request->'scopes', '[]'::jsonb),
(p_request->>'expires_at')::timestamptz::timestamp);
RETURN QUERY SELECT true, null::text, null::jsonb;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_get_consent(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_row jsonb;
BEGIN
SELECT jsonb_build_object('user_id', user_id, 'client_id', client_id, 'scopes', COALESCE(scopes, '[]'::jsonb), 'expires_at', expires_at)
INTO v_row
FROM oauth_consents
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id' AND expires_at > now();
IF v_row IS NULL THEN
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
ELSE
RETURN QUERY SELECT true, null::text, v_row;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_consent(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
DELETE FROM oauth_consents
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id';
RETURN QUERY SELECT true, null::text, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_refresh(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
INSERT INTO oauth_refresh_tokens (token_hash, family_id, client_id, user_id, session_token, scopes, extra, expires_at)
VALUES (p_request->>'token_hash', p_request->>'family_id', p_request->>'client_id', (p_request->>'user_id')::int,
p_request->>'session_token', p_request->'scopes', p_request->'extra',
(p_request->>'expires_at')::timestamptz::timestamp);
RETURN QUERY SELECT true, null::text, null::jsonb;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_rotate_refresh(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
r oauth_refresh_tokens%ROWTYPE;
v_next jsonb := p_request->'next';
v_old jsonb;
BEGIN
SELECT * INTO r FROM oauth_refresh_tokens WHERE token_hash = p_request->>'old_hash' FOR UPDATE;
IF NOT FOUND OR r.revoked_at IS NOT NULL OR r.expires_at <= now() THEN
RETURN QUERY SELECT false, 'refresh_invalid'::text, null::jsonb;
RETURN;
END IF;
v_old := jsonb_build_object('token_hash', r.token_hash, 'family_id', r.family_id, 'client_id', r.client_id,
'user_id', r.user_id, 'session_token', r.session_token,
'scopes', COALESCE(r.scopes, '[]'::jsonb), 'extra', COALESCE(r.extra, '{}'::jsonb),
'expires_at', r.expires_at);
IF r.used_at IS NOT NULL THEN
-- A rotated token came back: revoke the whole family. Returning (not raising) keeps the revoke.
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE family_id = r.family_id AND revoked_at IS NULL;
RETURN QUERY SELECT false, 'refresh_reused'::text, v_old;
RETURN;
END IF;
UPDATE oauth_refresh_tokens SET used_at = now() WHERE id = r.id;
INSERT INTO oauth_refresh_tokens (token_hash, family_id, client_id, user_id, session_token, scopes, extra, expires_at)
VALUES (v_next->>'token_hash', r.family_id, r.client_id, r.user_id,
COALESCE(NULLIF(v_next->>'session_token', ''), r.session_token),
v_next->'scopes', v_next->'extra', (v_next->>'expires_at')::timestamptz::timestamp);
RETURN QUERY SELECT true, null::text, v_old;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_peek_refresh(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_row jsonb;
BEGIN
SELECT jsonb_build_object('token_hash', token_hash, 'family_id', family_id, 'client_id', client_id,
'user_id', user_id, 'session_token', session_token,
'scopes', COALESCE(scopes, '[]'::jsonb), 'extra', COALESCE(extra, '{}'::jsonb),
'expires_at', expires_at)
INTO v_row
FROM oauth_refresh_tokens
WHERE token_hash = p_request->>'token_hash' AND revoked_at IS NULL AND expires_at > now();
IF v_row IS NULL THEN
RETURN QUERY SELECT false, 'refresh_invalid'::text, null::jsonb;
ELSE
RETURN QUERY SELECT true, null::text, v_row;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_refresh_family(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE family_id = p_request->>'family_id' AND revoked_at IS NULL;
RETURN QUERY SELECT true, null::text, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_refresh_session(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE session_token = p_request->>'session_token' AND revoked_at IS NULL;
RETURN QUERY SELECT true, null::text, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_create_device(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
INSERT INTO oauth_device_codes (device_hash, user_code, client_id, scopes, status, poll_interval, expires_at)
VALUES (p_request->>'device_hash', upper(p_request->>'user_code'), p_request->>'client_id', p_request->'scopes',
COALESCE(NULLIF(p_request->>'status', ''), 'pending'), COALESCE((p_request->>'interval')::int, 5),
(p_request->>'expires_at')::timestamptz::timestamp);
RETURN QUERY SELECT true, null::text, null::jsonb;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_by_user_code(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_row jsonb;
BEGIN
SELECT jsonb_build_object('device_hash', device_hash, 'user_code', user_code, 'client_id', client_id,
'scopes', COALESCE(scopes, '[]'::jsonb), 'status', status, 'interval', poll_interval,
'expires_at', expires_at)
INTO v_row
FROM oauth_device_codes
WHERE user_code = upper(p_request->>'user_code') AND status = 'pending' AND expires_at > now();
IF v_row IS NULL THEN
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
ELSE
RETURN QUERY SELECT true, null::text, v_row;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_decide(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_rows int;
v_approve boolean := COALESCE((p_request->>'approve')::boolean, false);
BEGIN
UPDATE oauth_device_codes
SET status = CASE WHEN v_approve THEN 'approved' ELSE 'denied' END,
user_id = CASE WHEN v_approve THEN (p_request->>'user_id')::int ELSE user_id END,
session_token = CASE WHEN v_approve THEN p_request->>'session_token' ELSE session_token END
WHERE user_code = upper(p_request->>'user_code') AND status = 'pending' AND expires_at > now();
GET DIAGNOSTICS v_rows = ROW_COUNT;
IF v_rows = 0 THEN
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
ELSE
RETURN QUERY SELECT true, null::text, null::jsonb;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_poll(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
d oauth_device_codes%ROWTYPE;
v_slow boolean;
BEGIN
SELECT * INTO d FROM oauth_device_codes WHERE device_hash = p_request->>'device_hash' FOR UPDATE;
IF NOT FOUND THEN
RETURN QUERY SELECT false, 'device_expired'::text, null::jsonb;
RETURN;
END IF;
IF d.expires_at <= now() THEN
DELETE FROM oauth_device_codes WHERE id = d.id;
RETURN QUERY SELECT false, 'device_expired'::text, null::jsonb;
RETURN;
END IF;
v_slow := d.last_polled_at IS NOT NULL AND (now() - d.last_polled_at) < make_interval(secs => d.poll_interval);
UPDATE oauth_device_codes SET last_polled_at = now() WHERE id = d.id;
IF v_slow THEN
RETURN QUERY SELECT false, 'device_slowdown'::text, null::jsonb;
ELSIF d.status = 'denied' THEN
DELETE FROM oauth_device_codes WHERE id = d.id;
RETURN QUERY SELECT false, 'device_denied'::text, null::jsonb;
ELSIF d.status = 'approved' THEN
DELETE FROM oauth_device_codes WHERE id = d.id;
RETURN QUERY SELECT true, null::text, jsonb_build_object('device_hash', d.device_hash, 'user_code', d.user_code,
'client_id', d.client_id, 'scopes', COALESCE(d.scopes, '[]'::jsonb), 'status', d.status,
'user_id', d.user_id, 'session_token', d.session_token, 'interval', d.poll_interval, 'expires_at', d.expires_at);
ELSE
RETURN QUERY SELECT false, 'device_pending'::text, null::jsonb;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_par(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
INSERT INTO oauth_par_requests (request_uri, client_id, params, expires_at)
VALUES (p_request->>'request_uri', p_request->>'client_id', p_request->'params',
(p_request->>'expires_at')::timestamptz::timestamp);
RETURN QUERY SELECT true, null::text, null::jsonb;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_consume_par(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_row jsonb;
BEGIN
DELETE FROM oauth_par_requests
WHERE request_uri = p_request->>'request_uri' AND expires_at > now()
RETURNING jsonb_build_object('request_uri', request_uri, 'client_id', client_id,
'params', COALESCE(params, '{}'::jsonb), 'expires_at', expires_at)
INTO v_row;
IF v_row IS NULL THEN
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
ELSE
RETURN QUERY SELECT true, null::text, v_row;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_seen_jti(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_rows int;
BEGIN
DELETE FROM oauth_jti WHERE expires_at < now();
INSERT INTO oauth_jti (jti_key, expires_at)
VALUES (p_request->>'key', (p_request->>'expires_at')::timestamptz::timestamp)
ON CONFLICT (jti_key) DO NOTHING;
GET DIAGNOSTICS v_rows = ROW_COUNT;
RETURN QUERY SELECT true, null::text, jsonb_build_object('seen', v_rows = 0);
END;
$$;
+84
View File
@@ -0,0 +1,84 @@
package lookup_test
import (
"database/sql"
"testing"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/sqlitedialect"
"github.com/uptrace/bun/driver/sqliteshim"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
func TestFromDatabase(t *testing.T) {
sqldb, err := sql.Open(sqliteshim.ShimName, "file::memory:?cache=shared")
if err != nil {
t.Fatal(err)
}
defer sqldb.Close()
gdb, err := gorm.Open(sqlite.Open("file::memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
cases := map[string]struct {
db common.Database
same *sql.DB // expected handle; nil = just must be non-nil
}{
"pgsql": {database.NewPgSQLAdapter(sqldb, "sqlite"), sqldb},
"bun": {database.NewBunAdapter(bun.NewDB(sqldb, sqlitedialect.New())), sqldb},
"gorm": {database.NewGormAdapter(gdb), nil},
}
for name, c := range cases {
got, dialectName, err := lookup.FromDatabase(c.db)
if err != nil {
t.Errorf("%s: %v", name, err)
continue
}
if got == nil || (c.same != nil && got != c.same) {
t.Errorf("%s: unexpected *sql.DB %v", name, got)
}
if dialectName != "sqlite" {
t.Errorf("%s: dialect = %q, want sqlite", name, dialectName)
}
if err := got.Ping(); err != nil {
t.Errorf("%s: handle not usable: %v", name, err)
}
}
}
func TestFromDatabaseRejects(t *testing.T) {
if _, _, err := lookup.FromDatabase(nil); err == nil {
t.Error("nil database should fail")
}
// A database that does not expose a *sql.DB (the embedded interface is nil; only the type matters).
if _, _, err := lookup.FromDatabase(struct{ common.Database }{}); err == nil {
t.Error("adapter without SQLDB should fail")
}
}
func TestResolveDialect(t *testing.T) {
sqldb, err := sql.Open(sqliteshim.ShimName, "file::memory:")
if err != nil {
t.Fatal(err)
}
defer sqldb.Close()
d, err := lookup.Config{}.ResolveDialect(sqldb)
if err != nil || d.Name() != "sqlite" {
t.Errorf("detected = %v, %v", d, err)
}
d, err = lookup.Config{Dialect: "mysql"}.ResolveDialect(sqldb)
if err != nil || d.Name() != "mysql" {
t.Errorf("explicit dialect should win: %v, %v", d, err)
}
if _, err := (lookup.Config{Dialect: "oracle"}).Resolve(); err == nil {
t.Error("unknown dialect should fail Resolve")
}
}
+50
View File
@@ -0,0 +1,50 @@
// Package ddl holds the reference table schemas for the lookup direct backend, one per
// dialect. They use the lookup.DefaultSchema table and column names; copy and adapt them
// when you override names through lookup.Config.Schema.
//
// The Postgres file creates tables only. The stored-procedure schema
// (lookup/database_schema.sql) is a separate script with native bytea / text[] columns and
// must not be combined with it.
package ddl
import (
"embed"
"fmt"
"strings"
)
//go:embed postgres.sql sqlite.sql mysql.sql mssql.sql
var files embed.FS
// SQL returns the schema script for a dialect name ("postgres", "sqlite", "mysql", "mssql").
func SQL(dialect string) (string, error) {
b, err := files.ReadFile(dialect + ".sql")
if err != nil {
return "", fmt.Errorf("ddl: no reference schema for dialect %q", dialect)
}
return string(b), nil
}
// Statements returns the schema as separate statements, for drivers that reject
// multi-statement execution. Comment-only lines are dropped.
func Statements(dialect string) ([]string, error) {
s, err := SQL(dialect)
if err != nil {
return nil, err
}
var out []string
var cur strings.Builder
for _, line := range strings.Split(s, "\n") {
t := strings.TrimSpace(line)
if t == "" || strings.HasPrefix(t, "--") {
continue
}
cur.WriteString(line)
cur.WriteString("\n")
if strings.HasSuffix(t, ";") {
out = append(out, strings.TrimSpace(cur.String()))
cur.Reset()
}
}
return out, nil
}
+58
View File
@@ -0,0 +1,58 @@
package ddl_test
import (
"database/sql"
"strings"
"testing"
_ "github.com/glebarez/go-sqlite"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl"
)
func TestStatementsAllDialects(t *testing.T) {
for _, d := range []string{"postgres", "sqlite", "mysql", "mssql"} {
st, err := ddl.Statements(d)
if err != nil {
t.Fatal(err)
}
if len(st) < 12 {
t.Errorf("%s: %d statements", d, len(st))
}
all := strings.Join(st, "\n")
for _, tbl := range []string{"users", "user_sessions", "user_keys", "oauth_codes", "sec_group_members", "sec_column_rules", "sec_row_rules"} {
if !strings.Contains(all, tbl+" (") {
t.Errorf("%s: missing table %s", d, tbl)
}
}
}
if _, err := ddl.SQL("oracle"); err == nil {
t.Error("unknown dialect must error")
}
}
// Every table and column the default schema names must exist in the sqlite reference DDL.
func TestSQLiteMatchesDefaultSchema(t *testing.T) {
db, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatal(err)
}
defer db.Close()
db.SetMaxOpenConns(1)
s, _ := ddl.SQL("sqlite")
if _, err := db.Exec(s); err != nil {
t.Fatal(err)
}
if _, err := db.Exec(s); err != nil {
t.Fatalf("script must be re-runnable: %v", err)
}
sc := lookup.DefaultSchema()
for _, tbl := range sc {
for _, c := range tbl.Columns {
if _, err := db.Exec("SELECT " + c + " FROM " + tbl.Name + " WHERE 1=0"); err != nil {
t.Errorf("%s.%s: %v", tbl.Name, c, err)
}
}
}
}
+344
View File
@@ -0,0 +1,344 @@
-- Reference schema for the lookup direct backend: Microsoft SQL Server 2016+. Direct backend only. Run each statement separately (see ddl.Statements).
-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema.
-- Generated by the project, edit freely for your deployment (types, collations, extra columns).
-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login
-- roles: comma-separated roles
IF OBJECT_ID(N'users', N'U') IS NULL
CREATE TABLE users (
id INT IDENTITY(1,1) PRIMARY KEY,
username NVARCHAR(255) NOT NULL UNIQUE,
email NVARCHAR(255) NOT NULL UNIQUE,
password NVARCHAR(255),
user_level INT DEFAULT 0,
roles NVARCHAR(500),
is_active BIT DEFAULT 1,
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
updated_at DATETIME2 DEFAULT SYSUTCDATETIME(),
last_login_at DATETIME2,
program_user_id INT DEFAULT 0,
program_user_table NVARCHAR(255) DEFAULT '',
remote_id NVARCHAR(255),
auth_provider NVARCHAR(50),
totp_secret NVARCHAR(255),
totp_enabled BIT DEFAULT 0,
totp_enabled_at DATETIME2
);
IF OBJECT_ID(N'user_sessions', N'U') IS NULL
CREATE TABLE user_sessions (
id INT IDENTITY(1,1) PRIMARY KEY,
session_token NVARCHAR(450) NOT NULL UNIQUE,
user_id INT NOT NULL,
expires_at DATETIME2 NOT NULL,
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
last_activity_at DATETIME2 DEFAULT SYSUTCDATETIME(),
ip_address NVARCHAR(45),
user_agent NVARCHAR(MAX),
access_token NVARCHAR(MAX),
refresh_token NVARCHAR(MAX),
token_type NVARCHAR(50) DEFAULT 'Bearer',
auth_provider NVARCHAR(50),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_sessions_user_id' AND object_id = OBJECT_ID(N'user_sessions'))
CREATE INDEX idx_user_sessions_user_id ON user_sessions(user_id);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_sessions_expires_at' AND object_id = OBJECT_ID(N'user_sessions'))
CREATE INDEX idx_user_sessions_expires_at ON user_sessions(expires_at);
IF OBJECT_ID(N'token_blacklist', N'U') IS NULL
CREATE TABLE token_blacklist (
id INT IDENTITY(1,1) PRIMARY KEY,
token NVARCHAR(500) NOT NULL,
user_id INT,
expires_at DATETIME2 NOT NULL,
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
-- code_hash: SHA-256 hex of the backup code
IF OBJECT_ID(N'user_totp_backup_codes', N'U') IS NULL
CREATE TABLE user_totp_backup_codes (
id INT IDENTITY(1,1) PRIMARY KEY,
user_id INT NOT NULL,
code_hash NVARCHAR(64) NOT NULL,
used BIT DEFAULT 0,
used_at DATETIME2,
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_totp_user_id' AND object_id = OBJECT_ID(N'user_totp_backup_codes'))
CREATE INDEX idx_totp_user_id ON user_totp_backup_codes(user_id);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_totp_code_hash' AND object_id = OBJECT_ID(N'user_totp_backup_codes'))
CREATE INDEX idx_totp_code_hash ON user_totp_backup_codes(code_hash);
-- credential_id: base64 text
-- public_key: base64 text
-- aaguid: base64 text
-- transports: JSON-encoded array
IF OBJECT_ID(N'user_passkey_credentials', N'U') IS NULL
CREATE TABLE user_passkey_credentials (
id INT IDENTITY(1,1) PRIMARY KEY,
user_id INT NOT NULL,
credential_id VARCHAR(900) NOT NULL UNIQUE,
public_key NVARCHAR(MAX) NOT NULL,
attestation_type NVARCHAR(50) DEFAULT 'none',
aaguid NVARCHAR(MAX),
sign_count INT DEFAULT 0,
clone_warning BIT DEFAULT 0,
transports NVARCHAR(MAX),
backup_eligible BIT DEFAULT 0,
backup_state BIT DEFAULT 0,
name NVARCHAR(255),
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
last_used_at DATETIME2 DEFAULT SYSUTCDATETIME(),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_passkey_user_id' AND object_id = OBJECT_ID(N'user_passkey_credentials'))
CREATE INDEX idx_passkey_user_id ON user_passkey_credentials(user_id);
-- token_hash: SHA-256 hex of the raw token
IF OBJECT_ID(N'user_password_resets', N'U') IS NULL
CREATE TABLE user_password_resets (
id INT IDENTITY(1,1) PRIMARY KEY,
user_id INT NOT NULL,
token_hash NVARCHAR(64) NOT NULL UNIQUE,
expires_at DATETIME2 NOT NULL,
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
used BIT DEFAULT 0,
used_at DATETIME2,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_pw_reset_user_id' AND object_id = OBJECT_ID(N'user_password_resets'))
CREATE INDEX idx_pw_reset_user_id ON user_password_resets(user_id);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_pw_reset_expires_at' AND object_id = OBJECT_ID(N'user_password_resets'))
CREATE INDEX idx_pw_reset_expires_at ON user_password_resets(expires_at);
-- redirect_uris: JSON-encoded array
-- grant_types: JSON-encoded array
-- allowed_scopes: JSON-encoded array
-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients
IF OBJECT_ID(N'oauth_clients', N'U') IS NULL
CREATE TABLE oauth_clients (
id INT IDENTITY(1,1) PRIMARY KEY,
client_id NVARCHAR(255) NOT NULL UNIQUE,
redirect_uris NVARCHAR(MAX) NOT NULL,
client_name NVARCHAR(255),
grant_types NVARCHAR(MAX),
allowed_scopes NVARCHAR(MAX),
client_secret_hash NVARCHAR(MAX),
token_endpoint_auth_method NVARCHAR(30) DEFAULT 'none',
is_active BIT DEFAULT 1,
metadata NVARCHAR(MAX),
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
);
-- scopes: JSON-encoded array
IF OBJECT_ID(N'oauth_codes', N'U') IS NULL
CREATE TABLE oauth_codes (
id INT IDENTITY(1,1) PRIMARY KEY,
code NVARCHAR(255) NOT NULL UNIQUE,
client_id NVARCHAR(255) NOT NULL,
redirect_uri NVARCHAR(MAX) NOT NULL,
client_state NVARCHAR(MAX),
code_challenge NVARCHAR(255) NOT NULL,
code_challenge_method NVARCHAR(10) DEFAULT 'S256',
session_token NVARCHAR(MAX) NOT NULL,
refresh_token NVARCHAR(MAX),
scopes NVARCHAR(MAX),
expires_at DATETIME2 NOT NULL,
extra NVARCHAR(MAX),
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_codes_expires' AND object_id = OBJECT_ID(N'oauth_codes'))
CREATE INDEX idx_oauth_codes_expires ON oauth_codes(expires_at);
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
IF OBJECT_ID(N'oauth_consents', N'U') IS NULL
CREATE TABLE oauth_consents (
id INT IDENTITY(1,1) PRIMARY KEY,
user_id INT NOT NULL,
client_id NVARCHAR(255) NOT NULL,
scopes NVARCHAR(MAX),
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
expires_at DATETIME2 NOT NULL
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_consents_user_client' AND object_id = OBJECT_ID(N'oauth_consents'))
CREATE INDEX idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
IF OBJECT_ID(N'oauth_refresh_tokens', N'U') IS NULL
CREATE TABLE oauth_refresh_tokens (
id INT IDENTITY(1,1) PRIMARY KEY,
token_hash NVARCHAR(64) NOT NULL UNIQUE,
family_id NVARCHAR(64) NOT NULL,
client_id NVARCHAR(255) NOT NULL,
user_id INT NOT NULL,
session_token NVARCHAR(255),
scopes NVARCHAR(MAX),
extra NVARCHAR(MAX),
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
expires_at DATETIME2 NOT NULL,
used_at DATETIME2,
revoked_at DATETIME2
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_refresh_family' AND object_id = OBJECT_ID(N'oauth_refresh_tokens'))
CREATE INDEX idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_refresh_session' AND object_id = OBJECT_ID(N'oauth_refresh_tokens'))
CREATE INDEX idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_refresh_expires' AND object_id = OBJECT_ID(N'oauth_refresh_tokens'))
CREATE INDEX idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
IF OBJECT_ID(N'oauth_device_codes', N'U') IS NULL
CREATE TABLE oauth_device_codes (
id INT IDENTITY(1,1) PRIMARY KEY,
device_hash NVARCHAR(64) NOT NULL UNIQUE,
user_code NVARCHAR(32) NOT NULL UNIQUE,
client_id NVARCHAR(255) NOT NULL,
scopes NVARCHAR(MAX),
status NVARCHAR(16) NOT NULL DEFAULT 'pending',
user_id INT,
session_token NVARCHAR(255),
poll_interval INT NOT NULL DEFAULT 5,
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
expires_at DATETIME2 NOT NULL,
last_polled_at DATETIME2
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_device_expires' AND object_id = OBJECT_ID(N'oauth_device_codes'))
CREATE INDEX idx_oauth_device_expires ON oauth_device_codes(expires_at);
IF OBJECT_ID(N'oauth_par_requests', N'U') IS NULL
CREATE TABLE oauth_par_requests (
id INT IDENTITY(1,1) PRIMARY KEY,
request_uri NVARCHAR(255) NOT NULL UNIQUE,
client_id NVARCHAR(255) NOT NULL,
params NVARCHAR(MAX),
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
expires_at DATETIME2 NOT NULL
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_par_expires' AND object_id = OBJECT_ID(N'oauth_par_requests'))
CREATE INDEX idx_oauth_par_expires ON oauth_par_requests(expires_at);
IF OBJECT_ID(N'oauth_jti', N'U') IS NULL
CREATE TABLE oauth_jti (
id INT IDENTITY(1,1) PRIMARY KEY,
jti_key NVARCHAR(255) NOT NULL UNIQUE,
expires_at DATETIME2 NOT NULL
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_jti_expires' AND object_id = OBJECT_ID(N'oauth_jti'))
CREATE INDEX idx_oauth_jti_expires ON oauth_jti(expires_at);
-- key_hash: SHA-256 hex
-- scopes: JSON-encoded array
-- meta: JSON-encoded object
IF OBJECT_ID(N'user_keys', N'U') IS NULL
CREATE TABLE user_keys (
id INT IDENTITY(1,1) PRIMARY KEY,
user_id INT NOT NULL,
key_type NVARCHAR(50) NOT NULL,
key_hash NVARCHAR(64) NOT NULL UNIQUE,
name NVARCHAR(255) NOT NULL DEFAULT '',
scopes NVARCHAR(MAX),
meta NVARCHAR(MAX),
expires_at DATETIME2,
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
last_used_at DATETIME2,
is_active BIT DEFAULT 1,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_keys_user_id' AND object_id = OBJECT_ID(N'user_keys'))
CREATE INDEX idx_user_keys_user_id ON user_keys(user_id);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_keys_key_type' AND object_id = OBJECT_ID(N'user_keys'))
CREATE INDEX idx_user_keys_key_type ON user_keys(key_type);
-- Optional: omit to use per-user rules only.
IF OBJECT_ID(N'sec_group_members', N'U') IS NULL
CREATE TABLE sec_group_members (
group_id INT NOT NULL,
user_id INT NOT NULL,
PRIMARY KEY (group_id, user_id),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
-- column_path: dot path under the table: col or col.sub.field
-- access_type: mask, hide, read, ...
-- extra_filters: JSON object
IF OBJECT_ID(N'sec_column_rules', N'U') IS NULL
CREATE TABLE sec_column_rules (
id INT IDENTITY(1,1) PRIMARY KEY,
user_id INT,
group_id INT,
schema_name NVARCHAR(255) NOT NULL,
table_name NVARCHAR(255) NOT NULL,
column_path NVARCHAR(255) NOT NULL,
access_type NVARCHAR(50) NOT NULL,
mask_start INT DEFAULT 0,
mask_end INT DEFAULT 0,
mask_invert BIT DEFAULT 0,
mask_char NVARCHAR(10) DEFAULT '*',
extra_filters NVARCHAR(MAX),
is_active BIT NOT NULL DEFAULT 1,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_sec_column_rules_table' AND object_id = OBJECT_ID(N'sec_column_rules'))
CREATE INDEX idx_sec_column_rules_table ON sec_column_rules(schema_name, table_name);
-- template: SQL fragment, e.g. user_id = {UserID}
IF OBJECT_ID(N'sec_row_rules', N'U') IS NULL
CREATE TABLE sec_row_rules (
id INT IDENTITY(1,1) PRIMARY KEY,
user_id INT,
group_id INT,
schema_name NVARCHAR(255) NOT NULL,
table_name NVARCHAR(255) NOT NULL,
template NVARCHAR(MAX),
has_block BIT NOT NULL DEFAULT 0,
is_active BIT NOT NULL DEFAULT 1,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_sec_row_rules_table' AND object_id = OBJECT_ID(N'sec_row_rules'))
CREATE INDEX idx_sec_row_rules_table ON sec_row_rules(schema_name, table_name);
+289
View File
@@ -0,0 +1,289 @@
-- Reference schema for the lookup direct backend: MySQL 8.0.16+ / MariaDB 10.2+. Direct backend only. Run each statement separately (see ddl.Statements) unless multiStatements=true.
-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema.
-- Generated by the project, edit freely for your deployment (types, collations, extra columns).
-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login
-- roles: comma-separated roles
CREATE TABLE IF NOT EXISTS users (
id INT AUTO_INCREMENT PRIMARY KEY,
username VARCHAR(255) NOT NULL UNIQUE,
email VARCHAR(255) NOT NULL UNIQUE,
password VARCHAR(255),
user_level INT DEFAULT 0,
roles VARCHAR(500),
is_active TINYINT(1) DEFAULT 1,
created_at DATETIME NULL,
updated_at DATETIME NULL,
last_login_at DATETIME,
program_user_id INT DEFAULT 0,
program_user_table VARCHAR(255) DEFAULT '',
remote_id VARCHAR(255),
auth_provider VARCHAR(50),
totp_secret VARCHAR(255),
totp_enabled TINYINT(1) DEFAULT 0,
totp_enabled_at DATETIME
);
CREATE TABLE IF NOT EXISTS user_sessions (
id INT AUTO_INCREMENT PRIMARY KEY,
session_token VARCHAR(500) NOT NULL UNIQUE,
user_id INT NOT NULL,
expires_at DATETIME NOT NULL,
created_at DATETIME NULL,
last_activity_at DATETIME NULL,
ip_address VARCHAR(45),
user_agent TEXT,
access_token TEXT,
refresh_token TEXT,
token_type VARCHAR(50) DEFAULT 'Bearer',
auth_provider VARCHAR(50),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
INDEX idx_user_sessions_user_id (user_id),
INDEX idx_user_sessions_expires_at (expires_at)
);
CREATE TABLE IF NOT EXISTS token_blacklist (
id INT AUTO_INCREMENT PRIMARY KEY,
token VARCHAR(500) NOT NULL,
user_id INT,
expires_at DATETIME NOT NULL,
created_at DATETIME NULL,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
-- code_hash: SHA-256 hex of the backup code
CREATE TABLE IF NOT EXISTS user_totp_backup_codes (
id INT AUTO_INCREMENT PRIMARY KEY,
user_id INT NOT NULL,
code_hash VARCHAR(64) NOT NULL,
used TINYINT(1) DEFAULT 0,
used_at DATETIME,
created_at DATETIME NULL,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
INDEX idx_totp_user_id (user_id),
INDEX idx_totp_code_hash (code_hash)
);
-- credential_id: base64 text
-- public_key: base64 text
-- aaguid: base64 text
-- transports: JSON-encoded array
CREATE TABLE IF NOT EXISTS user_passkey_credentials (
id INT AUTO_INCREMENT PRIMARY KEY,
user_id INT NOT NULL,
credential_id VARCHAR(1400) CHARACTER SET ascii NOT NULL UNIQUE,
public_key TEXT NOT NULL,
attestation_type VARCHAR(50) DEFAULT 'none',
aaguid TEXT,
sign_count INT DEFAULT 0,
clone_warning TINYINT(1) DEFAULT 0,
transports TEXT,
backup_eligible TINYINT(1) DEFAULT 0,
backup_state TINYINT(1) DEFAULT 0,
name VARCHAR(255),
created_at DATETIME NULL,
last_used_at DATETIME NULL,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
INDEX idx_passkey_user_id (user_id)
);
-- token_hash: SHA-256 hex of the raw token
CREATE TABLE IF NOT EXISTS user_password_resets (
id INT AUTO_INCREMENT PRIMARY KEY,
user_id INT NOT NULL,
token_hash VARCHAR(64) NOT NULL UNIQUE,
expires_at DATETIME NOT NULL,
created_at DATETIME NULL,
used TINYINT(1) DEFAULT 0,
used_at DATETIME,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
INDEX idx_pw_reset_user_id (user_id),
INDEX idx_pw_reset_expires_at (expires_at)
);
-- redirect_uris: JSON-encoded array
-- grant_types: JSON-encoded array
-- allowed_scopes: JSON-encoded array
-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients
CREATE TABLE IF NOT EXISTS oauth_clients (
id INT AUTO_INCREMENT PRIMARY KEY,
client_id VARCHAR(255) NOT NULL UNIQUE,
redirect_uris TEXT NOT NULL,
client_name VARCHAR(255),
grant_types TEXT,
allowed_scopes TEXT,
client_secret_hash TEXT,
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
is_active TINYINT(1) DEFAULT 1,
metadata TEXT,
created_at DATETIME NULL
);
-- scopes: JSON-encoded array
CREATE TABLE IF NOT EXISTS oauth_codes (
id INT AUTO_INCREMENT PRIMARY KEY,
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,
expires_at DATETIME NOT NULL,
extra TEXT,
created_at DATETIME NULL,
INDEX idx_oauth_codes_expires (expires_at)
);
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
CREATE TABLE IF NOT EXISTS oauth_consents (
id INT AUTO_INCREMENT PRIMARY KEY,
user_id INT NOT NULL,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
created_at DATETIME NULL,
expires_at DATETIME NOT NULL,
INDEX idx_oauth_consents_user_client (user_id, client_id)
);
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
id INT AUTO_INCREMENT PRIMARY KEY,
token_hash VARCHAR(64) NOT NULL UNIQUE,
family_id VARCHAR(64) NOT NULL,
client_id VARCHAR(255) NOT NULL,
user_id INT NOT NULL,
session_token VARCHAR(255),
scopes TEXT,
extra TEXT,
created_at DATETIME NULL,
expires_at DATETIME NOT NULL,
used_at DATETIME,
revoked_at DATETIME,
INDEX idx_oauth_refresh_family (family_id),
INDEX idx_oauth_refresh_session (session_token),
INDEX idx_oauth_refresh_expires (expires_at)
);
CREATE TABLE IF NOT EXISTS oauth_device_codes (
id INT AUTO_INCREMENT PRIMARY KEY,
device_hash VARCHAR(64) NOT NULL UNIQUE,
user_code VARCHAR(32) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
status VARCHAR(16) NOT NULL DEFAULT 'pending',
user_id INT,
session_token VARCHAR(255),
poll_interval INT NOT NULL DEFAULT 5,
created_at DATETIME NULL,
expires_at DATETIME NOT NULL,
last_polled_at DATETIME,
INDEX idx_oauth_device_expires (expires_at)
);
CREATE TABLE IF NOT EXISTS oauth_par_requests (
id INT AUTO_INCREMENT PRIMARY KEY,
request_uri VARCHAR(255) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
params TEXT,
created_at DATETIME NULL,
expires_at DATETIME NOT NULL,
INDEX idx_oauth_par_expires (expires_at)
);
CREATE TABLE IF NOT EXISTS oauth_jti (
id INT AUTO_INCREMENT PRIMARY KEY,
jti_key VARCHAR(255) NOT NULL UNIQUE,
expires_at DATETIME NOT NULL,
INDEX idx_oauth_jti_expires (expires_at)
);
-- key_hash: SHA-256 hex
-- scopes: JSON-encoded array
-- meta: JSON-encoded object
CREATE TABLE IF NOT EXISTS user_keys (
id INT AUTO_INCREMENT PRIMARY KEY,
user_id INT NOT NULL,
key_type VARCHAR(50) NOT NULL,
key_hash VARCHAR(64) NOT NULL UNIQUE,
name VARCHAR(255) NOT NULL DEFAULT '',
scopes TEXT,
meta TEXT,
expires_at DATETIME,
created_at DATETIME NULL,
last_used_at DATETIME,
is_active TINYINT(1) DEFAULT 1,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
INDEX idx_user_keys_user_id (user_id),
INDEX idx_user_keys_key_type (key_type)
);
-- Optional: omit to use per-user rules only.
CREATE TABLE IF NOT EXISTS sec_group_members (
group_id INT NOT NULL,
user_id INT NOT NULL,
PRIMARY KEY (group_id, user_id),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
-- column_path: dot path under the table: col or col.sub.field
-- access_type: mask, hide, read, ...
-- extra_filters: JSON object
CREATE TABLE IF NOT EXISTS sec_column_rules (
id INT AUTO_INCREMENT PRIMARY KEY,
user_id INT,
group_id INT,
schema_name VARCHAR(255) NOT NULL,
table_name VARCHAR(255) NOT NULL,
column_path VARCHAR(255) NOT NULL,
access_type VARCHAR(50) NOT NULL,
mask_start INT DEFAULT 0,
mask_end INT DEFAULT 0,
mask_invert TINYINT(1) DEFAULT 0,
mask_char VARCHAR(10) DEFAULT '*',
extra_filters TEXT,
is_active TINYINT(1) NOT NULL DEFAULT 1,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)),
INDEX idx_sec_column_rules_table (schema_name, table_name)
);
-- template: SQL fragment, e.g. user_id = {UserID}
CREATE TABLE IF NOT EXISTS sec_row_rules (
id INT AUTO_INCREMENT PRIMARY KEY,
user_id INT,
group_id INT,
schema_name VARCHAR(255) NOT NULL,
table_name VARCHAR(255) NOT NULL,
template TEXT,
has_block TINYINT(1) NOT NULL DEFAULT 0,
is_active TINYINT(1) NOT NULL DEFAULT 1,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)),
INDEX idx_sec_row_rules_table (schema_name, table_name)
);
+312
View File
@@ -0,0 +1,312 @@
-- Reference schema for the lookup direct backend: PostgreSQL (tables only, no stored procedures). Use with lookup.Config{Mode: lookup.ModeDirect}.
-- Do not combine with the procedure schema (database_schema.sql): that schema stores credential ids as bytea and
-- list columns as text[], this one stores base64 / JSON text which is what the direct backend reads and writes.
-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema.
-- Generated by the project, edit freely for your deployment (types, collations, extra columns).
-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login
-- roles: comma-separated roles
CREATE TABLE IF NOT EXISTS users (
id SERIAL PRIMARY KEY,
username VARCHAR(255) NOT NULL UNIQUE,
email VARCHAR(255) NOT NULL UNIQUE,
password VARCHAR(255),
user_level INTEGER DEFAULT 0,
roles VARCHAR(500),
is_active BOOLEAN DEFAULT true,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMP DEFAULT CURRENT_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 false,
totp_enabled_at TIMESTAMP
);
CREATE TABLE IF NOT EXISTS user_sessions (
id SERIAL PRIMARY KEY,
session_token VARCHAR(500) NOT NULL UNIQUE,
user_id INTEGER NOT NULL,
expires_at TIMESTAMP NOT NULL,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
last_activity_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
ip_address VARCHAR(45),
user_agent TEXT,
access_token TEXT,
refresh_token TEXT,
token_type VARCHAR(50) DEFAULT 'Bearer',
auth_provider VARCHAR(50),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_user_sessions_user_id ON user_sessions(user_id);
CREATE INDEX IF NOT EXISTS idx_user_sessions_expires_at ON user_sessions(expires_at);
CREATE INDEX IF NOT EXISTS idx_user_sessions_refresh_token ON user_sessions(refresh_token);
CREATE TABLE IF NOT EXISTS token_blacklist (
id SERIAL PRIMARY KEY,
token VARCHAR(500) NOT NULL,
user_id INTEGER,
expires_at TIMESTAMP NOT NULL,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
-- code_hash: SHA-256 hex of the backup code
CREATE TABLE IF NOT EXISTS user_totp_backup_codes (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL,
code_hash VARCHAR(64) NOT NULL,
used BOOLEAN DEFAULT false,
used_at TIMESTAMP,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
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);
-- credential_id: base64 text
-- public_key: base64 text
-- aaguid: base64 text
-- transports: JSON-encoded array
CREATE TABLE IF NOT EXISTS user_passkey_credentials (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL,
credential_id TEXT NOT NULL UNIQUE,
public_key TEXT NOT NULL,
attestation_type VARCHAR(50) DEFAULT 'none',
aaguid TEXT,
sign_count INTEGER DEFAULT 0,
clone_warning BOOLEAN DEFAULT false,
transports TEXT,
backup_eligible BOOLEAN DEFAULT false,
backup_state BOOLEAN DEFAULT false,
name VARCHAR(255),
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
last_used_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_passkey_user_id ON user_passkey_credentials(user_id);
-- token_hash: SHA-256 hex of the raw token
CREATE TABLE IF NOT EXISTS user_password_resets (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL,
token_hash VARCHAR(64) NOT NULL UNIQUE,
expires_at TIMESTAMP NOT NULL,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
used BOOLEAN DEFAULT false,
used_at TIMESTAMP,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_pw_reset_user_id ON user_password_resets(user_id);
CREATE INDEX IF NOT EXISTS idx_pw_reset_expires_at ON user_password_resets(expires_at);
-- redirect_uris: JSON-encoded array
-- grant_types: JSON-encoded array
-- allowed_scopes: JSON-encoded array
-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients
CREATE TABLE IF NOT EXISTS oauth_clients (
id SERIAL PRIMARY KEY,
client_id VARCHAR(255) NOT NULL UNIQUE,
redirect_uris TEXT NOT NULL,
client_name VARCHAR(255),
grant_types TEXT,
allowed_scopes TEXT,
client_secret_hash TEXT,
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
is_active BOOLEAN DEFAULT true,
metadata TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
-- scopes: JSON-encoded array
CREATE TABLE IF NOT EXISTS oauth_codes (
id SERIAL PRIMARY KEY,
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,
expires_at TIMESTAMP NOT NULL,
extra TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
CREATE TABLE IF NOT EXISTS oauth_consents (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
id SERIAL PRIMARY KEY,
token_hash VARCHAR(64) NOT NULL UNIQUE,
family_id VARCHAR(64) NOT NULL,
client_id VARCHAR(255) NOT NULL,
user_id INTEGER NOT NULL,
session_token VARCHAR(255),
scopes TEXT,
extra TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
used_at TIMESTAMP,
revoked_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
CREATE TABLE IF NOT EXISTS oauth_device_codes (
id SERIAL PRIMARY KEY,
device_hash VARCHAR(64) NOT NULL UNIQUE,
user_code VARCHAR(32) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
status VARCHAR(16) NOT NULL DEFAULT 'pending',
user_id INTEGER,
session_token VARCHAR(255),
poll_interval INTEGER NOT NULL DEFAULT 5,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
last_polled_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_device_expires ON oauth_device_codes(expires_at);
CREATE TABLE IF NOT EXISTS oauth_par_requests (
id SERIAL PRIMARY KEY,
request_uri VARCHAR(255) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
params TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_par_expires ON oauth_par_requests(expires_at);
CREATE TABLE IF NOT EXISTS oauth_jti (
id SERIAL PRIMARY KEY,
jti_key VARCHAR(255) NOT NULL UNIQUE,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_jti_expires ON oauth_jti(expires_at);
-- key_hash: SHA-256 hex
-- scopes: JSON-encoded array
-- meta: JSON-encoded object
CREATE TABLE IF NOT EXISTS user_keys (
id SERIAL PRIMARY KEY,
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,
meta TEXT,
expires_at TIMESTAMP,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
last_used_at TIMESTAMP,
is_active BOOLEAN DEFAULT true,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_user_keys_user_id ON user_keys(user_id);
CREATE INDEX IF NOT EXISTS idx_user_keys_key_type ON user_keys(key_type);
-- Optional: omit to use per-user rules only.
CREATE TABLE IF NOT EXISTS sec_group_members (
group_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
PRIMARY KEY (group_id, user_id),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);
-- column_path: dot path under the table: col or col.sub.field
-- access_type: mask, hide, read, ...
-- extra_filters: JSON object
CREATE TABLE IF NOT EXISTS sec_column_rules (
id SERIAL PRIMARY KEY,
user_id INTEGER,
group_id INTEGER,
schema_name TEXT NOT NULL,
table_name TEXT NOT NULL,
column_path TEXT NOT NULL,
access_type VARCHAR(50) NOT NULL,
mask_start INTEGER DEFAULT 0,
mask_end INTEGER DEFAULT 0,
mask_invert BOOLEAN DEFAULT false,
mask_char VARCHAR(10) DEFAULT '*',
extra_filters TEXT,
is_active BOOLEAN NOT NULL DEFAULT true,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
);
CREATE INDEX IF NOT EXISTS idx_sec_column_rules_table ON sec_column_rules(schema_name, table_name);
-- template: SQL fragment, e.g. user_id = {UserID}
CREATE TABLE IF NOT EXISTS sec_row_rules (
id SERIAL PRIMARY KEY,
user_id INTEGER,
group_id INTEGER,
schema_name TEXT NOT NULL,
table_name TEXT NOT NULL,
template TEXT,
has_block BOOLEAN NOT NULL DEFAULT false,
is_active BOOLEAN NOT NULL DEFAULT true,
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
);
CREATE INDEX IF NOT EXISTS idx_sec_row_rules_table ON sec_row_rules(schema_name, table_name);
+301
View File
@@ -0,0 +1,301 @@
-- Reference schema for the lookup direct backend: SQLite. Direct backend only.
-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema.
-- Generated by the project, edit freely for your deployment (types, collations, extra columns).
-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login
-- roles: comma-separated roles
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),
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_user_sessions_user_id ON user_sessions(user_id);
CREATE INDEX IF NOT EXISTS idx_user_sessions_expires_at ON user_sessions(expires_at);
CREATE INDEX IF NOT EXISTS idx_user_sessions_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
);
-- code_hash: SHA-256 hex of the backup code
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);
-- credential_id: base64 text
-- public_key: base64 text
-- aaguid: base64 text
-- transports: JSON-encoded array
CREATE TABLE IF NOT EXISTS user_passkey_credentials (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
credential_id TEXT NOT NULL UNIQUE,
public_key TEXT NOT NULL,
attestation_type VARCHAR(50) DEFAULT 'none',
aaguid TEXT,
sign_count INTEGER DEFAULT 0,
clone_warning BOOLEAN DEFAULT 0,
transports TEXT,
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);
-- token_hash: SHA-256 hex of the raw token
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 INDEX IF NOT EXISTS idx_pw_reset_user_id ON user_password_resets(user_id);
CREATE INDEX IF NOT EXISTS idx_pw_reset_expires_at ON user_password_resets(expires_at);
-- redirect_uris: JSON-encoded array
-- grant_types: JSON-encoded array
-- allowed_scopes: JSON-encoded array
-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients
CREATE TABLE IF NOT EXISTS oauth_clients (
id INTEGER PRIMARY KEY AUTOINCREMENT,
client_id VARCHAR(255) NOT NULL UNIQUE,
redirect_uris TEXT NOT NULL,
client_name VARCHAR(255),
grant_types TEXT,
allowed_scopes TEXT,
client_secret_hash TEXT,
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
is_active BOOLEAN DEFAULT 1,
metadata TEXT,
created_at TIMESTAMP
);
-- scopes: JSON-encoded array
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,
expires_at TIMESTAMP NOT NULL,
extra TEXT,
created_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
CREATE TABLE IF NOT EXISTS oauth_consents (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
created_at TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
id INTEGER PRIMARY KEY AUTOINCREMENT,
token_hash VARCHAR(64) NOT NULL UNIQUE,
family_id VARCHAR(64) NOT NULL,
client_id VARCHAR(255) NOT NULL,
user_id INTEGER NOT NULL,
session_token VARCHAR(255),
scopes TEXT,
extra TEXT,
created_at TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
used_at TIMESTAMP,
revoked_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
CREATE TABLE IF NOT EXISTS oauth_device_codes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
device_hash VARCHAR(64) NOT NULL UNIQUE,
user_code VARCHAR(32) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
status VARCHAR(16) NOT NULL DEFAULT 'pending',
user_id INTEGER,
session_token VARCHAR(255),
poll_interval INTEGER NOT NULL DEFAULT 5,
created_at TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
last_polled_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_device_expires ON oauth_device_codes(expires_at);
CREATE TABLE IF NOT EXISTS oauth_par_requests (
id INTEGER PRIMARY KEY AUTOINCREMENT,
request_uri VARCHAR(255) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
params TEXT,
created_at TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_par_expires ON oauth_par_requests(expires_at);
CREATE TABLE IF NOT EXISTS oauth_jti (
id INTEGER PRIMARY KEY AUTOINCREMENT,
jti_key VARCHAR(255) NOT NULL UNIQUE,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_jti_expires ON oauth_jti(expires_at);
-- key_hash: SHA-256 hex
-- scopes: JSON-encoded array
-- meta: JSON-encoded object
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,
meta TEXT,
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_type ON user_keys(key_type);
-- Optional: omit to use per-user rules only.
CREATE TABLE IF NOT EXISTS sec_group_members (
group_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
PRIMARY KEY (group_id, user_id)
);
-- column_path: dot path under the table: col or col.sub.field
-- access_type: mask, hide, read, ...
-- extra_filters: JSON object
CREATE TABLE IF NOT EXISTS sec_column_rules (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER,
group_id INTEGER,
schema_name TEXT NOT NULL,
table_name TEXT NOT NULL,
column_path TEXT NOT NULL,
access_type VARCHAR(50) NOT NULL,
mask_start INTEGER DEFAULT 0,
mask_end INTEGER DEFAULT 0,
mask_invert BOOLEAN DEFAULT 0,
mask_char VARCHAR(10) DEFAULT '*',
extra_filters TEXT,
is_active BOOLEAN NOT NULL DEFAULT 1,
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
);
CREATE INDEX IF NOT EXISTS idx_sec_column_rules_table ON sec_column_rules(schema_name, table_name);
-- template: SQL fragment, e.g. user_id = {UserID}
CREATE TABLE IF NOT EXISTS sec_row_rules (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER,
group_id INTEGER,
schema_name TEXT NOT NULL,
table_name TEXT NOT NULL,
template TEXT,
has_block BOOLEAN NOT NULL DEFAULT 0,
is_active BOOLEAN NOT NULL DEFAULT 1,
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
);
CREATE INDEX IF NOT EXISTS idx_sec_row_rules_table ON sec_row_rules(schema_name, table_name);
+34
View File
@@ -0,0 +1,34 @@
package lookup_test
import (
"os/exec"
"strings"
"testing"
)
// lookup and its backends sit above sectypes and below security: they must not import the
// core package (security imports lookup) or any sibling sub package.
func TestNoUpwardImports(t *testing.T) {
for _, pkg := range []string{".", "./procedure", "./direct", "./conformance"} {
checkNoUpwardImports(t, pkg)
}
}
func checkNoUpwardImports(t *testing.T, pkg string) {
t.Helper()
out, err := exec.Command("go", "list", "-deps", "-f", "{{.ImportPath}}", pkg).Output()
if err != nil {
t.Skipf("go list unavailable: %v", err)
}
for _, p := range strings.Fields(string(out)) {
if strings.Contains(p, "uptrace/bun") || strings.Contains(p, "gorm.io") {
t.Errorf("%s must not depend on an ORM, found %s", pkg, p)
}
if !strings.Contains(p, "/pkg/security") || strings.HasSuffix(p, "/lookup") ||
strings.HasSuffix(p, "/sectypes") || strings.HasSuffix(p, "/lookup/dialect") ||
strings.HasSuffix(p, "/lookup/procedure") || strings.HasSuffix(p, "/lookup/direct") || strings.HasSuffix(p, "/lookup/conformance") {
continue
}
t.Errorf("%s must not import %s", pkg, p)
}
}
+93
View File
@@ -0,0 +1,93 @@
package dialect
import (
"strconv"
"strings"
"time"
)
// --- postgres ---------------------------------------------------------------
type postgres struct{}
func (postgres) Name() string { return "postgres" }
func (postgres) Matches(driver string) bool {
return strings.Contains(driver, "pgx") || strings.Contains(driver, "lib/pq") ||
strings.Contains(driver, "postgres")
}
func (postgres) Placeholder(n int) string { return "$" + strconv.Itoa(n) }
func (postgres) Quote(ident string) string { return quoteWith(ident, `"`, `"`) }
func (postgres) Bool(v bool) any { return v }
func (postgres) ScanBool(src any) (bool, error) { return scanBool(src) }
func (postgres) ScanTime(src any) (time.Time, error) { return scanTime(src) }
func (postgres) EncodeJSON(v any) (any, error) { return encodeJSON(v) }
func (postgres) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) }
func (d postgres) InsertReturningID(table string, cols []string, idCol string) Insert {
return Insert{SQL: insertSQL(d, table, cols, "", "RETURNING "+d.Quote(idCol), "DEFAULT VALUES"), Strategy: ReturningQuery}
}
// --- sqlite -----------------------------------------------------------------
type sqlite struct{}
func (sqlite) Name() string { return "sqlite" }
func (sqlite) Matches(driver string) bool {
return strings.Contains(driver, "sqlite")
}
func (sqlite) Placeholder(int) string { return "?" }
func (sqlite) Quote(ident string) string { return quoteWith(ident, `"`, `"`) }
func (sqlite) Bool(v bool) any { return boolInt(v) }
func (sqlite) ScanBool(src any) (bool, error) { return scanBool(src) }
func (sqlite) ScanTime(src any) (time.Time, error) { return scanTime(src) }
func (sqlite) EncodeJSON(v any) (any, error) { return encodeJSON(v) }
func (sqlite) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) }
func (d sqlite) InsertReturningID(table string, cols []string, _ string) Insert {
return Insert{SQL: insertSQL(d, table, cols, "", "", "DEFAULT VALUES"), Strategy: LastInsertID}
}
// --- mysql / mariadb ----------------------------------------------------------
type mysql struct{}
func (mysql) Name() string { return "mysql" }
func (mysql) Matches(driver string) bool {
return strings.Contains(driver, "mysql") || strings.Contains(driver, "mariadb")
}
func (mysql) Placeholder(int) string { return "?" }
func (mysql) Quote(ident string) string { return quoteWith(ident, "`", "`") }
func (mysql) Bool(v bool) any { return boolInt(v) }
func (mysql) ScanBool(src any) (bool, error) { return scanBool(src) }
func (mysql) ScanTime(src any) (time.Time, error) { return scanTime(src) }
func (mysql) EncodeJSON(v any) (any, error) { return encodeJSON(v) }
func (mysql) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) }
func (d mysql) InsertReturningID(table string, cols []string, _ string) Insert {
// MySQL has no DEFAULT VALUES; an empty column list is spelled "() VALUES ()".
s := insertSQL(d, table, cols, "", "", "() VALUES ()")
return Insert{SQL: s, Strategy: LastInsertID}
}
// --- mssql (SQL Server) --------------------------------------------------------
type mssql struct{}
func (mssql) Name() string { return "mssql" }
func (mssql) Matches(driver string) bool {
return strings.Contains(driver, "mssql") || strings.Contains(driver, "sqlserver")
}
func (mssql) Placeholder(n int) string { return "@p" + strconv.Itoa(n) }
func (mssql) Quote(ident string) string { return quoteWith(ident, "[", "]") }
func (mssql) Bool(v bool) any { return v }
func (mssql) ScanBool(src any) (bool, error) { return scanBool(src) }
func (mssql) ScanTime(src any) (time.Time, error) { return scanTime(src) }
func (mssql) EncodeJSON(v any) (any, error) { return encodeJSON(v) }
func (mssql) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) }
func (d mssql) InsertReturningID(table string, cols []string, idCol string) Insert {
return Insert{SQL: insertSQL(d, table, cols, "OUTPUT INSERTED."+d.Quote(idCol), "", "DEFAULT VALUES"), Strategy: ReturningQuery}
}
func boolInt(v bool) int64 {
if v {
return 1
}
return 0
}
+312
View File
@@ -0,0 +1,312 @@
// Package dialect holds the per-database adaptors used by the lookup direct backend.
// An adaptor supplies only what differs between databases (placeholders, quoting,
// booleans, time and JSON handling, insert-returning-id); the backend builds queries from it.
// It imports only the standard library.
package dialect
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"reflect"
"sort"
"strings"
"sync"
"time"
)
// Dialect is the adaptor for one database type.
type Dialect interface {
// Name is the registry name, e.g. "postgres".
Name() string
// Matches reports whether a driver type identifier (lowercased "<pkgpath>.<Type>")
// belongs to this database. Used by Detect.
Matches(driver string) bool
// Placeholder returns the bind placeholder for the n-th (1-based) argument.
Placeholder(n int) string
// Quote quotes an identifier. A dotted name is quoted per part ("schema.table").
// Embedded quote characters are escaped, never interpolated raw.
Quote(ident string) string
// Bool converts a Go bool to the value bound as a boolean column argument.
Bool(v bool) any
// ScanBool reads a boolean column value that the driver returned as bool, integer, string or bytes.
ScanBool(src any) (bool, error)
// ScanTime reads a time column value that the driver returned as time.Time, string or bytes.
// NULL (nil) yields the zero time.
ScanTime(src any) (time.Time, error)
// EncodeJSON converts a value to the argument bound to a JSON/TEXT column. A nil
// value, map or slice yields nil (SQL NULL).
EncodeJSON(v any) (any, error)
// DecodeJSON reads a JSON/TEXT column value into dst. NULL and empty values leave dst untouched.
DecodeJSON(src any, dst any) error
// InsertReturningID builds an INSERT of cols into table and describes how to read the new id.
// Arguments are bound positionally in cols order.
InsertReturningID(table string, cols []string, idCol string) Insert
}
// InsertStrategy tells how the generated id is read after an Insert.
type InsertStrategy int
const (
// ReturningQuery means the statement returns the id as a single row (QueryRow + Scan).
ReturningQuery InsertStrategy = iota
// LastInsertID means the id is read from sql.Result.LastInsertId after Exec.
LastInsertID
)
// Insert is a generated INSERT statement and how to read the id it creates.
type Insert struct {
SQL string
Strategy InsertStrategy
}
// Querier is implemented by *sql.DB, *sql.Tx and *sql.Conn.
type Querier interface {
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
}
// Run executes the insert and returns the generated id.
func (i Insert) Run(ctx context.Context, q Querier, args ...any) (int64, error) {
switch i.Strategy {
case ReturningQuery:
var id int64
if err := q.QueryRowContext(ctx, i.SQL, args...).Scan(&id); err != nil {
return 0, err
}
return id, nil
case LastInsertID:
res, err := q.ExecContext(ctx, i.SQL, args...)
if err != nil {
return 0, err
}
return res.LastInsertId()
}
return 0, fmt.Errorf("dialect: unknown insert strategy %d", i.Strategy)
}
// Factory creates a Dialect.
type Factory func() Dialect
var (
regMu sync.RWMutex
registry = map[string]Factory{}
)
// Register adds a dialect under name. Registering a name twice replaces it, so applications
// can override a built-in. Adding a database = implementing Dialect and calling Register.
func Register(name string, f Factory) {
regMu.Lock()
defer regMu.Unlock()
registry[strings.ToLower(name)] = f
}
// Get returns the dialect registered under name.
func Get(name string) (Dialect, error) {
regMu.RLock()
f, ok := registry[strings.ToLower(name)]
regMu.RUnlock()
if !ok {
return nil, fmt.Errorf("dialect: unknown dialect %q (registered: %s)", name, strings.Join(Names(), ", "))
}
return f(), nil
}
// Names lists the registered dialect names, sorted.
func Names() []string {
regMu.RLock()
defer regMu.RUnlock()
names := make([]string, 0, len(registry))
for n := range registry {
names = append(names, n)
}
sort.Strings(names)
return names
}
// Detect picks the dialect for db from its driver type.
func Detect(db *sql.DB) (Dialect, error) {
if db == nil {
return nil, fmt.Errorf("dialect: nil database")
}
return DetectDriver(driverID(db.Driver()))
}
// DetectDriver picks the dialect for a driver type identifier (see Dialect.Matches).
func DetectDriver(driver string) (Dialect, error) {
driver = strings.ToLower(driver)
for _, n := range Names() {
d, _ := Get(n)
if d != nil && d.Matches(driver) {
return d, nil
}
}
return nil, fmt.Errorf("dialect: cannot detect a dialect for driver %q; set the dialect explicitly", driver)
}
// driverID builds "<pkgpath>.<Type>" for a driver value, lowercased.
func driverID(drv any) string {
t := reflect.TypeOf(drv)
for t != nil && t.Kind() == reflect.Pointer {
t = t.Elem()
}
if t == nil {
return ""
}
return strings.ToLower(t.PkgPath() + "." + t.Name())
}
func init() {
Register("postgres", func() Dialect { return postgres{} })
Register("sqlite", func() Dialect { return sqlite{} })
Register("mysql", func() Dialect { return mysql{} })
Register("mssql", func() Dialect { return mssql{} })
}
// --- shared helpers -------------------------------------------------------
// quoteWith quotes each dotted part of ident with open/close, doubling embedded close characters.
func quoteWith(ident, open, closeq string) string {
parts := strings.Split(ident, ".")
for i, p := range parts {
parts[i] = open + strings.ReplaceAll(p, closeq, closeq+closeq) + closeq
}
return strings.Join(parts, ".")
}
func scanBool(src any) (bool, error) {
switch v := src.(type) {
case nil:
return false, nil
case bool:
return v, nil
case int64:
return v != 0, nil
case int:
return v != 0, nil
case int32:
return v != 0, nil
case float64:
return v != 0, nil
case []byte:
return scanBool(string(v))
case string:
switch strings.ToLower(strings.TrimSpace(v)) {
case "1", "t", "true", "y", "yes", "on":
return true, nil
case "", "0", "f", "false", "n", "no", "off":
return false, nil
}
}
return false, fmt.Errorf("dialect: cannot read %T (%v) as bool", src, src)
}
var timeLayouts = []string{
time.RFC3339Nano,
"2006-01-02 15:04:05.999999999 -0700",
"2006-01-02 15:04:05.999999999-07:00",
"2006-01-02 15:04:05.999999999Z07:00",
"2006-01-02T15:04:05.999999999",
"2006-01-02 15:04:05.999999999",
"2006-01-02",
}
func scanTime(src any) (time.Time, error) {
switch v := src.(type) {
case nil:
return time.Time{}, nil
case time.Time:
return v, nil
case []byte:
return scanTime(string(v))
case string:
s := strings.TrimSpace(v)
if s == "" {
return time.Time{}, nil
}
// Some drivers append the Go monotonic/zone suffix ("+0000 UTC"); drop it.
if i := strings.Index(s, " m="); i >= 0 {
s = s[:i]
}
s = strings.TrimSuffix(s, " UTC")
s = strings.TrimSuffix(s, " +0000 +0000")
for _, l := range timeLayouts {
if t, err := time.Parse(l, s); err == nil {
return t, nil
}
}
}
return time.Time{}, fmt.Errorf("dialect: cannot read %T (%v) as time", src, src)
}
func encodeJSON(v any) (any, error) {
if v == nil {
return nil, nil
}
rv := reflect.ValueOf(v)
switch rv.Kind() {
case reflect.Map, reflect.Slice, reflect.Pointer, reflect.Interface:
if rv.IsNil() {
return nil, nil
}
}
b, err := json.Marshal(v)
if err != nil {
return nil, fmt.Errorf("dialect: encode json: %w", err)
}
return string(b), nil
}
func decodeJSON(src any, dst any) error {
var raw []byte
switch v := src.(type) {
case nil:
return nil
case string:
raw = []byte(v)
case []byte:
raw = v
default:
return fmt.Errorf("dialect: cannot read %T as json", src)
}
if len(strings.TrimSpace(string(raw))) == 0 {
return nil
}
if err := json.Unmarshal(raw, dst); err != nil {
return fmt.Errorf("dialect: decode json: %w", err)
}
return nil
}
// insertSQL assembles "INSERT INTO t (cols) <mid> VALUES (...) <tail>" for a dialect.
func insertSQL(d Dialect, table string, cols []string, mid, tail, defaults string) string {
var b strings.Builder
b.WriteString("INSERT INTO ")
b.WriteString(d.Quote(table))
if len(cols) == 0 {
if mid != "" {
b.WriteString(" " + mid)
}
b.WriteString(" " + defaults)
if tail != "" {
b.WriteString(" " + tail)
}
return b.String()
}
qc := make([]string, len(cols))
ph := make([]string, len(cols))
for i, c := range cols {
qc[i] = d.Quote(c)
ph[i] = d.Placeholder(i + 1)
}
b.WriteString(" (" + strings.Join(qc, ", ") + ")")
if mid != "" {
b.WriteString(" " + mid)
}
b.WriteString(" VALUES (" + strings.Join(ph, ", ") + ")")
if tail != "" {
b.WriteString(" " + tail)
}
return b.String()
}
+270
View File
@@ -0,0 +1,270 @@
package dialect_test
import (
"context"
"database/sql"
"strings"
"testing"
"time"
_ "github.com/glebarez/go-sqlite"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
)
func get(t *testing.T, name string) dialect.Dialect {
t.Helper()
d, err := dialect.Get(name)
if err != nil {
t.Fatal(err)
}
return d
}
func TestRegistryHasBuiltins(t *testing.T) {
got := strings.Join(dialect.Names(), ",")
if got != "mssql,mysql,postgres,sqlite" {
t.Errorf("names = %s", got)
}
if _, err := dialect.Get("oracle"); err == nil {
t.Error("unknown dialect should error")
}
}
func TestRegisterCustomDialect(t *testing.T) {
// A new database is added by registering a dialect; no core change needed.
dialect.Register("custom", func() dialect.Dialect { return customDialect{get(t, "sqlite")} })
d, err := dialect.Get("CUSTOM")
if err != nil || d.Name() != "custom" {
t.Fatalf("got %v, %v", d, err)
}
if got, _ := dialect.DetectDriver("example.com/driver.customdriver"); got == nil || got.Name() != "custom" {
t.Errorf("custom detection failed: %v", got)
}
}
type customDialect struct{ dialect.Dialect }
func (customDialect) Name() string { return "custom" }
func (customDialect) Matches(driver string) bool { return strings.Contains(driver, "customdriver") }
func TestPlaceholderAndQuote(t *testing.T) {
cases := []struct {
name, ph1, ph3, quote, quoteDotted, quoteEscape string
}{
{"postgres", "$1", "$3", `"users"`, `"auth"."users"`, `"a""b"`},
{"sqlite", "?", "?", `"users"`, `"auth"."users"`, `"a""b"`},
{"mysql", "?", "?", "`users`", "`auth`.`users`", "`a``b`"},
{"mssql", "@p1", "@p3", "[users]", "[auth].[users]", "[a]]b]"},
}
for _, c := range cases {
d := get(t, c.name)
if d.Placeholder(1) != c.ph1 || d.Placeholder(3) != c.ph3 {
t.Errorf("%s placeholders: %s %s", c.name, d.Placeholder(1), d.Placeholder(3))
}
if d.Quote("users") != c.quote {
t.Errorf("%s quote: %s", c.name, d.Quote("users"))
}
if d.Quote("auth.users") != c.quoteDotted {
t.Errorf("%s dotted: %s", c.name, d.Quote("auth.users"))
}
raw := map[string]string{"postgres": `a"b`, "sqlite": `a"b`, "mysql": "a`b", "mssql": "a]b"}[c.name]
q := c.quoteEscape
if got := d.Quote(raw); got != q {
t.Errorf("%s escape: got %s want %s", c.name, got, q)
}
}
}
func TestInsertReturningID(t *testing.T) {
cols := []string{"user_id", "name"}
cases := []struct {
name string
wantSQL string
strategy dialect.InsertStrategy
empty string
}{
{"postgres", `INSERT INTO "t" ("user_id", "name") VALUES ($1, $2) RETURNING "id"`, dialect.ReturningQuery, `INSERT INTO "t" DEFAULT VALUES RETURNING "id"`},
{"sqlite", `INSERT INTO "t" ("user_id", "name") VALUES (?, ?)`, dialect.LastInsertID, `INSERT INTO "t" DEFAULT VALUES`},
{"mysql", "INSERT INTO `t` (`user_id`, `name`) VALUES (?, ?)", dialect.LastInsertID, "INSERT INTO `t` () VALUES ()"},
{"mssql", `INSERT INTO [t] ([user_id], [name]) OUTPUT INSERTED.[id] VALUES (@p1, @p2)`, dialect.ReturningQuery, `INSERT INTO [t] OUTPUT INSERTED.[id] DEFAULT VALUES`},
}
for _, c := range cases {
ins := get(t, c.name).InsertReturningID("t", cols, "id")
if ins.SQL != c.wantSQL || ins.Strategy != c.strategy {
t.Errorf("%s: got %q (%d)", c.name, ins.SQL, ins.Strategy)
}
if e := get(t, c.name).InsertReturningID("t", nil, "id"); e.SQL != c.empty {
t.Errorf("%s empty: got %q", c.name, e.SQL)
}
}
}
func TestBoolRoundTrip(t *testing.T) {
for _, name := range dialect.Names() {
if name == "custom" {
continue
}
d := get(t, name)
for _, v := range []bool{true, false} {
got, err := d.ScanBool(d.Bool(v))
if err != nil || got != v {
t.Errorf("%s: Bool(%v) round trip = %v, %v", name, v, got, err)
}
}
for _, c := range []struct {
in any
want bool
}{{int64(1), true}, {int64(0), false}, {"1", true}, {"false", false}, {[]byte("t"), true}, {nil, false}} {
if got, err := d.ScanBool(c.in); err != nil || got != c.want {
t.Errorf("%s: ScanBool(%v) = %v, %v", name, c.in, got, err)
}
}
if _, err := d.ScanBool("maybe"); err == nil {
t.Errorf("%s: ScanBool(maybe) should fail", name)
}
}
}
func TestScanTime(t *testing.T) {
d := get(t, "sqlite")
want := time.Date(2026, 3, 4, 5, 6, 7, 0, time.UTC)
for _, in := range []any{
want,
"2026-03-04 05:06:07+00:00",
"2026-03-04T05:06:07Z",
"2026-03-04 05:06:07",
[]byte("2026-03-04 05:06:07"),
"2026-03-04 05:06:07 +0000 UTC",
} {
got, err := d.ScanTime(in)
if err != nil || !got.Equal(want) {
t.Errorf("ScanTime(%v) = %v, %v", in, got, err)
}
}
if got, err := d.ScanTime(nil); err != nil || !got.IsZero() {
t.Errorf("ScanTime(nil) = %v, %v", got, err)
}
if _, err := d.ScanTime("not a time"); err == nil {
t.Error("ScanTime(garbage) should fail")
}
}
func TestJSON(t *testing.T) {
for _, name := range []string{"postgres", "sqlite", "mysql", "mssql"} {
d := get(t, name)
enc, err := d.EncodeJSON([]string{"a", "b"})
if err != nil || enc != `["a","b"]` {
t.Errorf("%s encode = %v, %v", name, enc, err)
}
var nilSlice []string
var nilMap map[string]any
for _, v := range []any{nil, nilSlice, nilMap} {
if enc, err := d.EncodeJSON(v); err != nil || enc != nil {
t.Errorf("%s: EncodeJSON(%T nil) = %v, %v; want SQL NULL", name, v, enc, err)
}
}
var out []string
if err := d.DecodeJSON([]byte(`["x"]`), &out); err != nil || len(out) != 1 || out[0] != "x" {
t.Errorf("%s decode bytes = %v, %v", name, out, err)
}
out = []string{"keep"}
if err := d.DecodeJSON(nil, &out); err != nil || out[0] != "keep" {
t.Errorf("%s decode NULL should leave dst: %v, %v", name, out, err)
}
if err := d.DecodeJSON("", &out); err != nil || out[0] != "keep" {
t.Errorf("%s decode empty should leave dst: %v, %v", name, out, err)
}
if err := d.DecodeJSON("{bad", &out); err == nil {
t.Errorf("%s decode of invalid json should fail", name)
}
}
}
func TestDetectDriver(t *testing.T) {
cases := map[string]string{
"github.com/jackc/pgx/v5/stdlib.driver": "postgres",
"github.com/lib/pq.driver": "postgres",
"github.com/mattn/go-sqlite3.sqlitedriver": "sqlite",
"modernc.org/sqlite.driver": "sqlite",
"github.com/go-sql-driver/mysql.mysqldriver": "mysql",
"github.com/microsoft/go-mssqldb.driver": "mssql",
"github.com/denisenkom/go-mssqldb.driver": "mssql",
}
for drv, want := range cases {
d, err := dialect.DetectDriver(drv)
if err != nil || d.Name() != want {
t.Errorf("%s: got %v, %v; want %s", drv, d, err, want)
}
}
if _, err := dialect.DetectDriver("example.com/unknown.driver"); err == nil {
t.Error("unknown driver should fail with an explicit-dialect hint")
}
if _, err := dialect.Detect(nil); err == nil {
t.Error("nil db should fail")
}
}
// TestSQLiteRoundTrip runs the sqlite dialect against a real in-memory database.
func TestSQLiteRoundTrip(t *testing.T) {
db, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatal(err)
}
defer db.Close()
db.SetMaxOpenConns(1)
d, err := dialect.Detect(db)
if err != nil || d.Name() != "sqlite" {
t.Fatalf("Detect = %v, %v", d, err)
}
ctx := context.Background()
if _, err := db.ExecContext(ctx, `CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT, active BOOLEAN, scopes TEXT, at TIMESTAMP)`); err != nil {
t.Fatal(err)
}
now := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC)
scopes, _ := d.EncodeJSON([]string{"read", "write"})
ins := d.InsertReturningID("t", []string{"name", "active", "scopes", "at"}, "id")
id, err := ins.Run(ctx, db, "k1", d.Bool(true), scopes, now)
if err != nil || id != 1 {
t.Fatalf("insert = %d, %v", id, err)
}
id, err = ins.Run(ctx, db, "k2", d.Bool(false), nil, now)
if err != nil || id != 2 {
t.Fatalf("second insert = %d, %v", id, err)
}
q := "SELECT active, scopes, at FROM " + d.Quote("t") + " WHERE " + d.Quote("id") + " = " + d.Placeholder(1)
var active, scopesRaw, at any
if err := db.QueryRowContext(ctx, q, 1).Scan(&active, &scopesRaw, &at); err != nil {
t.Fatal(err)
}
if b, err := d.ScanBool(active); err != nil || !b {
t.Errorf("active = %v, %v", b, err)
}
var got []string
if err := d.DecodeJSON(scopesRaw, &got); err != nil || len(got) != 2 {
t.Errorf("scopes = %v, %v", got, err)
}
if ts, err := d.ScanTime(at); err != nil || !ts.Equal(now) {
t.Errorf("at = %v (%T), %v", ts, at, err)
}
if err := db.QueryRowContext(ctx, q, 2).Scan(&active, &scopesRaw, &at); err != nil {
t.Fatal(err)
}
if b, _ := d.ScanBool(active); b {
t.Error("second row should be inactive")
}
if scopesRaw != nil {
t.Errorf("nil scopes should be NULL, got %v", scopesRaw)
}
// Insert inside a transaction goes through the same Querier.
tx, _ := db.BeginTx(ctx, nil)
if id, err := d.InsertReturningID("t", nil, "id").Run(ctx, tx); err != nil || id != 3 {
t.Errorf("DEFAULT VALUES insert in tx = %d, %v", id, err)
}
_ = tx.Rollback()
}
+608
View File
@@ -0,0 +1,608 @@
package direct
import (
"context"
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/hex"
"errors"
"fmt"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
const sessionLifetime = 24 * time.Hour
// AuthOptions tunes Auth.
type AuthOptions struct {
// UpgradePasswordHash rewrites legacy cleartext passwords as bcrypt on a successful login.
UpgradePasswordHash bool
}
// Auth implements lookup.AuthStore on the tables. Passwords are verified with bcrypt;
// legacy cleartext rows are accepted at login and only rewritten when UpgradePasswordHash
// is set. Registration never honours client-supplied user_level or roles. Multi-step writes
// (login, register, refresh, reset) run in one transaction.
type Auth struct {
*Base
opts AuthOptions
}
var _ lookup.AuthStore = (*Auth)(nil)
// NewAuth creates the direct AuthStore.
func NewAuth(b *Base, opts AuthOptions) *Auth { return &Auth{Base: b, opts: opts} }
// GenerateSessionToken returns "sess_<64 hex>_<unix>".
func GenerateSessionToken() (string, error) {
buf := make([]byte, 32)
if _, err := rand.Read(buf); err != nil {
return "", err
}
return fmt.Sprintf("sess_%s_%d", hex.EncodeToString(buf), time.Now().Unix()), nil
}
// ParseRoles splits the comma-separated roles column.
func ParseRoles(s string) []string {
if s == "" {
return []string{}
}
return strings.Split(s, ",")
}
func claimStrings(claims map[string]any) (ip, ua string) {
if claims == nil {
return "", ""
}
if v, ok := claims["ip_address"].(string); ok {
ip = v
}
if v, ok := claims["user_agent"].(string); ok {
ua = v
}
return ip, ua
}
func sha256Hex(s string) string {
h := sha256.Sum256([]byte(s))
return hex.EncodeToString(h[:])
}
// userRow is the users columns every session-bearing response needs.
type userRow struct {
id int
username sql.NullString
email sql.NullString
roles sql.NullString
programUserTable sql.NullString
userLevel sql.NullInt64
programUserID sql.NullInt64
}
func (u *userRow) context(sessionID string) *sectypes.UserContext {
return &sectypes.UserContext{
UserID: u.id,
UserName: u.username.String,
Email: u.email.String,
UserLevel: int(u.userLevel.Int64),
SessionID: sessionID,
Roles: ParseRoles(u.roles.String),
ProgramUserID: int(u.programUserID.Int64),
ProgramUserTable: u.programUserTable.String,
}
}
// insertSession writes a session row and stamps the user's last login.
func (a *Auth) insertSession(ctx context.Context, q Querier, token string, userID int64, expiresAt time.Time, ip, ua string, now time.Time) error {
err := a.Insert(lookup.EntityUserSessions).Set(
Set(lookup.SessionsToken, token),
Set(lookup.SessionsUserID, userID),
Set(lookup.SessionsExpiresAt, expiresAt),
Set(lookup.SessionsIPAddress, ip),
Set(lookup.SessionsUserAgent, ua),
Set(lookup.SessionsLastActivityAt, now),
Set(lookup.SessionsCreatedAt, now),
).Exec(ctx, q)
if err != nil {
return err
}
return a.touchLastLogin(ctx, q, userID, now)
}
func (a *Auth) touchLastLogin(ctx context.Context, q Querier, userID int64, now time.Time) error {
_, err := a.Update(lookup.EntityUsers).Set(Set(lookup.UsersLastLoginAt, now)).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
return err
}
// Login implements lookup.AuthStore.
func (a *Auth) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
var userID int
var email, roles, programUserTable, storedPassword sql.NullString
var userLevel, programUserID sql.NullInt64
err := a.do(func(q Querier) error {
return a.From(lookup.EntityUsers).
Cols(lookup.UsersID, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles,
lookup.UsersProgramUserID, lookup.UsersProgramUserTable, lookup.UsersPassword).
Where(Eq(lookup.UsersUsername, req.Username), Eq(lookup.UsersIsActive, true)).
QueryRow(ctx, q, &userID, &email, &userLevel, &roles, &programUserID, &programUserTable, &storedPassword)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
BurnPasswordCheck(req.Password)
return nil, fmt.Errorf("invalid credentials")
}
return nil, fmt.Errorf("login query failed: %w", err)
}
ok, needsRehash := VerifyPassword(storedPassword.String, req.Password)
if !ok {
if storedPassword.String == "" {
BurnPasswordCheck(req.Password)
}
return nil, fmt.Errorf("invalid credentials")
}
if needsRehash && a.opts.UpgradePasswordHash {
a.upgradePasswordHash(ctx, userID, req.Password)
}
token, err := GenerateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
now := a.Now()
ip, ua := claimStrings(req.Claims)
err = a.tx(ctx, func(q Querier) error {
return a.insertSession(ctx, q, token, int64(userID), now.Add(sessionLifetime), ip, ua, now)
})
if err != nil {
return nil, fmt.Errorf("login query failed: %w", err)
}
return &sectypes.LoginResponse{
Token: token,
User: &sectypes.UserContext{
UserID: userID,
UserName: req.Username,
Email: email.String,
UserLevel: int(userLevel.Int64),
Roles: ParseRoles(roles.String),
SessionID: token,
ProgramUserID: int(programUserID.Int64),
ProgramUserTable: programUserTable.String,
},
ExpiresIn: int64(sessionLifetime.Seconds()),
}, nil
}
// upgradePasswordHash replaces a legacy cleartext password with a bcrypt hash. Failure is
// logged and ignored: the login itself already succeeded.
func (a *Auth) upgradePasswordHash(ctx context.Context, userID int, password string) {
h, err := HashPassword(password)
if err != nil {
return
}
err = a.do(func(q Querier) error {
_, err := a.Update(lookup.EntityUsers).
Set(Set(lookup.UsersPassword, h), Set(lookup.UsersUpdatedAt, a.Now())).
Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
return err
})
if err != nil {
logger.Warn("failed to upgrade legacy password hash for user %d: %v", userID, err)
}
}
// Register implements lookup.AuthStore.
func (a *Auth) Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) {
if req.Username == "" {
return nil, fmt.Errorf("username is required")
}
if req.Email == "" {
return nil, fmt.Errorf("email is required")
}
if req.Password == "" {
return nil, fmt.Errorf("password is required")
}
hash, err := HashPassword(req.Password)
if err != nil {
return nil, err
}
token, err := GenerateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
// Privileges are never taken from the request: self-registration always creates an
// unprivileged user.
const userLevel = 0
now := a.Now()
ip, ua := claimStrings(req.Claims)
var userID int64
err = a.tx(ctx, func(q Querier) error {
exists, err := a.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersUsername, req.Username)).Exists(ctx, q)
if err != nil {
return err
}
if exists {
return lookup.ErrUsernameExists
}
exists, err = a.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersEmail, req.Email)).Exists(ctx, q)
if err != nil {
return err
}
if exists {
return lookup.ErrEmailExists
}
userID, err = a.Insert(lookup.EntityUsers).Set(
Set(lookup.UsersUsername, req.Username),
Set(lookup.UsersEmail, req.Email),
Set(lookup.UsersPassword, hash),
Set(lookup.UsersUserLevel, userLevel),
Set(lookup.UsersRoles, ""),
Set(lookup.UsersIsActive, true),
Set(lookup.UsersCreatedAt, now),
Set(lookup.UsersUpdatedAt, now),
Set(lookup.UsersProgramUserID, 0),
Set(lookup.UsersProgramUserTable, ""),
).ExecID(ctx, q, lookup.UsersID)
if err != nil {
return err
}
return a.insertSession(ctx, q, token, userID, now.Add(sessionLifetime), ip, ua, now)
})
if err != nil {
if errors.Is(err, lookup.ErrUsernameExists) || errors.Is(err, lookup.ErrEmailExists) {
return nil, err
}
return nil, fmt.Errorf("register query failed: %w", err)
}
return &sectypes.LoginResponse{
Token: token,
User: &sectypes.UserContext{
UserID: int(userID),
UserName: req.Username,
Email: req.Email,
UserLevel: userLevel,
Roles: ParseRoles(""),
SessionID: token,
},
ExpiresIn: int64(sessionLifetime.Seconds()),
}, nil
}
// Logout implements lookup.AuthStore.
func (a *Auth) Logout(ctx context.Context, req sectypes.LogoutRequest) error {
token := strings.TrimPrefix(strings.TrimPrefix(req.Token, "Bearer "), "bearer ")
var rows int64
err := a.do(func(q Querier) error {
var err error
rows, err = a.Delete(lookup.EntityUserSessions).
Where(Eq(lookup.SessionsToken, token), Eq(lookup.SessionsUserID, req.UserID)).Exec(ctx, q)
return err
})
if err != nil {
return fmt.Errorf("logout query failed: %w", err)
}
if rows == 0 {
return fmt.Errorf("session not found")
}
return nil
}
// sessionUser selects the user behind a live session token.
func (a *Auth) sessionUser(ctx context.Context, q Querier, token string, extra ...lookup.Column) (*userRow, []any, error) {
var u userRow
dest := []any{&u.id, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable}
cols := []lookup.Column{lookup.SessionsUserID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel,
lookup.UsersRoles, lookup.UsersProgramUserID, lookup.UsersProgramUserTable}
extras := make([]any, len(extra))
for i, c := range extra {
cols = append(cols, c)
extras[i] = new(sql.NullString)
dest = append(dest, extras[i])
}
err := a.From(lookup.EntityUserSessions).Cols(cols...).
Join(lookup.EntityUsers, EqCol(lookup.SessionsUserID, lookup.UsersID)).
Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, a.Now()), Eq(lookup.UsersIsActive, true)).
QueryRow(ctx, q, dest...)
return &u, extras, err
}
// Session implements lookup.AuthStore. reference is only meaningful to the procedure backend.
func (a *Auth) Session(ctx context.Context, token, _ string) (*sectypes.UserContext, error) {
var u *userRow
err := a.do(func(q Querier) error {
var err error
u, _, err = a.sessionUser(ctx, q, token)
return err
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("invalid or expired session")
}
return nil, fmt.Errorf("session query failed: %w", err)
}
return u.context(token), nil
}
// TouchSession implements lookup.AuthStore.
func (a *Auth) TouchSession(ctx context.Context, token string, _ *sectypes.UserContext) error {
return a.do(func(q Querier) error {
now := a.Now()
_, err := a.Update(lookup.EntityUserSessions).Set(Set(lookup.SessionsLastActivityAt, now)).
Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, now)).Exec(ctx, q)
return err
})
}
// Refresh implements lookup.AuthStore: the old session is replaced by a new one.
func (a *Auth) Refresh(ctx context.Context, oldToken string) (*sectypes.LoginResponse, error) {
var u *userRow
var extras []any
err := a.do(func(q Querier) error {
var err error
u, extras, err = a.sessionUser(ctx, q, oldToken, lookup.SessionsIPAddress, lookup.SessionsUserAgent)
return err
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("invalid or expired refresh token")
}
return nil, fmt.Errorf("refresh token query failed: %w", err)
}
ip := extras[0].(*sql.NullString).String
ua := extras[1].(*sql.NullString).String
newToken, err := GenerateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
now := a.Now()
err = a.tx(ctx, func(q Querier) error {
err := a.Insert(lookup.EntityUserSessions).Set(
Set(lookup.SessionsToken, newToken),
Set(lookup.SessionsUserID, u.id),
Set(lookup.SessionsExpiresAt, now.Add(sessionLifetime)),
Set(lookup.SessionsIPAddress, ip),
Set(lookup.SessionsUserAgent, ua),
Set(lookup.SessionsLastActivityAt, now),
Set(lookup.SessionsCreatedAt, now),
).Exec(ctx, q)
if err != nil {
return err
}
_, err = a.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsToken, oldToken)).Exec(ctx, q)
return err
})
if err != nil {
return nil, fmt.Errorf("refresh token generation failed: %w", err)
}
return &sectypes.LoginResponse{
Token: newToken,
User: u.context(newToken),
ExpiresIn: int64(sessionLifetime.Seconds()),
}, nil
}
// apiKeyTypes are the key types accepted by LoginAPIKey.
var apiKeyTypes = []any{string(sectypes.KeyTypeHeaderAPI), string(sectypes.KeyTypeGenericAPI)}
// LoginAPIKey implements lookup.AuthStore. Unknown, expired, inactive and wrong-type keys
// (and inactive users) all return lookup.ErrInvalidAPIKey; the raw key is never logged.
func (a *Auth) LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error) {
if rawKey == "" {
return nil, lookup.ErrInvalidAPIKey
}
now := a.Now()
var keyID int64
var u userRow
err := a.do(func(q Querier) error {
return a.From(lookup.EntityUserKeys).
Cols(lookup.KeysID, lookup.UsersID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel,
lookup.UsersRoles, lookup.UsersProgramUserID, lookup.UsersProgramUserTable).
Join(lookup.EntityUsers, EqCol(lookup.KeysUserID, lookup.UsersID)).
Where(
Eq(lookup.KeysKeyHash, sectypes.HashKey(rawKey)),
In(lookup.KeysKeyType, apiKeyTypes...),
Eq(lookup.KeysIsActive, true),
Or(IsNull(lookup.KeysExpiresAt), Gt(lookup.KeysExpiresAt, now)),
Eq(lookup.UsersIsActive, true),
).QueryRow(ctx, q, &keyID, &u.id, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, lookup.ErrInvalidAPIKey
}
return nil, fmt.Errorf("api key login query failed: %w", err)
}
token, err := GenerateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
ip, ua := claimStrings(claims)
err = a.tx(ctx, func(q Querier) error {
if err := a.insertSession(ctx, q, token, int64(u.id), now.Add(sessionLifetime), ip, ua, now); err != nil {
return err
}
_, err := a.Update(lookup.EntityUserKeys).Set(Set(lookup.KeysLastUsedAt, now)).Where(Eq(lookup.KeysID, keyID)).Exec(ctx, q)
return err
})
if err != nil {
return nil, fmt.Errorf("api key login query failed: %w", err)
}
return &sectypes.LoginResponse{
Token: token,
User: u.context(token),
ExpiresIn: int64(sessionLifetime.Seconds()),
}, nil
}
// JWTLogin implements lookup.AuthStore (mirrors resolvespec_jwt_login). The token is a
// placeholder until JWT signing is wired in.
func (a *Auth) JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
var userID int
var email, roles, storedPassword sql.NullString
var userLevel sql.NullInt64
err := a.do(func(q Querier) error {
return a.From(lookup.EntityUsers).
Cols(lookup.UsersID, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles, lookup.UsersPassword).
Where(Eq(lookup.UsersUsername, req.Username), Eq(lookup.UsersIsActive, true)).
QueryRow(ctx, q, &userID, &email, &userLevel, &roles, &storedPassword)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
BurnPasswordCheck(req.Password)
return nil, fmt.Errorf("invalid credentials")
}
return nil, fmt.Errorf("login query failed: %w", err)
}
ok, needsRehash := VerifyPassword(storedPassword.String, req.Password)
if !ok {
if storedPassword.String == "" {
BurnPasswordCheck(req.Password)
}
return nil, fmt.Errorf("invalid credentials")
}
if needsRehash && a.opts.UpgradePasswordHash {
a.upgradePasswordHash(ctx, userID, req.Password)
}
expiresAt := a.Now().Add(sessionLifetime)
return &sectypes.LoginResponse{
Token: fmt.Sprintf("token_%d_%d", userID, expiresAt.Unix()),
User: &sectypes.UserContext{
UserID: userID,
UserName: req.Username,
Email: email.String,
UserLevel: int(userLevel.Int64),
Roles: ParseRoles(roles.String),
},
ExpiresIn: int64(sessionLifetime.Seconds()),
}, nil
}
// JWTLogout implements lookup.AuthStore: the token goes on the blacklist.
func (a *Auth) JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error {
now := a.Now()
err := a.do(func(q Querier) error {
return a.Insert(lookup.EntityTokenBlacklist).Set(
Set(lookup.BlacklistToken, req.Token),
Set(lookup.BlacklistUserID, req.UserID),
Set(lookup.BlacklistExpiresAt, now.Add(sessionLifetime)),
Set(lookup.BlacklistCreatedAt, now),
).Exec(ctx, q)
})
if err != nil {
return fmt.Errorf("logout query failed: %w", err)
}
return nil
}
// ResetRequest implements lookup.AuthStore. An unknown user yields a generic empty success
// so accounts cannot be enumerated.
func (a *Auth) ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) {
if req.Email == "" && req.Username == "" {
return nil, fmt.Errorf("email or username is required")
}
var userID int
err := a.do(func(q Querier) error {
lookupCol, val := lookup.UsersUsername, req.Username
if req.Email != "" {
lookupCol, val = lookup.UsersEmail, req.Email
}
return a.From(lookup.EntityUsers).Cols(lookup.UsersID).
Where(Eq(lookupCol, val), Eq(lookup.UsersIsActive, true)).QueryRow(ctx, q, &userID)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return &sectypes.PasswordResetResponse{Token: "", ExpiresIn: 0}, nil
}
return nil, fmt.Errorf("password reset request query failed: %w", err)
}
raw := make([]byte, 32)
if _, err := rand.Read(raw); err != nil {
return nil, fmt.Errorf("failed to generate reset token: %w", err)
}
rawToken := hex.EncodeToString(raw)
now := a.Now()
err = a.tx(ctx, func(q Querier) error {
if _, err := a.Delete(lookup.EntityUserPasswordResets).
Where(Eq(lookup.ResetsUserID, userID), Eq(lookup.ResetsUsed, false)).Exec(ctx, q); err != nil {
return err
}
return a.Insert(lookup.EntityUserPasswordResets).Set(
Set(lookup.ResetsUserID, userID),
Set(lookup.ResetsTokenHash, sha256Hex(rawToken)),
Set(lookup.ResetsExpiresAt, now.Add(time.Hour)),
Set(lookup.ResetsCreatedAt, now),
Set(lookup.ResetsUsed, false),
).Exec(ctx, q)
})
if err != nil {
return nil, fmt.Errorf("password reset request query failed: %w", err)
}
return &sectypes.PasswordResetResponse{Token: rawToken, ExpiresIn: 3600}, nil
}
// ResetComplete implements lookup.AuthStore: sets the new password, ends every session of
// the user and consumes the reset token, atomically.
func (a *Auth) ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error {
if req.Token == "" {
return fmt.Errorf("token is required")
}
if req.NewPassword == "" {
return fmt.Errorf("new_password is required")
}
newHash, err := HashPassword(req.NewPassword)
if err != nil {
return err
}
tokenHash := sha256Hex(req.Token)
now := a.Now()
var resetID, userID int
var expiresAt time.Time
err = a.do(func(q Querier) error {
return a.From(lookup.EntityUserPasswordResets).
Cols(lookup.ResetsID, lookup.ResetsUserID, lookup.ResetsExpiresAt).
Where(Eq(lookup.ResetsTokenHash, tokenHash), Eq(lookup.ResetsUsed, false)).
QueryRow(ctx, q, &resetID, &userID, a.timeDest(&expiresAt))
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return fmt.Errorf("invalid or expired token")
}
return fmt.Errorf("password reset complete query failed: %w", err)
}
if !expiresAt.After(now) {
return fmt.Errorf("invalid or expired token")
}
err = a.tx(ctx, func(q Querier) error {
if _, err := a.Update(lookup.EntityUsers).
Set(Set(lookup.UsersPassword, newHash), Set(lookup.UsersUpdatedAt, now)).
Where(Eq(lookup.UsersID, userID)).Exec(ctx, q); err != nil {
return err
}
if _, err := a.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsUserID, userID)).Exec(ctx, q); err != nil {
return err
}
_, err := a.Update(lookup.EntityUserPasswordResets).
Set(Set(lookup.ResetsUsed, true), Set(lookup.ResetsUsedAt, now)).
Where(Eq(lookup.ResetsID, resetID)).Exec(ctx, q)
return err
})
if err != nil {
return fmt.Errorf("password reset complete query failed: %w", err)
}
return nil
}
+268
View File
@@ -0,0 +1,268 @@
package direct
import (
"context"
"database/sql"
"errors"
"strings"
"testing"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
func newAuth(t *testing.T, opts AuthOptions) (*Auth, *sql.DB) {
db := newTestDB(t)
return NewAuth(newTestBase(t, db, nil), opts), db
}
func registerUser(t *testing.T, a *Auth, name string) *sectypes.LoginResponse {
t.Helper()
resp, err := a.Register(context.Background(), sectypes.RegisterRequest{Username: name, Email: name + "@x.io", Password: "pw-" + name})
if err != nil {
t.Fatal(err)
}
return resp
}
func TestRegisterLoginSessionFlow(t *testing.T) {
ctx := context.Background()
a, db := newAuth(t, AuthOptions{})
reg, err := a.Register(ctx, sectypes.RegisterRequest{
Username: "ann", Email: "ann@x.io", Password: "secret",
UserLevel: 99, Roles: []string{"admin"}, // must be ignored
})
if err != nil {
t.Fatal(err)
}
if reg.User.UserLevel != 0 || len(reg.User.Roles) != 0 {
t.Fatalf("register honoured privileges: %+v", reg.User)
}
var stored string
if err := db.QueryRow(`SELECT password FROM users WHERE username='ann'`).Scan(&stored); err != nil || !strings.HasPrefix(stored, "$2") {
t.Fatalf("password not bcrypt: %q %v", stored, err)
}
if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "ann", Email: "other@x.io", Password: "x"}); !errors.Is(err, lookup.ErrUsernameExists) {
t.Fatalf("got %v", err)
}
if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "bob", Email: "ann@x.io", Password: "x"}); !errors.Is(err, lookup.ErrEmailExists) {
t.Fatalf("got %v", err)
}
var n int
_ = db.QueryRow(`SELECT COUNT(*) FROM users`).Scan(&n)
if n != 1 {
t.Fatalf("failed register left a row: %d", n)
}
login, err := a.Login(ctx, sectypes.LoginRequest{Username: "ann", Password: "secret", Claims: map[string]any{"ip_address": "1.2.3.4", "user_agent": "ua"}})
if err != nil {
t.Fatal(err)
}
if !strings.HasPrefix(login.Token, "sess_") || login.ExpiresIn != 86400 || login.User.Email != "ann@x.io" {
t.Fatalf("login: %+v", login)
}
var ip string
_ = db.QueryRow(`SELECT ip_address FROM user_sessions WHERE session_token=?`, login.Token).Scan(&ip)
if ip != "1.2.3.4" {
t.Fatalf("ip %q", ip)
}
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "ann", Password: "wrong"}); err == nil || err.Error() != "invalid credentials" {
t.Fatalf("got %v", err)
}
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "nobody", Password: "x"}); err == nil || err.Error() != "invalid credentials" {
t.Fatalf("got %v", err)
}
u, err := a.Session(ctx, login.Token, "authenticate")
if err != nil || u.UserName != "ann" || u.SessionID != login.Token {
t.Fatalf("session: %+v %v", u, err)
}
if err := a.TouchSession(ctx, login.Token, u); err != nil {
t.Fatal(err)
}
if _, err := a.Session(ctx, "nope", "authenticate"); err == nil || err.Error() != "invalid or expired session" {
t.Fatalf("got %v", err)
}
ref, err := a.Refresh(ctx, login.Token)
if err != nil || ref.Token == login.Token {
t.Fatalf("refresh: %+v %v", ref, err)
}
if _, err := a.Session(ctx, login.Token, ""); err == nil {
t.Fatal("old session still valid after refresh")
}
if _, err := a.Refresh(ctx, login.Token); err == nil || err.Error() != "invalid or expired refresh token" {
t.Fatalf("got %v", err)
}
if err := a.Logout(ctx, sectypes.LogoutRequest{Token: "Bearer " + ref.Token, UserID: ref.User.UserID}); err != nil {
t.Fatal(err)
}
if err := a.Logout(ctx, sectypes.LogoutRequest{Token: ref.Token, UserID: ref.User.UserID}); err == nil || err.Error() != "session not found" {
t.Fatalf("got %v", err)
}
}
func TestExpiredSessionRejected(t *testing.T) {
ctx := context.Background()
a, _ := newAuth(t, AuthOptions{})
resp := registerUser(t, a, "eve")
a.Now = func() time.Time { return time.Now().Add(48 * time.Hour) }
if _, err := a.Session(ctx, resp.Token, ""); err == nil {
t.Fatal("expired session accepted")
}
}
func TestLegacyPasswordUpgradeIsOptIn(t *testing.T) {
ctx := context.Background()
for _, upgrade := range []bool{false, true} {
a, db := newAuth(t, AuthOptions{UpgradePasswordHash: upgrade})
_, err := db.Exec(`INSERT INTO users (username, email, password, user_level, roles, is_active) VALUES ('old','o@x.io','clear',1,'a,b',1)`)
if err != nil {
t.Fatal(err)
}
resp, err := a.Login(ctx, sectypes.LoginRequest{Username: "old", Password: "clear"})
if err != nil || len(resp.User.Roles) != 2 {
t.Fatalf("login: %+v %v", resp, err)
}
var stored string
_ = db.QueryRow(`SELECT password FROM users WHERE username='old'`).Scan(&stored)
if got := strings.HasPrefix(stored, "$2"); got != upgrade {
t.Fatalf("upgrade=%v stored=%q", upgrade, stored)
}
}
}
func TestInactiveUserCannotLogin(t *testing.T) {
ctx := context.Background()
a, db := newAuth(t, AuthOptions{})
resp := registerUser(t, a, "ian")
_, _ = db.Exec(`UPDATE users SET is_active = 0`)
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "ian", Password: "pw-ian"}); err == nil {
t.Fatal("inactive login accepted")
}
if _, err := a.Session(ctx, resp.Token, ""); err == nil {
t.Fatal("inactive session accepted")
}
}
func TestLoginAPIKey(t *testing.T) {
ctx := context.Background()
a, db := newAuth(t, AuthOptions{})
reg := registerUser(t, a, "kim")
insert := func(raw, typ string, active int, expires any) {
t.Helper()
_, err := db.Exec(`INSERT INTO user_keys (user_id, key_type, key_hash, name, is_active, expires_at) VALUES (?,?,?,?,?,?)`,
reg.User.UserID, typ, sectypes.HashKey(raw), "k", active, expires)
if err != nil {
t.Fatal(err)
}
}
insert("good", "header_api", 1, nil)
insert("generic", "api", 1, nil)
insert("jwt", "jwt_secret", 1, nil)
insert("off", "api", 0, nil)
insert("old", "api", 1, time.Now().UTC().Add(-time.Hour))
for _, k := range []string{"good", "generic"} {
resp, err := a.LoginAPIKey(ctx, k, map[string]any{"ip_address": "9.9.9.9"})
if err != nil || resp.User.UserName != "kim" || !strings.HasPrefix(resp.Token, "sess_") {
t.Fatalf("%s: %+v %v", k, resp, err)
}
if _, err := a.Session(ctx, resp.Token, ""); err != nil {
t.Fatal(err)
}
}
var used sql.NullString
_ = db.QueryRow(`SELECT last_used_at FROM user_keys WHERE key_hash = ?`, sectypes.HashKey("good")).Scan(&used)
if !used.Valid {
t.Fatal("last_used_at not stamped")
}
for _, k := range []string{"", "missing", "jwt", "off", "old"} {
if _, err := a.LoginAPIKey(ctx, k, nil); !errors.Is(err, lookup.ErrInvalidAPIKey) {
t.Fatalf("%q: got %v", k, err)
}
}
_, _ = db.Exec(`UPDATE users SET is_active = 0`)
if _, err := a.LoginAPIKey(ctx, "good", nil); !errors.Is(err, lookup.ErrInvalidAPIKey) {
t.Fatalf("inactive user: got %v", err)
}
}
func TestPasswordReset(t *testing.T) {
ctx := context.Background()
a, _ := newAuth(t, AuthOptions{})
reg := registerUser(t, a, "rae")
empty, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: "none@x.io"})
if err != nil || empty.Token != "" {
t.Fatalf("enumeration leak: %+v %v", empty, err)
}
if _, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{}); err == nil {
t.Fatal("expected error")
}
r1, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: "rae@x.io"})
if err != nil || r1.Token == "" {
t.Fatal(err)
}
r2, _ := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Username: "rae"})
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r1.Token, NewPassword: "n"}); err == nil {
t.Fatal("superseded token accepted")
}
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r2.Token, NewPassword: "newpw"}); err != nil {
t.Fatal(err)
}
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r2.Token, NewPassword: "again"}); err == nil {
t.Fatal("token reused")
}
if _, err := a.Session(ctx, reg.Token, ""); err == nil {
t.Fatal("sessions survived reset")
}
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "rae", Password: "newpw"}); err != nil {
t.Fatal(err)
}
}
func TestJWTLoginLogout(t *testing.T) {
ctx := context.Background()
a, db := newAuth(t, AuthOptions{})
reg := registerUser(t, a, "jay")
resp, err := a.JWTLogin(ctx, sectypes.LoginRequest{Username: "jay", Password: "pw-jay"})
if err != nil || !strings.HasPrefix(resp.Token, "token_") {
t.Fatalf("%+v %v", resp, err)
}
if err := a.JWTLogout(ctx, sectypes.LogoutRequest{Token: "tok", UserID: reg.User.UserID}); err != nil {
t.Fatal(err)
}
var n int
_ = db.QueryRow(`SELECT COUNT(*) FROM token_blacklist WHERE token='tok'`).Scan(&n)
if n != 1 {
t.Fatal("token not blacklisted")
}
}
func TestCustomSchemaNames(t *testing.T) {
db := newTestDB(t, `
CREATE TABLE app_users (uid INTEGER PRIMARY KEY AUTOINCREMENT, login TEXT, email TEXT, password TEXT,
user_level INTEGER, roles TEXT, is_active INTEGER, created_at DATETIME, updated_at DATETIME,
last_login_at DATETIME, program_user_id INTEGER, program_user_table TEXT, remote_id TEXT, auth_provider TEXT,
totp_secret TEXT, totp_enabled INTEGER, totp_enabled_at DATETIME);`)
schema := lookup.Schema{lookup.EntityUsers: {Name: "app_users", Columns: map[string]string{"id": "uid", "username": "login"}}}
a := NewAuth(newTestBase(t, db, schema), AuthOptions{})
ctx := context.Background()
if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "zed", Email: "z@x.io", Password: "p"}); err != nil {
t.Fatal(err)
}
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "zed", Password: "p"}); err != nil {
t.Fatal(err)
}
var login string
if err := db.QueryRow(`SELECT login FROM app_users`).Scan(&login); err != nil || login != "zed" {
t.Fatalf("%q %v", login, err)
}
}

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