mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 19:20:31 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0d3ad9e4fd | ||
|
|
f5d232d971 | ||
|
|
a4702161fb | ||
|
|
640faeeeaf | ||
|
|
f54b707040 | ||
|
|
ca89cb8a73 | ||
|
|
c1153522f2 | ||
|
|
155e04deea | ||
|
|
e49c3a916e | ||
|
|
276c3814d8 | ||
|
|
82f901a49c | ||
|
|
ad2f54693f | ||
|
|
7662d5055c | ||
|
|
ea6a2e705f | ||
|
|
2516fcb13d | ||
|
|
c9fa8c60f2 | ||
|
|
60bd0a6dd3 | ||
|
|
982c90bfdd |
@@ -1,6 +1,3 @@
|
||||
# This workflow will build a golang project
|
||||
# For more information see: https://docs.github.com/en/actions/automating-builds-and-tests/building-and-testing-go
|
||||
|
||||
name: Create Go Release (Tag Versioning)
|
||||
|
||||
on:
|
||||
@@ -26,7 +23,9 @@ jobs:
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v2
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Set up Git
|
||||
run: |
|
||||
@@ -38,7 +37,7 @@ jobs:
|
||||
run: |
|
||||
git fetch --tags
|
||||
latest_tag=$(git describe --tags `git rev-list --tags --max-count=1`)
|
||||
echo "::set-output name=tag::$latest_tag"
|
||||
echo "tag=${latest_tag}" >> "${GITHUB_OUTPUT}"
|
||||
|
||||
- name: Determine new tag version
|
||||
id: new_tag
|
||||
@@ -57,7 +56,7 @@ jobs:
|
||||
((minor++))
|
||||
patch=0
|
||||
;;
|
||||
"release")
|
||||
"major")
|
||||
((major++))
|
||||
minor=0
|
||||
patch=0
|
||||
@@ -68,15 +67,11 @@ jobs:
|
||||
;;
|
||||
esac
|
||||
new_tag="v$major.$minor.$patch"
|
||||
echo "::set-output name=tag::$new_tag"
|
||||
echo "tag=${new_tag}" >> "${GITHUB_OUTPUT}"
|
||||
|
||||
- name: Create tag
|
||||
run: |
|
||||
git tag -a ${{ steps.new_tag.outputs.tag }} -m "Tagging ${{ steps.new_tag.outputs.tag }} for release"
|
||||
|
||||
- name: Push changes
|
||||
uses: ad-m/github-push-action@master
|
||||
with:
|
||||
github_token: ${{ secrets.BITECH_GITHUB_TOKEN }}
|
||||
force: true
|
||||
tags: true
|
||||
- name: Push tag
|
||||
run: git push origin ${{ steps.new_tag.outputs.tag }}
|
||||
@@ -0,0 +1,264 @@
|
||||
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
|
||||
run: |
|
||||
rustup toolchain install stable --profile minimal
|
||||
rustup default 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
|
||||
|
||||
- name: Dry run
|
||||
if: ${{ env.PUBLISH != 'true' }}
|
||||
run: dart pub publish --dry-run
|
||||
|
||||
- name: Publish
|
||||
if: ${{ env.PUBLISH == 'true' }}
|
||||
run: |
|
||||
dart pub token add "${SERVER_URL}/api/packages/${OWNER}/pub" --env-var TOKEN
|
||||
dart pub publish --force
|
||||
@@ -9,9 +9,9 @@ jobs:
|
||||
name: Unit Tests
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v6
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version: "1.24"
|
||||
- name: Run unit tests
|
||||
@@ -22,7 +22,7 @@ jobs:
|
||||
go test ./pkg/resolvespec ./pkg/restheadspec -coverprofile=coverage.out
|
||||
go tool cover -html=coverage.out -o coverage.html
|
||||
- name: Upload coverage
|
||||
uses: actions/upload-artifact@v5
|
||||
uses: actions/upload-artifact@v3
|
||||
continue-on-error: true
|
||||
with:
|
||||
name: coverage-report
|
||||
@@ -31,15 +31,16 @@ jobs:
|
||||
name: Race Detector
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v6
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version: "1.24"
|
||||
- name: Run unit tests with the race detector
|
||||
run: go test -race -count=1 ./pkg/...
|
||||
integration-tests:
|
||||
name: Integration Tests
|
||||
if: false # disabled for now
|
||||
runs-on: ubuntu-latest
|
||||
services:
|
||||
postgres:
|
||||
@@ -56,46 +57,51 @@ jobs:
|
||||
ports:
|
||||
- 5432:5432
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v6
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version: "1.24"
|
||||
- name: Install PostgreSQL client
|
||||
run: |
|
||||
SUDO=""; [ "$(id -u)" -ne 0 ] && SUDO="sudo"
|
||||
$SUDO apt-get update -qq
|
||||
$SUDO apt-get install -y -qq postgresql-client
|
||||
- name: Create test databases
|
||||
env:
|
||||
PGPASSWORD: postgres
|
||||
run: |
|
||||
psql -h localhost -U postgres -c "CREATE DATABASE resolvespec_test;"
|
||||
psql -h localhost -U postgres -c "CREATE DATABASE restheadspec_test;"
|
||||
psql -h postgres -U postgres -c "CREATE DATABASE resolvespec_test;"
|
||||
psql -h postgres -U postgres -c "CREATE DATABASE restheadspec_test;"
|
||||
- name: Run resolvespec integration tests
|
||||
continue-on-error: true
|
||||
env:
|
||||
TEST_DATABASE_URL: "host=localhost user=postgres password=postgres dbname=resolvespec_test port=5432 sslmode=disable"
|
||||
TEST_DATABASE_URL: "host=postgres user=postgres password=postgres dbname=resolvespec_test port=5432 sslmode=disable"
|
||||
run: go test -tags=integration ./pkg/resolvespec -v -coverprofile=coverage-resolvespec-integration.out
|
||||
- name: Run restheadspec integration tests
|
||||
continue-on-error: true
|
||||
env:
|
||||
TEST_DATABASE_URL: "host=localhost user=postgres password=postgres dbname=restheadspec_test port=5432 sslmode=disable"
|
||||
TEST_DATABASE_URL: "host=postgres user=postgres password=postgres dbname=restheadspec_test port=5432 sslmode=disable"
|
||||
run: go test -tags=integration ./pkg/restheadspec -v -coverprofile=coverage-restheadspec-integration.out
|
||||
- name: Generate integration coverage
|
||||
continue-on-error: true
|
||||
env:
|
||||
TEST_DATABASE_URL: "host=localhost user=postgres password=postgres dbname=resolvespec_test port=5432 sslmode=disable"
|
||||
TEST_DATABASE_URL: "host=postgres user=postgres password=postgres dbname=resolvespec_test port=5432 sslmode=disable"
|
||||
run: |
|
||||
go tool cover -html=coverage-resolvespec-integration.out -o coverage-resolvespec-integration.html
|
||||
go tool cover -html=coverage-restheadspec-integration.out -o coverage-restheadspec-integration.html
|
||||
|
||||
- name: Upload resolvespec integration coverage
|
||||
uses: actions/upload-artifact@v5
|
||||
uses: actions/upload-artifact@v3
|
||||
continue-on-error: true
|
||||
with:
|
||||
name: resolvespec-integration-coverage-report
|
||||
path: coverage-resolvespec-integration.html
|
||||
|
||||
- name: Upload restheadspec integration coverage
|
||||
uses: actions/upload-artifact@v5
|
||||
uses: actions/upload-artifact@v3
|
||||
continue-on-error: true
|
||||
|
||||
with:
|
||||
name: integration-coverage-restheadspec-report
|
||||
path: coverage-restheadspec-integration
|
||||
path: coverage-restheadspec-integration.html
|
||||
@@ -28,6 +28,7 @@ All share the same core architecture and provide dynamic data querying, relation
|
||||
- [Testing](#testing)
|
||||
- [Additional Packages](#additional-packages)
|
||||
- [Security Considerations](#security-considerations)
|
||||
- [Breaking Changes](#breaking-changes)
|
||||
- [What's New](#whats-new)
|
||||
|
||||
## Features
|
||||
@@ -277,32 +278,28 @@ ResolveMCP exposes registered models as Model Context Protocol tools so AI model
|
||||
```go
|
||||
import "github.com/bitechdev/ResolveSpec/pkg/resolvemcp"
|
||||
|
||||
// Create handler
|
||||
handler := resolvemcp.NewHandlerWithGORM(db)
|
||||
handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{BaseURL: "http://localhost:8080", BasePath: "/mcp"})
|
||||
|
||||
securityList, _ := security.NewSecurityList(provider)
|
||||
resolvemcp.RegisterSecurityHooks(handler, securityList)
|
||||
|
||||
// Register models — must be done BEFORE Build()
|
||||
handler.RegisterModel("public", "users", &User{})
|
||||
handler.RegisterModel("public", "posts", &Post{})
|
||||
|
||||
// Finalize: registers MCP tools and resources
|
||||
handler.Build()
|
||||
|
||||
// Mount SSE transport on your existing router
|
||||
// Mount the guarded SSE transport (OAuth bearer, session token or API key required)
|
||||
router := mux.NewRouter()
|
||||
resolvemcp.SetupMuxRoutes(router, handler, "http://localhost:8080")
|
||||
resolvemcp.SetupMuxRoutes(router, handler, securityList)
|
||||
|
||||
// MCP clients connect to:
|
||||
// SSE stream: GET http://localhost:8080/mcp/sse
|
||||
// Messages: POST http://localhost:8080/mcp/message
|
||||
//
|
||||
// Auto-registered tools per model:
|
||||
// read_public_users — filter, sort, paginate, preload
|
||||
// create_public_users — insert a new record
|
||||
// update_public_users — update a record by ID
|
||||
// delete_public_users — delete a record by ID
|
||||
// Fixed meta tools (independent of the number of models):
|
||||
// list_tables, describe_table, select_table, insert_into_table,
|
||||
// update_table, delete_from_table, list_functions, call_function
|
||||
```
|
||||
|
||||
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
|
||||
|
||||
@@ -646,9 +643,11 @@ For documentation, see [pkg/cache/README.md](pkg/cache/README.md).
|
||||
|
||||
#### 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
|
||||
|
||||
@@ -732,6 +731,58 @@ 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.
|
||||
|
||||
## 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
|
||||
|
||||
### Unreleased
|
||||
|
||||
+2
-2
@@ -1,6 +1,6 @@
|
||||
# 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
|
||||
|
||||
@@ -72,7 +72,7 @@ Same rules as resolvespec CRUD, plus guardrails.
|
||||
|
||||
### 1. API key login (`pkg/security`)
|
||||
- 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`.
|
||||
- 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.
|
||||
|
||||
@@ -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` |
|
||||
| **Tests** | `tools_test.go` (34), `tx_test.go` (207); `go test` passes. No hostile-input tests, no `-race` |
|
||||
| **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 |
|
||||
| **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) |
|
||||
|
||||
@@ -9,6 +9,7 @@
|
||||
| **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 |
|
||||
| **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 |
|
||||
| **Threat model** | hostile internet client; request bodies, headers, query params, schema/table/column names and filter expressions all attacker-controlled |
|
||||
| **Depth** | deep |
|
||||
|
||||
@@ -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
@@ -10,7 +10,7 @@
|
||||
- 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.
|
||||
|
||||
## 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 |
|
||||
|---|---|---|---|---|
|
||||
| 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`.
|
||||
- `OnTxBegin` failure aborts the whole request, rolls back, returns an error with no detail leaked to the client.
|
||||
|
||||
## Open
|
||||
- Consumer's ResolveSpec version: confirm it is >= v1.1.28 (read/create already in tx). Not blocking.
|
||||
## Status summary
|
||||
**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
|
||||
| # | 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 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.
|
||||
- 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`.
|
||||
- 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).
|
||||
@@ -103,7 +113,7 @@
|
||||
|
||||
## 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.
|
||||
- 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).
|
||||
@@ -118,5 +128,5 @@
|
||||
- `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.
|
||||
- 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`.
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
# Changelog
|
||||
|
||||
## 0.1.0
|
||||
|
||||
- Initial release: ResolveSpec (JSON body) and FunctionSpec client.
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2026 Hein
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -1,6 +1,6 @@
|
||||
# resolvespec-go
|
||||
|
||||
Go client for ResolveSpec (JSON body) and FunctionSpec. Module: `github.com/bitechdev/ResolveSpec/clients/resolvespec-go`. Stdlib only.
|
||||
Go client for ResolveSpec (JSON body) and FunctionSpec. Module: `git.warky.dev/wdevs/ResolveSpec/clients/resolvespec-go`. Stdlib only.
|
||||
|
||||
## Clients
|
||||
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
module github.com/bitechdev/ResolveSpec/clients/resolvespec-go
|
||||
module git.warky.dev/wdevs/ResolveSpec/clients/resolvespec-go
|
||||
|
||||
go 1.22
|
||||
|
||||
@@ -275,6 +275,14 @@ func (b *BunAdapter) GetUnderlyingDB() interface{} {
|
||||
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 {
|
||||
// Normalize Bun's dialect name to match the project's canonical vocabulary.
|
||||
// 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
|
||||
}
|
||||
|
||||
// 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 {
|
||||
// Extract optional prefix from args
|
||||
// If the last arg is a string that looks like a table prefix, use it
|
||||
|
||||
@@ -2,6 +2,7 @@ package database
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
@@ -227,6 +228,20 @@ func (g *GormAdapter) GetUnderlyingDB() interface{} {
|
||||
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 {
|
||||
return normalizeGormDriverName(g.getDB())
|
||||
}
|
||||
@@ -362,6 +377,16 @@ func (g *GormSelectQuery) WhereOr(query string, args ...interface{}) common.Sele
|
||||
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 {
|
||||
// Extract optional prefix from args
|
||||
// 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)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -5,7 +5,9 @@ import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -223,6 +225,11 @@ func (p *PgSQLAdapter) GetUnderlyingDB() interface{} {
|
||||
return p.db
|
||||
}
|
||||
|
||||
// SQLDB implements common.SQLDBProvider.
|
||||
func (p *PgSQLAdapter) SQLDB() *sql.DB {
|
||||
return p.db
|
||||
}
|
||||
|
||||
func (p *PgSQLAdapter) DriverName() string {
|
||||
return p.driverName
|
||||
}
|
||||
@@ -318,6 +325,32 @@ func (p *PgSQLSelectQuery) WhereOr(query string, args ...interface{}) common.Sel
|
||||
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 {
|
||||
query = p.replacePlaceholders(query, len(args))
|
||||
p.joins = append(p.joins, "JOIN "+query)
|
||||
@@ -762,6 +795,9 @@ func (p *PgSQLInsertQuery) Scan(ctx context.Context, dest interface{}) (err erro
|
||||
return nil
|
||||
}
|
||||
|
||||
// placeholderRe matches a numbered SQL parameter such as $12.
|
||||
var placeholderRe = regexp.MustCompile(`\$\d+`)
|
||||
|
||||
// PgSQLUpdateQuery implements UpdateQuery for PostgreSQL
|
||||
type PgSQLUpdateQuery struct {
|
||||
db *sql.DB
|
||||
@@ -892,23 +928,17 @@ func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err err
|
||||
p.tableName,
|
||||
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 {
|
||||
shift := len(setArgs)
|
||||
updatedWhereClauses := make([]string, 0, len(p.whereClauses))
|
||||
for _, whereClause := range p.whereClauses {
|
||||
// Find and replace parameter placeholders
|
||||
updatedClause := whereClause
|
||||
paramNum := i
|
||||
// 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
|
||||
updatedWhereClauses = append(updatedWhereClauses, placeholderRe.ReplaceAllStringFunc(whereClause, func(m string) string {
|
||||
n, _ := strconv.Atoi(m[1:])
|
||||
return fmt.Sprintf("$%d", n+shift)
|
||||
}))
|
||||
}
|
||||
p.whereClauses = updatedWhereClauses
|
||||
}
|
||||
|
||||
@@ -627,3 +627,23 @@ func TestRawSQL(t *testing.T) {
|
||||
|
||||
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
@@ -20,7 +20,15 @@ func DefaultCORSConfig() CORSConfig {
|
||||
configManager := config.GetConfigManager()
|
||||
cfg, _ := configManager.GetConfig()
|
||||
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()
|
||||
|
||||
@@ -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) {
|
||||
// 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")
|
||||
if origin == "" {
|
||||
origin = "*"
|
||||
// Not a cross-origin browser request; nothing to protect.
|
||||
w.SetHeader("Access-Control-Allow-Origin", "*")
|
||||
} else {
|
||||
// Vary must be set so caches don't serve one origin's response to another
|
||||
httpW := w.UnderlyingResponseWriter()
|
||||
httpW.Header().Set("Vary", "Origin")
|
||||
w.UnderlyingResponseWriter().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
|
||||
if len(config.AllowedMethods) > 0 {
|
||||
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")
|
||||
if 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))
|
||||
}
|
||||
|
||||
// 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 != "*" {
|
||||
w.SetHeader("Access-Control-Allow-Credentials", "true")
|
||||
}
|
||||
|
||||
// Expose headers that clients can read
|
||||
exposeHeaders := config.AllowedHeaders
|
||||
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, ", "))
|
||||
}
|
||||
|
||||
@@ -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() }
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package common
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
@@ -38,6 +39,15 @@ type Database interface {
|
||||
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)
|
||||
type SelectQuery interface {
|
||||
Model(model interface{}) SelectQuery
|
||||
@@ -309,3 +319,11 @@ type QueryHandler interface {
|
||||
SpecHandler
|
||||
// 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
|
||||
}
|
||||
|
||||
@@ -148,6 +148,93 @@ func validateWhereClauseSecurity(where string) error {
|
||||
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
|
||||
// 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)
|
||||
|
||||
// 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)
|
||||
return ""
|
||||
}
|
||||
|
||||
@@ -102,25 +102,25 @@ func TestSanitizeWhereClause(t *testing.T) {
|
||||
name: "dangerous DELETE keyword - blocked",
|
||||
where: "status = 'active'; DELETE FROM users",
|
||||
tableName: "users",
|
||||
expected: "",
|
||||
expected: "(1=0)", // fail closed,
|
||||
},
|
||||
{
|
||||
name: "dangerous UPDATE keyword - blocked",
|
||||
where: "1=1; UPDATE users SET admin = true",
|
||||
tableName: "users",
|
||||
expected: "",
|
||||
expected: "(1=0)", // fail closed,
|
||||
},
|
||||
{
|
||||
name: "dangerous TRUNCATE keyword - blocked",
|
||||
where: "status = 'active' OR TRUNCATE TABLE users",
|
||||
tableName: "users",
|
||||
expected: "",
|
||||
expected: "(1=0)", // fail closed,
|
||||
},
|
||||
{
|
||||
name: "dangerous DROP keyword - blocked",
|
||||
where: "status = 'active'; DROP TABLE users",
|
||||
tableName: "users",
|
||||
expected: "",
|
||||
expected: "(1=0)", // fail closed,
|
||||
},
|
||||
{
|
||||
name: "subquery with table alias should not be modified",
|
||||
|
||||
@@ -27,12 +27,13 @@ var allowedPoolHookTx = map[string]int{
|
||||
"resolvespec/handler.go": 1,
|
||||
"websocketspec/handler.go": 1,
|
||||
"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.
|
||||
var allowedPoolQuery = map[string]int{
|
||||
"resolvemcp/annotation.go": 2, // tool annotations, not a data request
|
||||
}
|
||||
var allowedPoolQuery = map[string]int{}
|
||||
|
||||
func guardedFiles(t *testing.T) map[string][]string {
|
||||
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
|
||||
// unwired hook silently disables whatever is registered on it (resolvespec's
|
||||
// AfterRead skipped column-level security masking until it was wired).
|
||||
var unwiredHooks = map[string]string{
|
||||
"websocketspec/BeforeDisconnect": "connection close is not hooked yet",
|
||||
"websocketspec/AfterDisconnect": "connection close is not hooked yet",
|
||||
}
|
||||
var unwiredHooks = map[string]string{}
|
||||
|
||||
var hookConstRE = regexp.MustCompile(`(?m)^\s*([A-Z][A-Za-z0-9]*)\s+HookType\s*=`)
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ package common
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
@@ -105,7 +106,7 @@ func (v *ColumnValidator) ValidateColumn(column string) error {
|
||||
}
|
||||
|
||||
// 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
|
||||
}
|
||||
|
||||
@@ -275,8 +276,14 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
||||
validSorts = append(validSorts, sort)
|
||||
} else {
|
||||
foundJoin := false
|
||||
strictSort := Hardening().SortStrict
|
||||
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
|
||||
break
|
||||
}
|
||||
@@ -287,7 +294,7 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
||||
}
|
||||
if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||
// 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)
|
||||
} else {
|
||||
logger.Warn("Unsafe sort expression '%s' removed", sort.Column)
|
||||
@@ -376,6 +383,56 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
||||
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
|
||||
// and doesn't contain SQL injection attempts or dangerous commands
|
||||
func IsSafeSortExpression(expr string) bool {
|
||||
|
||||
@@ -435,19 +435,19 @@ func TestFilterRequestOptions_WithSortExpressions(t *testing.T) {
|
||||
|
||||
options := RequestOptions{
|
||||
Sort: []SortOption{
|
||||
{Column: "id", Direction: "ASC"}, // Valid column
|
||||
{Column: "(SELECT MAX(age) FROM users)", Direction: "DESC"}, // Safe expression
|
||||
{Column: "name", Direction: "ASC"}, // Valid column
|
||||
{Column: "(id); DROP TABLE users; --", Direction: "DESC"}, // Dangerous expression
|
||||
{Column: "invalid_col", Direction: "ASC"}, // Invalid column
|
||||
{Column: "id", Direction: "ASC"}, // Valid column
|
||||
{Column: "(SELECT MAX(age) FROM users)", Direction: "DESC"}, // Safe expression
|
||||
{Column: "name", Direction: "ASC"}, // Valid column
|
||||
{Column: "(id); DROP TABLE users; --", Direction: "DESC"}, // Dangerous expression
|
||||
{Column: "invalid_col", Direction: "ASC"}, // Invalid column
|
||||
{Column: "(CASE WHEN age > 18 THEN 1 ELSE 0 END)", Direction: "ASC"}, // Safe expression
|
||||
},
|
||||
}
|
||||
|
||||
filtered := validator.FilterRequestOptions(options)
|
||||
|
||||
// Should keep: id, safe expression, name, another safe expression
|
||||
// Should remove: dangerous expression, invalid column
|
||||
// Keeps: id, subquery expression, name, CASE expression
|
||||
// Removes: dangerous expression, invalid column
|
||||
expectedCount := 4
|
||||
if len(filtered.Sort) != expectedCount {
|
||||
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
|
||||
// the relation field is the name used in x-preload headers.
|
||||
type PreloadParentModel struct {
|
||||
ID int64 `bun:"id,pk"`
|
||||
Name string `bun:"name"`
|
||||
ID int64 `bun:"id,pk"`
|
||||
Name string `bun:"name"`
|
||||
RELATED *RelatedModel `json:"RELATED" bun:"rel:has-one,join:id=related_id"`
|
||||
}
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ type Config struct {
|
||||
EventBroker EventBrokerConfig `mapstructure:"event_broker"`
|
||||
DBManager DBManagerConfig `mapstructure:"dbmanager"`
|
||||
DBTrace DBTraceConfig `mapstructure:"db_trace"`
|
||||
Hardening HardeningConfig `mapstructure:"hardening"`
|
||||
Paths PathsConfig `mapstructure:"paths"`
|
||||
Extensions map[string]interface{} `mapstructure:"extensions"`
|
||||
}
|
||||
@@ -143,6 +144,25 @@ type CORSConfig struct {
|
||||
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).
|
||||
// Env: RESOLVESPEC_DB_TRACE_ENABLED, _MIN_CALLS, _MIN_DURATION, _POOL_LOG.
|
||||
type DBTraceConfig struct {
|
||||
|
||||
@@ -168,6 +168,7 @@ func (m *Manager) SetConfig(cfg *Config) error {
|
||||
m.v.Set("event_broker", cfg.EventBroker)
|
||||
m.v.Set("dbmanager", cfg.DBManager)
|
||||
m.v.Set("db_trace", cfg.DBTrace)
|
||||
m.v.Set("hardening", cfg.Hardening)
|
||||
m.v.Set("paths", cfg.Paths)
|
||||
m.v.Set("extensions", cfg.Extensions)
|
||||
|
||||
@@ -279,6 +280,12 @@ func setDefaults(v *viper.Viper) {
|
||||
v.SetDefault("cors.allowed_headers", []string{"*"})
|
||||
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
|
||||
v.SetDefault("database.url", "")
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@ Wire: `dbtrace.Configure(dbtrace.FromConfig(cfg.DBTrace))` and wrap handlers wit
|
||||
## Log fields
|
||||
- `tx` transactions begun · `tx_queries` adapter queries inside `RunInTransaction` (share the tx 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`
|
||||
|
||||
## Pool log
|
||||
|
||||
+122
-134
@@ -1,46 +1,53 @@
|
||||
# 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
|
||||
|
||||
```go
|
||||
import (
|
||||
"github.com/bitechdev/ResolveSpec/pkg/resolvemcp"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
"github.com/gorilla/mux"
|
||||
)
|
||||
|
||||
// 1. Create a handler
|
||||
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", "orders", &Order{})
|
||||
|
||||
// 3. Mount routes
|
||||
r := mux.NewRouter()
|
||||
resolvemcp.SetupMuxRoutes(r, handler)
|
||||
resolvemcp.SetupMuxRoutes(r, handler, securityList) // guarded
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Config
|
||||
|
||||
```go
|
||||
type Config struct {
|
||||
// BaseURL is the public-facing base URL of the server (e.g. "http://localhost:8080").
|
||||
// Sent to MCP clients during the SSE handshake so they know where to POST messages.
|
||||
// If empty, it is detected from each incoming request using the Host header and
|
||||
// TLS state (X-Forwarded-Proto is honoured for reverse-proxy deployments).
|
||||
BaseURL string
|
||||
| Field | Default | Purpose |
|
||||
|---|---|---|
|
||||
| `BaseURL` | request-detected | Public base URL sent to SSE clients |
|
||||
| `BasePath` | request-detected | Mount path (e.g. `/mcp`) |
|
||||
| `DefaultLimit` | 50 | Page size when a read gives no limit |
|
||||
| `MaxLimit` | 1000 | Larger limits are clamped |
|
||||
| `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
|
||||
|
||||
@@ -63,14 +70,36 @@ handler.RegisterModel(schema, entity string, model interface{}) error
|
||||
- `entity` — table/entity name (e.g. `"users"`).
|
||||
- `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
|
||||
|
||||
`Config.BasePath` is required and used for all route registration.
|
||||
`Config.BaseURL` is optional — when empty it is detected from each request.
|
||||
`Config.BasePath` is used for route registration. `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).
|
||||
|
||||
@@ -83,7 +112,7 @@ Two endpoints: `GET {BasePath}/sse` (subscribe) + `POST {BasePath}/message` (sen
|
||||
#### Gorilla Mux
|
||||
|
||||
```go
|
||||
resolvemcp.SetupMuxRoutes(r, handler)
|
||||
resolvemcp.SetupMuxRoutes(r, handler, securityList)
|
||||
```
|
||||
|
||||
| Route | Method | Description |
|
||||
@@ -94,13 +123,13 @@ resolvemcp.SetupMuxRoutes(r, handler)
|
||||
#### bunrouter
|
||||
|
||||
```go
|
||||
resolvemcp.SetupBunRouterRoutes(router, handler)
|
||||
resolvemcp.SetupBunRouterRoutes(router, handler, securityList)
|
||||
```
|
||||
|
||||
#### Gin / net/http / Echo
|
||||
|
||||
```go
|
||||
sse := handler.SSEServer()
|
||||
sse := resolvemcp.NewSSEServer(handler, securityList) // guarded
|
||||
|
||||
engine.Any("/mcp/*path", gin.WrapH(sse)) // Gin
|
||||
http.Handle("/mcp/", sse) // net/http
|
||||
@@ -116,7 +145,7 @@ Single endpoint at `{BasePath}`. Handles POST (client→server) and GET (server
|
||||
#### Gorilla Mux
|
||||
|
||||
```go
|
||||
resolvemcp.SetupMuxStreamableHTTPRoutes(r, handler)
|
||||
resolvemcp.SetupMuxStreamableHTTPRoutes(r, handler, securityList)
|
||||
```
|
||||
|
||||
Mounts the handler at `{BasePath}` (all methods).
|
||||
@@ -124,7 +153,7 @@ Mounts the handler at `{BasePath}` (all methods).
|
||||
#### bunrouter
|
||||
|
||||
```go
|
||||
resolvemcp.SetupBunRouterStreamableHTTPRoutes(router, handler)
|
||||
resolvemcp.SetupBunRouterStreamableHTTPRoutes(router, handler, securityList)
|
||||
```
|
||||
|
||||
Registers GET, POST, DELETE on `{BasePath}`.
|
||||
@@ -132,8 +161,7 @@ Registers GET, POST, DELETE on `{BasePath}`.
|
||||
#### Gin / net/http / Echo
|
||||
|
||||
```go
|
||||
h := handler.StreamableHTTPServer()
|
||||
// or: h := resolvemcp.NewStreamableHTTPHandler(handler)
|
||||
h := resolvemcp.NewStreamableHTTPHandler(handler, securityList) // guarded
|
||||
|
||||
engine.Any("/mcp", gin.WrapH(h)) // Gin
|
||||
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.)
|
||||
- **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
|
||||
|
||||
| Path | Spec | Purpose |
|
||||
@@ -187,7 +217,7 @@ handler.EnableOAuthServer(security.OAuthServerConfig{
|
||||
|
||||
provider, _ := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
|
||||
securityList, _ := security.NewSecurityList(provider)
|
||||
security.RegisterSecurityHooks(handler, securityList)
|
||||
resolvemcp.RegisterSecurityHooks(handler, securityList)
|
||||
|
||||
http.ListenAndServe(":8080", handler.HTTPHandler(securityList))
|
||||
```
|
||||
@@ -286,7 +316,10 @@ resolvemcp.SetupMuxRoutesWithAuth(r, handler, securityList)
|
||||
```go
|
||||
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)
|
||||
```
|
||||
|
||||
@@ -294,12 +327,17 @@ Call `RegisterSecurityHooks` **once**, after creating the handler and before reg
|
||||
|
||||
| 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 |
|
||||
| `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 |
|
||||
| `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 |
|
||||
|
||||
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
|
||||
|
||||
Use `RegisterModelWithRules` instead of `RegisterModel` to set access rules at registration time:
|
||||
@@ -356,121 +394,58 @@ handler.SetModelRules("public", "users", modelregistry.ModelRules{
|
||||
|
||||
## 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).
|
||||
|
||||
```
|
||||
{operation}_{schema}_{entity} // e.g. read_public_users
|
||||
{operation}_{entity} // e.g. read_users (when schema is empty)
|
||||
```
|
||||
| Tool | Purpose |
|
||||
|---|---|
|
||||
| `list_tables` | Tables the caller may use and the allowed operations |
|
||||
| `describe_table` | Columns, PK, relations, writable columns, operations, limits |
|
||||
| `select_table` | Read rows (filters, sort, columns, preloads, paging) |
|
||||
| `insert_into_table` | Insert one row or a capped batch |
|
||||
| `update_table` | Update by `id` or `filters` |
|
||||
| `delete_from_table` | Delete by `id` or `filters` |
|
||||
| `list_functions` | Registered functions the caller may call, with parameters |
|
||||
| `call_function` | Call a registered function |
|
||||
| `resolvespec_annotate` | Only with `EnableAnnotations` |
|
||||
|
||||
Operations: `read`, `create`, `update`, `delete`.
|
||||
|
||||
### Read Tool — `read_{schema}_{entity}`
|
||||
|
||||
Fetch one or many records.
|
||||
### `select_table`
|
||||
|
||||
| Argument | Type | Description |
|
||||
|---|---|---|
|
||||
| `id` | string | Primary key value. Omit to return multiple records. |
|
||||
| `limit` | number | Max records per page (recommended: 10–100). |
|
||||
| `offset` | number | Records to skip (offset-based pagination). |
|
||||
| `cursor_forward` | string | PK of the **last** record on the current page (next-page cursor). |
|
||||
| `cursor_backward` | string | PK of the **first** record on the current page (prev-page cursor). |
|
||||
| `columns` | array | Column names to include. Omit for all columns. |
|
||||
| `omit_columns` | array | Column names to exclude. |
|
||||
| `filters` | array | Filter objects (see [Filtering](#filtering)). |
|
||||
| `sort` | array | Sort objects (see [Sorting](#sorting)). |
|
||||
| `preloads` | array | Relation preload objects (see [Preloading](#preloading)). |
|
||||
| `table` | string (required) | `schema.entity` |
|
||||
| `id` | string | Primary key of one row |
|
||||
| `filters`, `sort` | array | See [Filtering](#filtering), [Sorting](#sorting) |
|
||||
| `columns`, `omit_columns` | array | Column selection |
|
||||
| `preloads` | array | Relations (validated against the model, max depth `MaxPreloadDepth`) |
|
||||
| `limit`, `offset` | number | Clamped to `MaxLimit` / rejected above `MaxOffset` |
|
||||
| `cursor_forward`, `cursor_backward` | string | PK cursor, requires `sort` |
|
||||
| `include_count` | boolean | Also compute totals (slower); otherwise `total`/`filtered` are 0 |
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"data": [...],
|
||||
"metadata": {
|
||||
"total": 100,
|
||||
"filtered": 100,
|
||||
"count": 10,
|
||||
"limit": 10,
|
||||
"offset": 0
|
||||
}
|
||||
}
|
||||
```
|
||||
Response: `{"success":true,"data":[...],"metadata":{"total","filtered","count","limit","offset"}}`
|
||||
|
||||
### 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 |
|
||||
|---|---|---|
|
||||
| `data` | object \| array | Single object or array of objects to insert. |
|
||||
### `update_table` / `delete_from_table`
|
||||
|
||||
Array input runs inside a single transaction — all succeed or all fail.
|
||||
Either `id` or `filters` is required.
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{ "success": true, "data": { ... } }
|
||||
```
|
||||
| Mode | Behaviour |
|
||||
|---|---|
|
||||
| `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 |
|
||||
|---|---|---|
|
||||
| `id` | string | Primary key of the record. Can also be included inside `data`. |
|
||||
| `data` | object (required) | Fields to update. |
|
||||
`name` and `arguments` (object). See [Functions](#functions).
|
||||
|
||||
**Response:**
|
||||
```json
|
||||
{ "success": true, "data": { ...merged record... } }
|
||||
```
|
||||
### `resolvespec_annotate`
|
||||
|
||||
### Delete Tool — `delete_{schema}_{entity}`
|
||||
|
||||
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`.
|
||||
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.
|
||||
|
||||
---
|
||||
|
||||
@@ -570,6 +545,9 @@ Hooks let you intercept and modify CRUD operations at well-defined lifecycle poi
|
||||
| `BeforeCreate` / `AfterCreate` | Around insert |
|
||||
| `BeforeUpdate` / `AfterUpdate` | Around update |
|
||||
| `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
|
||||
|
||||
@@ -599,7 +577,7 @@ handler.Hooks().RegisterMultiple(
|
||||
| `Entity` | `string` | Entity/table name |
|
||||
| `Model` | `interface{}` | Registered model instance |
|
||||
| `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) |
|
||||
| `Data` | `interface{}` | Input data (create/update — modifiable) |
|
||||
| `Result` | `interface{}` | Output data (set by After hooks) |
|
||||
@@ -633,7 +611,7 @@ registry.ClearAll() // remove all hooks
|
||||
|
||||
## 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
|
||||
schema := resolvemcp.GetSchema(ctx)
|
||||
@@ -653,7 +631,7 @@ ctx = resolvemcp.WithSchema(ctx, "tenant_a")
|
||||
|
||||
## 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
|
||||
mcpServer := handler.MCPServer()
|
||||
@@ -669,3 +647,13 @@ The handler resolves table names in priority order:
|
||||
1. `TableNameProvider` interface — `TableName() string` (can return `"schema.table"`)
|
||||
2. `SchemaProvider` interface — `SchemaName() string` (combined with entity name)
|
||||
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.
|
||||
|
||||
@@ -6,6 +6,8 @@ import (
|
||||
"fmt"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
)
|
||||
|
||||
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) {
|
||||
hookCtx, err := annotationGate(ctx, h, "annotate_set", toolName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
jsonBytes, err := json.Marshal(annotations)
|
||||
if 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 {
|
||||
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) {
|
||||
hookCtx, err := annotationGate(ctx, h, "annotate_get", toolName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
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 {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get annotation: %v", err)), nil
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -1,6 +1,11 @@
|
||||
package resolvemcp
|
||||
|
||||
import "context"
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
)
|
||||
|
||||
type contextKey string
|
||||
|
||||
@@ -69,3 +74,18 @@ func withRequestData(ctx context.Context, schema, entity, tableName string, mode
|
||||
ctx = WithModelPtr(ctx, modelPtr)
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
@@ -4,9 +4,12 @@ import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"runtime/debug"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
@@ -30,6 +33,8 @@ type Handler struct {
|
||||
version string
|
||||
oauth2Regs []oauth2Registration
|
||||
oauthSrv *security.OAuthServer
|
||||
functions functionRegistry
|
||||
confirms *confirmStore
|
||||
}
|
||||
|
||||
// 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,
|
||||
hooks: NewHookRegistry(),
|
||||
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"),
|
||||
config: cfg,
|
||||
config: cfg.withDefaults(),
|
||||
confirms: newConfirmStore(),
|
||||
name: "resolvemcp",
|
||||
version: "1.0.0",
|
||||
}
|
||||
registerAnnotationTool(h)
|
||||
registerMetaTools(h)
|
||||
if cfg.EnableAnnotations {
|
||||
registerAnnotationTool(h)
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
@@ -97,7 +106,20 @@ type dynamicSSEHandler struct {
|
||||
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) {
|
||||
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)
|
||||
|
||||
d.mu.Lock()
|
||||
@@ -106,6 +128,12 @@ func (d *dynamicSSEHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
s, ok := d.pool[baseURL]
|
||||
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)
|
||||
d.pool[baseURL] = s
|
||||
}
|
||||
@@ -114,6 +142,20 @@ func (d *dynamicSSEHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
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.
|
||||
// It honours the X-Forwarded-Proto header for deployments behind a proxy.
|
||||
func requestBaseURL(r *http.Request) string {
|
||||
@@ -127,13 +169,13 @@ func requestBaseURL(r *http.Request) string {
|
||||
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 {
|
||||
fullName := buildModelName(schema, entity)
|
||||
if err := h.registry.RegisterModel(fullName, model); err != nil {
|
||||
return err
|
||||
}
|
||||
registerModelTools(h, schema, entity, model)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -149,7 +191,6 @@ func (h *Handler) RegisterModelWithRules(schema, entity string, model interface{
|
||||
if err := reg.RegisterModelWithRules(fullName, model, rules); err != nil {
|
||||
return err
|
||||
}
|
||||
registerModelTools(h, schema, entity, model)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -200,22 +241,36 @@ func (h *Handler) getSchemaAndTable(defaultSchema, entity string, model interfac
|
||||
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.
|
||||
// Usage: defer recoverPanic(&returnedErr)
|
||||
func recoverPanic(err *error) {
|
||||
if r := recover(); r != nil {
|
||||
msg := fmt.Sprintf("%v", r)
|
||||
logger.Error("[resolvemcp] panic recovered: %s", msg)
|
||||
*err = fmt.Errorf("internal error: %s", msg)
|
||||
logger.Error("[resolvemcp] panic recovered: %v\n%s", r, debug.Stack())
|
||||
*err = errInternal
|
||||
}
|
||||
}
|
||||
|
||||
// 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)
|
||||
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)
|
||||
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)
|
||||
@@ -226,7 +281,7 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
|
||||
model = unwrapped.Model
|
||||
modelType := unwrapped.ModelType
|
||||
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)
|
||||
options = validator.FilterRequestOptions(options)
|
||||
@@ -253,7 +308,7 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
|
||||
var metadata *common.Metadata
|
||||
err = h.runInTx(ctx, hookCtx, func(common.Database) 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
|
||||
})
|
||||
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.
|
||||
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))
|
||||
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
|
||||
cursorFilter, err := getCursorFilter(tableName, pkName, modelColumns, options, nil)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("cursor error: %w", err)
|
||||
return nil, nil, invalidArg("invalid cursor")
|
||||
}
|
||||
|
||||
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.
|
||||
total, err := query.Count(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("error counting records: %w", err)
|
||||
total := 0
|
||||
if count {
|
||||
var err error
|
||||
total, err = query.Count(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("error counting records: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 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.
|
||||
if len(options.Preload) > 0 {
|
||||
if err := h.validatePreloads(model, options.Preload); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
var preloadErr error
|
||||
query, preloadErr = h.applyPreloads(model, query, options.Preload)
|
||||
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.
|
||||
if err := query.ScanModel(ctx); err != nil {
|
||||
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)
|
||||
}
|
||||
@@ -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.
|
||||
scannedResults := reflect.ValueOf(modelPtr).Elem()
|
||||
if scannedResults.Len() == 0 {
|
||||
return nil, nil, fmt.Errorf("record not found")
|
||||
return nil, nil, errRecordNotFound
|
||||
}
|
||||
data = scannedResults.Index(0).Interface()
|
||||
} else {
|
||||
@@ -421,9 +483,14 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
|
||||
// executeCreate inserts one or more records.
|
||||
func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data interface{}) (_ interface{}, retErr error) {
|
||||
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)
|
||||
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)
|
||||
@@ -433,7 +500,7 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
|
||||
|
||||
model = result.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{
|
||||
Context: ctx,
|
||||
@@ -455,11 +522,11 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
|
||||
modelType = modelType.Elem()
|
||||
}
|
||||
|
||||
// Transaction 1: BeforeCreate + inserts.
|
||||
var (
|
||||
single bool
|
||||
originals []map[string]interface{}
|
||||
insertedIDs []interface{}
|
||||
results []interface{}
|
||||
)
|
||||
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
||||
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 {
|
||||
itemMap, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
return fmt.Errorf("each item must be an object")
|
||||
return invalidArg("each item must be an object")
|
||||
}
|
||||
originals = append(originals, itemMap)
|
||||
}
|
||||
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))
|
||||
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)
|
||||
for key, value := range itemMap {
|
||||
for key, value := range cols {
|
||||
q = q.Value(key, value)
|
||||
}
|
||||
if pkName == "" {
|
||||
@@ -502,22 +576,11 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
|
||||
}
|
||||
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.
|
||||
results := make([]interface{}, 0, len(insertedIDs))
|
||||
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
||||
results = results[:0]
|
||||
// Re-fetch inside the same transaction to capture DB-generated defaults/triggers, then
|
||||
// AfterCreate: the write is only committed when the whole sequence succeeds, so a
|
||||
// failure here cannot leave a committed insert behind an error the client may retry.
|
||||
results = make([]interface{}, 0, len(insertedIDs))
|
||||
for i, pkVal := range insertedIDs {
|
||||
if pkVal == nil {
|
||||
results = append(results, originals[i])
|
||||
@@ -544,6 +607,12 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
|
||||
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
|
||||
}
|
||||
if single {
|
||||
@@ -555,9 +624,11 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
|
||||
// executeUpdate updates a record by ID.
|
||||
func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, data interface{}) (_ interface{}, retErr error) {
|
||||
defer recoverPanic(&retErr)
|
||||
ctx, cancel := h.callContext(ctx)
|
||||
defer cancel()
|
||||
model, err := h.registry.GetModelByEntity(schema, entity)
|
||||
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)
|
||||
@@ -567,11 +638,11 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
|
||||
|
||||
model = result.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{})
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("data must be an object")
|
||||
return nil, invalidArg("data must be an object")
|
||||
}
|
||||
|
||||
if id == "" {
|
||||
@@ -580,7 +651,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
|
||||
}
|
||||
}
|
||||
if id == "" {
|
||||
return nil, fmt.Errorf("update requires an ID")
|
||||
return nil, invalidArg("update requires an id")
|
||||
}
|
||||
|
||||
pkName := reflection.GetPrimaryKeyName(model)
|
||||
@@ -602,23 +673,47 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
|
||||
|
||||
var updateResult interface{}
|
||||
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)
|
||||
if modelType.Kind() == reflect.Pointer {
|
||||
modelType = modelType.Elem()
|
||||
}
|
||||
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)
|
||||
|
||||
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 {
|
||||
return fmt.Errorf("no records found to update")
|
||||
return errRecordNotFound
|
||||
}
|
||||
return fmt.Errorf("error fetching existing record: %w", err)
|
||||
}
|
||||
|
||||
// Convert to map
|
||||
existingMap := make(map[string]interface{})
|
||||
jsonData, err := json.Marshal(existingRecord)
|
||||
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 {
|
||||
return fmt.Errorf("error unmarshaling existing record: %w", err)
|
||||
}
|
||||
|
||||
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
|
||||
return err
|
||||
}
|
||||
if modifiedData, ok := hookCtx.Data.(map[string]interface{}); ok {
|
||||
updates = modifiedData
|
||||
for key, v := range updates {
|
||||
existingMap[key] = v
|
||||
}
|
||||
|
||||
// Merge non-nil, non-empty values
|
||||
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).
|
||||
q := tx.NewUpdate().Table(tableName).SetMap(setCols).
|
||||
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
|
||||
res, err := q.Exec(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("error updating record: %w", err)
|
||||
}
|
||||
if res.RowsAffected() == 0 {
|
||||
return fmt.Errorf("no records found to update")
|
||||
return errRecordNotFound
|
||||
}
|
||||
|
||||
updateResult = existingMap
|
||||
hookCtx.Result = updateResult
|
||||
return h.hooks.Execute(AfterUpdate, hookCtx)
|
||||
})
|
||||
hookCtx.Result = existingMap
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 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 {
|
||||
// Re-fetch inside the same transaction to capture DB-generated changes, then
|
||||
// AfterUpdate; see executeCreate.
|
||||
fetchedRecord := reflect.New(modelType).Interface()
|
||||
if err := tx.NewSelect().Model(fetchedRecord).
|
||||
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id).
|
||||
ScanModel(ctx); err == nil {
|
||||
jsonData, marshalErr := json.Marshal(fetchedRecord)
|
||||
if marshalErr == nil {
|
||||
if jsonData, marshalErr := json.Marshal(fetchedRecord); marshalErr == nil {
|
||||
var fetchedMap map[string]interface{}
|
||||
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 {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return updateResult, nil
|
||||
}
|
||||
|
||||
// executeDelete deletes a record by ID.
|
||||
func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string) (_ interface{}, retErr error) {
|
||||
defer recoverPanic(&retErr)
|
||||
ctx, cancel := h.callContext(ctx)
|
||||
defer cancel()
|
||||
if id == "" {
|
||||
return nil, fmt.Errorf("delete requires an ID")
|
||||
return nil, invalidArg("delete requires an id")
|
||||
}
|
||||
|
||||
model, err := h.registry.GetModelByEntity(schema, entity)
|
||||
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)
|
||||
@@ -711,7 +782,7 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
|
||||
|
||||
model = result.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)
|
||||
|
||||
@@ -741,11 +812,14 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
|
||||
return err
|
||||
}
|
||||
record := reflect.New(modelType).Interface()
|
||||
selectQuery := tx.NewSelect().Model(record).
|
||||
hookCtx.Query = tx.NewSelect().Model(record).
|
||||
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 {
|
||||
return fmt.Errorf("record not found")
|
||||
return errRecordNotFound
|
||||
}
|
||||
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)
|
||||
}
|
||||
if res.RowsAffected() == 0 {
|
||||
return fmt.Errorf("record not found or already deleted")
|
||||
return errRecordNotFound
|
||||
}
|
||||
|
||||
recordToDelete = record
|
||||
@@ -902,3 +976,56 @@ func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(t
|
||||
return h.hooks.Execute(OnTxBegin, hookCtx)
|
||||
}, 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
|
||||
}
|
||||
|
||||
@@ -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
@@ -3,6 +3,8 @@ package resolvemcp
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"runtime/debug"
|
||||
"sync"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
@@ -21,12 +23,23 @@ const (
|
||||
BeforeCreate HookType = "before_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"
|
||||
AfterUpdate HookType = "after_update"
|
||||
|
||||
BeforeDelete HookType = "before_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
|
||||
// (including the second short transaction for post-commit work). hookCtx.Tx is
|
||||
// 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
|
||||
type HookRegistry struct {
|
||||
mu sync.RWMutex
|
||||
hooks map[HookType][]HookFunc
|
||||
}
|
||||
|
||||
@@ -72,11 +86,14 @@ func NewHookRegistry() *HookRegistry {
|
||||
}
|
||||
|
||||
func (r *HookRegistry) Register(hookType HookType, hook HookFunc) {
|
||||
r.mu.Lock()
|
||||
if r.hooks == nil {
|
||||
r.hooks = make(map[HookType][]HookFunc)
|
||||
}
|
||||
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) {
|
||||
@@ -86,37 +103,57 @@ func (r *HookRegistry) RegisterMultiple(hookTypes []HookType, hook HookFunc) {
|
||||
}
|
||||
|
||||
func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
||||
hooks, exists := r.hooks[hookType]
|
||||
if !exists || len(hooks) == 0 {
|
||||
// Append-only slices: a snapshot of the slice header is safe to iterate without the lock.
|
||||
r.mu.RLock()
|
||||
hooks := r.hooks[hookType]
|
||||
r.mu.RUnlock()
|
||||
if len(hooks) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
logger.Debug("Executing %d resolvemcp hook(s) for %s", len(hooks), hookType)
|
||||
|
||||
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)
|
||||
return fmt.Errorf("hook execution failed: %w", err)
|
||||
}
|
||||
|
||||
if ctx.Abort {
|
||||
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
|
||||
}
|
||||
|
||||
// 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) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
delete(r.hooks, hookType)
|
||||
}
|
||||
|
||||
func (r *HookRegistry) ClearAll() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.hooks = make(map[HookType][]HookFunc)
|
||||
}
|
||||
|
||||
func (r *HookRegistry) HasHooks(hookType HookType) bool {
|
||||
hooks, exists := r.hooks[hookType]
|
||||
return exists && len(hooks) > 0
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return len(r.hooks[hookType]) > 0
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -114,25 +114,12 @@ func (h *Handler) mountOAuth2Routes(mux *http.ServeMux) {
|
||||
// context into the request context, making it available to BeforeHandle security hooks.
|
||||
// Unauthenticated requests receive 401 before reaching any MCP tool.
|
||||
func (h *Handler) AuthedSSEServer(securityList *security.SecurityList) http.Handler {
|
||||
return security.NewAuthMiddleware(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())
|
||||
return Guard(securityList)(h.SSEServer())
|
||||
}
|
||||
|
||||
// AuthedStreamableHTTPServer wraps StreamableHTTPServer with required authentication middleware.
|
||||
func (h *Handler) AuthedStreamableHTTPServer(securityList *security.SecurityList) http.Handler {
|
||||
return security.NewAuthMiddleware(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())
|
||||
return Guard(securityList)(h.StreamableHTTPServer())
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
@@ -243,22 +230,3 @@ func SetupMuxOAuth2Routes(muxRouter *mux.Router, auth *security.DatabaseAuthenti
|
||||
OAuth2CallbackHandler(auth, cfg.ProviderName, cfg.AfterLoginRedirect, cookieOpts...),
|
||||
).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
@@ -12,12 +12,13 @@
|
||||
// handler.RegisterModel("public", "users", &User{})
|
||||
//
|
||||
// r := mux.NewRouter()
|
||||
// resolvemcp.SetupMuxRoutes(r, handler)
|
||||
// resolvemcp.SetupMuxRoutes(r, handler, securityList) // requires an authenticated caller
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"runtime/debug"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/uptrace/bun"
|
||||
@@ -28,6 +29,7 @@ import (
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
)
|
||||
|
||||
// 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").
|
||||
// If empty, the path is detected from each incoming request automatically.
|
||||
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.
|
||||
@@ -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
|
||||
// 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)
|
||||
// - POST {basePath}/message — JSON-RPC message endpoint (client sends requests here)
|
||||
//
|
||||
// To protect these routes with authentication, wrap the mux router or apply middleware
|
||||
// before calling SetupMuxRoutes.
|
||||
func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler) {
|
||||
basePath := handler.config.BasePath
|
||||
h := handler.SSEServer()
|
||||
// Nothing is mounted (and an error is logged) when securityList has no provider.
|
||||
func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler, securityList *security.SecurityList) {
|
||||
if !requireGuard("SetupMuxRoutes", securityList) {
|
||||
return
|
||||
}
|
||||
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+"/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
|
||||
// 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
|
||||
// - 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() {
|
||||
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
|
||||
h := handler.SSEServer()
|
||||
|
||||
router.GET(basePath+"/sse", bunrouter.HTTPHandler(h))
|
||||
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)
|
||||
}
|
||||
|
||||
// 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
|
||||
// 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)
|
||||
func NewSSEServer(handler *Handler) http.Handler {
|
||||
return handler.SSEServer()
|
||||
func NewSSEServer(handler *Handler, securityList *security.SecurityList) http.Handler {
|
||||
return handler.AuthedSSEServer(securityList)
|
||||
}
|
||||
|
||||
// SetupMuxStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on the given Gorilla Mux router.
|
||||
// The streamable HTTP transport uses a single endpoint (Config.BasePath) for all communication:
|
||||
// POST for client→server messages, GET for server→client streaming.
|
||||
// SetupMuxStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on the given Gorilla Mux
|
||||
// router, behind Guard(securityList). The streamable HTTP transport uses a single endpoint
|
||||
// (Config.BasePath) for all communication: POST for client→server messages, GET for
|
||||
// server→client streaming.
|
||||
//
|
||||
// Example:
|
||||
//
|
||||
// resolvemcp.SetupMuxStreamableHTTPRoutes(r, handler) // mounts at Config.BasePath
|
||||
func SetupMuxStreamableHTTPRoutes(muxRouter *mux.Router, handler *Handler) {
|
||||
// Nothing is mounted (and an error is logged) when securityList has no provider.
|
||||
func SetupMuxStreamableHTTPRoutes(muxRouter *mux.Router, handler *Handler, securityList *security.SecurityList) {
|
||||
if !requireGuard("SetupMuxStreamableHTTPRoutes", securityList) {
|
||||
return
|
||||
}
|
||||
basePath := handler.config.BasePath
|
||||
h := handler.StreamableHTTPServer()
|
||||
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, h))
|
||||
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, handler.AuthedStreamableHTTPServer(securityList)))
|
||||
}
|
||||
|
||||
// SetupBunRouterStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on a bunrouter router.
|
||||
// The streamable HTTP transport uses a single endpoint (Config.BasePath).
|
||||
func SetupBunRouterStreamableHTTPRoutes(router *bunrouter.Router, handler *Handler) {
|
||||
// SetupMuxStreamableHTTPRoutesUnauthenticated is SetupMuxStreamableHTTPRoutes without the guard.
|
||||
// A warning is logged.
|
||||
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
|
||||
h := handler.StreamableHTTPServer()
|
||||
router.GET(basePath, bunrouter.HTTPHandler(h))
|
||||
router.POST(basePath, bunrouter.HTTPHandler(h))
|
||||
router.DELETE(basePath, bunrouter.HTTPHandler(h))
|
||||
}
|
||||
|
||||
// NewStreamableHTTPHandler returns an http.Handler that serves MCP over the streamable HTTP transport.
|
||||
// Mount it at the desired path; that path becomes the MCP endpoint.
|
||||
// NewStreamableHTTPHandler returns an http.Handler that serves MCP over the streamable HTTP
|
||||
// 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)
|
||||
// engine.Any("/mcp", gin.WrapH(h))
|
||||
func NewStreamableHTTPHandler(handler *Handler) http.Handler {
|
||||
return handler.StreamableHTTPServer()
|
||||
func NewStreamableHTTPHandler(handler *Handler, securityList *security.SecurityList) http.Handler {
|
||||
return handler.AuthedStreamableHTTPServer(securityList)
|
||||
}
|
||||
|
||||
@@ -31,7 +31,7 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
||||
hookCtx.Abort = true
|
||||
hookCtx.AbortMessage = err.Error()
|
||||
hookCtx.AbortCode = http.StatusUnauthorized
|
||||
return err
|
||||
return NewClientError(CodeForbidden, err.Error())
|
||||
}
|
||||
return nil
|
||||
})
|
||||
@@ -62,6 +62,15 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.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.
|
||||
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||
return security.ApplyColumnSecurity(newSecurityContext(hookCtx), securityList)
|
||||
@@ -72,14 +81,19 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
||||
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.
|
||||
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
|
||||
return security.CheckModelUpdateAllowed(newSecurityContext(hookCtx))
|
||||
return forbidden(security.CheckModelUpdateAllowed(newSecurityContext(hookCtx)))
|
||||
})
|
||||
|
||||
// BeforeDelete: enforce CanDelete rule.
|
||||
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")
|
||||
@@ -153,3 +167,11 @@ func (s *securityContext) GetResult() interface{} {
|
||||
func (s *securityContext) SetResult(result interface{}) {
|
||||
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())
|
||||
}
|
||||
|
||||
@@ -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
@@ -1,7 +1,6 @@
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"reflect"
|
||||
@@ -10,34 +9,9 @@ import (
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"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.
|
||||
type modelInfo struct {
|
||||
fullName string // e.g. "public.users"
|
||||
@@ -227,350 +201,8 @@ func buildSchemaDoc(info modelInfo) string {
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
// columnNameList returns a comma-separated list of JSON column names (for descriptions).
|
||||
func columnNameList(cols []columnInfo) string {
|
||||
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.
|
||||
// parseRequestOptions reads the paging, filter, sort, column and preload arguments shared
|
||||
// by the read tools.
|
||||
func parseRequestOptions(args map[string]interface{}) common.RequestOptions {
|
||||
options := common.RequestOptions{}
|
||||
|
||||
|
||||
+35
-28
@@ -128,14 +128,12 @@ func TestReadRunsInOneTransaction(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateSingleUsesTwoTransactions(t *testing.T) {
|
||||
func TestCreateSingleRunsInOneTransaction(t *testing.T) {
|
||||
h, mock, ctx := newTxHarness(t)
|
||||
tr := traceHooks(h, OnTxBegin, BeforeCreate, AfterCreate)
|
||||
|
||||
mock.ExpectBegin()
|
||||
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.ExpectCommit()
|
||||
|
||||
@@ -145,24 +143,19 @@ func TestCreateSingleUsesTwoTransactions(t *testing.T) {
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tr.assertOrder(t, "on_tx_begin", "before_create", "on_tx_begin", "after_create")
|
||||
if tr.txs["before_create"][0] != tr.txs["on_tx_begin"][0] {
|
||||
t.Fatal("BeforeCreate must run on the first 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")
|
||||
tr.assertOrder(t, "on_tx_begin", "before_create", "after_create")
|
||||
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, the re-fetch and AfterCreate must share one transaction")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateBatchRefetchOnSecondTransaction(t *testing.T) {
|
||||
func TestCreateBatchRefetchInSameTransaction(t *testing.T) {
|
||||
h, mock, ctx := newTxHarness(t)
|
||||
tr := traceHooks(h, OnTxBegin, AfterCreate)
|
||||
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
||||
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(2, "b"))
|
||||
mock.ExpectCommit()
|
||||
@@ -174,10 +167,10 @@ func TestCreateBatchRefetchOnSecondTransaction(t *testing.T) {
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
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)
|
||||
tr := traceHooks(h, OnTxBegin, BeforeUpdate, AfterUpdate)
|
||||
|
||||
@@ -185,25 +178,38 @@ func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) {
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a"))
|
||||
mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
mock.ExpectCommit()
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b"))
|
||||
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)
|
||||
}
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tr.assertOrder(t, "on_tx_begin", "before_update", "after_update", "on_tx_begin")
|
||||
if tr.txs["on_tx_begin"][0] == tr.txs["on_tx_begin"][1] {
|
||||
t.Fatal("re-fetch must run on a second transaction")
|
||||
if m, _ := res.(map[string]interface{}); m["name"] != "b" {
|
||||
t.Fatalf("result must be the re-fetched row, got %v", res)
|
||||
}
|
||||
for _, ht := range []string{"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)
|
||||
}
|
||||
tr.assertOrder(t, "on_tx_begin", "before_update", "after_update")
|
||||
}
|
||||
|
||||
// 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
|
||||
}
|
||||
|
||||
func (stubProvider) GetRowSecurity(context.Context, any, string, string) (security.RowSecurity, error) {
|
||||
return security.RowSecurity{}, nil
|
||||
}
|
||||
|
||||
func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) {
|
||||
h, mock, ctx := newTxHarness(t)
|
||||
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.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"}
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0))
|
||||
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a"))
|
||||
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.ExpectCommit()
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
+73
-58
@@ -646,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
|
||||
}
|
||||
|
||||
// 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]
|
||||
// Client-controlled conditions (filters + x-custom-sql-w) are built by a closure so they can
|
||||
// be wrapped in a single group together with x-custom-sql-or: the OR then only widens the
|
||||
// client's own conditions and can never escape the server-side filters ANDed around it.
|
||||
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
|
||||
castInfo := h.ValidateAndAdjustFilterForColumnType(filter, model)
|
||||
// Validate and adjust filter based on column type
|
||||
castInfo := h.ValidateAndAdjustFilterForColumnType(filter, model)
|
||||
|
||||
// Default to AND if LogicOperator is not set
|
||||
logicOp := filter.LogicOperator
|
||||
if logicOp == "" {
|
||||
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
|
||||
}
|
||||
// Default to AND if LogicOperator is not set
|
||||
logicOp := filter.LogicOperator
|
||||
if logicOp == "" {
|
||||
logicOp = "AND"
|
||||
}
|
||||
|
||||
// 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++
|
||||
// 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
|
||||
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)
|
||||
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)
|
||||
sanitizedOr := ""
|
||||
if options.CustomSQLOr != "" {
|
||||
logger.Debug("Applying custom SQL OR: %s", options.CustomSQLOr)
|
||||
customOr := common.AddTablePrefixToColumns(options.CustomSQLOr, reflection.ExtractTableNameOnly(tableName))
|
||||
// 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
|
||||
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 != "" {
|
||||
query = query.WhereOr(sanitizedOr)
|
||||
}
|
||||
|
||||
+18
-23
@@ -19,7 +19,7 @@ In-memory store seeded from a static list. Suitable for a small, fixed set of se
|
||||
|
||||
```go
|
||||
// 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,
|
||||
KeyType: security.KeyTypeGenericAPI,
|
||||
@@ -33,7 +33,7 @@ store := security.NewConfigKeyStore([]security.UserKey{
|
||||
|
||||
### 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
|
||||
db, _ := sql.Open("postgres", dsn)
|
||||
@@ -43,8 +43,8 @@ store := security.NewDatabaseKeyStore(db)
|
||||
// With options
|
||||
store = security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{
|
||||
CacheTTL: 5 * time.Minute,
|
||||
SQLNames: &security.KeyStoreSQLNames{
|
||||
ValidateKey: "myapp_keystore_validate", // override one procedure name
|
||||
Lookup: lookup.Config{
|
||||
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>`
|
||||
|
||||
```go
|
||||
auth := security.NewKeyStoreAuthenticator(store, "") // "" = accept any key type
|
||||
auth := providers.NewKeyStoreAuthenticator(store, "") // "" = accept any key type
|
||||
// Restrict to a specific type:
|
||||
auth = security.NewKeyStoreAuthenticator(store, security.KeyTypeGenericAPI)
|
||||
auth = providers.NewKeyStoreAuthenticator(store, security.KeyTypeGenericAPI)
|
||||
```
|
||||
|
||||
Plug it into a handler:
|
||||
@@ -109,10 +109,10 @@ On successful validation the request context receives a `UserContext` where:
|
||||
|
||||
## 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
|
||||
\i pkg/security/keystore_schema.sql
|
||||
\i pkg/security/lookup/keystore_schema.sql
|
||||
```
|
||||
|
||||
This creates:
|
||||
@@ -123,28 +123,23 @@ This creates:
|
||||
- `resolvespec_keystore_delete_key(p_user_id, p_key_id)`
|
||||
- `resolvespec_keystore_validate_key(p_key_hash, p_key_type)`
|
||||
|
||||
### Custom procedure names
|
||||
### Custom names and modes
|
||||
|
||||
```go
|
||||
store := security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{
|
||||
SQLNames: &security.KeyStoreSQLNames{
|
||||
GetUserKeys: "myschema_get_keys",
|
||||
CreateKey: "myschema_create_key",
|
||||
DeleteKey: "myschema_delete_key",
|
||||
ValidateKey: "myschema_validate_key",
|
||||
Lookup: lookup.Config{
|
||||
Procs: lookup.ProcNames{
|
||||
KeystoreGetUserKeys: "myschema_get_keys",
|
||||
KeystoreCreateKey: "myschema_create_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
|
||||
|
||||
- Raw keys are never stored. Only the SHA-256 hex digest is persisted.
|
||||
|
||||
@@ -4,6 +4,8 @@
|
||||
|
||||
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
|
||||
|
||||
- **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
|
||||
- **State Validation**: Built-in CSRF protection
|
||||
- **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
|
||||
|
||||
## Quick Start
|
||||
@@ -21,7 +24,7 @@ The security package provides OAuth2 authentication support for any OAuth2-compl
|
||||
### 1. Database Setup
|
||||
|
||||
```sql
|
||||
-- Run the schema from database_schema.sql
|
||||
-- Run the schema from lookup/database_schema.sql
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id SERIAL PRIMARY KEY,
|
||||
username VARCHAR(255) NOT NULL UNIQUE,
|
||||
@@ -53,7 +56,7 @@ CREATE TABLE IF NOT EXISTS user_sessions (
|
||||
);
|
||||
|
||||
-- OAuth2 stored procedures (7 functions)
|
||||
-- See database_schema.sql for full implementation
|
||||
-- See lookup/database_schema.sql for full implementation
|
||||
```
|
||||
|
||||
### 2. Google OAuth2
|
||||
@@ -397,7 +400,7 @@ UserInfoParser: func(userInfo map[string]any) (*security.UserContext, error) {
|
||||
|
||||
## 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_createsession` - Create OAuth2 session
|
||||
- `resolvespec_oauth_getsession` - Validate and retrieve session
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# 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)
|
||||
|
||||
### 1. Initialize Authenticator
|
||||
@@ -276,6 +278,6 @@ authURL += "&access_type=offline&prompt=consent"
|
||||
|
||||
## 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`.
|
||||
|
||||
@@ -43,16 +43,16 @@ CREATE TABLE IF NOT EXISTS user_sessions (
|
||||
**`resolvespec_oauth_getrefreshtoken(p_refresh_token)`**
|
||||
- Gets OAuth2 session data by refresh token
|
||||
- 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)`**
|
||||
- Updates session with new tokens after refresh
|
||||
- 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)`**
|
||||
- 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)
|
||||
```
|
||||
|
||||
**Location:** `pkg/security/oauth2_methods.go:375`
|
||||
**Location:** `pkg/security/oauth2_methods.go` (`OAuth2RefreshToken`)
|
||||
|
||||
### Implementation Flow
|
||||
|
||||
@@ -476,7 +476,7 @@ auth.OAuth2RefreshToken(ctx, token, "google") // Must match ProviderName
|
||||
|
||||
## 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.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -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.
|
||||
@@ -6,9 +6,9 @@ Passkey authentication (WebAuthn/FIDO2) is now integrated into the DatabaseAuthe
|
||||
## Setup
|
||||
|
||||
### 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
|
||||
- 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
|
||||
|
||||
@@ -6,7 +6,7 @@
|
||||
// Step 1: Create security providers
|
||||
auth := security.NewDatabaseAuthenticator(db) // Session-based (recommended)
|
||||
// OR: auth := security.NewJWTAuthenticator("secret-key", db)
|
||||
// OR: auth := security.NewHeaderAuthenticator()
|
||||
// OR: auth := providers.NewHeaderAuthenticator()
|
||||
// OR: auth := security.NewGoogleAuthenticator(clientID, secret, redirectURL, db) // OAuth2
|
||||
|
||||
colSec := security.NewDatabaseColumnSecurityProvider(db)
|
||||
@@ -16,7 +16,8 @@ rowSec := security.NewDatabaseRowSecurityProvider(db)
|
||||
provider, _ := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
|
||||
|
||||
// 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.SetSecurityMiddleware(securityList))
|
||||
```
|
||||
@@ -25,7 +26,7 @@ router.Use(security.SetSecurityMiddleware(securityList))
|
||||
|
||||
## 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
|
||||
```go
|
||||
@@ -55,7 +56,7 @@ All stored procedures return structured results:
|
||||
- Session/Login: `(p_success bool, p_error text, p_data 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:
|
||||
// - users (id, username, email, password, user_level, roles, is_active)
|
||||
// - 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:
|
||||
// - Login with username/password
|
||||
@@ -313,16 +314,15 @@ func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context,
|
||||
}
|
||||
|
||||
query := `
|
||||
SELECT control, accesstype, jsonvalue
|
||||
FROM core.secaccess
|
||||
WHERE rid_hub IN (
|
||||
SELECT rid_hub_parent FROM core.hub_link
|
||||
WHERE rid_hub_child = ? AND parent_hubtype = 'secgroup'
|
||||
)
|
||||
AND control ILIKE ?
|
||||
SELECT schema_name || '.' || table_name || '.' || column_path AS control,
|
||||
access_type AS accesstype, COALESCE(extra_filters, '') AS jsonvalue
|
||||
FROM sec_column_rules
|
||||
WHERE is_active = true
|
||||
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 = ?))
|
||||
`
|
||||
|
||||
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 {
|
||||
return nil, err
|
||||
}
|
||||
@@ -378,19 +378,19 @@ func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userID i
|
||||
|
||||
```go
|
||||
// Test Authenticator
|
||||
auth := security.NewHeaderAuthenticator()
|
||||
auth := providers.NewHeaderAuthenticator()
|
||||
req := httptest.NewRequest("GET", "/", nil)
|
||||
req.Header.Set("X-User-ID", "123")
|
||||
userCtx, err := auth.Authenticate(req)
|
||||
assert.Equal(t, 123, userCtx.UserID)
|
||||
|
||||
// Test ColumnSecurityProvider
|
||||
colSec := security.NewConfigColumnSecurityProvider(rules)
|
||||
colSec := providers.NewConfigColumnSecurityProvider(rules)
|
||||
cols, err := colSec.GetColumnSecurity(context.Background(), 123, "public", "employees")
|
||||
assert.Equal(t, "mask", cols[0].Accesstype)
|
||||
|
||||
// Test RowSecurityProvider
|
||||
rowSec := security.NewConfigRowSecurityProvider(templates, blocked)
|
||||
rowSec := providers.NewConfigRowSecurityProvider(templates, blocked)
|
||||
row, err := rowSec.GetRowSecurity(context.Background(), 123, "public", "orders")
|
||||
assert.Equal(t, "user_id = {UserID}", row.Template)
|
||||
```
|
||||
@@ -633,7 +633,8 @@ func main() {
|
||||
|
||||
// Setup security
|
||||
provider := &SimpleProvider{}
|
||||
securityList := security.SetupSecurityProvider(handler, provider)
|
||||
securityList, _ := security.NewSecurityList(provider)
|
||||
restheadspec.RegisterSecurityHooks(handler, securityList)
|
||||
|
||||
// Apply middleware
|
||||
router := mux.NewRouter()
|
||||
@@ -762,7 +763,8 @@ auth := security.NewJWTAuthenticator("secret", db)
|
||||
colSec := security.NewDatabaseColumnSecurityProvider(db)
|
||||
rowSec := security.NewDatabaseRowSecurityProvider(db)
|
||||
provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
|
||||
securityList := security.SetupSecurityProvider(handler, provider)
|
||||
securityList, _ := security.NewSecurityList(provider)
|
||||
restheadspec.RegisterSecurityHooks(handler, securityList)
|
||||
|
||||
// ===== INTERFACE METHODS =====
|
||||
Authenticate(r *http.Request) (*UserContext, error)
|
||||
|
||||
+113
-102
@@ -13,12 +13,12 @@ Type-safe, composable security system for ResolveSpec with support for authentic
|
||||
- ✅ **Extensible** - Implement custom providers for your needs
|
||||
- ✅ **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
|
||||
- ✅ **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
|
||||
|
||||
## 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
|
||||
|
||||
@@ -35,6 +35,7 @@ Type-safe, composable security system for ResolveSpec with support for authentic
|
||||
| `resolvespec_login` | Session-based login | DatabaseAuthenticator |
|
||||
| `resolvespec_logout` | Session invalidation | 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_refresh_token` | Token refresh | DatabaseAuthenticator |
|
||||
| `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` | 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).
|
||||
- **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.
|
||||
- **procedure** (`lookup/procedure`): calls the `resolvespec_*` stored procedures (`p_success` / `p_error` / `p_data` contract). Postgres only.
|
||||
- **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
|
||||
type QueryMode int
|
||||
|
||||
const (
|
||||
ModeAuto QueryMode = iota // default
|
||||
ModeProcedure
|
||||
ModeDirect
|
||||
)
|
||||
```
|
||||
|
||||
- **`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
|
||||
type Config struct {
|
||||
Dialect string // "postgres", "sqlite", "mysql", "mssql", or one you registered; empty = detect from the driver
|
||||
Mode lookup.Mode // default for every operation
|
||||
Overrides map[lookup.Op]lookup.Mode // per-operation mode, e.g. lookup.OpSession: lookup.ModeDirect
|
||||
Procs lookup.ProcNames // procedure names, empty fields keep the default
|
||||
Schema lookup.Schema // table/column names, missing entries keep the default
|
||||
}
|
||||
```
|
||||
|
||||
`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
|
||||
auth := security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{
|
||||
TableNames: &security.TableNames{Users: "app_users"}, // only override what differs
|
||||
// SQLite or MySQL: nothing to configure, direct SQL is the default.
|
||||
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
|
||||
|
||||
- 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`).
|
||||
- Session tokens generated by Direct mode use the same `sess_<hex>_<unix-timestamp>` shape as the plpgsql procedures.
|
||||
- `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.
|
||||
- Direct login, register, refresh, API-key login, password reset and passkey login run in one transaction.
|
||||
- Passwords are stored as bcrypt; legacy cleartext values are accepted at login and only rewritten when `UpgradePasswordHash` is enabled.
|
||||
- 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
|
||||
|
||||
@@ -297,7 +299,7 @@ type UserContext struct {
|
||||
|
||||
**HeaderAuthenticator** - Simple header-based authentication:
|
||||
```go
|
||||
auth := security.NewHeaderAuthenticator()
|
||||
auth := providers.NewHeaderAuthenticator()
|
||||
// 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
|
||||
// All operations use stored procedures: resolvespec_login, resolvespec_logout,
|
||||
// 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:
|
||||
@@ -318,19 +320,19 @@ auth := security.NewJWTAuthenticator("secret-key", db)
|
||||
// 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
|
||||
baseAuth := security.NewDatabaseAuthenticator(db)
|
||||
|
||||
// Use in-memory provider (for testing)
|
||||
tfaProvider := security.NewMemoryTwoFactorProvider(nil)
|
||||
tfaProvider := totp.NewMemoryProvider(nil)
|
||||
|
||||
// Or use database provider (for production)
|
||||
tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil)
|
||||
// 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
|
||||
// Compatible with Google Authenticator, Microsoft Authenticator, Authy, etc.
|
||||
```
|
||||
@@ -341,7 +343,7 @@ auth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil)
|
||||
```go
|
||||
colSec := security.NewDatabaseColumnSecurityProvider(db)
|
||||
// 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:
|
||||
@@ -351,7 +353,7 @@ rules := map[string][]security.ColumnSecurity{
|
||||
{Path: []string{"ssn"}, Accesstype: "mask", MaskStart: 5},
|
||||
},
|
||||
}
|
||||
colSec := security.NewConfigColumnSecurityProvider(rules)
|
||||
colSec := providers.NewConfigColumnSecurityProvider(rules)
|
||||
```
|
||||
|
||||
### Row Security Providers
|
||||
@@ -370,7 +372,7 @@ templates := map[string]string{
|
||||
blocked := map[string]bool{
|
||||
"public.admin_logs": true,
|
||||
}
|
||||
rowSec := security.NewConfigRowSecurityProvider(templates, blocked)
|
||||
rowSec := providers.NewConfigRowSecurityProvider(templates, blocked)
|
||||
```
|
||||
|
||||
## Usage Examples
|
||||
@@ -381,7 +383,7 @@ rowSec := security.NewConfigRowSecurityProvider(templates, blocked)
|
||||
func main() {
|
||||
db := setupDatabase()
|
||||
|
||||
// Run migrations (see database_schema.sql)
|
||||
// Run migrations (see lookup/database_schema.sql)
|
||||
// db.Exec("CREATE TABLE users ...")
|
||||
// db.Exec("CREATE TABLE user_sessions ...")
|
||||
|
||||
@@ -475,8 +477,8 @@ func handleRefresh(securityList *security.SecurityList) http.HandlerFunc {
|
||||
```go
|
||||
// 1. Wrap existing authenticator with 2FA support
|
||||
baseAuth := security.NewDatabaseAuthenticator(db)
|
||||
tfaProvider := security.NewMemoryTwoFactorProvider(nil) // Use custom DB implementation in production
|
||||
tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil)
|
||||
tfaProvider := totp.NewMemoryProvider(nil) // Use custom DB implementation in production
|
||||
tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, nil)
|
||||
|
||||
// 2. Use as normal authenticator
|
||||
provider := security.NewCompositeSecurityProvider(tfaAuth, colSec, rowSec)
|
||||
@@ -548,18 +550,18 @@ has2FA, err := tfaProvider.Get2FAStatus(userID)
|
||||
// Uses PostgreSQL stored procedures for all operations
|
||||
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
|
||||
// - Create user_totp_backup_codes table
|
||||
// - Create resolvespec_totp_* stored procedures
|
||||
|
||||
tfaProvider := security.NewDatabaseTwoFactorProvider(db, nil)
|
||||
tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, nil)
|
||||
tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, nil)
|
||||
```
|
||||
|
||||
**Option 2: Implement Custom Provider**
|
||||
|
||||
Implement `TwoFactorAuthProvider` for custom storage:
|
||||
Implement `totp.AuthProvider` for custom storage:
|
||||
|
||||
```go
|
||||
type DBTwoFactorProvider struct {
|
||||
@@ -585,15 +587,15 @@ func (p *DBTwoFactorProvider) Get2FASecret(userID int) (string, error) {
|
||||
### Configuration
|
||||
|
||||
```go
|
||||
config := &security.TwoFactorConfig{
|
||||
config := &totp.Config{
|
||||
Algorithm: "SHA256", // SHA1, SHA256, SHA512
|
||||
Digits: 8, // 6 or 8
|
||||
Period: 30, // Seconds per code
|
||||
SkewWindow: 2, // Accept codes ±2 periods
|
||||
}
|
||||
|
||||
totp := security.NewTOTPGenerator(config)
|
||||
tfaAuth := security.NewTwoFactorAuthenticator(baseAuth, tfaProvider, config)
|
||||
totp := totp.NewGenerator(config)
|
||||
tfaAuth := totp.NewAuthenticator(baseAuth, tfaProvider, config)
|
||||
```
|
||||
|
||||
### API Response Structure
|
||||
@@ -662,9 +664,9 @@ func main() {
|
||||
}
|
||||
|
||||
// Create providers
|
||||
auth := security.NewHeaderAuthenticator()
|
||||
colSec := security.NewConfigColumnSecurityProvider(columnRules)
|
||||
rowSec := security.NewConfigRowSecurityProvider(rowTemplates, nil)
|
||||
auth := providers.NewHeaderAuthenticator()
|
||||
colSec := providers.NewConfigColumnSecurityProvider(columnRules)
|
||||
rowSec := providers.NewConfigRowSecurityProvider(rowTemplates, nil)
|
||||
|
||||
// Combine providers and register hooks
|
||||
provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
|
||||
@@ -934,7 +936,8 @@ func TestMyHandler(t *testing.T) {
|
||||
&MockRowSecurity{},
|
||||
)
|
||||
|
||||
securityList := security.SetupSecurityProvider(handler, provider)
|
||||
securityList, _ := security.NewSecurityList(provider)
|
||||
restheadspec.RegisterSecurityHooks(handler, securityList)
|
||||
// ... test your handler
|
||||
}
|
||||
```
|
||||
@@ -1013,7 +1016,7 @@ restheadspec.RegisterSecurityHooks(handler, securityList) // or funcspec/resolve
|
||||
|
||||
### 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`)
|
||||
- `resolvespec_password_reset_request` 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
|
||||
- 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
|
||||
type SQLNames struct {
|
||||
// ...
|
||||
PasswordResetRequest string // default: "resolvespec_password_reset_request"
|
||||
PasswordResetComplete string // default: "resolvespec_password_reset"
|
||||
lookup.ProcNames{
|
||||
PasswordResetRequest: "resolvespec_password_reset_request", // default
|
||||
PasswordResetComplete: "resolvespec_password_reset", // default
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1064,6 +1068,8 @@ type SQLNames struct {
|
||||
|
||||
## 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`.
|
||||
|
||||
### 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.
|
||||
|
||||
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
|
||||
|
||||
@@ -1198,10 +1204,10 @@ auth.OAuthIntrospectToken(ctx, token) // RFC 7662 — returns OAuthTokenInfo
|
||||
auth.OAuthRevokeToken(ctx, token) // RFC 7009 — revoke session
|
||||
```
|
||||
|
||||
#### SQLNames Fields
|
||||
#### Procedure names
|
||||
|
||||
```go
|
||||
type SQLNames struct {
|
||||
type ProcNames struct {
|
||||
// ... existing fields ...
|
||||
OAuthRegisterClient string // default: "resolvespec_oauth_register_client"
|
||||
OAuthGetClient string // default: "resolvespec_oauth_get_client"
|
||||
@@ -1224,9 +1230,14 @@ The main changes:
|
||||
| File | Description |
|
||||
|------|-------------|
|
||||
| **QUICK_REFERENCE.md** | Quick reference guide with examples |
|
||||
| **INTERFACE_GUIDE.md** | Complete implementation guide |
|
||||
| **examples.go** | Working provider implementations |
|
||||
| **setup_example.go** | 6 complete integration examples |
|
||||
| **KEYSTORE.md** | Per-user auth keys and key stores |
|
||||
| **OAUTH2.md** | OAuth2 client login |
|
||||
| **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
|
||||
|
||||
|
||||
@@ -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`.
|
||||
@@ -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 {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -88,6 +88,14 @@ func (c *CompositeSecurityProvider) RefreshToken(ctx context.Context, refreshTok
|
||||
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
|
||||
func (c *CompositeSecurityProvider) ValidateToken(ctx context.Context, token string) (bool, error) {
|
||||
if validatable, ok := c.auth.(Validatable); ok {
|
||||
|
||||
@@ -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);
|
||||
@@ -5,37 +5,41 @@ import (
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"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 {
|
||||
return time.Now().Add(1 * time.Hour)
|
||||
}
|
||||
|
||||
// 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.
|
||||
func newDirectTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
|
||||
db, err := sql.Open("sqlite3", "file::memory:?cache=shared")
|
||||
db, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
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
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
|
||||
schemaPath := filepath.Join("database_schema_sqlite.sql")
|
||||
schema, err := os.ReadFile(schemaPath)
|
||||
schema, err := ddl.SQL("sqlite")
|
||||
if err != nil {
|
||||
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)
|
||||
}
|
||||
return db
|
||||
@@ -49,7 +53,7 @@ func authenticatedRequest(token string) *http.Request {
|
||||
|
||||
func TestDirectMode_RegisterThenLogin(t *testing.T) {
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||
ctx := context.Background()
|
||||
|
||||
regResp, err := auth.Register(ctx, RegisterRequest{
|
||||
@@ -105,7 +109,7 @@ func TestDirectMode_RegisterThenLogin(t *testing.T) {
|
||||
|
||||
func TestDirectMode_SessionLifecycle(t *testing.T) {
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||
ctx := context.Background()
|
||||
|
||||
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) {
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||
ctx := context.Background()
|
||||
|
||||
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) {
|
||||
db := newDirectTestDB(t)
|
||||
jwtAuth := NewJWTAuthenticator("secret", db).WithQueryMode(ModeDirect)
|
||||
directAuth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
||||
jwtAuth := NewJWTAuthenticator("secret", db).WithLookup(directConfig)
|
||||
directAuth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||
ctx := context.Background()
|
||||
|
||||
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) {
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||
ctx := context.Background()
|
||||
|
||||
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
|
||||
|
||||
totp := NewDatabaseTwoFactorProvider(db, nil).WithQueryMode(ModeDirect)
|
||||
totp := NewDatabaseTwoFactorProvider(db, nil).WithLookup(directConfig)
|
||||
|
||||
if err := totp.Enable2FA(userID, "SECRET123", []string{"code1", "code2"}); err != nil {
|
||||
t.Fatalf("Enable2FA() error = %v", err)
|
||||
@@ -248,7 +252,7 @@ func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) {
|
||||
|
||||
func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) {
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||
ctx := context.Background()
|
||||
|
||||
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
|
||||
|
||||
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{
|
||||
@@ -317,7 +321,7 @@ func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) {
|
||||
|
||||
func TestDirectMode_OAuthGetOrCreateUserAndSession(t *testing.T) {
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||
ctx := context.Background()
|
||||
|
||||
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) {
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||
ctx := context.Background()
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
ks := NewDatabaseKeyStore(db, DatabaseKeyStoreOptions{QueryMode: ModeDirect})
|
||||
ks := NewDatabaseKeyStore(db, DatabaseKeyStoreOptions{Lookup: directConfig})
|
||||
|
||||
createResp, err := ks.CreateKey(ctx, CreateKeyRequest{
|
||||
UserID: regResp.User.UserID,
|
||||
@@ -399,7 +403,7 @@ func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) {
|
||||
|
||||
func TestDirectMode_OAuthServerClientAndCode(t *testing.T) {
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig})
|
||||
ctx := context.Background()
|
||||
|
||||
client := &OAuthServerClient{
|
||||
@@ -489,7 +493,7 @@ func TestDirectMode_OAuthServerClientAndCode(t *testing.T) {
|
||||
func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) {
|
||||
for _, enabled := range []bool{false, true} {
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect, UpgradePasswordHash: enabled})
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{Lookup: directConfig, UpgradePasswordHash: enabled})
|
||||
ctx := context.Background()
|
||||
|
||||
if _, err := db.Exec(`DELETE FROM users`); err != nil {
|
||||
@@ -518,18 +522,6 @@ func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyPasswordEdgeCases(t *testing.T) {
|
||||
h, _ := hashPassword("pw")
|
||||
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")
|
||||
}
|
||||
func isBcryptHash(s string) bool {
|
||||
return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$")
|
||||
}
|
||||
|
||||
@@ -441,6 +441,19 @@ func resolveModelRules(secCtx SecurityContext) (modelregistry.ModelRules, bool)
|
||||
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.
|
||||
func CheckModelUpdateAllowed(secCtx SecurityContext) error {
|
||||
return checkModelUpdateAllowed(secCtx)
|
||||
|
||||
@@ -5,81 +5,6 @@ import (
|
||||
"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
|
||||
type Authenticator interface {
|
||||
// Login authenticates credentials and returns a token
|
||||
@@ -144,6 +69,13 @@ type Refreshable interface {
|
||||
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
|
||||
type Validatable interface {
|
||||
// ValidateToken checks if a token is valid without extracting full user context
|
||||
|
||||
@@ -2,64 +2,12 @@ package security
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// hashSHA256Hex returns the lowercase hex SHA-256 digest of the given string.
|
||||
// Used by all keystore implementations to hash raw keys before storage or lookup.
|
||||
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
|
||||
}
|
||||
// hashSHA256Hex is kept as a short alias for sectypes.HashKey inside this package.
|
||||
func hashSHA256Hex(raw string) string { return sectypes.HashKey(raw) }
|
||||
|
||||
// KeyStore manages per-user auth keys with pluggable storage backends.
|
||||
// Implementations: ConfigKeyStore (static list) and DatabaseKeyStore (stored procedures).
|
||||
|
||||
@@ -5,16 +5,16 @@ import (
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sync/singleflight"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/backends"
|
||||
)
|
||||
|
||||
// DatabaseKeyStoreOptions configures DatabaseKeyStore.
|
||||
@@ -24,36 +24,28 @@ type DatabaseKeyStoreOptions struct {
|
||||
// CacheTTL is the duration to cache ValidateKey results.
|
||||
// Default: 2 minutes.
|
||||
CacheTTL time.Duration
|
||||
// SQLNames provides custom procedure names. If nil, uses DefaultKeyStoreSQLNames().
|
||||
SQLNames *KeyStoreSQLNames
|
||||
// TableNames provides custom table names for Direct mode. If nil, uses DefaultKeyStoreTableNames().
|
||||
TableNames *KeyStoreTableNames
|
||||
// QueryMode selects stored-procedure vs Direct-mode SQL. Default (zero value) is ModeAuto.
|
||||
QueryMode QueryMode
|
||||
// Lookup selects dialect, query mode and procedure/table/column names.
|
||||
// The zero value uses stored procedures on Postgres and direct SQL elsewhere.
|
||||
Lookup lookup.Config
|
||||
// LookupProvider, when set, is used instead of building one from Lookup and the db.
|
||||
LookupProvider *lookup.Provider
|
||||
// DBFactory is called to obtain a fresh *sql.DB when the existing connection is closed.
|
||||
// If nil, reconnection is disabled.
|
||||
DBFactory func() (*sql.DB, error)
|
||||
}
|
||||
|
||||
// DatabaseKeyStore is a KeyStore backed by PostgreSQL stored procedures.
|
||||
// All DB operations go through configurable procedure names; the raw key is
|
||||
// never passed to the database.
|
||||
// DatabaseKeyStore is a KeyStore backed by the lookup package (stored procedures on
|
||||
// Postgres by default, direct SQL elsewhere). The raw key is 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
|
||||
// cache TTL, a deleted key may continue to authenticate for up to CacheTTL
|
||||
// (default 2 minutes) if the cache entry cannot be invalidated.
|
||||
type DatabaseKeyStore struct {
|
||||
db *sql.DB
|
||||
dbMu sync.RWMutex
|
||||
dbFactory func() (*sql.DB, error)
|
||||
sqlNames *KeyStoreSQLNames
|
||||
tableNames *KeyStoreTableNames
|
||||
queryMode QueryMode
|
||||
capability *dbCapability
|
||||
cache *cache.Cache
|
||||
cacheTTL time.Duration
|
||||
src *lookupSource
|
||||
cache *cache.Cache
|
||||
cacheTTL time.Duration
|
||||
|
||||
// validateLoads collapses concurrent key lookups for the same key
|
||||
validateLoads singleflight.Group
|
||||
@@ -72,42 +64,14 @@ func NewDatabaseKeyStore(db *sql.DB, opts ...DatabaseKeyStoreOptions) *DatabaseK
|
||||
if c == nil {
|
||||
c = cache.GetDefaultCache()
|
||||
}
|
||||
names := MergeKeyStoreSQLNames(DefaultKeyStoreSQLNames(), o.SQLNames)
|
||||
tableNames := resolveKeyStoreTableNames(o.TableNames)
|
||||
return &DatabaseKeyStore{
|
||||
db: db,
|
||||
dbFactory: o.DBFactory,
|
||||
sqlNames: names,
|
||||
tableNames: tableNames,
|
||||
queryMode: o.QueryMode,
|
||||
capability: newDBCapability(),
|
||||
cache: c,
|
||||
cacheTTL: o.CacheTTL,
|
||||
}
|
||||
src := newLookupSource(db)
|
||||
src.cfg = o.Lookup
|
||||
src.provider = o.LookupProvider
|
||||
src.opts = backends.Options{DBFactory: o.DBFactory}
|
||||
return &DatabaseKeyStore{src: src, cache: c, cacheTTL: o.CacheTTL}
|
||||
}
|
||||
|
||||
func (ks *DatabaseKeyStore) getDB() *sql.DB {
|
||||
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
|
||||
}
|
||||
func (ks *DatabaseKeyStore) keys() lookup.KeyStore { return ks.src.get().Keys }
|
||||
|
||||
// CreateKey generates a raw key, stores its SHA-256 hash via the create procedure,
|
||||
// and returns the raw key once.
|
||||
@@ -119,110 +83,29 @@ func (ks *DatabaseKeyStore) CreateKey(ctx context.Context, req CreateKeyRequest)
|
||||
rawKey := base64.RawURLEncoding.EncodeToString(rawBytes)
|
||||
hash := hashSHA256Hex(rawKey)
|
||||
|
||||
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.CreateKey) {
|
||||
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,
|
||||
})
|
||||
key, err := ks.keys().Create(ctx, req, hash)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal create key request: %w", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
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
|
||||
return &CreateKeyResponse{Key: *key, RawKey: rawKey}, nil
|
||||
}
|
||||
|
||||
// GetUserKeys returns all active, non-expired keys for the given user.
|
||||
// Pass an empty KeyType to return all types.
|
||||
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.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
|
||||
return ks.keys().List(ctx, userID, keyType)
|
||||
}
|
||||
|
||||
// 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.
|
||||
// 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 {
|
||||
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.DeleteKey) {
|
||||
return ks.deleteKeyDirect(ctx, userID, keyID)
|
||||
keyHash, err := ks.keys().Delete(ctx, userID, keyID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var success bool
|
||||
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))
|
||||
if keyHash != "" && ks.cache != nil {
|
||||
_ = ks.cache.Delete(ctx, keystoreCacheKey(keyHash))
|
||||
}
|
||||
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.
|
||||
func (ks *DatabaseKeyStore) validateKeyLoad(ctx context.Context, hash, cacheKey string, keyType KeyType) (*UserKey, error) {
|
||||
dbtrace.Raw(ctx, "keystore.validate")
|
||||
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.ValidateKey) {
|
||||
key, err := ks.validateKeyDirect(ctx, hash, keyType)
|
||||
if err != nil {
|
||||
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)
|
||||
key, err := ks.keys().Validate(ctx, hash, keyType)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
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 {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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, §ypes.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 := §ypes.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, §ypes.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, §ypes.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, §ypes.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, §ypes.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)
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
$$ 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
|
||||
-- Input: user_id (int), schema (text), table_name (text)
|
||||
-- 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)
|
||||
RETURNS TABLE(p_success boolean, p_error text, p_rules jsonb) AS $$
|
||||
DECLARE
|
||||
v_rules jsonb;
|
||||
BEGIN
|
||||
-- Query column security rules from core.secaccess
|
||||
SELECT jsonb_agg(
|
||||
jsonb_build_object(
|
||||
'control', control,
|
||||
'accesstype', accesstype,
|
||||
'jsonvalue', jsonvalue
|
||||
'control', r.schema_name || '.' || r.table_name || '.' || r.column_path,
|
||||
'accesstype', r.access_type,
|
||||
'jsonvalue', COALESCE(r.extra_filters, '')
|
||||
)
|
||||
)
|
||||
INTO v_rules
|
||||
FROM core.secaccess
|
||||
WHERE rid_hub IN (
|
||||
SELECT rid_hub_parent
|
||||
FROM core.hub_link
|
||||
WHERE rid_hub_child = p_user_id AND parent_hubtype = 'secgroup'
|
||||
)
|
||||
AND control ILIKE (p_schema || '.' || p_table_name || '%');
|
||||
FROM sec_column_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)
|
||||
);
|
||||
|
||||
IF v_rules IS NULL THEN
|
||||
v_rules := '[]'::jsonb;
|
||||
@@ -464,20 +512,36 @@ EXCEPTION
|
||||
END;
|
||||
$$ 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)
|
||||
-- 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)
|
||||
RETURNS TABLE(p_template text, p_block boolean) AS $$
|
||||
DECLARE
|
||||
v_block boolean;
|
||||
v_template text;
|
||||
BEGIN
|
||||
-- Call the existing core function if it exists, or implement your own logic
|
||||
-- This is a placeholder that you should customize based on your core.api_sec_rowtemplate logic
|
||||
RETURN QUERY SELECT ''::text, false;
|
||||
SELECT COALESCE(bool_or(r.has_block), false),
|
||||
COALESCE(string_agg('(' || r.template || ')', ' AND ' ORDER BY r.id)
|
||||
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:
|
||||
-- RETURN QUERY SELECT template, has_block
|
||||
-- FROM core.row_security_config
|
||||
-- WHERE schema_name = p_schema AND table_name = p_table_name AND user_id = p_user_id;
|
||||
IF v_block THEN
|
||||
v_template := '';
|
||||
END IF;
|
||||
|
||||
RETURN QUERY SELECT v_template, v_block;
|
||||
END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
|
||||
@@ -650,7 +714,7 @@ BEGIN
|
||||
v_auth_provider := COALESCE(p_user_data->>'auth_provider', 'oauth2');
|
||||
|
||||
-- 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;
|
||||
|
||||
-- Try to find existing user by email
|
||||
@@ -701,7 +765,7 @@ BEGIN
|
||||
v_access_token := p_session_data->>'access_token';
|
||||
v_refresh_token := p_session_data->>'refresh_token';
|
||||
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');
|
||||
|
||||
-- Insert or update session
|
||||
@@ -857,7 +921,7 @@ BEGIN
|
||||
v_new_session_token := p_update_data->>'new_session_token';
|
||||
v_new_access_token := p_update_data->>'new_access_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 user_sessions
|
||||
@@ -1214,7 +1278,7 @@ BEGIN
|
||||
|
||||
-- Convert transports array
|
||||
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;
|
||||
END IF;
|
||||
|
||||
@@ -1304,12 +1368,11 @@ BEGIN
|
||||
'name', name,
|
||||
'created_at', created_at,
|
||||
'last_used_at', last_used_at
|
||||
)
|
||||
) ORDER BY created_at DESC
|
||||
), '[]'::jsonb)
|
||||
INTO v_credentials
|
||||
FROM user_passkey_credentials
|
||||
WHERE user_id = p_user_id
|
||||
ORDER BY created_at DESC;
|
||||
WHERE user_id = p_user_id;
|
||||
|
||||
RETURN QUERY SELECT true, NULL::text, v_credentials;
|
||||
EXCEPTION
|
||||
@@ -1451,6 +1514,64 @@ EXCEPTION
|
||||
END;
|
||||
$$ 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
|
||||
-- ============================================
|
||||
@@ -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
|
||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||
is_active BOOLEAN DEFAULT true,
|
||||
metadata jsonb, -- every other RFC 7591 field (see sectypes.OAuthServerClient)
|
||||
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)
|
||||
-- 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,
|
||||
scopes TEXT[],
|
||||
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
|
||||
);
|
||||
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_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
|
||||
-- ============================================
|
||||
|
||||
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)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_client_id text;
|
||||
v_row jsonb;
|
||||
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 (
|
||||
v_client_id,
|
||||
ARRAY(SELECT jsonb_array_elements_text(p_data->'redirect_uris')),
|
||||
p_data->>'client_name',
|
||||
COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'grant_types')), ARRAY['authorization_code']),
|
||||
COALESCE(ARRAY(SELECT jsonb_array_elements_text(p_data->'allowed_scopes')), ARRAY['openid','profile','email']),
|
||||
NULLIF(p_data->>'client_secret_hash', ''),
|
||||
COALESCE(NULLIF(p_data->>'token_endpoint_auth_method', ''), 'none')
|
||||
ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)),
|
||||
p_request->>'client_name',
|
||||
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,
|
||||
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_request->>'client_secret_hash', ''),
|
||||
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;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
@@ -1704,7 +1830,7 @@ LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_row jsonb;
|
||||
BEGIN
|
||||
SELECT to_jsonb(oauth_clients.*)
|
||||
SELECT (to_jsonb(oauth_clients.*) - 'metadata') || COALESCE(metadata, '{}'::jsonb)
|
||||
INTO v_row
|
||||
FROM oauth_clients
|
||||
WHERE client_id = p_client_id AND is_active = true;
|
||||
@@ -1717,22 +1843,23 @@ BEGIN
|
||||
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)
|
||||
LANGUAGE plpgsql AS $$
|
||||
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 (
|
||||
p_data->>'code',
|
||||
p_data->>'client_id',
|
||||
p_data->>'redirect_uri',
|
||||
p_data->>'client_state',
|
||||
p_data->>'code_challenge',
|
||||
COALESCE(p_data->>'code_challenge_method', 'S256'),
|
||||
p_data->>'session_token',
|
||||
p_data->>'refresh_token',
|
||||
ARRAY(SELECT jsonb_array_elements_text(p_data->'scopes')),
|
||||
(p_data->>'expires_at')::timestamp
|
||||
p_request->>'code',
|
||||
p_request->>'client_id',
|
||||
p_request->>'redirect_uri',
|
||||
p_request->>'client_state',
|
||||
p_request->>'code_challenge',
|
||||
COALESCE(p_request->>'code_challenge_method', 'S256'),
|
||||
p_request->>'session_token',
|
||||
p_request->>'refresh_token',
|
||||
ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'scopes') = 'array' THEN p_request->'scopes' ELSE '[]'::jsonb END)),
|
||||
(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;
|
||||
@@ -1758,7 +1885,7 @@ BEGIN
|
||||
'session_token', session_token,
|
||||
'refresh_token', refresh_token,
|
||||
'scopes', to_jsonb(scopes)
|
||||
) INTO v_row;
|
||||
) || COALESCE(extra, '{}'::jsonb) INTO v_row;
|
||||
|
||||
IF v_row IS NULL THEN
|
||||
RETURN QUERY SELECT false, 'invalid or expired code'::text, null::jsonb;
|
||||
@@ -1809,3 +1936,374 @@ BEGIN
|
||||
RETURN QUERY SELECT true, null::text;
|
||||
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;
|
||||
$$;
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
|
||||
@@ -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)
|
||||
);
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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 §ypes.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 §ypes.LoginResponse{
|
||||
Token: token,
|
||||
User: §ypes.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 §ypes.LoginResponse{
|
||||
Token: token,
|
||||
User: §ypes.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 §ypes.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 §ypes.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 §ypes.LoginResponse{
|
||||
Token: fmt.Sprintf("token_%d_%d", userID, expiresAt.Unix()),
|
||||
User: §ypes.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 §ypes.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 §ypes.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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,478 @@
|
||||
// Package direct is the table-backed implementation of the lookup stores. SQL is built
|
||||
// from the configured Schema (table and column names) and Dialect (placeholders, quoting,
|
||||
// booleans, insert-returning-id); no statement is written per database and no ORM is used.
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
|
||||
)
|
||||
|
||||
// Runner runs a database operation, reconnecting once when the *sql.DB has been closed.
|
||||
// procedure.Runner (and procedure.DB) satisfy it.
|
||||
type Runner interface {
|
||||
Run(run func(*sql.DB) error) error
|
||||
}
|
||||
|
||||
// Querier is implemented by *sql.DB and *sql.Tx.
|
||||
type Querier interface {
|
||||
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
|
||||
QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error)
|
||||
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
|
||||
}
|
||||
|
||||
// Base is the state shared by every direct store: the runner, dialect, schema and clock.
|
||||
type Base struct {
|
||||
run Runner
|
||||
d dialect.Dialect
|
||||
schema lookup.Schema
|
||||
// Now is the clock; tests replace it.
|
||||
Now func() time.Time
|
||||
}
|
||||
|
||||
// NewBase creates the shared state. The schema is merged with the defaults and validated.
|
||||
func NewBase(run Runner, d dialect.Dialect, schema lookup.Schema) (*Base, error) {
|
||||
if run == nil {
|
||||
return nil, fmt.Errorf("direct: nil runner")
|
||||
}
|
||||
if d == nil {
|
||||
return nil, fmt.Errorf("direct: nil dialect")
|
||||
}
|
||||
merged := lookup.DefaultSchema().Merge(schema)
|
||||
if err := merged.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Base{run: run, d: d, schema: merged, Now: time.Now}, nil
|
||||
}
|
||||
|
||||
// Dialect returns the dialect in use.
|
||||
func (b *Base) Dialect() dialect.Dialect { return b.d }
|
||||
|
||||
// do runs fn against the database without a transaction.
|
||||
func (b *Base) do(fn func(q Querier) error) error {
|
||||
return b.run.Run(func(db *sql.DB) error { return fn(db) })
|
||||
}
|
||||
|
||||
// tx runs fn in one transaction; an error rolls back.
|
||||
func (b *Base) tx(ctx context.Context, fn func(q Querier) error) error {
|
||||
return b.run.Run(func(db *sql.DB) error {
|
||||
tx, err := db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := fn(tx); err != nil {
|
||||
_ = tx.Rollback()
|
||||
return err
|
||||
}
|
||||
return tx.Commit()
|
||||
})
|
||||
}
|
||||
|
||||
// tableRef returns the (possibly schema-qualified) physical table name of an entity.
|
||||
func (b *Base) tableRef(e lookup.Entity) string {
|
||||
t := b.schema[e]
|
||||
name := t.Name
|
||||
if name == "" {
|
||||
name = string(e)
|
||||
}
|
||||
if t.Schema != "" {
|
||||
return t.Schema + "." + name
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
// colName returns the physical column name of a logical column.
|
||||
func (b *Base) colName(c lookup.Column) string {
|
||||
if t, ok := b.schema[c.Entity]; ok {
|
||||
if n := t.Columns[c.Name]; n != "" {
|
||||
return n
|
||||
}
|
||||
}
|
||||
return c.Name
|
||||
}
|
||||
|
||||
// arg converts a Go value to a bind argument (booleans go through the dialect).
|
||||
func (b *Base) arg(v any) any {
|
||||
if bv, ok := v.(bool); ok {
|
||||
return b.d.Bool(bv)
|
||||
}
|
||||
// Timestamp columns carry no zone: bind every instant as UTC so drivers that send an
|
||||
// offset (SQL Server) and ones that drop it agree on the stored wall clock.
|
||||
if tv, ok := v.(time.Time); ok {
|
||||
return tv.UTC()
|
||||
}
|
||||
if tp, ok := v.(*time.Time); ok {
|
||||
if tp == nil {
|
||||
return nil
|
||||
}
|
||||
return tp.UTC()
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
// timeDest scans a time column through the dialect, so drivers returning strings work.
|
||||
type timeDest struct {
|
||||
d dialect.Dialect
|
||||
v *time.Time
|
||||
}
|
||||
|
||||
func (t timeDest) Scan(src any) error {
|
||||
v, err := t.d.ScanTime(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
*t.v = v
|
||||
return nil
|
||||
}
|
||||
|
||||
type boolDest struct {
|
||||
d dialect.Dialect
|
||||
v *bool
|
||||
}
|
||||
|
||||
func (t boolDest) Scan(src any) error {
|
||||
v, err := t.d.ScanBool(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
*t.v = v
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *Base) timeDest(v *time.Time) sql.Scanner { return timeDest{d: b.d, v: v} }
|
||||
func (b *Base) boolDest(v *bool) sql.Scanner { return boolDest{d: b.d, v: v} }
|
||||
|
||||
// --- query builder --------------------------------------------------------
|
||||
|
||||
// builder accumulates bind arguments and renders column references.
|
||||
type builder struct {
|
||||
b *Base
|
||||
args []any
|
||||
aliases map[lookup.Entity]string
|
||||
nalias int
|
||||
}
|
||||
|
||||
func (bl *builder) ph(v any) string {
|
||||
bl.args = append(bl.args, bl.b.arg(v))
|
||||
return bl.b.d.Placeholder(len(bl.args))
|
||||
}
|
||||
|
||||
// col renders a column; with aliases set (select queries) it is qualified by its table alias.
|
||||
func (bl *builder) col(c lookup.Column) string {
|
||||
name := bl.b.d.Quote(bl.b.colName(c))
|
||||
if bl.aliases != nil {
|
||||
if a, ok := bl.aliases[c.Entity]; ok {
|
||||
return a + "." + name
|
||||
}
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
// Cond renders one boolean condition.
|
||||
type Cond func(*builder) string
|
||||
|
||||
// Eq is `col = value`.
|
||||
func Eq(c lookup.Column, v any) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " = " + bl.ph(v) }
|
||||
}
|
||||
|
||||
// EqFold is a case-insensitive `LOWER(col) = value` match (the value is lowered in Go).
|
||||
func EqFold(c lookup.Column, v string) Cond {
|
||||
return func(bl *builder) string { return "LOWER(" + bl.col(c) + ") = " + bl.ph(strings.ToLower(v)) }
|
||||
}
|
||||
|
||||
// Ne is `col <> value`.
|
||||
func Ne(c lookup.Column, v any) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " <> " + bl.ph(v) }
|
||||
}
|
||||
|
||||
// Gt is `col > value`.
|
||||
func Gt(c lookup.Column, v any) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " > " + bl.ph(v) }
|
||||
}
|
||||
|
||||
// Lt is `col < value`.
|
||||
func Lt(c lookup.Column, v any) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " < " + bl.ph(v) }
|
||||
}
|
||||
|
||||
// IsNull is `col IS NULL`.
|
||||
func IsNull(c lookup.Column) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " IS NULL" }
|
||||
}
|
||||
|
||||
// EqCol is `a = b` between two columns (join conditions).
|
||||
func EqCol(a, c lookup.Column) Cond {
|
||||
return func(bl *builder) string { return bl.col(a) + " = " + bl.col(c) }
|
||||
}
|
||||
|
||||
// In is `col IN (v...)`; an empty list renders a condition that is never true.
|
||||
func In(c lookup.Column, vs ...any) Cond {
|
||||
return func(bl *builder) string {
|
||||
if len(vs) == 0 {
|
||||
return "1 = 0"
|
||||
}
|
||||
ph := make([]string, len(vs))
|
||||
for i, v := range vs {
|
||||
ph[i] = bl.ph(v)
|
||||
}
|
||||
return bl.col(c) + " IN (" + strings.Join(ph, ", ") + ")"
|
||||
}
|
||||
}
|
||||
|
||||
// InSelect is `col IN (subselect)`; the subselect's arguments share the outer numbering.
|
||||
func InSelect(c lookup.Column, sub *Select) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " IN (" + sub.render(bl) + ")" }
|
||||
}
|
||||
|
||||
// Or joins conditions with OR inside parentheses.
|
||||
func Or(cs ...Cond) Cond { return joinConds("OR", cs) }
|
||||
|
||||
// And joins conditions with AND inside parentheses.
|
||||
func And(cs ...Cond) Cond { return joinConds("AND", cs) }
|
||||
|
||||
func joinConds(op string, cs []Cond) Cond {
|
||||
return func(bl *builder) string {
|
||||
parts := make([]string, len(cs))
|
||||
for i, c := range cs {
|
||||
parts[i] = c(bl)
|
||||
}
|
||||
return "(" + strings.Join(parts, " "+op+" ") + ")"
|
||||
}
|
||||
}
|
||||
|
||||
func (bl *builder) where(cs []Cond) string {
|
||||
if len(cs) == 0 {
|
||||
return ""
|
||||
}
|
||||
parts := make([]string, len(cs))
|
||||
for i, c := range cs {
|
||||
parts[i] = c(bl)
|
||||
}
|
||||
return " WHERE " + strings.Join(parts, " AND ")
|
||||
}
|
||||
|
||||
// Stmt is a rendered statement.
|
||||
type Stmt struct {
|
||||
SQL string
|
||||
Args []any
|
||||
}
|
||||
|
||||
// Select builds a SELECT.
|
||||
type Select struct {
|
||||
b *Base
|
||||
from lookup.Entity
|
||||
joins []join
|
||||
cols []lookup.Column
|
||||
conds []Cond
|
||||
order []lookup.Column
|
||||
}
|
||||
|
||||
type join struct {
|
||||
e lookup.Entity
|
||||
on Cond
|
||||
}
|
||||
|
||||
// From starts a SELECT on e.
|
||||
func (b *Base) From(e lookup.Entity) *Select { return &Select{b: b, from: e} }
|
||||
|
||||
// Cols sets the selected columns.
|
||||
func (s *Select) Cols(cs ...lookup.Column) *Select { s.cols = cs; return s }
|
||||
|
||||
// Join adds `JOIN e ON on`.
|
||||
func (s *Select) Join(e lookup.Entity, on Cond) *Select {
|
||||
s.joins = append(s.joins, join{e: e, on: on})
|
||||
return s
|
||||
}
|
||||
|
||||
// Where adds AND-ed conditions.
|
||||
func (s *Select) Where(cs ...Cond) *Select { s.conds = append(s.conds, cs...); return s }
|
||||
|
||||
// OrderBy adds ascending order columns.
|
||||
func (s *Select) OrderBy(cs ...lookup.Column) *Select { s.order = append(s.order, cs...); return s }
|
||||
|
||||
// Build renders the statement.
|
||||
func (s *Select) Build() Stmt {
|
||||
bl := &builder{b: s.b}
|
||||
sqlText := s.render(bl)
|
||||
return Stmt{SQL: sqlText, Args: bl.args}
|
||||
}
|
||||
|
||||
// render writes the select into bl, giving every table a fresh alias so a subselect cannot
|
||||
// clash with the statement around it.
|
||||
func (s *Select) render(bl *builder) string {
|
||||
saved := bl.aliases
|
||||
defer func() { bl.aliases = saved }()
|
||||
bl.aliases = map[lookup.Entity]string{}
|
||||
alias := func() string { a := fmt.Sprintf("t%d", bl.nalias); bl.nalias++; return a }
|
||||
bl.aliases[s.from] = alias()
|
||||
for _, j := range s.joins {
|
||||
bl.aliases[j.e] = alias()
|
||||
}
|
||||
sel := make([]string, len(s.cols))
|
||||
for i, c := range s.cols {
|
||||
sel[i] = bl.col(c)
|
||||
}
|
||||
var sb strings.Builder
|
||||
sb.WriteString("SELECT " + strings.Join(sel, ", "))
|
||||
sb.WriteString(" FROM " + s.b.d.Quote(s.b.tableRef(s.from)) + " " + bl.aliases[s.from])
|
||||
for _, j := range s.joins {
|
||||
sb.WriteString(" JOIN " + s.b.d.Quote(s.b.tableRef(j.e)) + " " + bl.aliases[j.e] + " ON " + j.on(bl))
|
||||
}
|
||||
sb.WriteString(bl.where(s.conds))
|
||||
if len(s.order) > 0 {
|
||||
o := make([]string, len(s.order))
|
||||
for i, c := range s.order {
|
||||
o[i] = bl.col(c)
|
||||
}
|
||||
sb.WriteString(" ORDER BY " + strings.Join(o, ", "))
|
||||
}
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
// QueryRow runs the select and scans the first row into dest.
|
||||
func (s *Select) QueryRow(ctx context.Context, q Querier, dest ...any) error {
|
||||
st := s.Build()
|
||||
return q.QueryRowContext(ctx, st.SQL, st.Args...).Scan(dest...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
|
||||
}
|
||||
|
||||
// Query runs the select.
|
||||
func (s *Select) Query(ctx context.Context, q Querier) (*sql.Rows, error) {
|
||||
st := s.Build()
|
||||
return q.QueryContext(ctx, st.SQL, st.Args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
|
||||
}
|
||||
|
||||
// Exists reports whether the select returns at least one row.
|
||||
func (s *Select) Exists(ctx context.Context, q Querier) (bool, error) {
|
||||
s.cols = []lookup.Column{s.firstCol()}
|
||||
rows, err := s.Query(ctx, q)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
ok := rows.Next()
|
||||
return ok, rows.Err()
|
||||
}
|
||||
|
||||
func (s *Select) firstCol() lookup.Column {
|
||||
if len(s.cols) > 0 {
|
||||
return s.cols[0]
|
||||
}
|
||||
return lookup.Column{Entity: s.from, Name: lookup.FirstColumn(s.from)}
|
||||
}
|
||||
|
||||
// Assignment is one `col = value` of an UPDATE or INSERT.
|
||||
type Assignment struct {
|
||||
Col lookup.Column
|
||||
Val any
|
||||
}
|
||||
|
||||
// Set builds an Assignment.
|
||||
func Set(c lookup.Column, v any) Assignment { return Assignment{Col: c, Val: v} }
|
||||
|
||||
// Update builds an UPDATE.
|
||||
type Update struct {
|
||||
b *Base
|
||||
e lookup.Entity
|
||||
sets []Assignment
|
||||
conds []Cond
|
||||
}
|
||||
|
||||
// Update starts an UPDATE of e.
|
||||
func (b *Base) Update(e lookup.Entity) *Update { return &Update{b: b, e: e} }
|
||||
|
||||
// Set adds assignments.
|
||||
func (u *Update) Set(as ...Assignment) *Update { u.sets = append(u.sets, as...); return u }
|
||||
|
||||
// Where adds AND-ed conditions.
|
||||
func (u *Update) Where(cs ...Cond) *Update { u.conds = append(u.conds, cs...); return u }
|
||||
|
||||
// Exec runs the update and returns the affected row count.
|
||||
func (u *Update) Exec(ctx context.Context, q Querier) (int64, error) {
|
||||
bl := &builder{b: u.b}
|
||||
set := make([]string, len(u.sets))
|
||||
for i, a := range u.sets {
|
||||
set[i] = bl.col(a.Col) + " = " + bl.ph(a.Val)
|
||||
}
|
||||
sqlText := "UPDATE " + u.b.d.Quote(u.b.tableRef(u.e)) + " SET " + strings.Join(set, ", ") + bl.where(u.conds)
|
||||
res, err := q.ExecContext(ctx, sqlText, bl.args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.RowsAffected()
|
||||
}
|
||||
|
||||
// Delete builds a DELETE.
|
||||
type Delete struct {
|
||||
b *Base
|
||||
e lookup.Entity
|
||||
conds []Cond
|
||||
}
|
||||
|
||||
// Delete starts a DELETE on e.
|
||||
func (b *Base) Delete(e lookup.Entity) *Delete { return &Delete{b: b, e: e} }
|
||||
|
||||
// Where adds AND-ed conditions.
|
||||
func (d *Delete) Where(cs ...Cond) *Delete { d.conds = append(d.conds, cs...); return d }
|
||||
|
||||
// Exec runs the delete and returns the affected row count.
|
||||
func (d *Delete) Exec(ctx context.Context, q Querier) (int64, error) {
|
||||
bl := &builder{b: d.b}
|
||||
sqlText := "DELETE FROM " + d.b.d.Quote(d.b.tableRef(d.e)) + bl.where(d.conds)
|
||||
res, err := q.ExecContext(ctx, sqlText, bl.args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.RowsAffected()
|
||||
}
|
||||
|
||||
// Insert builds an INSERT.
|
||||
type Insert struct {
|
||||
b *Base
|
||||
e lookup.Entity
|
||||
sets []Assignment
|
||||
}
|
||||
|
||||
// Insert starts an INSERT into e.
|
||||
func (b *Base) Insert(e lookup.Entity) *Insert { return &Insert{b: b, e: e} }
|
||||
|
||||
// Set adds assignments.
|
||||
func (i *Insert) Set(as ...Assignment) *Insert { i.sets = append(i.sets, as...); return i }
|
||||
|
||||
func (i *Insert) colsAndArgs() (cols []string, args []any) {
|
||||
cols = make([]string, len(i.sets))
|
||||
args = make([]any, len(i.sets))
|
||||
for n, a := range i.sets {
|
||||
cols[n] = i.b.colName(a.Col)
|
||||
args[n] = i.b.arg(a.Val)
|
||||
}
|
||||
return cols, args
|
||||
}
|
||||
|
||||
// Exec runs the insert.
|
||||
func (i *Insert) Exec(ctx context.Context, q Querier) error {
|
||||
cols, args := i.colsAndArgs()
|
||||
ph := make([]string, len(cols))
|
||||
qc := make([]string, len(cols))
|
||||
for n, c := range cols {
|
||||
qc[n] = i.b.d.Quote(c)
|
||||
ph[n] = i.b.d.Placeholder(n + 1)
|
||||
}
|
||||
sqlText := "INSERT INTO " + i.b.d.Quote(i.b.tableRef(i.e)) + " (" + strings.Join(qc, ", ") + ") VALUES (" + strings.Join(ph, ", ") + ")"
|
||||
_, err := q.ExecContext(ctx, sqlText, args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
|
||||
return err
|
||||
}
|
||||
|
||||
// ExecID runs the insert and returns the generated value of idCol, using the dialect's
|
||||
// insert-returning-id strategy.
|
||||
func (i *Insert) ExecID(ctx context.Context, q Querier, idCol lookup.Column) (int64, error) {
|
||||
cols, args := i.colsAndArgs()
|
||||
ins := i.b.d.InsertReturningID(i.b.tableRef(i.e), cols, i.b.colName(idCol))
|
||||
return ins.Run(ctx, q, args...)
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
|
||||
)
|
||||
|
||||
type noRun struct{}
|
||||
|
||||
func (noRun) Run(func(*sql.DB) error) error { return nil }
|
||||
|
||||
func TestBuilderRendersPerDialect(t *testing.T) {
|
||||
schema := lookup.Schema{lookup.EntityUserSessions: {Schema: "auth", Name: "sessions"}}
|
||||
cases := map[string]string{
|
||||
"postgres": `SELECT t0."session_token" FROM "auth"."sessions" t0`,
|
||||
"mysql": "SELECT t0.`session_token` FROM `auth`.`sessions` t0",
|
||||
"mssql": `SELECT t0.[session_token] FROM [auth].[sessions] t0`,
|
||||
}
|
||||
tails := map[string]string{
|
||||
"postgres": ` WHERE t0."user_id" = $1 AND t0."session_token" IN ($2, $3)`,
|
||||
"mysql": " WHERE t0.`user_id` = ? AND t0.`session_token` IN (?, ?)",
|
||||
"mssql": ` WHERE t0.[user_id] = @p1 AND t0.[session_token] IN (@p2, @p3)`,
|
||||
}
|
||||
for name, head := range cases {
|
||||
d, err := dialect.Get(name)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := NewBase(noRun{}, d, schema)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
st := b.From(lookup.EntityUserSessions).Cols(lookup.SessionsToken).
|
||||
Where(Eq(lookup.SessionsUserID, 7), In(lookup.SessionsToken, "a", "b")).Build()
|
||||
if st.SQL != head+tails[name] || len(st.Args) != 3 {
|
||||
t.Errorf("%s:\n got %s\n want %s", name, st.SQL, head+tails[name])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuilderBoolsGoThroughDialect(t *testing.T) {
|
||||
d, _ := dialect.Get("sqlite")
|
||||
b, _ := NewBase(noRun{}, d, nil)
|
||||
st := b.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersIsActive, true)).Build()
|
||||
if st.Args[0] != d.Bool(true) {
|
||||
t.Fatalf("bool not converted: %#v", st.Args[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuilderSubselectSharesArguments(t *testing.T) {
|
||||
d, _ := dialect.Get("postgres")
|
||||
b, _ := NewBase(noRun{}, d, nil)
|
||||
sub := b.From(lookup.EntitySecGroupMembers).Cols(lookup.GroupMembersGroupID).Where(Eq(lookup.GroupMembersUserID, 5))
|
||||
st := b.From(lookup.EntitySecRowRules).Cols(lookup.RowRulesID).
|
||||
Where(Eq(lookup.RowRulesTableName, "t"), InSelect(lookup.RowRulesGroupID, sub), Eq(lookup.RowRulesSchemaName, "s")).Build()
|
||||
want := `SELECT t0."id" FROM "sec_row_rules" t0 WHERE t0."table_name" = $1 AND t0."group_id" IN (SELECT t1."group_id" FROM "sec_group_members" t1 WHERE t1."user_id" = $2) AND t0."schema_name" = $3`
|
||||
if st.SQL != want || len(st.Args) != 3 {
|
||||
t.Fatalf("got %s\nwant %s", st.SQL, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaRejectsUnsafeIdentifiers(t *testing.T) {
|
||||
d, _ := dialect.Get("postgres")
|
||||
bad := lookup.Schema{lookup.EntityUsers: {Name: `users"; DROP TABLE x; --`}}
|
||||
if _, err := NewBase(noRun{}, d, bad); err == nil {
|
||||
t.Fatal("unsafe table name accepted")
|
||||
}
|
||||
bad = lookup.Schema{lookup.EntityUsers: {Columns: map[string]string{"username": "a b"}}}
|
||||
if _, err := NewBase(noRun{}, d, bad); err == nil {
|
||||
t.Fatal("unsafe column name accepted")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
|
||||
_ "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/lookup/dialect"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/procedure"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T, extraDDL ...string) *sql.DB {
|
||||
t.Helper()
|
||||
db, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db.SetMaxOpenConns(1)
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
ref, err := ddl.SQL("sqlite")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, s := range append([]string{ref}, extraDDL...) {
|
||||
if _, err := db.Exec(s); err != nil {
|
||||
t.Fatalf("ddl: %v", err)
|
||||
}
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
func newTestBase(t *testing.T, db *sql.DB, schema lookup.Schema) *Base {
|
||||
t.Helper()
|
||||
d, err := dialect.Detect(db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := NewBase(procedure.NewDB(db, nil, nil), d, schema)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return b
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user