Compare commits

...

11 Commits

Author SHA1 Message Date
Hein 7c737afc5a feat(handler): implement default sort configuration for models
* Add SetDefaultSort method to configure default sort order
* Implement getDefaultSort method to retrieve configured defaults
* Update handleRead to apply default sort when none specified
* Add tests for default sort functionality
2026-07-29 10:05:24 +02:00
Hein a70e3e02d0 feat(security): complete OAuth2/OIDC spec coverage in OAuthServer
Add client_secret_basic/client_secret_post client authentication, the
client_credentials grant (RFC 6749 §4.4, backed by a synthetic
service-account user so it reuses the existing session/introspection/RLS
pipeline unchanged), RFC 9728 protected resource metadata, and OIDC
discovery + JWKS + id_token/userinfo support.

Remove plan_oauth.md, which was only meant as a working handoff doc.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-29 09:45:09 +02:00
Hein cec8eb5c0f Merge branch 'main' of https://github.com/bitechdev/ResolveSpec
Tests / Integration Tests (push) Failing after 16s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 51s
Tests / Unit Tests (push) Failing after 11s
Build , Vet Test, and Lint / Build (push) Successful in 1m23s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m35s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m35s
2026-07-28 17:48:14 +02:00
Hein 06fa3198f2 feat(security): implement OAuth2 client authentication and grants
* Add client authentication methods and client_credentials grant support
* Introduce RFC 9728 Protected Resource Metadata endpoint
* Implement OIDC discovery and id_token issuance
* Update database schema for new client fields
* Add new HTTP handlers for metadata and userinfo
2026-07-28 17:48:09 +02:00
warkanum 52d3dca1fa feat(hooks): add BeforeOp hook and fix BeforeScan row-security gap
Tests / Unit Tests (push) Failing after 11s
Tests / Integration Tests (push) Failing after 15s
Build , Vet Test, and Lint / Build (push) Successful in 1m45s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m0s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m0s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m4s
Adds a BeforeOp HookType across resolvespec, restheadspec, websocketspec,
mqttspec, and funcspec that fires before every SQL operation (read,
create, update, delete, scan/query) via a new ExecuteBeforeOp helper.

Also closes a row-level-security gap: resolvespec registered a BeforeScan
hook for ApplyRowSecurity but never fired it, and websocketspec/mqttspec
had no BeforeScan hook point at all, so row security was never applied
to their queries. BeforeScan now fires right before the actual scan in
all three, with the (possibly hook-modified) query used for execution.
2026-07-25 11:39:31 +02:00
Hein 873e8925d4 fix(handler): ensure target ID is provided for updates
Tests / Unit Tests (push) Failing after 9s
Tests / Integration Tests (push) Failing after 12s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 48s
Build , Vet Test, and Lint / Build (push) Successful in 1m18s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m27s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m28s
2026-07-24 16:35:22 +02:00
Hein b23916048a fix: ordering in before hook 2026-07-24 15:56:23 +02:00
Hein 47708fc87a fix(security): run RLS-scoping hooks in the same transaction as their queries
Tests / Unit Tests (push) Failing after 12s
Tests / Integration Tests (push) Failing after 16s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 51s
Build , Vet Test, and Lint / Build (push) Successful in 1m24s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m31s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m4s
BeforeRead/BeforeCreate hooks (used to set session-scoped RLS GUCs) were
firing against the pooled db handle while the actual queries ran as
separate calls to the same pool. Under connection pooling these could
land on different physical connections, silently bypassing row-level
security on creates and reads. handleUpdate already did this correctly;
handleRead/handleCreate in both resolvespec and restheadspec now wrap
hook execution and queries in a single RunInTransaction call.
2026-07-24 12:40:27 +02:00
Hein a85e572732 fix(hooks): reset Tx to pooled connection for post-commit hook calls
Tests / Unit Tests (push) Failing after 2m40s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m9s
Tests / Integration Tests (push) Failing after 2m47s
Build , Vet Test, and Lint / Build (push) Successful in 5m56s
Build , Vet Test, and Lint / Lint Code (push) Failing after 6m37s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 6m48s
BeforeScan (restheadspec handleUpdate) and BeforeResponse (funcspec
list/single query handlers) fire after RunInTransaction commits, but
hookCtx.Tx still pointed at the now-dead transaction. Any hook that
executed a query against Tx (e.g. setUserViaContext) failed with
"sql: transaction has already been committed or rolled back".
2026-07-20 13:45:13 +02:00
Hein 598fd687f6 feat(security): add GetUserRef method for opaque user identifiers
Tests / Integration Tests (push) Failing after 1s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m6s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m21s
Tests / Unit Tests (push) Failing after 21s
Build , Vet Test, and Lint / Build (push) Successful in 50s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 54s
2026-07-10 13:47:32 +02:00
Hein eee83f9dc6 feat(security): add query mode handling for database operations
Tests / Integration Tests (push) Failing after 1s
Tests / Unit Tests (push) Failing after 24s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 59s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 58s
Build , Vet Test, and Lint / Build (push) Successful in 4m37s
Build , Vet Test, and Lint / Lint Code (push) Failing after 5m35s
* Introduced QueryMode to select between stored procedure and direct SQL execution.
* Implemented dbCapability to probe for stored procedure existence.
* Added table names configuration for direct SQL operations.
* Updated DatabaseTwoFactorProvider to support query mode and table names.
* Implemented direct SQL methods mirroring stored procedures for TOTP operations.
* Added tests for query mode logic and table names validation.
2026-07-07 15:29:29 +02:00
53 changed files with 5779 additions and 1221 deletions
+2 -2
View File
@@ -524,9 +524,9 @@ For documentation, see [pkg/cache/README.md](pkg/cache/README.md).
#### Security
Authentication and authorization framework with hooks integration.
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.
For documentation, see [pkg/security/README.md](pkg/security/README.md).
For documentation, see [pkg/security/README.md](pkg/security/README.md) (see "Direct Mode" for the SQLite/portable-SQL path).
#### Middleware
+2 -2
View File
@@ -8,6 +8,7 @@ require (
github.com/eclipse/paho.mqtt.golang v1.5.1
github.com/getsentry/sentry-go v0.46.2
github.com/glebarez/sqlite v1.11.0
github.com/golang-jwt/jwt/v5 v5.3.1
github.com/google/uuid v1.6.0
github.com/gorilla/mux v1.8.1
github.com/gorilla/websocket v1.5.3
@@ -61,7 +62,6 @@ require (
github.com/containerd/platforms v0.2.1 // indirect
github.com/cpuguy83/dockercfg v0.3.2 // indirect
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
github.com/distribution/reference v0.6.0 // indirect
github.com/docker/docker v28.5.1+incompatible // indirect
github.com/docker/go-connections v0.6.0 // indirect
@@ -144,12 +144,12 @@ require (
go.uber.org/multierr v1.11.0 // indirect
go.yaml.in/yaml/v2 v2.4.4 // indirect
go.yaml.in/yaml/v3 v3.0.4 // indirect
golang.org/x/exp v0.0.0-20260508232706-74f9aab9d74a // indirect
golang.org/x/mod v0.36.0 // indirect
golang.org/x/net v0.54.0 // indirect
golang.org/x/sync v0.20.0 // indirect
golang.org/x/sys v0.44.0 // indirect
golang.org/x/text v0.37.0 // indirect
golang.org/x/tools v0.45.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20260519071638-aa98bba5eb94 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260519071638-aa98bba5eb94 // indirect
google.golang.org/grpc v1.81.1 // indirect
+23 -143
View File
@@ -5,43 +5,35 @@ github.com/AdaLogics/go-fuzz-headers v0.0.0-20240806141605-e8a1dd7889d6/go.mod h
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.0/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.1/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.11.1/go.mod h1:a6xsAQUZg+VsS3TJ05SRp524Hs4pZ/AeFSr5ENf0Yjo=
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.18.0 h1:Gt0j3wceWMwPmiazCa8MzMA0MfhmPIz0Qp0FJ6qcM0U=
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.18.0/go.mod h1:Ot/6aikWnKWi4l9QB7qVSwa8iMphQNqkWALMoNT3rzM=
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1 h1:jHb/wfvRikGdxMXYV3QG/SzUOPYN9KEUUuC0Yd0/vC0=
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1/go.mod h1:pzBXCYn05zvYIrwLgtK8Ap8QcjRg+0i76tMQdWN6wOk=
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.3.1/go.mod h1:uE9zaUfEQT/nbQjVi2IblCG9iaLtZsuYZ8ne+PuQ02M=
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.6.0/go.mod h1:9kIvujWAA58nmPmWB1m23fyWic1kYZMxD9CxaWn4Qpg=
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.10.1 h1:B+blDbyVIG3WaikNxPnhPiJ1MThR03b3vKGtER95TP4=
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.10.1/go.mod h1:JdM5psgjfBf5fo2uWOZhflPWyDBZ/O/CNAH9CtsuZE4=
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1 h1:Hk5QBxZQC1jb2Fwj6mpzme37xbCDdNTxU7O9eb5+LB4=
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1/go.mod h1:IYus9qsFobWIc2YVwe/WPjcnyCkPKtnHAqUYeebc8z0=
github.com/Azure/azure-sdk-for-go/sdk/internal v1.3.0/go.mod h1:okt5dMMTOFjX/aovMlrjvvXoPMBVSPzk9185BT0+eZM=
github.com/Azure/azure-sdk-for-go/sdk/internal v1.5.2/go.mod h1:yInRyqWXAuaPrgI7p70+lDDgh3mlBohis29jGMISnmc=
github.com/Azure/azure-sdk-for-go/sdk/internal v1.8.0/go.mod h1:4OG6tQ9EOP/MT0NMjDlRzWoVFxfu9rN9B2X+tlSVktg=
github.com/Azure/azure-sdk-for-go/sdk/internal v1.11.1 h1:FPKJS1T+clwv+OLGt13a8UjqeRuh0O4SJ3lUriThc+4=
github.com/Azure/azure-sdk-for-go/sdk/internal v1.11.1/go.mod h1:j2chePtV91HrC22tGoRX3sGY42uF13WzmmV80/OdVAA=
github.com/Azure/azure-sdk-for-go/sdk/internal v1.12.0 h1:fhqpLE3UEXi9lPaBRpQ6XuRW0nU7hgg4zlmZZa+a9q4=
github.com/Azure/azure-sdk-for-go/sdk/internal v1.12.0/go.mod h1:7dCRMLwisfRH3dBupKeNCioWYUZ4SS09Z14H+7i8ZoY=
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.0.1/go.mod h1:GpPjLhVR9dnUoJMyHWSPy71xY9/lcmpzIPZXmF0FCVY=
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.3.1 h1:Wgf5rZba3YZqeTNJPtvqZoBu1sBN/L4sry+u2U3Y75w=
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.3.1/go.mod h1:xxCBG/f/4Vbmh2XQJBsOmNdxWUY5j/s27jujKPbQf14=
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.4.0 h1:E4MgwLBGeVB5f2MdcIVD3ELVAWpr+WD6MUe1i+tM/PA=
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.4.0/go.mod h1:Y2b/1clN4zsAoUd/pgNAQHjLDnTis/6ROkUfyob6psM=
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.0.0/go.mod h1:bTSOgj05NGRuHHhQwAdPnYr9TOdNmKlZTgGLL6nyAdI=
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.1.1 h1:bFWuoEKg+gImo7pvkiQEFAc8ocibADgXeiLAxWhWmkI=
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.1.1/go.mod h1:Vih/3yc6yac2JzU4hzpaDupBJP0Flaia9rXXrU8xyww=
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.2.0 h1:nCYfgcSyHZXJI8J0IWE5MsCGlb2xp9fJiXyxWgmOFg4=
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.2.0/go.mod h1:ucUjca2JtSZboY8IoUqyQyuuXvwbMBVwFOm0vdQPNhA=
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1 h1:UQHMgLO+TxOElx5B5HZ4hJQsoJ/PvUvKRhJHDQXO8P8=
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E=
github.com/AzureAD/microsoft-authentication-library-for-go v1.1.1/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
github.com/AzureAD/microsoft-authentication-library-for-go v1.2.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
github.com/AzureAD/microsoft-authentication-library-for-go v1.4.2 h1:oygO0locgZJe7PpYPXT5A29ZkwJaPqcva7BVeemZOZs=
github.com/AzureAD/microsoft-authentication-library-for-go v1.4.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
github.com/AzureAD/microsoft-authentication-library-for-go v1.6.0 h1:XRzhVemXdgvJqCH0sFfrBUTnUJSBrBf7++ypk+twtRs=
github.com/AzureAD/microsoft-authentication-library-for-go v1.6.0/go.mod h1:HKpQxkWaGLJ+D/5H8QRpyQXA1eKjxkFlOMwck5+33Jk=
github.com/DATA-DOG/go-sqlmock v1.5.2 h1:OcvFkGmslmlZibjAjaHm3L//6LiuBgolP7OputlJIzU=
github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU=
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf h1:TqhNAT4zKbTdLa62d2HDBFdvgSbIGB3eJE8HqhgiL9I=
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf/go.mod h1:r5xuitiExdLAJ09PR7vBVENGvp4ZuTBeWTGtxuX3K+c=
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c h1:6Gpm9YYUEQx2T9zMsYolQhr6sjwwGtFitSA0pQsa7a8=
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c/go.mod h1:r5xuitiExdLAJ09PR7vBVENGvp4ZuTBeWTGtxuX3K+c=
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
@@ -68,14 +60,13 @@ github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ3
github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY=
github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM=
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78=
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc=
github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk=
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
github.com/dlclark/regexp2 v1.11.0 h1:G/nrcoOa7ZXlpoa/91N3X7mM3r8eIlMBBJZvsz/mxKI=
github.com/dlclark/regexp2 v1.11.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8=
github.com/dnaeon/go-vcr v1.1.0/go.mod h1:M7tiix8f0r6mKKJ3Yq/kqU1OYf3MnfmBWVbPx/yU9ko=
github.com/dnaeon/go-vcr v1.2.0/go.mod h1:R4UdLID7HZT3taECzJs4YgbbH6PIGXB6W/sc5OLb6RQ=
github.com/docker/docker v28.5.1+incompatible h1:Bm8DchhSD2J6PsFzxC35TZo4TLGR2PdW/E69rU45NhM=
@@ -94,12 +85,8 @@ github.com/felixge/httpsnoop v1.0.4 h1:NFTV2Zj1bL4mc9sqWACXbQFVBBg2W3GPvqp8/ESS2
github.com/felixge/httpsnoop v1.0.4/go.mod h1:m8KPJKqk1gH5J9DgRY2ASl2lWCfGKXixSwevea8zH2U=
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k=
github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0=
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
github.com/getsentry/sentry-go v0.40.0 h1:VTJMN9zbTvqDqPwheRVLcp0qcUcM+8eFivvGocAaSbo=
github.com/getsentry/sentry-go v0.40.0/go.mod h1:eRXCoh3uvmjQLY6qu63BjUZnaBu5L5WhMV1RwYO8W5s=
github.com/getsentry/sentry-go v0.46.2 h1:1jhYwrKGa3sIpo/y5iDNXS5wDoT7I1KNzMHrnK6ojns=
github.com/getsentry/sentry-go v0.46.2/go.mod h1:evVbw2qotNUdYG8KxXbAdjOQWWvWIwKxpjdZZIvcIPw=
github.com/glebarez/go-sqlite v1.22.0 h1:uAcMJhaA6r3LHMTFgP0SifzgXg46yJkgxqyuyec+ruQ=
@@ -115,16 +102,13 @@ github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
github.com/go-ole/go-ole v1.2.6 h1:/Fpf6oFPoeFik9ty7siob0G6Ke8QvQEuVcuChpwXzpY=
github.com/go-ole/go-ole v1.2.6/go.mod h1:pprOEPIfldk/42T2oK7lQ4v4JSDwmV0As9GaiUsvbm0=
github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs=
github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro=
github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
github.com/golang-jwt/jwt/v5 v5.0.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo=
github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0=
github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A=
@@ -137,8 +121,6 @@ github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/google/jsonschema-go v0.4.2 h1:tmrUohrwoLZZS/P3x7ex0WAVknEkBZM46iALbcqoRA8=
github.com/google/jsonschema-go v0.4.2/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE=
github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0=
github.com/google/jsonschema-go v0.4.3/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE=
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
@@ -152,8 +134,6 @@ github.com/gorilla/securecookie v1.1.1/go.mod h1:ra0sb63/xPlUeL+yeDciTfxMRAA+MP+
github.com/gorilla/sessions v1.2.1/go.mod h1:dk2InVEVJ0sfLlnXv9EAgkf6ecYs/i80K/zI+bUmuGM=
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.2 h1:8Tjv8EJ+pM1xP8mK6egEbD1OgnVTyacbefKhmbLhIhU=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.2/go.mod h1:pkJQ2tZHJ0aFOVEEot6oZmaVEZcRme73eIFmhiVuRWs=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 h1:5VipnvEpbqr2gA2VbM+nYVbkIF28c5ZQfqCBQ5g2xfk=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0/go.mod h1:Hyl3n6Twe1hvtd9XUXDec4pTvgMSEixRuQKPTMH2bNs=
github.com/hashicorp/go-uuid v1.0.2/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro=
@@ -164,8 +144,6 @@ github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsI
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
github.com/jackc/pgx/v5 v5.8.0 h1:TYPDoleBBme0xGSAX3/+NujXXtpZn9HBONkQC7IEZSo=
github.com/jackc/pgx/v5 v5.8.0/go.mod h1:QVeDInX2m9VyzvNeiCJVjCkNFqzsNb43204HshNSZKw=
github.com/jackc/pgx/v5 v5.9.2 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw=
github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
@@ -183,10 +161,10 @@ github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkr
github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
github.com/kisielk/sqlstruct v0.0.0-20201105191214-5f3e10d3ab46/go.mod h1:yyMNCyc/Ib3bDTKd379tNMpB/7/H5TjM2Y9QJ5THLbE=
github.com/klauspost/compress v1.18.2 h1:iiPHWW0YrcFgpBYhsA6D1+fqHssJscY/Tm/y2Uqnapk=
github.com/klauspost/compress v1.18.2/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4=
github.com/klauspost/compress v1.18.6 h1:2jupLlAwFm95+YDR+NwD2MEfFO9d4z4Prjl1XXDjuao=
github.com/klauspost/compress v1.18.6/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE=
github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
@@ -200,21 +178,13 @@ github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0 h1:6E+4a0GO5zZEnZ
github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0/go.mod h1:zJYVVT2jmtg6P3p1VtQj7WsuWi/y4VnjVBn7F8KPB3I=
github.com/magiconair/properties v1.8.10 h1:s31yESBquKXCV9a/ScB3ESkOjUYYv+X0rg8SYxI99mE=
github.com/magiconair/properties v1.8.10/go.mod h1:Dhd985XPs7jluiymwWYZ0G4Z61jb3vdS329zhj2hYo0=
github.com/mark3labs/mcp-go v0.46.0 h1:8KRibF4wcKejbLsHxCA/QBVUr5fQ9nwz/n8lGqmaALo=
github.com/mark3labs/mcp-go v0.46.0/go.mod h1:JKTC7R2LLVagkEWK7Kwu7DbmA6iIvnNAod6yrHiQMag=
github.com/mark3labs/mcp-go v0.54.0 h1:PZhQvd+5xrT43cUoiaKn/hDcvLUhcLc1twSEKYPTcTA=
github.com/mark3labs/mcp-go v0.54.0/go.mod h1:+8WclSK1ZUweCP3hvktSji8n8ABG/95QaEkeVE/Uwas=
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/mattn/go-isatty v0.0.22 h1:j8l17JJ9i6VGPUFUYoTUKPSgKe/83EYU2zBC7YNKMw4=
github.com/mattn/go-isatty v0.0.22/go.mod h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4=
github.com/mattn/go-sqlite3 v1.14.33 h1:A5blZ5ulQo2AtayQ9/limgHEkFreKj1Dv226a1K73s0=
github.com/mattn/go-sqlite3 v1.14.33/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y=
github.com/mattn/go-sqlite3 v1.14.44 h1:3VSe+xafpbzsLbdr2AWlAZk9yRHiBhTBakioXaCKTF8=
github.com/mattn/go-sqlite3 v1.14.44/go.mod h1:pjEuOr8IwzLJP2MfGeTb0A35jauH+C2kbHKBr7yXKVQ=
github.com/microsoft/go-mssqldb v1.8.2/go.mod h1:vp38dT33FGfVotRiTmDo3bFyaHq+p3LektQrjTULowo=
github.com/microsoft/go-mssqldb v1.9.5 h1:orwya0X/5bsL1o+KasupTkk2eNTNFkTQG0BEe/HxCn0=
github.com/microsoft/go-mssqldb v1.9.5/go.mod h1:VCP2a0KEZZtGLRHd1PsLavLFYy/3xX2yJUPycv3Sr2Q=
github.com/microsoft/go-mssqldb v1.10.0 h1:pHEt+Qz6YFPWqREq10mqSE524QQo+/QremwTCQht7TY=
github.com/microsoft/go-mssqldb v1.10.0/go.mod h1:mnG7lGa9iYJbzJqGCXyuQCegStKMr3kogDLD6+bmggg=
github.com/moby/docker-image-spec v1.3.1 h1:jMKff3w6PgbfSa69GfNg+zN/XLhfXJGnEx3Nl2EsFP0=
@@ -237,20 +207,14 @@ github.com/mochi-mqtt/server/v2 v2.7.9 h1:y0g4vrSLAag7T07l2oCzOa/+nKVLoazKEWAArw
github.com/mochi-mqtt/server/v2 v2.7.9/go.mod h1:lZD3j35AVNqJL5cezlnSkuG05c0FCHSsfAKSPBOSbqc=
github.com/modocache/gover v0.0.0-20171022184752-b58185e213c5/go.mod h1:caMODM3PzxT8aQXRPkAt8xlV/e7d7w8GM5g0fa5F0D8=
github.com/montanaflynn/stats v0.7.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
github.com/montanaflynn/stats v0.7.1 h1:etflOAAHORrCC44V+aR6Ftzort912ZU+YLiSTuV8eaE=
github.com/montanaflynn/stats v0.7.1/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
github.com/montanaflynn/stats v0.9.0 h1:tsBJ0RXwph9BmAuFoCmqGv6e8xa0MENQ8m0ptKq29mQ=
github.com/montanaflynn/stats v0.9.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A=
github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
github.com/nats-io/nats.go v1.48.0 h1:pSFyXApG+yWU/TgbKCjmm5K4wrHu86231/w84qRVR+U=
github.com/nats-io/nats.go v1.48.0/go.mod h1:iRWIPokVIFbVijxuMQq4y9ttaBTMe0SFdlZfMDd+33g=
github.com/nats-io/nats.go v1.52.0 h1:n3avV4VBsCgsdwh71TppsTwtv+QdPs7ntSKM8qJLGsc=
github.com/nats-io/nats.go v1.52.0/go.mod h1:26HypzazeOkyO3/mqd1zZd53STJN0EjCYF9Uy2ZOBno=
github.com/nats-io/nkeys v0.4.11 h1:q44qGV008kYd9W1b1nEBkNzvnWxtRSQ7A8BoqRrcfa0=
github.com/nats-io/nkeys v0.4.11/go.mod h1:szDimtgmfOi9n25JpfIdGw12tZFYXqhGxjhVxsatHVE=
github.com/nats-io/nkeys v0.4.15 h1:JACV5jRVO9V856KOapQ7x+EY8Jo3qw1vJt/9Jpwzkk4=
github.com/nats-io/nkeys v0.4.15/go.mod h1:CpMchTXC9fxA5zrMo4KpySxNjiDVvr8ANOSZdiNfUrs=
github.com/nats-io/nuid v1.0.1 h1:5iA8DT8V7q8WK2EScv2padNa/rTESc1KdnPw4TC2paw=
@@ -261,8 +225,6 @@ github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8
github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040=
github.com/opencontainers/image-spec v1.1.1/go.mod h1:qpqAh3Dmcf36wStyyWU+kCeDgrGnAve2nCC8+7h8Q0M=
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4=
github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
github.com/pelletier/go-toml/v2 v2.3.1 h1:MYEvvGnQjeNkRF1qUuGolNtNExTDwct51yp7olPtrEc=
github.com/pelletier/go-toml/v2 v2.3.1/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
github.com/pingcap/errors v0.11.4 h1:lFuQV/oaUMGcD2tqt+01ROSmJs75VG1ToEOkZIZ4nE4=
@@ -281,29 +243,20 @@ github.com/prometheus/client_golang v1.23.2 h1:Je96obch5RDVy3FDMndoUsjAhG5Edi49h
github.com/prometheus/client_golang v1.23.2/go.mod h1:Tb1a6LWHB3/SPIzCoaDXI4I8UHKeFTEQ1YCr+0Gyqmg=
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
github.com/prometheus/common v0.67.4 h1:yR3NqWO1/UyO1w2PhUvXlGQs/PtFmoveVO0KZ4+Lvsc=
github.com/prometheus/common v0.67.4/go.mod h1:gP0fq6YjjNCLssJCQp0yk4M8W6ikLURwkdd/YKtTbyI=
github.com/prometheus/common v0.67.5 h1:pIgK94WWlQt1WLwAC5j2ynLaBRDiinoAb86HZHTUGI4=
github.com/prometheus/common v0.67.5/go.mod h1:SjE/0MzDEEAyrdr5Gqc6G+sXI67maCxzaT3A2+HqjUw=
github.com/prometheus/procfs v0.19.2 h1:zUMhqEW66Ex7OXIiDkll3tl9a1ZdilUOd/F6ZXw4Vws=
github.com/prometheus/procfs v0.19.2/go.mod h1:M0aotyiemPhBCM0z5w87kL22CxfcH05ZpYlu+b4J7mw=
github.com/prometheus/procfs v0.20.1 h1:XwbrGOIplXW/AU3YhIhLODXMJYyC1isLFfYCsTEycfc=
github.com/prometheus/procfs v0.20.1/go.mod h1:o9EMBZGRyvDrSPH1RqdxhojkuXstoe4UlK79eF5TGGo=
github.com/puzpuzpuz/xsync/v3 v3.5.1 h1:GJYJZwO6IdxN/IKbneznS6yPkVC+c3zyY/j19c++5Fg=
github.com/puzpuzpuz/xsync/v3 v3.5.1/go.mod h1:VjzYrABPabuM4KyBh1Ftq6u8nhwY5tBPKP9jpmh0nnA=
github.com/redis/go-redis/v9 v9.17.2 h1:P2EGsA4qVIM3Pp+aPocCJ7DguDHhqrXNhVcEp4ViluI=
github.com/redis/go-redis/v9 v9.17.2/go.mod h1:u410H11HMLoB+TP67dz8rL9s6QW2j76l0//kSOd3370=
github.com/redis/go-redis/v9 v9.19.0 h1:XPVaaPSnG6RhYf7p+rmSa9zZfeVAnWsH5h3lxthOm/k=
github.com/redis/go-redis/v9 v9.19.0/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4=
github.com/rogpeppe/go-internal v1.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII=
github.com/rogpeppe/go-internal v1.13.1/go.mod h1:uMEvuHeurkdAXX61udpOXGD/AzZDWNMNyH2VO9fmH0o=
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
github.com/rs/xid v1.4.0 h1:qd7wPTDkN6KQx2VmMBLrpHkiyQwgFXRnkOLacUiaSNY=
github.com/rs/xid v1.4.0/go.mod h1:trrq9SKmegXys3aeAKXMUTdJsYXVwGY3RLcfgqegfbg=
github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc=
github.com/rs/xid v1.6.0 h1:fV591PaemRlL6JfRxGDEPl69wICngIQ3shQtzfy2gxU=
github.com/rs/xid v1.6.0/go.mod h1:7XoLgs4eV+QndskICGsho+ADou8ySMSjJKDIan90Nz0=
github.com/sagikazarmark/locafero v0.12.0 h1:/NQhBAkUb4+fH1jivKHWusDYFjMOOKU88eegjfxfHb4=
@@ -344,8 +297,6 @@ github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSW
github.com/testcontainers/testcontainers-go v0.40.0 h1:pSdJYLOVgLE8YdUY2FHQ1Fxu+aMnb6JfVz1mxk7OeMU=
github.com/testcontainers/testcontainers-go v0.40.0/go.mod h1:FSXV5KQtX2HAMlm7U3APNyLkkap35zNLxukw9oBi/MY=
github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
github.com/tidwall/gjson v1.18.0 h1:FIDeeyB800efLX89e5a8Y0BNH+LOngJyGrIWxG2FKQY=
github.com/tidwall/gjson v1.18.0/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
github.com/tidwall/gjson v1.19.0 h1:xwxm7n691Uf3u5OFjzngavjGTh55KX5q/9w9xHW88JU=
github.com/tidwall/gjson v1.19.0/go.mod h1:V37/opeE/JbLUOfH0QTXiNez2l0RUjYUhpT4szFQAfc=
github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
@@ -364,22 +315,10 @@ github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc h1:9lRDQMhESg+zvGYm
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc/go.mod h1:bciPuU6GHm1iF1pBvUfxfsH0Wmnc2VbpgvbI9ZWuIRs=
github.com/uptrace/bun/dialect/mssqldialect v1.2.16 h1:rKv0cKPNBviXadB/+2Y/UedA/c1JnwGzUWZkdN5FdSQ=
github.com/uptrace/bun/dialect/mssqldialect v1.2.16/go.mod h1:J5U7tGKWDsx2Q7MwDZF2417jCdpD6yD/ZMFJcCR80bk=
github.com/uptrace/bun/dialect/mssqldialect v1.2.17 h1:xEUH4WamuY9rXT9d8wHVZanhmLJCrc4s4v7frDH/PMc=
github.com/uptrace/bun/dialect/mssqldialect v1.2.17/go.mod h1:i1NRx/5cz1nivwtV7FEb/gP3CIbRTj4AQC9/Q0lNVno=
github.com/uptrace/bun/dialect/mssqldialect v1.2.18 h1:nYzHoyJKJlIyl5i95Exi8ZTK8ooKWG+o3z3f404d/yQ=
github.com/uptrace/bun/dialect/mssqldialect v1.2.18/go.mod h1:Su45Je7z66sfeZ3d1ZsnOQEK8xfzGgaMzBvtoE8yFhk=
github.com/uptrace/bun/dialect/pgdialect v1.2.16 h1:KFNZ0LxAyczKNfK/IJWMyaleO6eI9/Z5tUv3DE1NVL4=
github.com/uptrace/bun/dialect/pgdialect v1.2.16/go.mod h1:IJdMeV4sLfh0LDUZl7TIxLI0LipF1vwTK3hBC7p5qLo=
github.com/uptrace/bun/dialect/pgdialect v1.2.17 h1:DFmhOollvbYHvooxoS8ZIbiGC0wXIzstKeFUmWs+TP4=
github.com/uptrace/bun/dialect/pgdialect v1.2.17/go.mod h1:ej8ZDsvLETvyELlRDfUtIoA57sWnATv1GhOEVsuVG/k=
github.com/uptrace/bun/dialect/pgdialect v1.2.18 h1:IZ6nM2+OYrL8lkEAy7UkSEZvoa3vluTAUlZfPtlRB2k=
github.com/uptrace/bun/dialect/pgdialect v1.2.18/go.mod h1:Tqdf4QP1okrGYpXfodXvCOK6Ob1OOTwSaoAzCgBB3IU=
github.com/uptrace/bun/dialect/sqlitedialect v1.2.16 h1:6wVAiYLj1pMibRthGwy4wDLa3D5AQo32Y8rvwPd8CQ0=
github.com/uptrace/bun/dialect/sqlitedialect v1.2.16/go.mod h1:Z7+5qK8CGZkDQiPMu+LSdVuDuR1I5jcwtkB1Pi3F82E=
github.com/uptrace/bun/dialect/sqlitedialect v1.2.17 h1:ZipEoNr+wQJQleGy2poKSSoaQDavzc+nXTDp3ZzkA0E=
github.com/uptrace/bun/dialect/sqlitedialect v1.2.17/go.mod h1:phXmrxxeYqUhMU09FgazbfNxq9LlArdqjZqHc1ILy9U=
github.com/uptrace/bun/dialect/sqlitedialect v1.2.18 h1:Z33SY/U++XK9uGWqS4h8OZVxfCXguIG+sU9cYq2PGFQ=
github.com/uptrace/bun/dialect/sqlitedialect v1.2.18/go.mod h1:1MVOS/Ncy4FZbkJcgUFH6OqYoQinYNjkEwsmNQEXz2A=
github.com/uptrace/bun/driver/sqliteshim v1.2.16 h1:M6Dh5kkDWFbUWBrOsIE1g1zdZ5JbSytTD4piFRBOUAI=
github.com/uptrace/bun/driver/sqliteshim v1.2.16/go.mod h1:iKdJ06P3XS+pwKcONjSIK07bbhksH3lWsw3mpfr0+bY=
github.com/uptrace/bunrouter v1.0.23 h1:Bi7NKw3uCQkcA/GUCtDNPq5LE5UdR9pe+UyWbjHB/wU=
@@ -403,47 +342,30 @@ github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78/go.mod h1:aL8wCCfTfS
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
go.mongodb.org/mongo-driver v1.17.6 h1:87JUG1wZfWsr6rIz3ZmpH90rL5tea7O3IHuSwHUpsss=
go.mongodb.org/mongo-driver v1.17.6/go.mod h1:Hy04i7O2kC4RS06ZrhPRqj/u4DTYkFDAAccj+rVKqgQ=
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s=
go.mongodb.org/mongo-driver v1.17.9 h1:IexDdCuuNJ3BHrELgBlyaH9p60JXAvdzWR128q+U5tU=
go.mongodb.org/mongo-driver v1.17.9/go.mod h1:LlOhpH5NUEfhxcAwG0UEkMqwYcc4JU18gtCdGudk/tQ=
go.opentelemetry.io/auto/sdk v1.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA=
go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A=
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 h1:jq9TW8u3so/bN+JPT166wjOI6/vQPF6Xe7nMNIltagk=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0/go.mod h1:p8pYQP+m5XfbZm9fxtSKAbM6oIllS7s2AfxrChvc7iw=
go.opentelemetry.io/otel v1.38.0 h1:RkfdswUDRimDg0m2Az18RKOsnI8UDzppJAtj01/Ymk8=
go.opentelemetry.io/otel v1.38.0/go.mod h1:zcmtmQ1+YmQM9wrNsTGV/q/uyusom3P8RxwExxkZhjM=
go.opentelemetry.io/otel v1.43.0 h1:mYIM03dnh5zfN7HautFE4ieIig9amkNANT+xcVxAj9I=
go.opentelemetry.io/otel v1.43.0/go.mod h1:JuG+u74mvjvcm8vj8pI5XiHy1zDeoCS2LB1spIq7Ay0=
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.38.0 h1:GqRJVj7UmLjCVyVJ3ZFLdPRmhDUp2zFmQe3RHIOsw24=
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.38.0/go.mod h1:ri3aaHSmCTVYu2AWv44YMauwAQc0aqI9gHKIcSbI1pU=
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.43.0 h1:88Y4s2C8oTui1LGM6bTWkw0ICGcOLCAI5l6zsD1j20k=
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.43.0/go.mod h1:Vl1/iaggsuRlrHf/hfPJPvVag77kKyvrLeD10kpMl+A=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.38.0 h1:lwI4Dc5leUqENgGuQImwLo4WnuXFPetmPpkLi2IrX54=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.38.0/go.mod h1:Kz/oCE7z5wuyhPxsXDuaPteSWqjSBD5YaSdbxZYGbGk=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.43.0 h1:RAE+JPfvEmvy+0LzyUA25/SGawPwIUbZ6u0Wug54sLc=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.43.0/go.mod h1:AGmbycVGEsRx9mXMZ75CsOyhSP6MFIcj/6dnG+vhVjk=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.19.0 h1:IeMeyr1aBvBiPVYihXIaeIZba6b8E1bYp7lbdxK8CQg=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.19.0/go.mod h1:oVdCUtjq9MK9BlS7TtucsQwUcXcymNiEDjgDD2jMtZU=
go.opentelemetry.io/otel/metric v1.38.0 h1:Kl6lzIYGAh5M159u9NgiRkmoMKjvbsKtYRwgfrA6WpA=
go.opentelemetry.io/otel/metric v1.38.0/go.mod h1:kB5n/QoRM8YwmUahxvI3bO34eVtQf2i4utNVLr9gEmI=
go.opentelemetry.io/otel/metric v1.43.0 h1:d7638QeInOnuwOONPp4JAOGfbCEpYb+K6DVWvdxGzgM=
go.opentelemetry.io/otel/metric v1.43.0/go.mod h1:RDnPtIxvqlgO8GRW18W6Z/4P462ldprJtfxHxyKd2PY=
go.opentelemetry.io/otel/sdk v1.38.0 h1:l48sr5YbNf2hpCUj/FoGhW9yDkl+Ma+LrVl8qaM5b+E=
go.opentelemetry.io/otel/sdk v1.38.0/go.mod h1:ghmNdGlVemJI3+ZB5iDEuk4bWA3GkTpW+DOoZMYBVVg=
go.opentelemetry.io/otel/sdk v1.43.0 h1:pi5mE86i5rTeLXqoF/hhiBtUNcrAGHLKQdhg4h4V9Dg=
go.opentelemetry.io/otel/sdk v1.43.0/go.mod h1:P+IkVU3iWukmiit/Yf9AWvpyRDlUeBaRg6Y+C58QHzg=
go.opentelemetry.io/otel/sdk/metric v1.38.0 h1:aSH66iL0aZqo//xXzQLYozmWrXxyFkBJ6qT5wthqPoM=
go.opentelemetry.io/otel/sdk/metric v1.38.0/go.mod h1:dg9PBnW9XdQ1Hd6ZnRz689CbtrUp0wMMs9iPcgT9EZA=
go.opentelemetry.io/otel/sdk/metric v1.43.0 h1:S88dyqXjJkuBNLeMcVPRFXpRw2fuwdvfCGLEo89fDkw=
go.opentelemetry.io/otel/trace v1.38.0 h1:Fxk5bKrDZJUH+AMyyIXGcFAPah0oRcT+LuNtJrmcNLE=
go.opentelemetry.io/otel/trace v1.38.0/go.mod h1:j1P9ivuFsTceSWe1oY+EeW3sc+Pp42sO++GHkg4wwhs=
go.opentelemetry.io/otel/sdk/metric v1.43.0/go.mod h1:C/RJtwSEJ5hzTiUz5pXF1kILHStzb9zFlIEe85bhj6A=
go.opentelemetry.io/otel/trace v1.43.0 h1:BkNrHpup+4k4w+ZZ86CZoHHEkohws8AY+WTX09nk+3A=
go.opentelemetry.io/otel/trace v1.43.0/go.mod h1:/QJhyVBUUswCphDVxq+8mld+AvhXZLhe+8WVFxiFff0=
go.opentelemetry.io/proto/otlp v1.7.1 h1:gTOMpGDb0WTBOP8JaO72iL3auEZhVmAQg4ipjOVAtj4=
go.opentelemetry.io/proto/otlp v1.7.1/go.mod h1:b2rVh6rfI/s2pHWNlB7ILJcRALpcNDzKhACevjI+ZnE=
go.opentelemetry.io/proto/otlp v1.10.0 h1:IQRWgT5srOCYfiWnpqUYz9CVmbO8bFmKcwYxpuCSL2g=
go.opentelemetry.io/proto/otlp v1.10.0/go.mod h1:/CV4QoCR/S9yaPj8utp3lvQPoqMtxXdzn7ozvvozVqk=
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
@@ -452,12 +374,8 @@ go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0=
go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
go.uber.org/zap v1.27.1 h1:08RqriUEv8+ArZRYSTXy1LeBScaMpVSTBhCeaZYfMYc=
go.uber.org/zap v1.27.1/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E=
go.uber.org/zap v1.28.0 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo=
go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q=
go.yaml.in/yaml/v2 v2.4.3 h1:6gvOSjQoTB3vt1l+CU+tSyi/HOjfOjRLJ4YwYZGwRO0=
go.yaml.in/yaml/v2 v2.4.3/go.mod h1:zSxWcmIDjOzPXpjlTTbAsKokqkDNAVtZO0WOMiT90s8=
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
@@ -474,22 +392,14 @@ golang.org/x/crypto v0.21.0/go.mod h1:0BP7YvVV9gBbVKyeTG0Gyn+gZm94bibOW5BjDEYAOM
golang.org/x/crypto v0.22.0/go.mod h1:vr6Su+7cTlO45qkww3VDJlzDn0ctJvRgYbC2NvXHt+M=
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
golang.org/x/crypto v0.24.0/go.mod h1:Z1PMYSOR5nyMcyAVAIQSKCDwalqy85Aqn1x3Ws4L5DM=
golang.org/x/crypto v0.46.0 h1:cKRW/pmt1pKAfetfu+RCEvjvZkA9RimPbh7bhFjGVBU=
golang.org/x/crypto v0.46.0/go.mod h1:Evb/oLKmMraqjZ2iQTwDwvCtJkczlDuTmdJXoZVzqU0=
golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI=
golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8=
golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93 h1:fQsdNF2N+/YewlRZiricy4P1iimyPKZ/xwniHj8Q2a0=
golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93/go.mod h1:EPRbTFwzwjXj9NpYyyrvenVh9Y+GFeEvMNh7Xuz7xgU=
golang.org/x/exp v0.0.0-20260508232706-74f9aab9d74a h1:+3jdDGGB8NGb1Zktc737jlt3/A5f6UlwSzmvqUuufxw=
golang.org/x/exp v0.0.0-20260508232706-74f9aab9d74a/go.mod h1:d2fgXJLVs4dYDHUk5lwMIfzRzSrWCfGZb0ZqeLa/Vcw=
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
golang.org/x/mod v0.9.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
golang.org/x/mod v0.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
golang.org/x/mod v0.31.0/go.mod h1:43JraMp9cGx1Rx3AqioxrbrhNsLl2l/iNAvuBkrezpg=
golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4=
golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
@@ -509,12 +419,8 @@ golang.org/x/net v0.22.0/go.mod h1:JKghWKKOSdJwpW2GEx0Ja7fmaKnMsbu+MWVZTokSYmg=
golang.org/x/net v0.24.0/go.mod h1:2Q7sJY5mzlzWjKtYUEXSlBWCdyaioyXzRB2RtU8KVE8=
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
golang.org/x/net v0.26.0/go.mod h1:5YKkiSynbBIh3p6iOc/vibscux0x38BZDkn8sCUPxHE=
golang.org/x/net v0.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU=
golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY=
golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w=
golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ=
golang.org/x/oauth2 v0.34.0 h1:hqK/t4AKgbqWkdkcAeI8XLmbK+4m4G5YeQRrmiotGlw=
golang.org/x/oauth2 v0.34.0/go.mod h1:lzm5WQJQwKZ3nwavOZ3IS5Aulzxi68dUSgRHujetwEA=
golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
@@ -524,8 +430,6 @@ golang.org/x/sync v0.3.0/go.mod h1:FU7BRWz2tNW+3quACPkgCx/L+uEAv1htQ0V83Z9Rj+Y=
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.9.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
@@ -551,8 +455,6 @@ golang.org/x/sys v0.18.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/sys v0.19.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/sys v0.21.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk=
golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ=
golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
@@ -570,9 +472,8 @@ golang.org/x/term v0.18.0/go.mod h1:ILwASektA3OnRv7amZ1xhE/KTR+u50pbXfZ03+6Nx58=
golang.org/x/term v0.19.0/go.mod h1:2CuTdWZ7KHSQwUzKva0cbMg6q2DMI3Mmxp+gKJbskEk=
golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
golang.org/x/term v0.21.0/go.mod h1:ooXLefLobQVslOqselCNF4SxFAaoS6KujMbsGzSDmX0=
golang.org/x/term v0.38.0 h1:PQ5pkm/rLO6HnxFR7N2lJHOZX6Kez5Y1gDSJla6jo7Q=
golang.org/x/term v0.38.0/go.mod h1:bSEAKrOT1W+VSu9TSCMtoGEOUcKxOKgl3LE5QEF/xVg=
golang.org/x/term v0.43.0 h1:S4RLU2sB31O/NCl+zFN9Aru9A/Cq2aqKpTZJ6B+DwT4=
golang.org/x/term v0.43.0/go.mod h1:lrhlHNdQJHO+1qVYiHfFKVuVioJIheAc3fBSMFYEIsk=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
@@ -587,12 +488,8 @@ golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
golang.org/x/text v0.16.0/go.mod h1:GhwF1Be+LQoKShO3cGOHzqOgRrGaYc9AvblQOmPVHnI=
golang.org/x/text v0.20.0/go.mod h1:D4IsuqiFMhST5bX19pQ9ikHC2GsaKyk/oF+pn3ducp4=
golang.org/x/text v0.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU=
golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY=
golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc=
golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38=
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
@@ -601,24 +498,16 @@ golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc
golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58=
golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
golang.org/x/tools v0.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA=
golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc=
golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8=
golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk=
gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E=
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
google.golang.org/genproto/googleapis/api v0.0.0-20250825161204-c5933d9347a5 h1:BIRfGDEjiHRrk0QKZe3Xv2ieMhtgRGeLcZQ0mIVn4EY=
google.golang.org/genproto/googleapis/api v0.0.0-20250825161204-c5933d9347a5/go.mod h1:j3QtIyytwqGr1JUDtYXwtMXWPKsEa5LtzIFN1Wn5WvE=
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
google.golang.org/genproto/googleapis/api v0.0.0-20260519071638-aa98bba5eb94 h1:DddG61lE5LkX6144z22i0gma9BMBs5aZ9B8lZLobxyw=
google.golang.org/genproto/googleapis/api v0.0.0-20260519071638-aa98bba5eb94/go.mod h1:1dCETSCY2YKZNXQE3h4fun3TYwF5p8jejRKZgfWAgAY=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250825161204-c5933d9347a5 h1:eaY8u2EuxbRv7c3NiGK0/NedzVsCcV6hDuU5qPX5EGE=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250825161204-c5933d9347a5/go.mod h1:M4/wBTSeyLxupu3W3tJtOgB14jILAS/XWPSSa3TAlJc=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260519071638-aa98bba5eb94 h1:eZCjr/aAF8c5ccm5pb6T4EXgIei5MlAAPWPJk+5ArfY=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260519071638-aa98bba5eb94/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.75.0 h1:+TW+dqTd2Biwe6KKfhE5JpiYIBWq865PhKGSXiivqt4=
google.golang.org/grpc v1.75.0/go.mod h1:JtPAzKiq4v1xcAB2hydNlWI2RnF85XXcV0mhKXr2ecQ=
google.golang.org/grpc v1.81.1 h1:VnnIIZ88UzOOKLukQi+ImGz8O1Wdp8nAGGnvOfEIWQQ=
google.golang.org/grpc v1.81.1/go.mod h1:xGH9GfzOyMTGIOXBJmXt+BX/V0kcdQbdcuwQ/zNw42I=
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
@@ -644,37 +533,28 @@ gorm.io/gorm v1.31.1 h1:7CA8FTFz/gRfgqgpeKIBcervUn3xSyPUmr6B2WXJ7kg=
gorm.io/gorm v1.31.1/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs=
gotest.tools/v3 v3.5.2 h1:7koQfIKdy+I8UTetycgUqXWSDwpgv193Ka+qRsmBY8Q=
gotest.tools/v3 v3.5.2/go.mod h1:LtdLGcnqToBH83WByAAi/wiwSFCArdFIUV/xxN4pcjA=
modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis=
modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
modernc.org/cc/v4 v4.28.2 h1:3tQ0lf2ADtoby2EtSP+J7IE2SHwEJdP8ioR59wx7XpY=
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
modernc.org/ccgo/v4 v4.30.1/go.mod h1:bIOeI1JL54Utlxn+LwrFyjCx2n2RDiYEaJVSrgdrRfM=
modernc.org/cc/v4 v4.28.2/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI=
modernc.org/ccgo/v4 v4.34.0 h1:yRLPFZieg532OT4rp4JFNIVcquwalMX26G95WQDqwCQ=
modernc.org/fileutil v1.3.40 h1:ZGMswMNc9JOCrcrakF1HrvmergNLAmxOPjizirpfqBA=
modernc.org/fileutil v1.3.40/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc=
modernc.org/ccgo/v4 v4.34.0/go.mod h1:AS5WYMyBakQ+fhsHhtP8mWB82KTGPkNNJDGfGQCe0/A=
modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM=
modernc.org/fileutil v1.4.0/go.mod h1:EqdKFDxiByqxLk8ozOxObDSfcVOv/54xDs/DUHdvCUU=
modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI=
modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
modernc.org/gc/v3 v3.1.1 h1:k8T3gkXWY9sEiytKhcgyiZ2L0DTyCQ/nvX+LoCljoRE=
modernc.org/gc/v3 v3.1.1/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY=
modernc.org/gc/v3 v3.1.2 h1:ZtDCnhonXSZexk/AYsegNRV1lJGgaNZJuKjJSWKyEqo=
modernc.org/gc/v3 v3.1.2/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY=
modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks=
modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI=
modernc.org/libc v1.67.4 h1:zZGmCMUVPORtKv95c2ReQN5VDjvkoRm9GWPTEPuvlWg=
modernc.org/libc v1.67.4/go.mod h1:QvvnnJ5P7aitu0ReNpVIEyesuhmDLQ8kaEoyMjIFZJA=
modernc.org/libc v1.72.3 h1:ZnDF4tXn4NBXFutMMQC4vtbTFSXhhKzR73fv0beZEAU=
modernc.org/libc v1.72.3/go.mod h1:dn0dZNnnn1clLyvRxLxYExxiKRZIRENOfqQ8XEeg4Qs=
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8=
modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg=
modernc.org/opt v0.2.0/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=
modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
modernc.org/sqlite v1.42.2 h1:7hkZUNJvJFN2PgfUdjni9Kbvd4ef4mNLOu0B9FGxM74=
modernc.org/sqlite v1.42.2/go.mod h1:+VkC6v3pLOAE0A0uVucQEcbVW0I5nHCeDaBf+DpsQT8=
modernc.org/sqlite v1.50.1 h1:l+cQvn0sd0zJJtfygGHuQJ5AjlrwXmWPw4KP3ZMwr9w=
modernc.org/sqlite v1.50.1/go.mod h1:tcNzv5p84E0skkmJn038y+hWJbLQXQqEnQfeh5r2JLM=
modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
+24
View File
@@ -90,6 +90,30 @@ type SortOption struct {
Direction string `json:"direction"`
}
// PrimaryKeySortColumn is a sentinel SortOption.Column value that gets
// resolved to a model's actual primary key column name at query time.
// This lets a single global default sort (e.g. set once via
// Handler.SetDefaultSort) work across models with different primary keys,
// e.g. common.SortOption{Column: common.PrimaryKeySortColumn, Direction: "asc"}.
const PrimaryKeySortColumn = "$pk"
// ResolveSortColumns returns a copy of sort with any PrimaryKeySortColumn
// entries replaced by pkName. If pkName is empty, matching entries are
// dropped since there is no column to sort by.
func ResolveSortColumns(sort []SortOption, pkName string) []SortOption {
resolved := make([]SortOption, 0, len(sort))
for _, s := range sort {
if s.Column == PrimaryKeySortColumn {
if pkName == "" {
continue
}
s.Column = pkName
}
resolved = append(resolved, s)
}
return resolved
}
type CustomOperator struct {
Name string `json:"name"`
SQL string `json:"sql"`
+46
View File
@@ -0,0 +1,46 @@
package common
import (
"reflect"
"testing"
)
func TestResolveSortColumns(t *testing.T) {
sort := []SortOption{
{Column: PrimaryKeySortColumn, Direction: "asc"},
{Column: "name", Direction: "desc"},
}
got := ResolveSortColumns(sort, "user_id")
want := []SortOption{
{Column: "user_id", Direction: "asc"},
{Column: "name", Direction: "desc"},
}
if !reflect.DeepEqual(got, want) {
t.Errorf("ResolveSortColumns() = %v, want %v", got, want)
}
// Original slice must not be mutated
if sort[0].Column != PrimaryKeySortColumn {
t.Errorf("ResolveSortColumns() mutated input slice: %v", sort)
}
}
func TestResolveSortColumnsEmptyPK(t *testing.T) {
sort := []SortOption{
{Column: PrimaryKeySortColumn, Direction: "asc"},
{Column: "name", Direction: "desc"},
}
got := ResolveSortColumns(sort, "")
want := []SortOption{{Column: "name", Direction: "desc"}}
if !reflect.DeepEqual(got, want) {
t.Errorf("ResolveSortColumns() with empty pk = %v, want %v", got, want)
}
}
func TestResolveSortColumnsNil(t *testing.T) {
if got := ResolveSortColumns(nil, "id"); len(got) != 0 {
t.Errorf("ResolveSortColumns(nil) = %v, want empty", got)
}
}
+12 -6
View File
@@ -197,7 +197,7 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
hookCtx.Tx = tx
// Execute BeforeQueryList hook (inside transaction)
if err := h.hooks.Execute(BeforeQueryList, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeQueryList, hookCtx); err != nil {
logger.Error("BeforeQueryList hook failed: %v", err)
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
return err
@@ -261,7 +261,7 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
// Execute BeforeSQLExec hook
hookCtx.SQLQuery = sqlquery
if err := h.hooks.Execute(BeforeSQLExec, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
logger.Error("BeforeSQLExec hook failed: %v", err)
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
return err
@@ -331,7 +331,10 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
w.Header().Set("Content-Range", fmt.Sprintf("items %d-%d/%d", respOffset, respOffset+len(dbobjlist), total))
logger.Info("Serving: Records %d of %d", len(dbobjlist), total)
// Execute BeforeResponse hook
// Execute BeforeResponse hook. The transaction has already committed by
// this point, so hooks must use the pooled connection rather than the
// now-dead tx.
hookCtx.Tx = h.db
hookCtx.Result = dbobjlist
hookCtx.Total = total
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
@@ -560,7 +563,7 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
hookCtx.Tx = tx
// Execute BeforeQuery hook (inside transaction)
if err := h.hooks.Execute(BeforeQuery, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeQuery, hookCtx); err != nil {
logger.Error("BeforeQuery hook failed: %v", err)
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
return err
@@ -579,7 +582,7 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
sqlquery = hookCtx.SQLQuery
// Execute BeforeSQLExec hook
if err := h.hooks.Execute(BeforeSQLExec, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
logger.Error("BeforeSQLExec hook failed: %v", err)
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
return err
@@ -631,7 +634,10 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
return
}
// Execute BeforeResponse hook
// Execute BeforeResponse hook. The transaction has already committed by
// this point, so hooks must use the pooled connection rather than the
// now-dead tx.
hookCtx.Tx = h.db
hookCtx.Result = dbobj
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
logger.Error("BeforeResponse hook failed: %v", err)
+14
View File
@@ -28,6 +28,10 @@ const (
// Response hooks (before response is sent)
BeforeResponse HookType = "before_response"
// BeforeOp fires immediately before every SQL operation (query, query list, SQL exec).
// It fires at each individual SQL-operation hook point, so it runs once per statement executed.
BeforeOp HookType = "before_op"
)
// HookContext contains all the data available to a hook
@@ -130,6 +134,16 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
return nil
}
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
// (BeforeQuery, BeforeQueryList, or BeforeSQLExec). BeforeOp always runs first so it can
// observe/veto every SQL operation regardless of type.
func (r *HookRegistry) ExecuteBeforeOp(hookType HookType, ctx *HookContext) error {
if err := r.Execute(BeforeOp, ctx); err != nil {
return err
}
return r.Execute(hookType, ctx)
}
// Clear removes all hooks for the specified type
func (r *HookRegistry) Clear(hookType HookType) {
delete(r.hooks, hookType)
+10
View File
@@ -71,6 +71,16 @@ func (f *funcSpecSecurityContext) GetUserID() (int, bool) {
return int(f.ctx.UserContext.UserID), true
}
// GetUserRef returns an opaque user identifier for row security lookups.
// It returns the full *security.UserContext so providers can read JWT claims
// (e.g. a UUID subject) instead of relying on the int user ID.
func (f *funcSpecSecurityContext) GetUserRef() (any, bool) {
if f.ctx.UserContext == nil {
return nil, false
}
return f.ctx.UserContext, true
}
func (f *funcSpecSecurityContext) GetSchema() string {
// funcspec doesn't have a schema concept, extract from SQL query or use default
return "public"
+28 -4
View File
@@ -313,7 +313,7 @@ func (h *Handler) handleRequest(client *Client, msg *Message) {
// handleRead processes a read operation
func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext) {
// Execute before hook
if err := h.hooks.Execute(BeforeRead, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil {
logger.Error("[MQTTSpec] BeforeRead hook failed: %v", err)
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
return
@@ -356,7 +356,7 @@ func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext)
// handleCreate processes a create operation
func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContext) {
// Execute before hook
if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
logger.Error("[MQTTSpec] BeforeCreate hook failed: %v", err)
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
return
@@ -390,7 +390,7 @@ func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContex
// handleUpdate processes an update operation
func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContext) {
// Execute before hook
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil {
logger.Error("[MQTTSpec] BeforeUpdate hook failed: %v", err)
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
return
@@ -424,7 +424,7 @@ func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContex
// handleDelete processes a delete operation
func (h *Handler) handleDelete(client *Client, msg *Message, hookCtx *HookContext) {
// Execute before hook
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil {
logger.Error("[MQTTSpec] BeforeDelete hook failed: %v", err)
h.sendError(client.ID, msg.ID, "hook_error", err.Error())
return
@@ -686,6 +686,18 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) {
}
}
// Execute BeforeScan hooks - pass query so hooks (e.g. row-level security) can modify it
if hookCtx.Metadata == nil {
hookCtx.Metadata = make(map[string]interface{})
}
hookCtx.Metadata["query"] = query
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
return nil, fmt.Errorf("BeforeScan hook failed: %w", err)
}
if modifiedQuery, ok := hookCtx.Metadata["query"].(common.SelectQuery); ok {
query = modifiedQuery
}
// Execute query
if err := query.ScanModel(hookCtx.Context); err != nil {
return nil, fmt.Errorf("failed to read record: %w", err)
@@ -738,6 +750,18 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
}
}
// Execute BeforeScan hooks - pass query so hooks (e.g. row-level security) can modify it
if hookCtx.Metadata == nil {
hookCtx.Metadata = make(map[string]interface{})
}
hookCtx.Metadata["query"] = query
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
return nil, nil, fmt.Errorf("BeforeScan hook failed: %w", err)
}
if modifiedQuery, ok := hookCtx.Metadata["query"].(common.SelectQuery); ok {
query = modifiedQuery
}
// Execute query
if err := query.ScanModel(hookCtx.Context); err != nil {
return nil, nil, fmt.Errorf("failed to read records: %w", err)
+4
View File
@@ -34,6 +34,7 @@ const (
AfterUpdate = websocketspec.AfterUpdate
BeforeDelete = websocketspec.BeforeDelete
AfterDelete = websocketspec.AfterDelete
BeforeScan = websocketspec.BeforeScan
// Subscription hooks
BeforeSubscribe = websocketspec.BeforeSubscribe
@@ -46,6 +47,9 @@ const (
AfterConnect = websocketspec.AfterConnect
BeforeDisconnect = websocketspec.BeforeDisconnect
AfterDisconnect = websocketspec.AfterDisconnect
// BeforeOp fires immediately before every SQL operation (read, create, update, delete)
BeforeOp = websocketspec.BeforeOp
)
// NewHookRegistry creates a new hook registry
+21 -4
View File
@@ -27,25 +27,31 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
return security.LoadSecurityRules(secCtx, securityList)
})
// Hook 2: AfterRead - Apply column-level security (masking)
// Hook 2: BeforeScan - Apply row-level security filters
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
return security.ApplyRowSecurity(secCtx, securityList)
})
// Hook 3: AfterRead - Apply column-level security (masking)
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
return security.ApplyColumnSecurity(secCtx, securityList)
})
// Hook 3 (Optional): Audit logging
// Hook 4 (Optional): Audit logging
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
return security.LogDataAccess(secCtx)
})
// Hook 4: BeforeUpdate - enforce CanUpdate rule from context/registry
// Hook 5: BeforeUpdate - enforce CanUpdate rule from context/registry
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
return security.CheckModelUpdateAllowed(secCtx)
})
// Hook 5: BeforeDelete - enforce CanDelete rule from context/registry
// Hook 6: BeforeDelete - enforce CanDelete rule from context/registry
handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
return security.CheckModelDeleteAllowed(secCtx)
@@ -71,6 +77,17 @@ func (s *securityContext) GetUserID() (int, bool) {
return security.GetUserID(s.ctx.Context)
}
// GetUserRef returns an opaque user identifier for row security lookups.
// It prefers the full *security.UserContext (so providers can read JWT claims,
// e.g. a UUID subject) and falls back to the int user ID.
func (s *securityContext) GetUserRef() (any, bool) {
if userCtx, ok := security.GetUserContext(s.ctx.Context); ok {
return userCtx, true
}
userID, ok := security.GetUserID(s.ctx.Context)
return userID, ok
}
func (s *securityContext) GetSchema() string {
return s.ctx.Schema
}
+11
View File
@@ -84,6 +84,17 @@ func (s *securityContext) GetUserID() (int, bool) {
return security.GetUserID(s.ctx.Context)
}
// GetUserRef returns an opaque user identifier for row security lookups.
// It prefers the full *security.UserContext (so providers can read JWT claims,
// e.g. a UUID subject) and falls back to the int user ID.
func (s *securityContext) GetUserRef() (any, bool) {
if userCtx, ok := security.GetUserContext(s.ctx.Context); ok {
return userCtx, true
}
userID, ok := security.GetUserID(s.ctx.Context)
return userID, ok
}
func (s *securityContext) GetSchema() string {
return s.ctx.Schema
}
+550 -351
View File
File diff suppressed because it is too large Load Diff
+30
View File
@@ -41,6 +41,36 @@ func TestSetFallbackHandler(t *testing.T) {
}
}
func TestSetDefaultSort(t *testing.T) {
handler := NewHandler(nil, nil)
// No default configured yet
if got := handler.getDefaultSort("public", "users"); got != nil {
t.Errorf("Expected no default sort, got %v", got)
}
// Global default
global := []common.SortOption{{Column: "created_at", Direction: "desc"}}
handler.SetDefaultSort("", "", global...)
if got := handler.getDefaultSort("public", "users"); !reflect.DeepEqual(got, global) {
t.Errorf("Expected global default sort %v, got %v", global, got)
}
// Per-model default overrides the global default
perModel := []common.SortOption{{Column: "name", Direction: "asc"}}
handler.SetDefaultSort("public", "users", perModel...)
if got := handler.getDefaultSort("public", "users"); !reflect.DeepEqual(got, perModel) {
t.Errorf("Expected per-model default sort %v, got %v", perModel, got)
}
// Other models still fall back to the global default
if got := handler.getDefaultSort("public", "orders"); !reflect.DeepEqual(got, global) {
t.Errorf("Expected global default sort %v for unrelated model, got %v", global, got)
}
}
func TestGetDatabase(t *testing.T) {
handler := NewHandler(nil, nil)
db := handler.GetDatabase()
+15
View File
@@ -34,6 +34,11 @@ const (
// Scan/Execute operation hooks (for query building)
BeforeScan HookType = "before_scan"
// BeforeOp fires immediately before every SQL operation (read, create, update, delete, scan).
// Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each
// individual SQL-operation hook point, so it runs once per statement executed.
BeforeOp HookType = "before_op"
)
// HookContext contains all the data available to a hook
@@ -128,6 +133,16 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
return nil
}
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
// (BeforeRead, BeforeCreate, BeforeUpdate, BeforeDelete, or BeforeScan). BeforeOp always
// runs first so it can observe/veto every SQL operation regardless of type.
func (r *HookRegistry) ExecuteBeforeOp(hookType HookType, ctx *HookContext) error {
if err := r.Execute(BeforeOp, ctx); err != nil {
return err
}
return r.Execute(hookType, ctx)
}
// Clear removes all hooks for the specified type
func (r *HookRegistry) Clear(hookType HookType) {
delete(r.hooks, hookType)
+11
View File
@@ -78,6 +78,17 @@ func (s *securityContext) GetUserID() (int, bool) {
return security.GetUserID(s.ctx.Context)
}
// GetUserRef returns an opaque user identifier for row security lookups.
// It prefers the full *security.UserContext (so providers can read JWT claims,
// e.g. a UUID subject) and falls back to the int user ID.
func (s *securityContext) GetUserRef() (any, bool) {
if userCtx, ok := security.GetUserContext(s.ctx.Context); ok {
return userCtx, true
}
userID, ok := security.GetUserID(s.ctx.Context)
return userID, ok
}
func (s *securityContext) GetSchema() string {
return s.ctx.Schema
}
+38
View File
@@ -0,0 +1,38 @@
package restheadspec
import (
"reflect"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/common"
)
func TestSetDefaultSort(t *testing.T) {
handler := NewHandler(nil, nil)
// No default configured yet
if got := handler.getDefaultSort("public", "users"); got != nil {
t.Errorf("Expected no default sort, got %v", got)
}
// Global default
global := []common.SortOption{{Column: "created_at", Direction: "desc"}}
handler.SetDefaultSort("", "", global...)
if got := handler.getDefaultSort("public", "users"); !reflect.DeepEqual(got, global) {
t.Errorf("Expected global default sort %v, got %v", global, got)
}
// Per-model default overrides the global default
perModel := []common.SortOption{{Column: "name", Direction: "asc"}}
handler.SetDefaultSort("public", "users", perModel...)
if got := handler.getDefaultSort("public", "users"); !reflect.DeepEqual(got, perModel) {
t.Errorf("Expected per-model default sort %v, got %v", perModel, got)
}
// Other models still fall back to the global default
if got := handler.getDefaultSort("public", "orders"); !reflect.DeepEqual(got, global) {
t.Errorf("Expected global default sort %v for unrelated model, got %v", global, got)
}
}
File diff suppressed because it is too large Load Diff
+15
View File
@@ -34,6 +34,11 @@ const (
// Scan/Execute operation hooks
BeforeScan HookType = "before_scan"
// BeforeOp fires immediately before every SQL operation (read, create, update, delete, scan).
// Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each
// individual SQL-operation hook point, so it runs once per statement executed.
BeforeOp HookType = "before_op"
)
// HookContext contains all the data available to a hook
@@ -137,6 +142,16 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
return nil
}
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
// (BeforeRead, BeforeCreate, BeforeUpdate, BeforeDelete, or BeforeScan). BeforeOp always
// runs first so it can observe/veto every SQL operation regardless of type.
func (r *HookRegistry) ExecuteBeforeOp(hookType HookType, ctx *HookContext) error {
if err := r.Execute(BeforeOp, ctx); err != nil {
return err
}
return r.Execute(hookType, ctx)
}
// Clear removes all hooks for the specified type
func (r *HookRegistry) Clear(hookType HookType) {
delete(r.hooks, hookType)
+11
View File
@@ -77,6 +77,17 @@ func (s *securityContext) GetUserID() (int, bool) {
return security.GetUserID(s.ctx.Context)
}
// GetUserRef returns an opaque user identifier for row security lookups.
// It prefers the full *security.UserContext (so providers can read JWT claims,
// e.g. a UUID subject) and falls back to the int user ID.
func (s *securityContext) GetUserRef() (any, bool) {
if userCtx, ok := security.GetUserContext(s.ctx.Context); ok {
return userCtx, true
}
userID, ok := security.GetUserID(s.ctx.Context)
return userID, ok
}
func (s *securityContext) GetSchema() string {
return s.ctx.Schema
}
+90 -1
View File
@@ -11,7 +11,8 @@ Type-safe, composable security system for ResolveSpec with support for authentic
-**No Global State** - Each handler has its own security configuration
-**Testable** - Easy to mock and test
-**Extensible** - Implement custom providers for your needs
-**Stored Procedures** - All database operations use PostgreSQL stored procedures for security and maintainability
-**Stored Procedures** - Database operations use PostgreSQL stored procedures where available, for security and maintainability
-**Direct Mode** - Portable Go/SQL fallback for SQLite, MySQL, or Postgres without the stored procedures installed — no code changes required
-**OAuth2 Authorization Server** - Built-in OAuth 2.1 + PKCE server (RFC 8414, 7591, 7009, 7662) with login form and external provider federation
-**Password Reset** - Self-service password reset with secure token generation and session invalidation
@@ -51,6 +52,94 @@ Type-safe, composable security system for ResolveSpec with support for authentic
See `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.
## Direct Mode (portable SQL without stored procedures)
Every database-backed provider (`DatabaseAuthenticator`, `JWTAuthenticator`, `DatabaseTwoFactorProvider`, `DatabasePasskeyProvider`, the OAuth2 methods/server, `DatabaseKeyStore`) has two code paths:
- **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.
### QueryMode
Selection is controlled per-provider by a `QueryMode`:
```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
}
```
`DefaultTableNames()` / `MergeTableNames()` / `ValidateTableNames()` mirror `DefaultSQLNames()` / `MergeSQLNames()` / `ValidateSQLNames()`. Set custom names via the same `Options`/`With...` surface as `QueryMode`:
```go
auth := security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{
TableNames: &security.TableNames{Users: "app_users"}, // only override what differs
})
```
`oauth2_methods.go` and `oauth_server_db.go` are methods on `*DatabaseAuthenticator` and reuse its `TableNames`/`QueryMode`; there's no separate config for them.
### Schema
`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.
### What's NOT covered
`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.
### 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.
## Quick Start
```go
+2 -2
View File
@@ -74,8 +74,8 @@ func (c *CompositeSecurityProvider) GetColumnSecurity(ctx context.Context, userI
}
// GetRowSecurity delegates to the row security provider
func (c *CompositeSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
return c.rowSec.GetRowSecurity(ctx, userID, schema, table)
func (c *CompositeSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
return c.rowSec.GetRowSecurity(ctx, userRef, schema, table)
}
// Optional interface implementations (if wrapped providers support them)
+1 -1
View File
@@ -79,7 +79,7 @@ type mockRowSec struct {
supportsCache bool
}
func (m *mockRowSec) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
func (m *mockRowSec) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
return m.rowSec, m.err
}
+6 -2
View File
@@ -1597,6 +1597,8 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
client_name VARCHAR(255),
grant_types TEXT[] DEFAULT ARRAY['authorization_code'],
allowed_scopes TEXT[] DEFAULT ARRAY['openid','profile','email'],
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,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
@@ -1634,13 +1636,15 @@ DECLARE
BEGIN
v_client_id := p_data->>'client_id';
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes)
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method)
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'])
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')
)
RETURNING to_jsonb(oauth_clients.*) INTO v_row;
+143
View File
@@ -0,0 +1,143 @@
-- 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),
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);
+467
View File
@@ -0,0 +1,467 @@
package security
import (
"context"
"database/sql"
"encoding/base64"
"net/http"
"os"
"path/filepath"
"testing"
"time"
_ "github.com/mattn/go-sqlite3"
)
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
// 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")
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)
if err != nil {
t.Fatalf("failed to read schema: %v", err)
}
if _, err := db.Exec(string(schema)); err != nil {
t.Fatalf("failed to apply schema: %v", err)
}
return db
}
func authenticatedRequest(token string) *http.Request {
req, _ := http.NewRequest(http.MethodGet, "/", nil)
req.Header.Set("Authorization", "Bearer "+token)
return req
}
func TestDirectMode_RegisterThenLogin(t *testing.T) {
db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
ctx := context.Background()
regResp, err := auth.Register(ctx, RegisterRequest{
Username: "alice",
Password: "hunter2",
Email: "alice@example.com",
Roles: []string{"user", "admin"},
Claims: map[string]any{"ip_address": "127.0.0.1", "user_agent": "test-agent"},
})
if err != nil {
t.Fatalf("Register() error = %v", err)
}
if regResp.Token == "" || regResp.User == nil {
t.Fatalf("Register() returned incomplete response: %+v", regResp)
}
if len(regResp.User.Roles) != 2 {
t.Errorf("expected 2 roles, got %v", regResp.User.Roles)
}
loginResp, err := auth.Login(ctx, LoginRequest{Username: "alice", Password: "hunter2"})
if err != nil {
t.Fatalf("Login() error = %v", err)
}
if loginResp.User.UserName != "alice" {
t.Errorf("Login() user = %q, want alice", loginResp.User.UserName)
}
// Duplicate registration should fail.
if _, err := auth.Register(ctx, RegisterRequest{Username: "alice", Password: "x", Email: "other@example.com"}); err == nil {
t.Error("expected duplicate username registration to fail")
}
}
func TestDirectMode_SessionLifecycle(t *testing.T) {
db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
ctx := context.Background()
loginResp, err := auth.Register(ctx, RegisterRequest{Username: "bob", Password: "p", Email: "bob@example.com"})
if err != nil {
t.Fatalf("Register() error = %v", err)
}
userCtx, err := auth.Authenticate(authenticatedRequest(loginResp.Token))
if err != nil {
t.Fatalf("Authenticate() error = %v", err)
}
if userCtx.UserName != "bob" {
t.Errorf("Authenticate() user = %q, want bob", userCtx.UserName)
}
refreshResp, err := auth.RefreshToken(ctx, loginResp.Token)
if err != nil {
t.Fatalf("RefreshToken() error = %v", err)
}
if refreshResp.Token == "" || refreshResp.Token == loginResp.Token {
t.Errorf("RefreshToken() should return a new token, got %q", refreshResp.Token)
}
if err := auth.Logout(ctx, LogoutRequest{Token: refreshResp.Token, UserID: refreshResp.User.UserID}); err != nil {
t.Fatalf("Logout() error = %v", err)
}
if _, err := auth.RefreshToken(ctx, refreshResp.Token); err == nil {
t.Error("expected RefreshToken() after logout to fail")
}
}
func TestDirectMode_PasswordReset(t *testing.T) {
db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
ctx := context.Background()
if _, err := auth.Register(ctx, RegisterRequest{Username: "carol", Password: "old", Email: "carol@example.com"}); err != nil {
t.Fatalf("Register() error = %v", err)
}
resetResp, err := auth.RequestPasswordReset(ctx, PasswordResetRequest{Email: "carol@example.com"})
if err != nil {
t.Fatalf("RequestPasswordReset() error = %v", err)
}
if resetResp.Token == "" {
t.Fatal("expected a non-empty reset token")
}
if err := auth.CompletePasswordReset(ctx, PasswordResetCompleteRequest{Token: resetResp.Token, NewPassword: "new"}); err != nil {
t.Fatalf("CompletePasswordReset() error = %v", err)
}
// Reusing the same token should now fail.
if err := auth.CompletePasswordReset(ctx, PasswordResetCompleteRequest{Token: resetResp.Token, NewPassword: "again"}); err == nil {
t.Error("expected reusing a consumed reset token to fail")
}
}
func TestDirectMode_JWTLoginAndLogout(t *testing.T) {
db := newDirectTestDB(t)
jwtAuth := NewJWTAuthenticator("secret", db).WithQueryMode(ModeDirect)
directAuth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
ctx := context.Background()
if _, err := directAuth.Register(ctx, RegisterRequest{Username: "dave", Password: "p", Email: "dave@example.com"}); err != nil {
t.Fatalf("Register() error = %v", err)
}
resp, err := jwtAuth.Login(ctx, LoginRequest{Username: "dave", Password: "p"})
if err != nil {
t.Fatalf("JWTAuthenticator.Login() error = %v", err)
}
if resp.User.UserName != "dave" {
t.Errorf("JWT login user = %q, want dave", resp.User.UserName)
}
if err := jwtAuth.Logout(ctx, LogoutRequest{Token: resp.Token, UserID: resp.User.UserID}); err != nil {
t.Fatalf("JWTAuthenticator.Logout() error = %v", err)
}
}
func TestDirectMode_TOTPEnableAndValidateBackupCode(t *testing.T) {
db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
ctx := context.Background()
regResp, err := auth.Register(ctx, RegisterRequest{Username: "erin", Password: "p", Email: "erin@example.com"})
if err != nil {
t.Fatalf("Register() error = %v", err)
}
userID := regResp.User.UserID
totp := NewDatabaseTwoFactorProvider(db, nil).WithQueryMode(ModeDirect)
if err := totp.Enable2FA(userID, "SECRET123", []string{"code1", "code2"}); err != nil {
t.Fatalf("Enable2FA() error = %v", err)
}
enabled, err := totp.Get2FAStatus(userID)
if err != nil {
t.Fatalf("Get2FAStatus() error = %v", err)
}
if !enabled {
t.Error("expected 2FA to be enabled")
}
secret, err := totp.Get2FASecret(userID)
if err != nil {
t.Fatalf("Get2FASecret() error = %v", err)
}
if secret != "SECRET123" {
t.Errorf("Get2FASecret() = %q, want SECRET123", secret)
}
valid, err := totp.ValidateBackupCode(userID, "code1")
if err != nil {
t.Fatalf("ValidateBackupCode() error = %v", err)
}
if !valid {
t.Error("expected backup code to be valid")
}
// Reusing the same backup code should fail.
if _, err := totp.ValidateBackupCode(userID, "code1"); err == nil {
t.Error("expected reusing a consumed backup code to fail")
}
if err := totp.Disable2FA(userID); err != nil {
t.Fatalf("Disable2FA() error = %v", err)
}
enabled, err = totp.Get2FAStatus(userID)
if err != nil {
t.Fatalf("Get2FAStatus() error = %v", err)
}
if enabled {
t.Error("expected 2FA to be disabled")
}
}
func TestDirectMode_PasskeyStoreAndFetch(t *testing.T) {
db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
ctx := context.Background()
regResp, err := auth.Register(ctx, RegisterRequest{Username: "frank", Password: "p", Email: "frank@example.com"})
if err != nil {
t.Fatalf("Register() error = %v", err)
}
userID := regResp.User.UserID
passkeys := NewDatabasePasskeyProvider(db, DatabasePasskeyProviderOptions{
RPID: "example.com", RPName: "Example", RPOrigin: "https://example.com", QueryMode: ModeDirect,
})
cred, err := passkeys.CompleteRegistration(ctx, userID, PasskeyRegistrationResponse{
RawID: []byte("credential-1"),
Response: PasskeyAuthenticatorAttestationResponse{
AttestationObject: []byte("public-key-bytes"),
},
Transports: []string{"internal", "usb"},
}, nil)
if err != nil {
t.Fatalf("CompleteRegistration() error = %v", err)
}
if cred.UserID != userID {
t.Errorf("CompleteRegistration() UserID = %d, want %d", cred.UserID, userID)
}
creds, err := passkeys.GetCredentials(ctx, userID)
if err != nil {
t.Fatalf("GetCredentials() error = %v", err)
}
if len(creds) != 1 {
t.Fatalf("expected 1 credential, got %d", len(creds))
}
if string(creds[0].CredentialID) != "credential-1" {
t.Errorf("GetCredentials() CredentialID = %q, want credential-1", creds[0].CredentialID)
}
if len(creds[0].Transports) != 2 {
t.Errorf("expected 2 transports, got %v", creds[0].Transports)
}
credentialIDB64 := base64.StdEncoding.EncodeToString(creds[0].CredentialID)
if err := passkeys.UpdateCredentialName(ctx, userID, credentialIDB64, "My Phone"); err != nil {
t.Fatalf("UpdateCredentialName() error = %v", err)
}
updated, err := passkeys.GetCredentials(ctx, userID)
if err != nil {
t.Fatalf("GetCredentials() error = %v", err)
}
if updated[0].Name != "My Phone" {
t.Errorf("expected updated name 'My Phone', got %q", updated[0].Name)
}
if err := passkeys.DeleteCredential(ctx, userID, credentialIDB64); err != nil {
t.Fatalf("DeleteCredential() error = %v", err)
}
remaining, err := passkeys.GetCredentials(ctx, userID)
if err != nil {
t.Fatalf("GetCredentials() error = %v", err)
}
if len(remaining) != 0 {
t.Errorf("expected 0 credentials after delete, got %d", len(remaining))
}
}
func TestDirectMode_OAuthGetOrCreateUserAndSession(t *testing.T) {
db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
ctx := context.Background()
userCtx := &UserContext{UserName: "gina", Email: "gina@example.com", Roles: []string{"user"}}
userID, err := auth.oauth2GetOrCreateUser(ctx, userCtx, "google")
if err != nil {
t.Fatalf("oauth2GetOrCreateUser() error = %v", err)
}
if userID == 0 {
t.Fatal("expected non-zero user ID")
}
// Calling again with the same email should return the same user, not create a duplicate.
userID2, err := auth.oauth2GetOrCreateUser(ctx, userCtx, "google")
if err != nil {
t.Fatalf("oauth2GetOrCreateUser() second call error = %v", err)
}
if userID2 != userID {
t.Errorf("expected same user ID on repeat call, got %d and %d", userID, userID2)
}
}
func TestDirectMode_KeyStoreCreateAndValidate(t *testing.T) {
db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
ctx := context.Background()
regResp, err := auth.Register(ctx, RegisterRequest{Username: "henry", Password: "p", Email: "henry@example.com"})
if err != nil {
t.Fatalf("Register() error = %v", err)
}
ks := NewDatabaseKeyStore(db, DatabaseKeyStoreOptions{QueryMode: ModeDirect})
createResp, err := ks.CreateKey(ctx, CreateKeyRequest{
UserID: regResp.User.UserID,
KeyType: KeyTypeGenericAPI,
Name: "test key",
Scopes: []string{"read", "write"},
Meta: map[string]any{"note": "test"},
})
if err != nil {
t.Fatalf("CreateKey() error = %v", err)
}
if createResp.RawKey == "" {
t.Fatal("expected a non-empty raw key")
}
validated, err := ks.ValidateKey(ctx, createResp.RawKey, KeyTypeGenericAPI)
if err != nil {
t.Fatalf("ValidateKey() error = %v", err)
}
if validated.UserID != regResp.User.UserID {
t.Errorf("ValidateKey() UserID = %d, want %d", validated.UserID, regResp.User.UserID)
}
if len(validated.Scopes) != 2 {
t.Errorf("expected 2 scopes, got %v", validated.Scopes)
}
keys, err := ks.GetUserKeys(ctx, regResp.User.UserID, "")
if err != nil {
t.Fatalf("GetUserKeys() error = %v", err)
}
if len(keys) != 1 {
t.Fatalf("expected 1 key, got %d", len(keys))
}
if err := ks.DeleteKey(ctx, regResp.User.UserID, keys[0].ID); err != nil {
t.Fatalf("DeleteKey() error = %v", err)
}
remaining, err := ks.GetUserKeys(ctx, regResp.User.UserID, "")
if err != nil {
t.Fatalf("GetUserKeys() error = %v", err)
}
if len(remaining) != 0 {
t.Errorf("expected 0 keys after delete, got %d", len(remaining))
}
}
func TestDirectMode_OAuthServerClientAndCode(t *testing.T) {
db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
ctx := context.Background()
client := &OAuthServerClient{
ClientID: "client-1",
RedirectURIs: []string{"https://app.example.com/callback"},
ClientName: "Example App",
}
registered, err := auth.OAuthRegisterClient(ctx, client)
if err != nil {
t.Fatalf("OAuthRegisterClient() error = %v", err)
}
if len(registered.GrantTypes) != 1 || registered.GrantTypes[0] != "authorization_code" {
t.Errorf("expected default grant types, got %v", registered.GrantTypes)
}
if len(registered.AllowedScopes) != 3 {
t.Errorf("expected default allowed scopes, got %v", registered.AllowedScopes)
}
fetched, err := auth.OAuthGetClient(ctx, "client-1")
if err != nil {
t.Fatalf("OAuthGetClient() error = %v", err)
}
if fetched.ClientName != "Example App" {
t.Errorf("OAuthGetClient() ClientName = %q, want %q", fetched.ClientName, "Example App")
}
if len(fetched.RedirectURIs) != 1 || fetched.RedirectURIs[0] != "https://app.example.com/callback" {
t.Errorf("OAuthGetClient() RedirectURIs = %v", fetched.RedirectURIs)
}
regResp, err := auth.Register(ctx, RegisterRequest{Username: "ivan", Password: "p", Email: "ivan@example.com"})
if err != nil {
t.Fatalf("Register() error = %v", err)
}
if err := auth.OAuthSaveCode(ctx, &OAuthCode{
Code: "auth-code-2",
ClientID: "client-1",
RedirectURI: "https://app.example.com/callback",
CodeChallenge: "challenge",
SessionToken: regResp.Token,
Scopes: []string{"openid", "profile"},
ExpiresAt: futureTime(),
}); err != nil {
t.Fatalf("OAuthSaveCode() error = %v", err)
}
exchanged, err := auth.OAuthExchangeCode(ctx, "auth-code-2")
if err != nil {
t.Fatalf("OAuthExchangeCode() error = %v", err)
}
if exchanged.ClientID != "client-1" {
t.Errorf("OAuthExchangeCode() ClientID = %q, want client-1", exchanged.ClientID)
}
if len(exchanged.Scopes) != 2 {
t.Errorf("expected 2 scopes, got %v", exchanged.Scopes)
}
// Code should be single-use.
if _, err := auth.OAuthExchangeCode(ctx, "auth-code-2"); err == nil {
t.Error("expected re-exchanging a used code to fail")
}
info, err := auth.OAuthIntrospectToken(ctx, regResp.Token)
if err != nil {
t.Fatalf("OAuthIntrospectToken() error = %v", err)
}
if !info.Active {
t.Error("expected token to be active")
}
if info.Username != "ivan" {
t.Errorf("OAuthIntrospectToken() Username = %q, want ivan", info.Username)
}
if err := auth.OAuthRevokeToken(ctx, regResp.Token); err != nil {
t.Fatalf("OAuthRevokeToken() error = %v", err)
}
info, err = auth.OAuthIntrospectToken(ctx, regResp.Token)
if err != nil {
t.Fatalf("OAuthIntrospectToken() after revoke error = %v", err)
}
if info.Active {
t.Error("expected token to be inactive after revoke")
}
}
+23 -9
View File
@@ -14,6 +14,11 @@ import (
type SecurityContext interface {
GetContext() context.Context
GetUserID() (int, bool)
// GetUserRef returns an opaque user identifier for row security lookups.
// Unlike GetUserID, it is not required to be an integer: implementations backed by
// non-integer identifiers (e.g. UUIDs) can return a string, or the full
// *security.UserContext so a RowSecurityProvider can read JWT claims directly.
GetUserRef() (any, bool)
GetSchema() string
GetEntity() string
GetModel() interface{}
@@ -45,8 +50,13 @@ func loadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error
// return err
}
// Load row security rules using the provider
_, err = securityList.LoadRowSecurity(secCtx.GetContext(), userID, schema, tablename, false)
// Load row security rules using the provider. Row security uses the opaque
// user ref (not the int-only user ID) so non-integer user identifiers work.
userRef, refOK := secCtx.GetUserRef()
if !refOK {
userRef = userID
}
_, err = securityList.LoadRowSecurity(secCtx.GetContext(), userRef, schema, tablename, false)
if err != nil {
logger.Warn("Failed to load row security: %v", err)
// Don't fail the request if no security rules exist
@@ -58,25 +68,29 @@ func loadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error
// applyRowSecurity applies row-level security filters to the query (generic version)
func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
userID, ok := secCtx.GetUserID()
userRef, ok := secCtx.GetUserRef()
if !ok {
return nil // No user context, skip
userID, idOK := secCtx.GetUserID()
if !idOK {
return nil // No user context, skip
}
userRef = userID
}
schema := secCtx.GetSchema()
tablename := secCtx.GetEntity()
// Get row security template
rowSec, err := securityList.GetRowSecurityTemplate(userID, schema, tablename)
rowSec, err := securityList.GetRowSecurityTemplate(userRef, schema, tablename)
if err != nil {
// No row security defined, allow query to proceed
logger.Debug("No row security for %s.%s@%d: %v", schema, tablename, userID, err)
logger.Debug("No row security for %s.%s@%v: %v", schema, tablename, userRef, err)
return nil
}
// Check if user has a blocking rule
if rowSec.HasBlock {
logger.Warn("User %d blocked from accessing %s.%s", userID, schema, tablename)
logger.Warn("User %v blocked from accessing %s.%s", userRef, schema, tablename)
return fmt.Errorf("access denied to %s", tablename)
}
@@ -112,8 +126,8 @@ func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error
// Generate the WHERE clause from template
whereClause := rowSec.GetTemplate(pkName, modelType)
logger.Info("Applying row security filter for user %d on %s.%s: %s",
userID, schema, tablename, whereClause)
logger.Info("Applying row security filter for user %v on %s.%s: %s",
userRef, schema, tablename, whereClause)
// Apply the WHERE clause to the query
query := secCtx.GetQuery()
+4
View File
@@ -26,6 +26,10 @@ func (m *mockSecurityContext) GetUserID() (int, bool) {
return m.userID, m.hasUser
}
func (m *mockSecurityContext) GetUserRef() (any, bool) {
return m.userID, m.hasUser
}
func (m *mockSecurityContext) GetSchema() string {
return m.schema
}
+6 -2
View File
@@ -121,8 +121,12 @@ type ColumnSecurityProvider interface {
// RowSecurityProvider handles row-level security (filtering)
type RowSecurityProvider interface {
// GetRowSecurity loads row security rules for a user and entity
GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error)
// GetRowSecurity loads row security rules for a user and entity.
// userRef identifies the user and is opaque to the caller: it may be an int ID,
// a string/UUID, or the full *security.UserContext (see SecurityContext.GetUserRef),
// so providers backed by non-integer user identifiers (e.g. UUIDs) or that need
// access to JWT claims can implement row security without relying on a numeric ID.
GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error)
}
// SecurityProvider is the main interface combining all security concerns
+52 -11
View File
@@ -23,6 +23,10 @@ type DatabaseKeyStoreOptions struct {
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
// 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)
@@ -38,12 +42,15 @@ type DatabaseKeyStoreOptions struct {
// 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
cache *cache.Cache
cacheTTL time.Duration
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
}
// NewDatabaseKeyStore creates a DatabaseKeyStore with optional configuration.
@@ -60,12 +67,16 @@ func NewDatabaseKeyStore(db *sql.DB, opts ...DatabaseKeyStoreOptions) *DatabaseK
c = cache.GetDefaultCache()
}
names := MergeKeyStoreSQLNames(DefaultKeyStoreSQLNames(), o.SQLNames)
tableNames := resolveKeyStoreTableNames(o.TableNames)
return &DatabaseKeyStore{
db: db,
dbFactory: o.DBFactory,
sqlNames: names,
cache: c,
cacheTTL: o.CacheTTL,
db: db,
dbFactory: o.DBFactory,
sqlNames: names,
tableNames: tableNames,
queryMode: o.QueryMode,
capability: newDBCapability(),
cache: c,
cacheTTL: o.CacheTTL,
}
}
@@ -86,6 +97,9 @@ func (ks *DatabaseKeyStore) reconnectDB() error {
ks.dbMu.Lock()
ks.db = newDB
ks.dbMu.Unlock()
if ks.capability != nil {
ks.capability.reset()
}
return nil
}
@@ -99,6 +113,14 @@ 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"`
@@ -145,6 +167,10 @@ func (ks *DatabaseKeyStore) CreateKey(ctx context.Context, req CreateKeyRequest)
// 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
@@ -173,6 +199,10 @@ func (ks *DatabaseKeyStore) GetUserKeys(ctx context.Context, userID int, keyType
// 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)
}
var success bool
var errorMsg sql.NullString
var keyHash sql.NullString
@@ -207,6 +237,17 @@ func (ks *DatabaseKeyStore) ValidateKey(ctx context.Context, rawKey string, keyT
}
}
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
+216
View File
@@ -0,0 +1,216 @@
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)
}
+44
View File
@@ -0,0 +1,44 @@
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)
}
+15 -83
View File
@@ -226,6 +226,10 @@ func (a *DatabaseAuthenticator) getOAuth2Provider(providerName string) (*OAuth2P
// oauth2GetOrCreateUser finds or creates a user based on OAuth2 info using stored procedure
func (a *DatabaseAuthenticator) oauth2GetOrCreateUser(ctx context.Context, userCtx *UserContext, providerName string) (int, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthGetOrCreateUser) {
return a.oauth2GetOrCreateUserDirect(ctx, userCtx, providerName)
}
userData := map[string]interface{}{
"username": userCtx.UserName,
"email": userCtx.Email,
@@ -269,6 +273,10 @@ func (a *DatabaseAuthenticator) oauth2GetOrCreateUser(ctx context.Context, userC
// oauth2CreateSession creates a new OAuth2 session using stored procedure
func (a *DatabaseAuthenticator) oauth2CreateSession(ctx context.Context, sessionToken string, userID int, token *oauth2.Token, expiresAt time.Time, providerName string) error {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthCreateSession) {
return a.oauth2CreateSessionDirect(ctx, sessionToken, userID, token, expiresAt, providerName)
}
sessionData := map[string]interface{}{
"session_token": sessionToken,
"user_id": userID,
@@ -381,35 +389,9 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
}
// Get session by refresh token from database
var success bool
var errMsg *string
var sessionData []byte
err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error, p_data::text
FROM %s($1)
`, a.sqlNames.OAuthGetRefreshToken), refreshToken).Scan(&success, &errMsg, &sessionData)
session, err := a.oauthGetByRefreshToken(ctx, refreshToken)
if err != nil {
return nil, fmt.Errorf("failed to get session by refresh token: %w", err)
}
if !success {
if errMsg != nil {
return nil, fmt.Errorf("%s", *errMsg)
}
return nil, fmt.Errorf("invalid or expired refresh token")
}
// Parse session data
var session struct {
UserID int `json:"user_id"`
AccessToken string `json:"access_token"`
TokenType string `json:"token_type"`
Expiry time.Time `json:"expiry"`
}
if err := json.Unmarshal(sessionData, &session); err != nil {
return nil, fmt.Errorf("failed to parse session data: %w", err)
return nil, err
}
// Create oauth2.Token from stored data
@@ -434,64 +416,14 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
}
// Update session in database with new tokens
updateData := map[string]interface{}{
"user_id": session.UserID,
"old_refresh_token": refreshToken,
"new_session_token": newSessionToken,
"new_access_token": newToken.AccessToken,
"new_refresh_token": newToken.RefreshToken,
"expires_at": newToken.Expiry,
}
updateJSON, err := json.Marshal(updateData)
if err != nil {
return nil, fmt.Errorf("failed to marshal update data: %w", err)
}
var updateSuccess bool
var updateErrMsg *string
err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error
FROM %s($1::jsonb)
`, a.sqlNames.OAuthUpdateRefreshToken), updateJSON).Scan(&updateSuccess, &updateErrMsg)
if err != nil {
return nil, fmt.Errorf("failed to update session: %w", err)
}
if !updateSuccess {
if updateErrMsg != nil {
return nil, fmt.Errorf("%s", *updateErrMsg)
}
return nil, fmt.Errorf("failed to update session")
if err := a.oauthUpdateRefreshTokenRecord(ctx, session.UserID, refreshToken, newSessionToken, newToken.AccessToken, newToken.RefreshToken, newToken.Expiry); err != nil {
return nil, err
}
// Get user data
var userSuccess bool
var userErrMsg *string
var userData []byte
err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error, p_data::text
FROM %s($1)
`, a.sqlNames.OAuthGetUser), session.UserID).Scan(&userSuccess, &userErrMsg, &userData)
userCtx, err := a.oauthGetUserByID(ctx, session.UserID)
if err != nil {
return nil, fmt.Errorf("failed to get user data: %w", err)
}
if !userSuccess {
if userErrMsg != nil {
return nil, fmt.Errorf("%s", *userErrMsg)
}
return nil, fmt.Errorf("failed to get user data")
}
// Parse user context
var userCtx UserContext
if err := json.Unmarshal(userData, &userCtx); err != nil {
return nil, fmt.Errorf("failed to parse user context: %w", err)
return nil, err
}
userCtx.SessionID = newSessionToken
@@ -499,7 +431,7 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
return &LoginResponse{
Token: newSessionToken,
RefreshToken: newToken.RefreshToken,
User: &userCtx,
User: userCtx,
ExpiresIn: int64(time.Until(newToken.Expiry).Seconds()),
}, nil
}
+242
View File
@@ -0,0 +1,242 @@
package security
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
"golang.org/x/oauth2"
)
// oauthRefreshSession is the session data needed to refresh an OAuth2 token,
// shared by both the stored-procedure and Direct-mode code paths.
type oauthRefreshSession struct {
UserID int `json:"user_id"`
AccessToken string `json:"access_token"`
TokenType string `json:"token_type"`
Expiry time.Time `json:"expiry"`
}
// oauth2GetOrCreateUserDirect mirrors resolvespec_oauth_getorcreateuser.
func (a *DatabaseAuthenticator) oauth2GetOrCreateUserDirect(ctx context.Context, userCtx *UserContext, providerName string) (int, error) {
rolesStr := strings.Join(userCtx.Roles, ",")
var userID int
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT id FROM %s WHERE email = ?`, a.tableNames.Users))
err := db.QueryRowContext(ctx, query, userCtx.Email).Scan(&userID)
if err == nil {
now := time.Now()
updQuery := rewritePlaceholders(db, fmt.Sprintf(
`UPDATE %s SET last_login_at = ?, updated_at = ?, remote_id = COALESCE(remote_id, ?), auth_provider = COALESCE(auth_provider, ?) WHERE id = ?`,
a.tableNames.Users))
_, err := db.ExecContext(ctx, updQuery, now, now, userCtx.RemoteID, providerName, userID)
return err
}
if !errors.Is(err, sql.ErrNoRows) {
return err
}
now := time.Now()
insQuery := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (username, email, password, user_level, roles, is_active, created_at, updated_at, last_login_at, remote_id, auth_provider) VALUES (?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?)`,
a.tableNames.Users))
res, err := db.ExecContext(ctx, insQuery, userCtx.UserName, userCtx.Email, userCtx.UserLevel, rolesStr, true, now, now, now, userCtx.RemoteID, providerName)
if err != nil {
return err
}
id, err := res.LastInsertId()
if err != nil {
return err
}
userID = int(id)
return nil
})
if err != nil {
return 0, fmt.Errorf("failed to get or create user: %w", err)
}
return userID, nil
}
// oauth2CreateSessionDirect mirrors resolvespec_oauth_createsession (insert-or-update by session_token).
func (a *DatabaseAuthenticator) oauth2CreateSessionDirect(ctx context.Context, sessionToken string, userID int, token *oauth2.Token, expiresAt time.Time, providerName string) error {
return a.runDBOpWithReconnect(func(db *sql.DB) error {
var exists int
checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT 1 FROM %s WHERE session_token = ?`, a.tableNames.UserSessions))
err := db.QueryRowContext(ctx, checkQuery, sessionToken).Scan(&exists)
now := time.Now()
if err == nil {
updQuery := rewritePlaceholders(db, fmt.Sprintf(
`UPDATE %s SET access_token = ?, refresh_token = ?, token_type = ?, expires_at = ?, last_activity_at = ? WHERE session_token = ?`,
a.tableNames.UserSessions))
_, err := db.ExecContext(ctx, updQuery, token.AccessToken, token.RefreshToken, token.TokenType, expiresAt, now, sessionToken)
return err
}
if !errors.Is(err, sql.ErrNoRows) {
return err
}
insQuery := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (session_token, user_id, expires_at, created_at, last_activity_at, access_token, refresh_token, token_type, auth_provider) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
a.tableNames.UserSessions))
_, err = db.ExecContext(ctx, insQuery, sessionToken, userID, expiresAt, now, now, token.AccessToken, token.RefreshToken, token.TokenType, providerName)
return err
})
}
// oauthGetByRefreshToken retrieves the session for a refresh token, dispatching between
// the resolvespec_oauth_getrefreshtoken stored procedure and Direct-mode SQL.
func (a *DatabaseAuthenticator) oauthGetByRefreshToken(ctx context.Context, refreshToken string) (*oauthRefreshSession, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthGetRefreshToken) {
var session oauthRefreshSession
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`SELECT user_id, access_token, token_type, expires_at FROM %s WHERE refresh_token = ? AND expires_at > ?`,
a.tableNames.UserSessions))
return db.QueryRowContext(ctx, query, refreshToken, time.Now()).Scan(&session.UserID, &session.AccessToken, &session.TokenType, &session.Expiry)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("refresh token not found or expired")
}
return nil, fmt.Errorf("failed to get session by refresh token: %w", err)
}
return &session, nil
}
var success bool
var errMsg *string
var sessionData []byte
err := a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error, p_data::text
FROM %s($1)
`, a.sqlNames.OAuthGetRefreshToken), refreshToken).Scan(&success, &errMsg, &sessionData)
if err != nil {
return nil, fmt.Errorf("failed to get session by refresh token: %w", err)
}
if !success {
if errMsg != nil {
return nil, fmt.Errorf("%s", *errMsg)
}
return nil, fmt.Errorf("invalid or expired refresh token")
}
var session oauthRefreshSession
if err := json.Unmarshal(sessionData, &session); err != nil {
return nil, fmt.Errorf("failed to parse session data: %w", err)
}
return &session, nil
}
// oauthUpdateRefreshTokenRecord updates a session with new tokens, dispatching between
// the resolvespec_oauth_updaterefreshtoken stored procedure and Direct-mode SQL.
func (a *DatabaseAuthenticator) oauthUpdateRefreshTokenRecord(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthUpdateRefreshToken) {
var rows int64
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`UPDATE %s SET session_token = ?, access_token = ?, refresh_token = ?, expires_at = ?, last_activity_at = ? WHERE user_id = ? AND refresh_token = ?`,
a.tableNames.UserSessions))
res, err := db.ExecContext(ctx, query, newSessionToken, newAccessToken, newRefreshToken, expiresAt, time.Now(), userID, oldRefreshToken)
if err != nil {
return err
}
rows, err = res.RowsAffected()
return err
})
if err != nil {
return fmt.Errorf("failed to update session: %w", err)
}
if rows == 0 {
return fmt.Errorf("session not found")
}
return nil
}
updateData := map[string]interface{}{
"user_id": userID,
"old_refresh_token": oldRefreshToken,
"new_session_token": newSessionToken,
"new_access_token": newAccessToken,
"new_refresh_token": newRefreshToken,
"expires_at": expiresAt,
}
updateJSON, err := json.Marshal(updateData)
if err != nil {
return fmt.Errorf("failed to marshal update data: %w", err)
}
var updateSuccess bool
var updateErrMsg *string
err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error
FROM %s($1::jsonb)
`, a.sqlNames.OAuthUpdateRefreshToken), updateJSON).Scan(&updateSuccess, &updateErrMsg)
if err != nil {
return fmt.Errorf("failed to update session: %w", err)
}
if !updateSuccess {
if updateErrMsg != nil {
return fmt.Errorf("%s", *updateErrMsg)
}
return fmt.Errorf("failed to update session")
}
return nil
}
// oauthGetUserByID retrieves user data by ID, dispatching between the
// resolvespec_oauth_getuser stored procedure and Direct-mode SQL.
func (a *DatabaseAuthenticator) oauthGetUserByID(ctx context.Context, userID int) (*UserContext, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthGetUser) {
var username, email, roles, programUserTable sql.NullString
var userLevel, programUserID sql.NullInt64
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`SELECT username, email, user_level, roles, program_user_id, program_user_table FROM %s WHERE id = ? AND is_active = ?`,
a.tableNames.Users))
return db.QueryRowContext(ctx, query, userID, true).Scan(&username, &email, &userLevel, &roles, &programUserID, &programUserTable)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("user not found")
}
return nil, fmt.Errorf("failed to get user data: %w", err)
}
return &UserContext{
UserID: userID,
UserName: username.String,
Email: email.String,
UserLevel: int(userLevel.Int64),
Roles: parseRoles(roles.String),
ProgramUserID: int(programUserID.Int64),
ProgramUserTable: programUserTable.String,
}, nil
}
var userSuccess bool
var userErrMsg *string
var userData []byte
err := a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error, p_data::text
FROM %s($1)
`, a.sqlNames.OAuthGetUser), userID).Scan(&userSuccess, &userErrMsg, &userData)
if err != nil {
return nil, fmt.Errorf("failed to get user data: %w", err)
}
if !userSuccess {
if userErrMsg != nil {
return nil, fmt.Errorf("%s", *userErrMsg)
}
return nil, fmt.Errorf("failed to get user data")
}
var userCtx UserContext
if err := json.Unmarshal(userData, &userCtx); err != nil {
return nil, fmt.Errorf("failed to parse user context: %w", err)
}
return &userCtx, nil
}
+424 -35
View File
@@ -3,8 +3,11 @@ package security
import (
"context"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"crypto/subtle"
"encoding/base64"
"encoding/hex"
"encoding/json"
"fmt"
"net/http"
@@ -12,6 +15,9 @@ import (
"strings"
"sync"
"time"
"github.com/golang-jwt/jwt/v5"
"golang.org/x/oauth2"
)
// OAuthServerConfig configures the MCP-standard OAuth2 authorization server.
@@ -44,15 +50,32 @@ type OAuthServerConfig struct {
// AuthCodeTTL is the auth code lifetime. Defaults to 2 minutes.
AuthCodeTTL time.Duration
// ResourceIdentifier is this server's protected-resource identifier, advertised in
// RFC 9728 metadata. Defaults to Issuer.
ResourceIdentifier string
// SigningKey signs id_tokens (RS256) and is exposed via the JWKS endpoint. If nil, an
// RSA-2048 key is generated in memory when the server starts. Supply a persistent key
// for multi-instance deployments so id_tokens remain verifiable across restarts/instances.
SigningKey *rsa.PrivateKey
}
// oauthClient is a dynamically registered OAuth2 client (RFC 7591).
type oauthClient struct {
ClientID string `json:"client_id"`
RedirectURIs []string `json:"redirect_uris"`
ClientName string `json:"client_name,omitempty"`
GrantTypes []string `json:"grant_types"`
AllowedScopes []string `json:"allowed_scopes,omitempty"`
ClientID string `json:"client_id"`
RedirectURIs []string `json:"redirect_uris"`
ClientName string `json:"client_name,omitempty"`
GrantTypes []string `json:"grant_types"`
AllowedScopes []string `json:"allowed_scopes,omitempty"`
ClientSecretHash string `json:"client_secret_hash,omitempty"`
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
}
// isConfidential reports whether the client has a registered secret and must
// authenticate itself at the token endpoint.
func (c *oauthClient) isConfidential() bool {
return c.ClientSecretHash != ""
}
// pendingAuth tracks an in-progress authorization code exchange.
@@ -85,13 +108,25 @@ type externalProvider struct {
// The server exposes these RFC-compliant endpoints:
//
// GET /.well-known/oauth-authorization-server RFC 8414 — server metadata discovery
// GET /.well-known/openid-configuration OIDC discovery (superset of the above)
// GET /.well-known/oauth-protected-resource RFC 9728 — protected resource metadata
// POST /oauth/register RFC 7591 — dynamic client registration
// GET /oauth/authorize OAuth 2.1 + PKCE — start authorization
// POST /oauth/authorize Direct login form submission
// POST /oauth/token Token exchange and refresh
// POST /oauth/token Token exchange: authorization_code,
// refresh_token, client_credentials (RFC 6749 §4.4)
// POST /oauth/revoke RFC 7009 — token revocation
// POST /oauth/introspect RFC 7662 — token introspection
// GET /oauth/userinfo OIDC UserInfo endpoint
// GET /oauth/jwks.json JWKS — id_token verification keys
// GET {ProviderCallbackPath} Internal — external provider callback
//
// Confidential clients (registered with token_endpoint_auth_method other than "none", or
// any grant_types including client_credentials) authenticate at /oauth/token via
// client_secret_basic or client_secret_post. Public clients keep relying on PKCE alone.
//
// When the granted scope includes "openid", authorization_code and refresh_token responses
// include an RS256-signed id_token (see OAuthServerConfig.SigningKey).
type OAuthServer struct {
cfg OAuthServerConfig
auth *DatabaseAuthenticator // nil = only external providers
@@ -102,6 +137,9 @@ type OAuthServer struct {
pending map[string]*pendingAuth // provider_state → pending (external flow)
codes map[string]*pendingAuth // auth_code → pending (post-auth)
signingKey *rsa.PrivateKey
signingKeyID string
done chan struct{} // closed by Close() to stop background goroutines
}
@@ -130,13 +168,32 @@ func NewOAuthServer(cfg OAuthServerConfig, auth *DatabaseAuthenticator) *OAuthSe
}
// Normalize issuer: remove trailing slash to ensure consistent endpoint URL construction.
cfg.Issuer = strings.TrimSuffix(cfg.Issuer, "/")
if cfg.ResourceIdentifier == "" {
cfg.ResourceIdentifier = cfg.Issuer
}
signingKey := cfg.SigningKey
if signingKey == nil {
var err error
signingKey, err = rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
// Signing keys are only required for id_token issuance (OIDC "openid" scope);
// leaving signingKey nil degrades gracefully by omitting id_token/JWKS support.
signingKey = nil
}
}
s := &OAuthServer{
cfg: cfg,
auth: auth,
clients: make(map[string]*oauthClient),
pending: make(map[string]*pendingAuth),
codes: make(map[string]*pendingAuth),
done: make(chan struct{}),
cfg: cfg,
auth: auth,
clients: make(map[string]*oauthClient),
pending: make(map[string]*pendingAuth),
codes: make(map[string]*pendingAuth),
signingKey: signingKey,
done: make(chan struct{}),
}
if signingKey != nil {
s.signingKeyID = rsaKeyID(&signingKey.PublicKey)
}
go s.cleanupExpired()
return s
@@ -178,11 +235,15 @@ func (s *OAuthServer) ProviderCallbackPath() string {
func (s *OAuthServer) HTTPHandler() http.Handler {
mux := http.NewServeMux()
mux.HandleFunc("/.well-known/oauth-authorization-server", s.metadataHandler)
mux.HandleFunc("/.well-known/openid-configuration", s.openIDConfigurationHandler)
mux.HandleFunc("/.well-known/oauth-protected-resource", s.protectedResourceHandler)
mux.HandleFunc("/oauth/register", s.registerHandler)
mux.HandleFunc("/oauth/authorize", s.authorizeHandler)
mux.HandleFunc("/oauth/token", s.tokenHandler)
mux.HandleFunc("/oauth/revoke", s.revokeHandler)
mux.HandleFunc("/oauth/introspect", s.introspectHandler)
mux.HandleFunc("/oauth/userinfo", s.userinfoHandler)
mux.HandleFunc("/oauth/jwks.json", s.jwksHandler)
mux.HandleFunc(s.cfg.ProviderCallbackPath, s.providerCallbackHandler)
return mux
}
@@ -217,25 +278,127 @@ func (s *OAuthServer) cleanupExpired() {
// RFC 8414 — Server metadata
// --------------------------------------------------------------------------
func (s *OAuthServer) metadataHandler(w http.ResponseWriter, r *http.Request) {
// serverMetadata builds the fields shared by RFC 8414 authorization-server
// metadata and OIDC discovery metadata.
func (s *OAuthServer) serverMetadata() map[string]interface{} {
issuer := s.cfg.Issuer
meta := map[string]interface{}{
grantTypes := []string{"authorization_code", "refresh_token"}
if s.auth != nil {
grantTypes = append(grantTypes, "client_credentials")
}
return map[string]interface{}{
"issuer": issuer,
"authorization_endpoint": issuer + "/oauth/authorize",
"token_endpoint": issuer + "/oauth/token",
"registration_endpoint": issuer + "/oauth/register",
"revocation_endpoint": issuer + "/oauth/revoke",
"introspection_endpoint": issuer + "/oauth/introspect",
"userinfo_endpoint": issuer + "/oauth/userinfo",
"jwks_uri": issuer + "/oauth/jwks.json",
"scopes_supported": s.cfg.DefaultScopes,
"response_types_supported": []string{"code"},
"grant_types_supported": []string{"authorization_code", "refresh_token"},
"grant_types_supported": grantTypes,
"code_challenge_methods_supported": []string{"S256"},
"token_endpoint_auth_methods_supported": []string{"none"},
"token_endpoint_auth_methods_supported": []string{"none", "client_secret_basic", "client_secret_post"},
}
}
func (s *OAuthServer) metadataHandler(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(s.serverMetadata()) //nolint:errcheck
}
// --------------------------------------------------------------------------
// OIDC discovery — GET /.well-known/openid-configuration
// --------------------------------------------------------------------------
func (s *OAuthServer) openIDConfigurationHandler(w http.ResponseWriter, r *http.Request) {
meta := s.serverMetadata()
meta["subject_types_supported"] = []string{"public"}
meta["id_token_signing_alg_values_supported"] = []string{"RS256"}
meta["claims_supported"] = []string{"sub", "iss", "aud", "exp", "iat", "email", "preferred_username"}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(meta) //nolint:errcheck
}
// --------------------------------------------------------------------------
// RFC 9728 — Protected Resource Metadata
// --------------------------------------------------------------------------
func (s *OAuthServer) protectedResourceHandler(w http.ResponseWriter, r *http.Request) {
meta := map[string]interface{}{
"resource": s.cfg.ResourceIdentifier,
"authorization_servers": []string{s.cfg.Issuer},
"scopes_supported": s.cfg.DefaultScopes,
"bearer_methods_supported": []string{"header"},
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(meta) //nolint:errcheck
}
// --------------------------------------------------------------------------
// JWKS — GET /oauth/jwks.json
// --------------------------------------------------------------------------
func (s *OAuthServer) jwksHandler(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
if s.signingKey == nil {
json.NewEncoder(w).Encode(map[string]interface{}{"keys": []interface{}{}}) //nolint:errcheck
return
}
pub := s.signingKey.PublicKey
jwk := map[string]interface{}{
"kty": "RSA",
"use": "sig",
"alg": "RS256",
"kid": s.signingKeyID,
"n": base64.RawURLEncoding.EncodeToString(pub.N.Bytes()),
"e": base64.RawURLEncoding.EncodeToString(bigEndianBytes(pub.E)),
}
json.NewEncoder(w).Encode(map[string]interface{}{"keys": []interface{}{jwk}}) //nolint:errcheck
}
// --------------------------------------------------------------------------
// Userinfo — GET/POST /oauth/userinfo
// --------------------------------------------------------------------------
func (s *OAuthServer) userinfoHandler(w http.ResponseWriter, r *http.Request) {
auth := r.Header.Get("Authorization")
token := strings.TrimPrefix(auth, "Bearer ")
if token == "" || token == auth {
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`)
writeOAuthError(w, "invalid_token", "missing bearer token", http.StatusUnauthorized)
return
}
authToUse := s.auth
if authToUse == nil {
s.mu.RLock()
if len(s.providers) > 0 {
authToUse = s.providers[0].auth
}
s.mu.RUnlock()
}
if authToUse == nil {
writeOAuthError(w, "invalid_token", "no authenticator configured", http.StatusUnauthorized)
return
}
info, err := authToUse.OAuthIntrospectToken(r.Context(), token)
if err != nil || !info.Active {
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`)
writeOAuthError(w, "invalid_token", "token is inactive or invalid", http.StatusUnauthorized)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]interface{}{ //nolint:errcheck
"sub": info.Sub,
"preferred_username": info.Username,
"email": info.Email,
})
}
// --------------------------------------------------------------------------
// RFC 7591 — Dynamic client registration
// --------------------------------------------------------------------------
@@ -246,10 +409,11 @@ func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
return
}
var req struct {
RedirectURIs []string `json:"redirect_uris"`
ClientName string `json:"client_name"`
GrantTypes []string `json:"grant_types"`
AllowedScopes []string `json:"allowed_scopes"`
RedirectURIs []string `json:"redirect_uris"`
ClientName string `json:"client_name"`
GrantTypes []string `json:"grant_types"`
AllowedScopes []string `json:"allowed_scopes"`
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"`
}
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
writeOAuthError(w, "invalid_request", "malformed JSON", http.StatusBadRequest)
@@ -272,21 +436,48 @@ func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
// client_credentials is a machine-to-machine grant and requires a confidential
// client (RFC 6749 §4.4), so it always forces secret issuance regardless of the
// requested auth method.
authMethod := req.TokenEndpointAuthMethod
if authMethod == "" {
authMethod = "none"
}
if oauthSliceContains(grantTypes, "client_credentials") && authMethod == "none" {
authMethod = "client_secret_basic"
}
var plaintextSecret string
var secretHash string
if authMethod != "none" {
plaintextSecret, err = randomOAuthToken()
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
secretHash = hashClientSecret(plaintextSecret)
}
client := &oauthClient{
ClientID: clientID,
RedirectURIs: req.RedirectURIs,
ClientName: req.ClientName,
GrantTypes: grantTypes,
AllowedScopes: allowedScopes,
ClientID: clientID,
RedirectURIs: req.RedirectURIs,
ClientName: req.ClientName,
GrantTypes: grantTypes,
AllowedScopes: allowedScopes,
ClientSecretHash: secretHash,
TokenEndpointAuthMethod: authMethod,
}
if s.cfg.PersistClients && s.auth != nil {
dbClient := &OAuthServerClient{
ClientID: client.ClientID,
RedirectURIs: client.RedirectURIs,
ClientName: client.ClientName,
GrantTypes: client.GrantTypes,
AllowedScopes: client.AllowedScopes,
ClientID: client.ClientID,
RedirectURIs: client.RedirectURIs,
ClientName: client.ClientName,
GrantTypes: client.GrantTypes,
AllowedScopes: client.AllowedScopes,
ClientSecretHash: client.ClientSecretHash,
TokenEndpointAuthMethod: client.TokenEndpointAuthMethod,
}
if _, err := s.auth.OAuthRegisterClient(r.Context(), dbClient); err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
@@ -298,9 +489,25 @@ func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
s.clients[clientID] = client
s.mu.Unlock()
// RFC 7591 registration response: the plaintext secret is returned exactly once here
// and never persisted or served again — only its hash (client.ClientSecretHash) is
// stored, and that hash is deliberately excluded from this response.
resp := map[string]interface{}{
"client_id": client.ClientID,
"redirect_uris": client.RedirectURIs,
"client_name": client.ClientName,
"grant_types": client.GrantTypes,
"allowed_scopes": client.AllowedScopes,
"token_endpoint_auth_method": client.TokenEndpointAuthMethod,
}
if plaintextSecret != "" {
resp["client_secret"] = plaintextSecret
resp["client_secret_expires_at"] = 0
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusCreated)
json.NewEncoder(w).Encode(client) //nolint:errcheck
json.NewEncoder(w).Encode(resp) //nolint:errcheck
}
// --------------------------------------------------------------------------
@@ -577,6 +784,8 @@ func (s *OAuthServer) tokenHandler(w http.ResponseWriter, r *http.Request) {
s.handleAuthCodeGrant(w, r)
case "refresh_token":
s.handleRefreshGrant(w, r)
case "client_credentials":
s.handleClientCredentialsGrant(w, r)
default:
writeOAuthError(w, "unsupported_grant_type", "", http.StatusBadRequest)
}
@@ -593,6 +802,15 @@ func (s *OAuthServer) handleAuthCodeGrant(w http.ResponseWriter, r *http.Request
return
}
// Confidential clients (those registered with a client_secret) must authenticate;
// public clients keep relying on PKCE alone, unchanged from prior behavior.
if client, ok := s.lookupOrFetchClient(r.Context(), clientID); ok && client.isConfidential() {
if _, err := s.authenticateClient(r); err != nil {
writeOAuthError(w, "invalid_client", err.Error(), http.StatusUnauthorized)
return
}
}
var sessionToken string
var refreshToken string
var scopes []string
@@ -647,12 +865,13 @@ func (s *OAuthServer) handleAuthCodeGrant(w http.ResponseWriter, r *http.Request
scopes = pending.Scopes
}
s.writeOAuthToken(w, sessionToken, refreshToken, scopes)
s.writeOAuthToken(w, r, sessionToken, refreshToken, clientID, scopes, true)
}
func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request) {
refreshToken := r.FormValue("refresh_token")
providerName := r.FormValue("provider")
clientID := r.FormValue("client_id")
if refreshToken == "" {
writeOAuthError(w, "invalid_request", "refresh_token required", http.StatusBadRequest)
return
@@ -666,7 +885,7 @@ func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request)
writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest)
return
}
s.writeOAuthToken(w, loginResp.Token, loginResp.RefreshToken, nil)
s.writeOAuthToken(w, r, loginResp.Token, loginResp.RefreshToken, clientID, nil, false)
return
}
@@ -676,13 +895,86 @@ func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request)
writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest)
return
}
s.writeOAuthToken(w, loginResp.Token, loginResp.RefreshToken, nil)
s.writeOAuthToken(w, r, loginResp.Token, loginResp.RefreshToken, clientID, nil, false)
return
}
writeOAuthError(w, "invalid_grant", "no provider available for refresh", http.StatusBadRequest)
}
// --------------------------------------------------------------------------
// RFC 6749 §4.4 — Client credentials grant
// --------------------------------------------------------------------------
func (s *OAuthServer) handleClientCredentialsGrant(w http.ResponseWriter, r *http.Request) {
if s.auth == nil {
writeOAuthError(w, "unsupported_grant_type", "client_credentials requires a local user store", http.StatusBadRequest)
return
}
client, err := s.authenticateClient(r)
if err != nil {
w.Header().Set("WWW-Authenticate", `Basic realm="oauth"`)
writeOAuthError(w, "invalid_client", err.Error(), http.StatusUnauthorized)
return
}
if !oauthSliceContains(client.GrantTypes, "client_credentials") {
writeOAuthError(w, "unauthorized_client", "client is not authorized for client_credentials", http.StatusBadRequest)
return
}
requested := strings.Fields(r.FormValue("scope"))
effectiveScopes := client.AllowedScopes
if len(requested) > 0 {
effectiveScopes = nil
for _, sc := range requested {
if oauthSliceContains(client.AllowedScopes, sc) {
effectiveScopes = append(effectiveScopes, sc)
}
}
if len(effectiveScopes) == 0 {
writeOAuthError(w, "invalid_scope", "no requested scope is allowed for this client", http.StatusBadRequest)
return
}
}
// client_credentials tokens have no end user, but the rest of the stack (RLS-scoping
// hooks, introspection) expects every access token to resolve to a user_sessions row
// with a user_id. Represent the client as a deterministic synthetic "service account"
// user so the existing get-or-create/create-session/introspection pipeline handles it
// unchanged — no new tables or code paths required.
userCtx := &UserContext{
UserName: "client:" + client.ClientID,
Email: "oauth-client-" + client.ClientID + "@service.internal",
RemoteID: client.ClientID,
Roles: effectiveScopes,
}
userID, err := s.auth.oauth2GetOrCreateUser(r.Context(), userCtx, "oauth2_client")
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
sessionToken, err := randomOAuthToken()
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
expiresAt := time.Now().Add(s.cfg.AccessTokenTTL)
err = s.auth.oauth2CreateSession(r.Context(), sessionToken, userID, &oauth2.Token{
AccessToken: sessionToken,
TokenType: "Bearer",
}, expiresAt, "oauth2_client")
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
// No refresh token per RFC 6749 §4.4.3, and no id_token — client_credentials has no
// end-user subject to represent in OIDC terms.
s.writeOAuthToken(w, r, sessionToken, "", client.ClientID, effectiveScopes, false)
}
// --------------------------------------------------------------------------
// RFC 7009 — Token revocation
// --------------------------------------------------------------------------
@@ -879,7 +1171,10 @@ func oauthSliceContains(slice []string, s string) bool {
return false
}
func (s *OAuthServer) writeOAuthToken(w http.ResponseWriter, accessToken, refreshToken string, scopes []string) {
// writeOAuthToken writes the token response. When issueIDToken is true and the granted
// scopes include "openid", an RS256-signed id_token is included (OIDC); client_credentials
// responses always pass issueIDToken=false since that grant has no end-user subject.
func (s *OAuthServer) writeOAuthToken(w http.ResponseWriter, r *http.Request, accessToken, refreshToken, clientID string, scopes []string, issueIDToken bool) {
expiresIn := int64(s.cfg.AccessTokenTTL.Seconds())
resp := map[string]interface{}{
"access_token": accessToken,
@@ -892,12 +1187,106 @@ func (s *OAuthServer) writeOAuthToken(w http.ResponseWriter, accessToken, refres
if len(scopes) > 0 {
resp["scope"] = strings.Join(scopes, " ")
}
if issueIDToken && oauthSliceContains(scopes, "openid") {
if idToken, err := s.buildIDToken(r.Context(), accessToken, clientID, scopes); err == nil {
resp["id_token"] = idToken
}
}
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Pragma", "no-cache")
json.NewEncoder(w).Encode(resp) //nolint:errcheck
}
// buildIDToken issues an OIDC id_token for the just-issued access token by reusing the
// existing introspection pipeline to resolve the subject's claims.
func (s *OAuthServer) buildIDToken(ctx context.Context, accessToken, clientID string, scopes []string) (string, error) {
if s.signingKey == nil {
return "", fmt.Errorf("no signing key configured")
}
authToUse := s.auth
if authToUse == nil {
s.mu.RLock()
if len(s.providers) > 0 {
authToUse = s.providers[0].auth
}
s.mu.RUnlock()
}
if authToUse == nil {
return "", fmt.Errorf("no authenticator configured")
}
info, err := authToUse.OAuthIntrospectToken(ctx, accessToken)
if err != nil || !info.Active {
return "", fmt.Errorf("token not active")
}
now := time.Now()
claims := jwt.MapClaims{
"iss": s.cfg.Issuer,
"sub": info.Sub,
"aud": clientID,
"exp": now.Add(s.cfg.AccessTokenTTL).Unix(),
"iat": now.Unix(),
}
if oauthSliceContains(scopes, "profile") && info.Username != "" {
claims["preferred_username"] = info.Username
}
if oauthSliceContains(scopes, "email") && info.Email != "" {
claims["email"] = info.Email
}
token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims)
token.Header["kid"] = s.signingKeyID
return token.SignedString(s.signingKey)
}
// authenticateClient validates client_secret_basic (Authorization: Basic) or
// client_secret_post (client_id/client_secret form fields) credentials against a
// registered confidential client's stored secret hash.
func (s *OAuthServer) authenticateClient(r *http.Request) (*oauthClient, error) {
clientID, clientSecret, ok := r.BasicAuth()
if !ok {
clientID = r.FormValue("client_id")
clientSecret = r.FormValue("client_secret")
}
if clientID == "" || clientSecret == "" {
return nil, fmt.Errorf("client authentication required")
}
client, ok := s.lookupOrFetchClient(r.Context(), clientID)
if !ok || !client.isConfidential() {
return nil, fmt.Errorf("invalid client credentials")
}
if subtle.ConstantTimeCompare([]byte(hashClientSecret(clientSecret)), []byte(client.ClientSecretHash)) != 1 {
return nil, fmt.Errorf("invalid client credentials")
}
return client, nil
}
func hashClientSecret(secret string) string {
sum := sha256.Sum256([]byte(secret))
return hex.EncodeToString(sum[:])
}
// rsaKeyID derives a stable JWKS "kid" from an RSA public key's modulus.
func rsaKeyID(pub *rsa.PublicKey) string {
sum := sha256.Sum256(pub.N.Bytes())
return base64.RawURLEncoding.EncodeToString(sum[:8])
}
// bigEndianBytes encodes a small positive int (e.g. an RSA public exponent) as
// minimal big-endian bytes for JWK "e" encoding.
func bigEndianBytes(n int) []byte {
if n == 0 {
return []byte{0}
}
var b []byte
for n > 0 {
b = append([]byte{byte(n & 0xff)}, b...)
n >>= 8
}
return b
}
func writeOAuthError(w http.ResponseWriter, errCode, description string, status int) {
resp := map[string]string{"error": errCode}
if description != "" {
+31 -5
View File
@@ -9,11 +9,13 @@ import (
// OAuthServerClient is a persisted RFC 7591 registered OAuth2 client.
type OAuthServerClient struct {
ClientID string `json:"client_id"`
RedirectURIs []string `json:"redirect_uris"`
ClientName string `json:"client_name,omitempty"`
GrantTypes []string `json:"grant_types"`
AllowedScopes []string `json:"allowed_scopes,omitempty"`
ClientID string `json:"client_id"`
RedirectURIs []string `json:"redirect_uris"`
ClientName string `json:"client_name,omitempty"`
GrantTypes []string `json:"grant_types"`
AllowedScopes []string `json:"allowed_scopes,omitempty"`
ClientSecretHash string `json:"client_secret_hash,omitempty"`
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
}
// OAuthCode is a short-lived authorization code.
@@ -44,6 +46,10 @@ type OAuthTokenInfo struct {
// OAuthRegisterClient persists an OAuth2 client registration.
func (a *DatabaseAuthenticator) OAuthRegisterClient(ctx context.Context, client *OAuthServerClient) (*OAuthServerClient, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthRegisterClient) {
return a.oauthRegisterClientDirect(ctx, client)
}
input, err := json.Marshal(client)
if err != nil {
return nil, fmt.Errorf("failed to marshal client: %w", err)
@@ -76,6 +82,10 @@ func (a *DatabaseAuthenticator) OAuthRegisterClient(ctx context.Context, client
// OAuthGetClient retrieves a registered client by ID.
func (a *DatabaseAuthenticator) OAuthGetClient(ctx context.Context, clientID string) (*OAuthServerClient, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthGetClient) {
return a.oauthGetClientDirect(ctx, clientID)
}
var success bool
var errMsg *string
var data []byte
@@ -103,6 +113,10 @@ func (a *DatabaseAuthenticator) OAuthGetClient(ctx context.Context, clientID str
// OAuthSaveCode persists an authorization code.
func (a *DatabaseAuthenticator) OAuthSaveCode(ctx context.Context, code *OAuthCode) error {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthSaveCode) {
return a.oauthSaveCodeDirect(ctx, code)
}
input, err := json.Marshal(code)
if err != nil {
return fmt.Errorf("failed to marshal code: %w", err)
@@ -129,6 +143,10 @@ func (a *DatabaseAuthenticator) OAuthSaveCode(ctx context.Context, code *OAuthCo
// OAuthExchangeCode retrieves and deletes an authorization code (single use).
func (a *DatabaseAuthenticator) OAuthExchangeCode(ctx context.Context, code string) (*OAuthCode, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthExchangeCode) {
return a.oauthExchangeCodeDirect(ctx, code)
}
var success bool
var errMsg *string
var data []byte
@@ -157,6 +175,10 @@ func (a *DatabaseAuthenticator) OAuthExchangeCode(ctx context.Context, code stri
// OAuthIntrospectToken validates a token and returns its metadata (RFC 7662).
func (a *DatabaseAuthenticator) OAuthIntrospectToken(ctx context.Context, token string) (*OAuthTokenInfo, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthIntrospect) {
return a.oauthIntrospectTokenDirect(ctx, token)
}
var success bool
var errMsg *string
var data []byte
@@ -184,6 +206,10 @@ func (a *DatabaseAuthenticator) OAuthIntrospectToken(ctx context.Context, token
// OAuthRevokeToken revokes a token by deleting the session (RFC 7009).
func (a *DatabaseAuthenticator) OAuthRevokeToken(ctx context.Context, token string) error {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.OAuthRevoke) {
return a.oauthRevokeTokenDirect(ctx, token)
}
var success bool
var errMsg *string
+209
View File
@@ -0,0 +1,209 @@
package security
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"time"
)
// Direct-mode implementations mirroring the OAuth2 server stored procedures
// (resolvespec_oauth_register_client, etc.) in database_schema.sql, using
// plain SQL against TableNames.OAuthClients / TableNames.OAuthCodes.
// Array columns (redirect_uris, grant_types, allowed_scopes, scopes) are
// JSON-encoded TEXT instead of native Postgres arrays.
// nullIfEmpty converts an empty string to a SQL NULL so optional TEXT columns
// (e.g. client_secret_hash for public clients) stay unset rather than "".
func nullIfEmpty(s string) any {
if s == "" {
return nil
}
return s
}
func (a *DatabaseAuthenticator) oauthRegisterClientDirect(ctx context.Context, client *OAuthServerClient) (*OAuthServerClient, error) {
grantTypes := client.GrantTypes
if len(grantTypes) == 0 {
grantTypes = []string{"authorization_code"}
}
allowedScopes := client.AllowedScopes
if len(allowedScopes) == 0 {
allowedScopes = []string{"openid", "profile", "email"}
}
redirectURIsJSON, err := json.Marshal(client.RedirectURIs)
if err != nil {
return nil, fmt.Errorf("failed to marshal redirect_uris: %w", err)
}
grantTypesJSON, err := json.Marshal(grantTypes)
if err != nil {
return nil, fmt.Errorf("failed to marshal grant_types: %w", err)
}
allowedScopesJSON, err := json.Marshal(allowedScopes)
if err != nil {
return nil, fmt.Errorf("failed to marshal allowed_scopes: %w", err)
}
authMethod := client.TokenEndpointAuthMethod
if authMethod == "" {
authMethod = "none"
}
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method, is_active, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
a.tableNames.OAuthClients))
_, err := db.ExecContext(ctx, query, client.ClientID, string(redirectURIsJSON), client.ClientName, string(grantTypesJSON), string(allowedScopesJSON), nullIfEmpty(client.ClientSecretHash), authMethod, true, time.Now())
return err
})
if err != nil {
return nil, fmt.Errorf("failed to register client: %w", err)
}
return &OAuthServerClient{
ClientID: client.ClientID,
RedirectURIs: client.RedirectURIs,
ClientName: client.ClientName,
GrantTypes: grantTypes,
AllowedScopes: allowedScopes,
ClientSecretHash: client.ClientSecretHash,
TokenEndpointAuthMethod: authMethod,
}, nil
}
func (a *DatabaseAuthenticator) oauthGetClientDirect(ctx context.Context, clientID string) (*OAuthServerClient, error) {
var redirectURIsJSON, grantTypesJSON, allowedScopesJSON sql.NullString
var clientName, clientSecretHash, authMethod sql.NullString
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`SELECT redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method FROM %s WHERE client_id = ? AND is_active = ?`,
a.tableNames.OAuthClients))
return db.QueryRowContext(ctx, query, clientID, true).Scan(&redirectURIsJSON, &clientName, &grantTypesJSON, &allowedScopesJSON, &clientSecretHash, &authMethod)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("client not found")
}
return nil, fmt.Errorf("failed to get client: %w", err)
}
result := &OAuthServerClient{
ClientID: clientID,
ClientName: clientName.String,
ClientSecretHash: clientSecretHash.String,
TokenEndpointAuthMethod: authMethod.String,
}
if redirectURIsJSON.Valid {
_ = json.Unmarshal([]byte(redirectURIsJSON.String), &result.RedirectURIs)
}
if grantTypesJSON.Valid {
_ = json.Unmarshal([]byte(grantTypesJSON.String), &result.GrantTypes)
}
if allowedScopesJSON.Valid {
_ = json.Unmarshal([]byte(allowedScopesJSON.String), &result.AllowedScopes)
}
return result, nil
}
func (a *DatabaseAuthenticator) oauthSaveCodeDirect(ctx context.Context, code *OAuthCode) error {
scopesJSON, err := json.Marshal(code.Scopes)
if err != nil {
return fmt.Errorf("failed to marshal scopes: %w", err)
}
method := code.CodeChallengeMethod
if method == "" {
method = "S256"
}
return a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, a.tableNames.OAuthCodes))
_, err := db.ExecContext(ctx, query, code.Code, code.ClientID, code.RedirectURI, code.ClientState, code.CodeChallenge,
method, code.SessionToken, code.RefreshToken, string(scopesJSON), code.ExpiresAt, time.Now())
return err
})
}
func (a *DatabaseAuthenticator) oauthExchangeCodeDirect(ctx context.Context, code string) (*OAuthCode, error) {
var result OAuthCode
var clientState, refreshToken sql.NullString
var scopesJSON sql.NullString
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`SELECT client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes
FROM %s WHERE code = ? AND expires_at > ?`, a.tableNames.OAuthCodes))
err := db.QueryRowContext(ctx, query, code, time.Now()).Scan(
&result.ClientID, &result.RedirectURI, &clientState, &result.CodeChallenge, &result.CodeChallengeMethod, &result.SessionToken, &refreshToken, &scopesJSON)
if err != nil {
return err
}
delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE code = ?`, a.tableNames.OAuthCodes))
_, err = db.ExecContext(ctx, delQuery, code)
return err
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("invalid or expired code")
}
return nil, fmt.Errorf("failed to exchange code: %w", err)
}
result.Code = code
result.ClientState = clientState.String
result.RefreshToken = refreshToken.String
if scopesJSON.Valid {
_ = json.Unmarshal([]byte(scopesJSON.String), &result.Scopes)
}
return &result, nil
}
func (a *DatabaseAuthenticator) oauthIntrospectTokenDirect(ctx context.Context, token string) (*OAuthTokenInfo, error) {
var info OAuthTokenInfo
var roles sql.NullString
var exp, iat sql.NullTime
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`SELECT u.id, u.username, u.email, u.user_level, u.roles, s.expires_at, s.created_at
FROM %s s JOIN %s u ON u.id = s.user_id
WHERE s.session_token = ? AND s.expires_at > ? AND u.is_active = ?`,
a.tableNames.UserSessions, a.tableNames.Users))
var userID int
err := db.QueryRowContext(ctx, query, token, time.Now(), true).Scan(&userID, &info.Username, &info.Email, &info.UserLevel, &roles, &exp, &iat)
if err != nil {
return err
}
info.Sub = fmt.Sprintf("%d", userID)
return nil
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return &OAuthTokenInfo{Active: false}, nil
}
return nil, fmt.Errorf("failed to introspect token: %w", err)
}
info.Active = true
info.Roles = parseRoles(roles.String)
if exp.Valid {
info.Exp = exp.Time.Unix()
}
if iat.Valid {
info.Iat = iat.Time.Unix()
}
return &info, nil
}
func (a *DatabaseAuthenticator) oauthRevokeTokenDirect(ctx context.Context, token string) error {
return a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE session_token = ?`, a.tableNames.UserSessions))
_, err := db.ExecContext(ctx, query, token)
return err
})
}
+483
View File
@@ -0,0 +1,483 @@
package security
import (
"context"
"crypto/rsa"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"math/big"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"github.com/golang-jwt/jwt/v5"
)
func newTestOAuthServer(t *testing.T) (*OAuthServer, *DatabaseAuthenticator) {
t.Helper()
db := newDirectTestDB(t)
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect})
srv := NewOAuthServer(OAuthServerConfig{Issuer: "https://auth.example.com", PersistCodes: true}, auth)
t.Cleanup(srv.Close)
return srv, auth
}
// s256Challenge computes the PKCE S256 code_challenge for a given verifier,
// matching validatePKCESHA256 in oauth_server.go.
func s256Challenge(verifier string) string {
h := sha256.Sum256([]byte(verifier))
return base64.RawURLEncoding.EncodeToString(h[:])
}
func doJSON(t *testing.T, mux http.Handler, method, path string, body map[string]interface{}) (*httptest.ResponseRecorder, map[string]interface{}) {
t.Helper()
var reqBody *strings.Reader
if body != nil {
b, err := json.Marshal(body)
if err != nil {
t.Fatalf("marshal body: %v", err)
}
reqBody = strings.NewReader(string(b))
} else {
reqBody = strings.NewReader("")
}
req := httptest.NewRequest(method, path, reqBody)
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
mux.ServeHTTP(rec, req)
var parsed map[string]interface{}
if rec.Body.Len() > 0 {
_ = json.Unmarshal(rec.Body.Bytes(), &parsed)
}
return rec, parsed
}
func doForm(t *testing.T, mux http.Handler, path string, form url.Values, basicUser, basicPass string) (*httptest.ResponseRecorder, map[string]interface{}) {
t.Helper()
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
if basicUser != "" {
req.SetBasicAuth(basicUser, basicPass)
}
rec := httptest.NewRecorder()
mux.ServeHTTP(rec, req)
var parsed map[string]interface{}
if rec.Body.Len() > 0 {
_ = json.Unmarshal(rec.Body.Bytes(), &parsed)
}
return rec, parsed
}
func doGet(mux http.Handler, path, bearer string) (*httptest.ResponseRecorder, map[string]interface{}) {
req := httptest.NewRequest(http.MethodGet, path, nil)
if bearer != "" {
req.Header.Set("Authorization", "Bearer "+bearer)
}
rec := httptest.NewRecorder()
mux.ServeHTTP(rec, req)
var parsed map[string]interface{}
if rec.Body.Len() > 0 {
_ = json.Unmarshal(rec.Body.Bytes(), &parsed)
}
return rec, parsed
}
func TestOAuthServer_RegisterConfidentialClient_IssuesSecretOnce(t *testing.T) {
srv, _ := newTestOAuthServer(t)
mux := srv.HTTPHandler()
rec, resp := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{"https://app.example.com/callback"},
"grant_types": []string{"client_credentials"},
})
if rec.Code != http.StatusCreated {
t.Fatalf("register status = %d, body = %s", rec.Code, rec.Body.String())
}
secret, _ := resp["client_secret"].(string)
if secret == "" {
t.Fatal("expected client_secret to be present in registration response")
}
if _, ok := resp["client_secret_hash"]; ok {
t.Error("client_secret_hash must never be returned in the registration response")
}
if resp["token_endpoint_auth_method"] != "client_secret_basic" {
t.Errorf("token_endpoint_auth_method = %v, want client_secret_basic", resp["token_endpoint_auth_method"])
}
if resp["client_id"] == "" || resp["client_id"] == nil {
t.Fatal("expected non-empty client_id")
}
// Fetching the client back (e.g. via a later authorize/token call) must never leak the hash.
clientID := resp["client_id"].(string)
fetched, ok := srv.lookupOrFetchClient(context.Background(), clientID)
if !ok {
t.Fatal("expected client to be found")
}
if fetched.ClientSecretHash == "" {
t.Error("expected ClientSecretHash to be stored internally")
}
if fetched.ClientSecretHash == secret {
t.Error("stored hash must not equal the plaintext secret")
}
// Public registration (no client_credentials) should stay a public client with no secret.
recPublic, respPublic := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{"https://app.example.com/callback"},
})
if recPublic.Code != http.StatusCreated {
t.Fatalf("public register status = %d", recPublic.Code)
}
if _, ok := respPublic["client_secret"]; ok {
t.Error("public client registration should not receive a client_secret")
}
if respPublic["token_endpoint_auth_method"] != "none" {
t.Errorf("public client token_endpoint_auth_method = %v, want none", respPublic["token_endpoint_auth_method"])
}
}
func TestOAuthServer_ClientCredentialsGrant(t *testing.T) {
srv, _ := newTestOAuthServer(t)
mux := srv.HTTPHandler()
_, reg := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{"https://app.example.com/callback"},
"grant_types": []string{"client_credentials"},
"allowed_scopes": []string{"read", "write"},
})
clientID := reg["client_id"].(string)
clientSecret := reg["client_secret"].(string)
t.Run("valid credentials", func(t *testing.T) {
rec, resp := doForm(t, mux, "/oauth/token", url.Values{
"grant_type": {"client_credentials"},
"scope": {"read"},
}, clientID, clientSecret)
if rec.Code != http.StatusOK {
t.Fatalf("token status = %d, body = %s", rec.Code, rec.Body.String())
}
if resp["access_token"] == "" || resp["access_token"] == nil {
t.Error("expected non-empty access_token")
}
if _, ok := resp["refresh_token"]; ok {
t.Error("client_credentials must not issue a refresh_token (RFC 6749 §4.4.3)")
}
if resp["scope"] != "read" {
t.Errorf("scope = %v, want read", resp["scope"])
}
})
t.Run("scope not allowed for client", func(t *testing.T) {
rec, resp := doForm(t, mux, "/oauth/token", url.Values{
"grant_type": {"client_credentials"},
"scope": {"admin"},
}, clientID, clientSecret)
if rec.Code != http.StatusBadRequest {
t.Fatalf("status = %d, want 400", rec.Code)
}
if resp["error"] != "invalid_scope" {
t.Errorf("error = %v, want invalid_scope", resp["error"])
}
})
t.Run("wrong secret", func(t *testing.T) {
rec, resp := doForm(t, mux, "/oauth/token", url.Values{
"grant_type": {"client_credentials"},
}, clientID, "wrong-secret")
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401", rec.Code)
}
if resp["error"] != "invalid_client" {
t.Errorf("error = %v, want invalid_client", resp["error"])
}
})
t.Run("public client cannot use client_credentials", func(t *testing.T) {
_, pubReg := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{"https://app.example.com/callback"},
})
pubClientID := pubReg["client_id"].(string)
rec, _ := doForm(t, mux, "/oauth/token", url.Values{
"grant_type": {"client_credentials"},
"client_id": {pubClientID},
}, "", "")
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401 for public client attempting client_credentials", rec.Code)
}
})
}
func TestOAuthServer_ConfidentialClient_AuthCodeGrantRequiresSecret(t *testing.T) {
srv, auth := newTestOAuthServer(t)
mux := srv.HTTPHandler()
regResp, err := auth.Register(context.Background(), RegisterRequest{
Username: "nadia", Password: "p", Email: "nadia@example.com",
})
if err != nil {
t.Fatalf("Register() error = %v", err)
}
_, reg := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{"https://app.example.com/callback"},
"token_endpoint_auth_method": "client_secret_basic",
})
clientID := reg["client_id"].(string)
clientSecret := reg["client_secret"].(string)
verifier := "verifier-nadia-1234567890"
challenge := s256Challenge(verifier)
if err := auth.OAuthSaveCode(context.Background(), &OAuthCode{
Code: "code-no-auth",
ClientID: clientID,
RedirectURI: "https://app.example.com/callback",
CodeChallenge: challenge,
SessionToken: regResp.Token,
Scopes: []string{"profile"},
ExpiresAt: futureTime(),
}); err != nil {
t.Fatalf("OAuthSaveCode() error = %v", err)
}
// No client credentials supplied -> must be rejected for a confidential client.
rec, resp := doForm(t, mux, "/oauth/token", url.Values{
"grant_type": {"authorization_code"},
"code": {"code-no-auth"},
"redirect_uri": {"https://app.example.com/callback"},
"client_id": {clientID},
"code_verifier": {verifier},
}, "", "")
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401 without client auth, body=%s", rec.Code, rec.Body.String())
}
if resp["error"] != "invalid_client" {
t.Errorf("error = %v, want invalid_client", resp["error"])
}
// Same code, now with correct client credentials -> succeeds.
if err := auth.OAuthSaveCode(context.Background(), &OAuthCode{
Code: "code-with-auth",
ClientID: clientID,
RedirectURI: "https://app.example.com/callback",
CodeChallenge: challenge,
SessionToken: regResp.Token,
Scopes: []string{"profile"},
ExpiresAt: futureTime(),
}); err != nil {
t.Fatalf("OAuthSaveCode() error = %v", err)
}
rec2, _ := doForm(t, mux, "/oauth/token", url.Values{
"grant_type": {"authorization_code"},
"code": {"code-with-auth"},
"redirect_uri": {"https://app.example.com/callback"},
"client_id": {clientID},
"code_verifier": {verifier},
}, clientID, clientSecret)
if rec2.Code != http.StatusOK {
t.Fatalf("status = %d, want 200 with client auth, body=%s", rec2.Code, rec2.Body.String())
}
}
func TestOAuthServer_PublicClient_AuthCodeGrantUnaffected(t *testing.T) {
srv, auth := newTestOAuthServer(t)
mux := srv.HTTPHandler()
regResp, err := auth.Register(context.Background(), RegisterRequest{
Username: "oscar", Password: "p", Email: "oscar@example.com",
})
if err != nil {
t.Fatalf("Register() error = %v", err)
}
_, reg := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{"https://app.example.com/callback"},
})
clientID := reg["client_id"].(string)
if _, ok := reg["client_secret"]; ok {
t.Fatal("expected no client_secret for a default (public) registration")
}
verifier := "verifier-oscar-1234567890"
challenge := s256Challenge(verifier)
if err := auth.OAuthSaveCode(context.Background(), &OAuthCode{
Code: "public-code",
ClientID: clientID,
RedirectURI: "https://app.example.com/callback",
CodeChallenge: challenge,
SessionToken: regResp.Token,
Scopes: []string{"profile"},
ExpiresAt: futureTime(),
}); err != nil {
t.Fatalf("OAuthSaveCode() error = %v", err)
}
// No credentials needed for a public client — PKCE alone is sufficient, unchanged.
rec, _ := doForm(t, mux, "/oauth/token", url.Values{
"grant_type": {"authorization_code"},
"code": {"public-code"},
"redirect_uri": {"https://app.example.com/callback"},
"client_id": {clientID},
"code_verifier": {verifier},
}, "", "")
if rec.Code != http.StatusOK {
t.Fatalf("status = %d, want 200 for public client without credentials, body=%s", rec.Code, rec.Body.String())
}
}
func TestOAuthServer_ProtectedResourceMetadata(t *testing.T) {
srv, _ := newTestOAuthServer(t)
mux := srv.HTTPHandler()
rec, resp := doGet(mux, "/.well-known/oauth-protected-resource", "")
if rec.Code != http.StatusOK {
t.Fatalf("status = %d", rec.Code)
}
if resp["resource"] != "https://auth.example.com" {
t.Errorf("resource = %v", resp["resource"])
}
servers, ok := resp["authorization_servers"].([]interface{})
if !ok || len(servers) != 1 || servers[0] != "https://auth.example.com" {
t.Errorf("authorization_servers = %v", resp["authorization_servers"])
}
}
func TestOAuthServer_OIDCDiscoveryAndIDToken(t *testing.T) {
srv, auth := newTestOAuthServer(t)
mux := srv.HTTPHandler()
// Discovery document.
rec, disc := doGet(mux, "/.well-known/openid-configuration", "")
if rec.Code != http.StatusOK {
t.Fatalf("discovery status = %d", rec.Code)
}
if disc["jwks_uri"] != "https://auth.example.com/oauth/jwks.json" {
t.Errorf("jwks_uri = %v", disc["jwks_uri"])
}
grantTypes, _ := disc["grant_types_supported"].([]interface{})
found := false
for _, g := range grantTypes {
if g == "client_credentials" {
found = true
}
}
if !found {
t.Errorf("expected client_credentials in grant_types_supported, got %v", grantTypes)
}
// JWKS.
rec, _ = doGet(mux, "/oauth/jwks.json", "")
var jwks struct {
Keys []struct {
Kid string `json:"kid"`
N string `json:"n"`
E string `json:"e"`
} `json:"keys"`
}
_ = json.Unmarshal(rec.Body.Bytes(), &jwks)
if len(jwks.Keys) != 1 {
t.Fatalf("expected 1 JWKS key, got %d", len(jwks.Keys))
}
// End-to-end authorization_code flow with scope=openid, verifying the id_token
// signature against the published JWKS key.
regResp, err := auth.Register(context.Background(), RegisterRequest{
Username: "olivia", Password: "p", Email: "olivia@example.com",
})
if err != nil {
t.Fatalf("Register() error = %v", err)
}
_, clientReg := doJSON(t, mux, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{"https://app.example.com/callback"},
})
clientID := clientReg["client_id"].(string)
verifier := "test-code-verifier-1234567890"
challenge := s256Challenge(verifier)
if err := auth.OAuthSaveCode(context.Background(), &OAuthCode{
Code: "oidc-code",
ClientID: clientID,
RedirectURI: "https://app.example.com/callback",
CodeChallenge: challenge,
SessionToken: regResp.Token,
Scopes: []string{"openid", "profile", "email"},
ExpiresAt: futureTime(),
}); err != nil {
t.Fatalf("OAuthSaveCode() error = %v", err)
}
tokRec, tokenResp := doForm(t, mux, "/oauth/token", url.Values{
"grant_type": {"authorization_code"},
"code": {"oidc-code"},
"redirect_uri": {"https://app.example.com/callback"},
"client_id": {clientID},
"code_verifier": {verifier},
}, "", "")
if tokRec.Code != http.StatusOK {
t.Fatalf("token exchange status = %d, body = %s", tokRec.Code, tokRec.Body.String())
}
idTokenStr, _ := tokenResp["id_token"].(string)
if idTokenStr == "" {
t.Fatal("expected id_token in response for scope containing openid")
}
nBytes, err := base64.RawURLEncoding.DecodeString(jwks.Keys[0].N)
if err != nil {
t.Fatalf("decode n: %v", err)
}
eBytes, err := base64.RawURLEncoding.DecodeString(jwks.Keys[0].E)
if err != nil {
t.Fatalf("decode e: %v", err)
}
pub := &rsa.PublicKey{N: new(big.Int).SetBytes(nBytes), E: int(new(big.Int).SetBytes(eBytes).Int64())}
parsed, err := jwt.Parse(idTokenStr, func(token *jwt.Token) (interface{}, error) {
return pub, nil
}, jwt.WithValidMethods([]string{"RS256"}))
if err != nil || !parsed.Valid {
t.Fatalf("id_token did not verify against JWKS key: %v", err)
}
claims := parsed.Claims.(jwt.MapClaims)
if claims["iss"] != "https://auth.example.com" {
t.Errorf("iss claim = %v", claims["iss"])
}
if claims["aud"] != clientID {
t.Errorf("aud claim = %v, want %v", claims["aud"], clientID)
}
if claims["preferred_username"] != "olivia" {
t.Errorf("preferred_username claim = %v", claims["preferred_username"])
}
if claims["email"] != "olivia@example.com" {
t.Errorf("email claim = %v", claims["email"])
}
}
func TestOAuthServer_Userinfo(t *testing.T) {
srv, auth := newTestOAuthServer(t)
mux := srv.HTTPHandler()
regResp, err := auth.Register(context.Background(), RegisterRequest{
Username: "pete", Password: "p", Email: "pete@example.com",
})
if err != nil {
t.Fatalf("Register() error = %v", err)
}
rec, info := doGet(mux, "/oauth/userinfo", regResp.Token)
if rec.Code != http.StatusOK {
t.Fatalf("userinfo status = %d, body = %s", rec.Code, rec.Body.String())
}
if info["preferred_username"] != "pete" {
t.Errorf("preferred_username = %v", info["preferred_username"])
}
rec2, _ := doGet(mux, "/oauth/userinfo", "not-a-real-token")
if rec2.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401 for invalid token", rec2.Code)
}
}
+132 -39
View File
@@ -14,14 +14,17 @@ import (
// DatabasePasskeyProvider implements PasskeyProvider using database storage
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
type DatabasePasskeyProvider struct {
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
rpID string // Relying Party ID (domain)
rpName string // Relying Party display name
rpOrigin string // Expected origin for WebAuthn
timeout int64 // Timeout in milliseconds (default: 60000)
sqlNames *SQLNames
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
rpID string // Relying Party ID (domain)
rpName string // Relying Party display name
rpOrigin string // Expected origin for WebAuthn
timeout int64 // Timeout in milliseconds (default: 60000)
sqlNames *SQLNames
tableNames *TableNames
queryMode QueryMode
capability *dbCapability
}
// DatabasePasskeyProviderOptions configures the passkey provider
@@ -36,6 +39,10 @@ type DatabasePasskeyProviderOptions struct {
Timeout int64
// SQLNames provides custom SQL procedure/function names. If nil, uses DefaultSQLNames().
SQLNames *SQLNames
// TableNames provides custom table names for Direct mode. If nil, uses DefaultTableNames().
TableNames *TableNames
// QueryMode selects stored-procedure vs Direct-mode SQL. Default (zero value) is ModeAuto.
QueryMode QueryMode
// 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)
@@ -48,15 +55,19 @@ func NewDatabasePasskeyProvider(db *sql.DB, opts DatabasePasskeyProviderOptions)
}
sqlNames := MergeSQLNames(DefaultSQLNames(), opts.SQLNames)
tableNames := resolveTableNames(opts.TableNames)
return &DatabasePasskeyProvider{
db: db,
dbFactory: opts.DBFactory,
rpID: opts.RPID,
rpName: opts.RPName,
rpOrigin: opts.RPOrigin,
timeout: opts.Timeout,
sqlNames: sqlNames,
db: db,
dbFactory: opts.DBFactory,
rpID: opts.RPID,
rpName: opts.RPName,
rpOrigin: opts.RPOrigin,
timeout: opts.Timeout,
sqlNames: sqlNames,
tableNames: tableNames,
queryMode: opts.QueryMode,
capability: newDBCapability(),
}
}
@@ -77,9 +88,26 @@ func (p *DatabasePasskeyProvider) reconnectDB() error {
p.dbMu.Lock()
p.db = newDB
p.dbMu.Unlock()
if p.capability != nil {
p.capability.reset()
}
return nil
}
func (p *DatabasePasskeyProvider) runDBOpWithReconnect(run func(*sql.DB) error) error {
db := p.getDB()
if db == nil {
return fmt.Errorf("database connection is nil")
}
err := run(db)
if isDBClosed(err) {
if reconnErr := p.reconnectDB(); reconnErr == nil {
err = run(p.getDB())
}
}
return err
}
// BeginRegistration creates registration options for a new passkey
func (p *DatabasePasskeyProvider) BeginRegistration(ctx context.Context, userID int, username, displayName string) (*PasskeyRegistrationOptions, error) {
// Generate challenge
@@ -145,10 +173,40 @@ func (p *DatabasePasskeyProvider) CompleteRegistration(ctx context.Context, user
// For now, this is a placeholder that stores the credential data
// In production, you MUST use a proper WebAuthn library
credIDB64 := base64.StdEncoding.EncodeToString(response.RawID)
pubKeyB64 := base64.StdEncoding.EncodeToString(response.Response.AttestationObject)
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyStoreCredential) {
credentialID, err := p.storeCredentialDirect(ctx, storeCredentialParams{
UserID: userID,
CredentialID: credIDB64,
PublicKey: pubKeyB64,
AttestationType: "none",
SignCount: 0,
Transports: response.Transports,
BackupEligible: false,
BackupState: false,
Name: "Passkey",
})
if err != nil {
return nil, err
}
return &PasskeyCredential{
ID: fmt.Sprintf("%d", credentialID),
UserID: userID,
CredentialID: response.RawID,
PublicKey: response.Response.AttestationObject,
AttestationType: "none",
Transports: response.Transports,
CreatedAt: time.Now(),
LastUsedAt: time.Now(),
}, nil
}
credData := map[string]any{
"user_id": userID,
"credential_id": base64.StdEncoding.EncodeToString(response.RawID),
"public_key": base64.StdEncoding.EncodeToString(response.Response.AttestationObject),
"credential_id": credIDB64,
"public_key": pubKeyB64,
"attestation_type": "none",
"sign_count": 0,
"transports": response.Transports,
@@ -202,31 +260,36 @@ func (p *DatabasePasskeyProvider) BeginAuthentication(ctx context.Context, usern
// If username is provided, get user's credentials
var allowCredentials []PasskeyCredentialDescriptor
if username != "" {
var success bool
var errorMsg sql.NullString
var userID sql.NullInt64
var credentialsJSON sql.NullString
var creds []passkeyCredential
query := fmt.Sprintf(`SELECT p_success, p_error, p_user_id, p_credentials::text FROM %s($1)`, p.sqlNames.PasskeyGetCredsByUsername)
err := p.getDB().QueryRowContext(ctx, query, username).Scan(&success, &errorMsg, &userID, &credentialsJSON)
if err != nil {
return nil, fmt.Errorf("failed to get credentials: %w", err)
}
if !success {
if errorMsg.Valid {
return nil, fmt.Errorf("%s", errorMsg.String)
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyGetCredsByUsername) {
_, directCreds, err := p.getCredsByUsernameDirect(ctx, username)
if err != nil {
return nil, err
}
return nil, fmt.Errorf("failed to get credentials")
}
creds = directCreds
} else {
var success bool
var errorMsg sql.NullString
var userID sql.NullInt64
var credentialsJSON sql.NullString
// Parse credentials
var creds []struct {
ID string `json:"credential_id"`
Transports []string `json:"transports"`
}
if err := json.Unmarshal([]byte(credentialsJSON.String), &creds); err != nil {
return nil, fmt.Errorf("failed to parse credentials: %w", err)
query := fmt.Sprintf(`SELECT p_success, p_error, p_user_id, p_credentials::text FROM %s($1)`, p.sqlNames.PasskeyGetCredsByUsername)
err := p.getDB().QueryRowContext(ctx, query, username).Scan(&success, &errorMsg, &userID, &credentialsJSON)
if err != nil {
return nil, fmt.Errorf("failed to get credentials: %w", err)
}
if !success {
if errorMsg.Valid {
return nil, fmt.Errorf("%s", errorMsg.String)
}
return nil, fmt.Errorf("failed to get credentials")
}
if err := json.Unmarshal([]byte(credentialsJSON.String), &creds); err != nil {
return nil, fmt.Errorf("failed to parse credentials: %w", err)
}
}
allowCredentials = make([]PasskeyCredentialDescriptor, 0, len(creds))
@@ -262,6 +325,24 @@ func (p *DatabasePasskeyProvider) CompleteAuthentication(ctx context.Context, re
// 3. Verify signature using stored public key
// 4. Update sign counter and check for cloning
credIDB64 := base64.StdEncoding.EncodeToString(response.RawID)
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyGetCredential) {
userID, signCount, err := p.getCredentialDirect(ctx, credIDB64)
if err != nil {
return 0, err
}
newCounter := signCount + 1
cloneWarning, err := p.updateCounterDirect(ctx, credIDB64, newCounter)
if err != nil {
return 0, fmt.Errorf("failed to update counter: %w", err)
}
if cloneWarning {
return 0, fmt.Errorf("credential cloning detected")
}
return userID, nil
}
// Get credential from database
var success bool
var errorMsg sql.NullString
@@ -321,6 +402,10 @@ func (p *DatabasePasskeyProvider) CompleteAuthentication(ctx context.Context, re
// GetCredentials returns all passkey credentials for a user
func (p *DatabasePasskeyProvider) GetCredentials(ctx context.Context, userID int) ([]PasskeyCredential, error) {
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyGetUserCredentials) {
return p.getUserCredentialsDirect(ctx, userID)
}
var success bool
var errorMsg sql.NullString
var credentialsJSON sql.NullString
@@ -401,6 +486,10 @@ func (p *DatabasePasskeyProvider) DeleteCredential(ctx context.Context, userID i
return fmt.Errorf("invalid credential ID: %w", err)
}
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyDeleteCredential) {
return p.deleteCredentialDirect(ctx, userID, credentialID)
}
var success bool
var errorMsg sql.NullString
@@ -427,6 +516,10 @@ func (p *DatabasePasskeyProvider) UpdateCredentialName(ctx context.Context, user
return fmt.Errorf("invalid credential ID: %w", err)
}
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.PasskeyUpdateName) {
return p.updateNameDirect(ctx, userID, credentialID, name)
}
var success bool
var errorMsg sql.NullString
+256
View File
@@ -0,0 +1,256 @@
package security
import (
"context"
"database/sql"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"time"
)
// Direct-mode implementations mirroring the resolvespec_passkey_* stored
// procedures in database_schema.sql using plain SQL against TableNames.
// credential_id/public_key/aaguid are stored as base64 TEXT (not native
// bytea) and transports as JSON-encoded TEXT, so the same schema works on
// SQLite, MySQL, and Postgres.
type storeCredentialParams struct {
UserID int
CredentialID string // base64
PublicKey string // base64
AttestationType string
SignCount int
Transports []string
BackupEligible bool
BackupState bool
Name string
}
func (p *DatabasePasskeyProvider) storeCredentialDirect(ctx context.Context, params storeCredentialParams) (int64, error) {
transportsJSON, err := json.Marshal(params.Transports)
if err != nil {
return 0, fmt.Errorf("failed to marshal transports: %w", err)
}
var credentialID int64
err = p.runDBOpWithReconnect(func(db *sql.DB) error {
var exists int
checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT 1 FROM %s WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials))
if err := db.QueryRowContext(ctx, checkQuery, params.CredentialID).Scan(&exists); err == nil {
return fmt.Errorf("credential already exists")
} else if !errors.Is(err, sql.ErrNoRows) {
return err
}
var userExists int
userCheckQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT 1 FROM %s WHERE id = ?`, p.tableNames.Users))
if err := db.QueryRowContext(ctx, userCheckQuery, params.UserID).Scan(&userExists); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return fmt.Errorf("user not found")
}
return err
}
now := time.Now()
insertQuery := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (user_id, credential_id, public_key, attestation_type, aaguid, sign_count, transports, backup_eligible, backup_state, name, created_at, last_used_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, p.tableNames.UserPasskeyCredentials))
res, err := db.ExecContext(ctx, insertQuery, params.UserID, params.CredentialID, params.PublicKey, params.AttestationType,
"", params.SignCount, string(transportsJSON), params.BackupEligible, params.BackupState, params.Name, now, now)
if err != nil {
return err
}
credentialID, err = res.LastInsertId()
return err
})
if err != nil {
return 0, err
}
return credentialID, nil
}
func (p *DatabasePasskeyProvider) getCredentialDirect(ctx context.Context, credentialIDB64 string) (userID int, signCount uint32, err error) {
err = p.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT user_id, sign_count FROM %s WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials))
return db.QueryRowContext(ctx, query, credentialIDB64).Scan(&userID, &signCount)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return 0, 0, fmt.Errorf("credential not found")
}
return 0, 0, fmt.Errorf("failed to get credential: %w", err)
}
return userID, signCount, nil
}
func (p *DatabasePasskeyProvider) updateCounterDirect(ctx context.Context, credentialIDB64 string, newCounter uint32) (cloneWarning bool, err error) {
err = p.runDBOpWithReconnect(func(db *sql.DB) error {
var oldCounter int
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT sign_count FROM %s WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials))
if err := db.QueryRowContext(ctx, query, credentialIDB64).Scan(&oldCounter); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return fmt.Errorf("credential not found")
}
return err
}
if int(newCounter) <= oldCounter {
cloneWarning = true
updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET clone_warning = ? WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials))
_, err := db.ExecContext(ctx, updQuery, true, credentialIDB64)
return err
}
updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET sign_count = ?, last_used_at = ? WHERE credential_id = ?`, p.tableNames.UserPasskeyCredentials))
_, err := db.ExecContext(ctx, updQuery, newCounter, time.Now(), credentialIDB64)
return err
})
return cloneWarning, err
}
func (p *DatabasePasskeyProvider) getUserCredentialsDirect(ctx context.Context, userID int) ([]PasskeyCredential, error) {
var credentials []PasskeyCredential
err := p.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`SELECT id, user_id, credential_id, public_key, attestation_type, aaguid, sign_count, clone_warning, transports, backup_eligible, backup_state, name, created_at, last_used_at
FROM %s WHERE user_id = ? ORDER BY created_at DESC`, p.tableNames.UserPasskeyCredentials))
rows, err := db.QueryContext(ctx, query, userID)
if err != nil {
return err
}
defer rows.Close()
credentials = make([]PasskeyCredential, 0)
for rows.Next() {
var id, uid int
var credIDB64, pubKeyB64, attestationType, aaguidB64, name string
var signCount uint32
var cloneWarning, backupEligible, backupState bool
var transportsJSON sql.NullString
var createdAt, lastUsedAt time.Time
if err := rows.Scan(&id, &uid, &credIDB64, &pubKeyB64, &attestationType, &aaguidB64, &signCount,
&cloneWarning, &transportsJSON, &backupEligible, &backupState, &name, &createdAt, &lastUsedAt); err != nil {
return err
}
credID, err := base64.StdEncoding.DecodeString(credIDB64)
if err != nil {
continue
}
pubKey, err := base64.StdEncoding.DecodeString(pubKeyB64)
if err != nil {
continue
}
aaguid, _ := base64.StdEncoding.DecodeString(aaguidB64)
var transports []string
if transportsJSON.Valid && transportsJSON.String != "" {
_ = json.Unmarshal([]byte(transportsJSON.String), &transports)
}
credentials = append(credentials, PasskeyCredential{
ID: fmt.Sprintf("%d", id),
UserID: uid,
CredentialID: credID,
PublicKey: pubKey,
AttestationType: attestationType,
AAGUID: aaguid,
SignCount: signCount,
CloneWarning: cloneWarning,
Transports: transports,
BackupEligible: backupEligible,
BackupState: backupState,
Name: name,
CreatedAt: createdAt,
LastUsedAt: lastUsedAt,
})
}
return rows.Err()
})
if err != nil {
return nil, fmt.Errorf("failed to get credentials: %w", err)
}
return credentials, nil
}
func (p *DatabasePasskeyProvider) deleteCredentialDirect(ctx context.Context, userID int, credentialIDB64 string) error {
return p.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ? AND credential_id = ?`, p.tableNames.UserPasskeyCredentials))
res, err := db.ExecContext(ctx, query, userID, credentialIDB64)
if err != nil {
return err
}
rows, err := res.RowsAffected()
if err != nil {
return err
}
if rows == 0 {
return fmt.Errorf("credential not found")
}
return nil
})
}
func (p *DatabasePasskeyProvider) updateNameDirect(ctx context.Context, userID int, credentialIDB64, name string) error {
return p.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET name = ? WHERE user_id = ? AND credential_id = ?`, p.tableNames.UserPasskeyCredentials))
res, err := db.ExecContext(ctx, query, name, userID, credentialIDB64)
if err != nil {
return err
}
rows, err := res.RowsAffected()
if err != nil {
return err
}
if rows == 0 {
return fmt.Errorf("credential not found")
}
return nil
})
}
type passkeyCredential struct {
ID string `json:"credential_id"`
Transports []string `json:"transports"`
}
func (p *DatabasePasskeyProvider) getCredsByUsernameDirect(ctx context.Context, username string) (userID int, creds []passkeyCredential, err error) {
err = p.runDBOpWithReconnect(func(db *sql.DB) error {
userQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT id FROM %s WHERE username = ? AND is_active = ?`, p.tableNames.Users))
if err := db.QueryRowContext(ctx, userQuery, username, true).Scan(&userID); err != nil {
return err
}
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT credential_id, transports FROM %s WHERE user_id = ?`, p.tableNames.UserPasskeyCredentials))
rows, err := db.QueryContext(ctx, query, userID)
if err != nil {
return err
}
defer rows.Close()
creds = make([]passkeyCredential, 0)
for rows.Next() {
var credID string
var transportsJSON sql.NullString
if err := rows.Scan(&credID, &transportsJSON); err != nil {
return err
}
var transports []string
if transportsJSON.Valid && transportsJSON.String != "" {
_ = json.Unmarshal([]byte(transportsJSON.String), &transports)
}
creds = append(creds, passkeyCredential{ID: credID, Transports: transports})
}
return rows.Err()
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return 0, nil, fmt.Errorf("user not found")
}
return 0, nil, fmt.Errorf("failed to get credentials: %w", err)
}
return userID, creds, nil
}
+10 -7
View File
@@ -34,7 +34,10 @@ type RowSecurity struct {
Tablename string `json:"tablename"`
Template string `json:"template"`
HasBlock bool `json:"has_block"`
UserID int `json:"user_id"`
// UserID is the opaque user reference the security rules were loaded for.
// It may be an int, a string/UUID, or a *UserContext, depending on what the
// RowSecurityProvider/SecurityContext.GetUserRef implementation returns.
UserID any `json:"user_id"`
}
func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Type) string {
@@ -42,7 +45,7 @@ func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Typ
str = strings.ReplaceAll(str, "{PrimaryKeyName}", pPrimaryKeyName)
str = strings.ReplaceAll(str, "{TableName}", m.Tablename)
str = strings.ReplaceAll(str, "{SchemaName}", m.Schema)
str = strings.ReplaceAll(str, "{UserID}", fmt.Sprintf("%d", m.UserID))
str = strings.ReplaceAll(str, "{UserID}", fmt.Sprintf("%v", m.UserID))
return str
}
@@ -413,7 +416,7 @@ func (m *SecurityList) ClearSecurity(pUserID int, pSchema, pTablename string) er
return nil
}
func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserID int, pSchema, pTablename string, pOverwrite bool) (RowSecurity, error) {
func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserRef any, pSchema, pTablename string, pOverwrite bool) (RowSecurity, error) {
if m.provider == nil {
return RowSecurity{}, fmt.Errorf("security provider not set")
}
@@ -424,10 +427,10 @@ func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserID int, pSchema
if m.RowSecurity == nil {
m.RowSecurity = make(map[string]RowSecurity, 0)
}
secKey := fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)
secKey := fmt.Sprintf("%s.%s@%v", pSchema, pTablename, pUserRef)
// Call the provider to load security rules
record, err := m.provider.GetRowSecurity(ctx, pUserID, pSchema, pTablename)
record, err := m.provider.GetRowSecurity(ctx, pUserRef, pSchema, pTablename)
if err != nil {
return RowSecurity{}, fmt.Errorf("GetRowSecurity failed: %v", err)
}
@@ -436,7 +439,7 @@ func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserID int, pSchema
return record, nil
}
func (m *SecurityList) GetRowSecurityTemplate(pUserID int, pSchema, pTablename string) (RowSecurity, error) {
func (m *SecurityList) GetRowSecurityTemplate(pUserRef any, pSchema, pTablename string) (RowSecurity, error) {
defer logger.CatchPanic("GetRowSecurityTemplate")()
if m.RowSecurity == nil {
@@ -446,7 +449,7 @@ func (m *SecurityList) GetRowSecurityTemplate(pUserID int, pSchema, pTablename s
m.RowSecurityMutex.RLock()
defer m.RowSecurityMutex.RUnlock()
rowSec, ok := m.RowSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
rowSec, ok := m.RowSecurity[fmt.Sprintf("%s.%s@%v", pSchema, pTablename, pUserRef)]
if !ok {
return RowSecurity{}, fmt.Errorf("no row security data")
}
+1 -1
View File
@@ -44,7 +44,7 @@ func (m *mockSecurityProvider) GetColumnSecurity(ctx context.Context, userID int
return m.columnSecurity, nil
}
func (m *mockSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
func (m *mockSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
return m.rowSecurity, nil
}
+131 -30
View File
@@ -71,12 +71,15 @@ func (a *HeaderAuthenticator) Authenticate(r *http.Request) (*UserContext, error
// Also supports multiple OAuth2 providers configured with WithOAuth2()
// Also supports passkey authentication configured with WithPasskey()
type DatabaseAuthenticator struct {
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
cache *cache.Cache
cacheTTL time.Duration
sqlNames *SQLNames
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
cache *cache.Cache
cacheTTL time.Duration
sqlNames *SQLNames
tableNames *TableNames
queryMode QueryMode
capability *dbCapability
// Cookie session support (optional, gated by enableCookieSession)
enableCookieSession bool
@@ -105,6 +108,10 @@ type DatabaseAuthenticatorOptions struct {
// SQLNames provides custom SQL procedure/function names. If nil, uses DefaultSQLNames().
// Partial overrides are supported: only set the fields you want to change.
SQLNames *SQLNames
// TableNames provides custom table names for Direct mode. If nil, uses DefaultTableNames().
TableNames *TableNames
// QueryMode selects stored-procedure vs Direct-mode SQL. Default (zero value) is ModeAuto.
QueryMode QueryMode
// 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)
@@ -139,6 +146,7 @@ func NewDatabaseAuthenticatorWithOptions(db *sql.DB, opts DatabaseAuthenticatorO
}
sqlNames := MergeSQLNames(DefaultSQLNames(), opts.SQLNames)
tableNames := resolveTableNames(opts.TableNames)
return &DatabaseAuthenticator{
db: db,
@@ -146,6 +154,9 @@ func NewDatabaseAuthenticatorWithOptions(db *sql.DB, opts DatabaseAuthenticatorO
cache: cacheInstance,
cacheTTL: opts.CacheTTL,
sqlNames: sqlNames,
tableNames: tableNames,
queryMode: opts.QueryMode,
capability: newDBCapability(),
passkeyProvider: opts.PasskeyProvider,
enableCookieSession: opts.EnableCookieSession,
cookieOptions: opts.CookieOptions,
@@ -170,6 +181,9 @@ func (a *DatabaseAuthenticator) reconnectDB() error {
a.dbMu.Lock()
a.db = newDB
a.dbMu.Unlock()
if a.capability != nil {
a.capability.reset()
}
return nil
}
@@ -194,6 +208,9 @@ func (a *DatabaseAuthenticator) SetAuthenticateCallback(fn func(r *http.Request)
}
func (a *DatabaseAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.Login) {
return a.loginDirect(ctx, req)
}
// Convert LoginRequest to JSON
reqJSON, err := json.Marshal(req)
if err != nil {
@@ -230,6 +247,9 @@ func (a *DatabaseAuthenticator) Login(ctx context.Context, req LoginRequest) (*L
// Register implements Registrable interface
func (a *DatabaseAuthenticator) Register(ctx context.Context, req RegisterRequest) (*LoginResponse, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.Register) {
return a.registerDirect(ctx, req)
}
// Convert RegisterRequest to JSON
reqJSON, err := json.Marshal(req)
if err != nil {
@@ -265,6 +285,9 @@ func (a *DatabaseAuthenticator) Register(ctx context.Context, req RegisterReques
}
func (a *DatabaseAuthenticator) Logout(ctx context.Context, req LogoutRequest) error {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.Logout) {
return a.logoutDirect(ctx, req)
}
// Convert LogoutRequest to JSON
reqJSON, err := json.Marshal(req)
if err != nil {
@@ -378,6 +401,10 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err
var userCtx UserContext
err := a.cache.GetOrSet(r.Context(), cacheKey, &userCtx, a.cacheTTL, func() (any, error) {
// This function is called only if cache miss
if !a.capability.ShouldUseProcedure(r.Context(), a.queryMode, a.getDB(), a.sqlNames.Session) {
return a.sessionDirect(r.Context(), token)
}
var success bool
var errorMsg sql.NullString
var userJSON sql.NullString
@@ -453,6 +480,11 @@ func (a *DatabaseAuthenticator) ClearUserCache(userID int) error {
// updateSessionActivity updates the last activity timestamp for the session
func (a *DatabaseAuthenticator) updateSessionActivity(ctx context.Context, sessionToken string, userCtx *UserContext) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.SessionUpdate) {
_ = a.updateSessionActivityDirect(ctx, sessionToken)
return
}
// Convert UserContext to JSON
userJSON, err := json.Marshal(userCtx)
if err != nil {
@@ -471,6 +503,9 @@ func (a *DatabaseAuthenticator) updateSessionActivity(ctx context.Context, sessi
// RefreshToken implements Refreshable interface
func (a *DatabaseAuthenticator) RefreshToken(ctx context.Context, refreshToken string) (*LoginResponse, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.RefreshToken) {
return a.refreshTokenDirect(ctx, refreshToken)
}
// First, we need to get the current user context for the refresh token
var success bool
var errorMsg sql.NullString
@@ -528,18 +563,23 @@ func (a *DatabaseAuthenticator) RefreshToken(ctx context.Context, refreshToken s
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
// NOTE: JWT signing/verification requires github.com/golang-jwt/jwt/v5 to be installed and imported
type JWTAuthenticator struct {
secretKey []byte
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
sqlNames *SQLNames
secretKey []byte
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
sqlNames *SQLNames
tableNames *TableNames
queryMode QueryMode
capability *dbCapability
}
func NewJWTAuthenticator(secretKey string, db *sql.DB, names ...*SQLNames) *JWTAuthenticator {
return &JWTAuthenticator{
secretKey: []byte(secretKey),
db: db,
sqlNames: resolveSQLNames(names...),
secretKey: []byte(secretKey),
db: db,
sqlNames: resolveSQLNames(names...),
tableNames: DefaultTableNames(),
capability: newDBCapability(),
}
}
@@ -549,6 +589,18 @@ func (a *JWTAuthenticator) WithDBFactory(factory func() (*sql.DB, error)) *JWTAu
return a
}
// WithTableNames configures Direct-mode table names. If names is nil, defaults are used.
func (a *JWTAuthenticator) WithTableNames(names *TableNames) *JWTAuthenticator {
a.tableNames = resolveTableNames(names)
return a
}
// WithQueryMode selects stored-procedure vs Direct-mode SQL (default ModeAuto).
func (a *JWTAuthenticator) WithQueryMode(mode QueryMode) *JWTAuthenticator {
a.queryMode = mode
return a
}
func (a *JWTAuthenticator) getDB() *sql.DB {
a.dbMu.RLock()
defer a.dbMu.RUnlock()
@@ -566,10 +618,17 @@ func (a *JWTAuthenticator) reconnectDB() error {
a.dbMu.Lock()
a.db = newDB
a.dbMu.Unlock()
if a.capability != nil {
a.capability.reset()
}
return nil
}
func (a *JWTAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.JWTLogin) {
return a.jwtLoginDirect(ctx, req)
}
var success bool
var errorMsg sql.NullString
var userJSON []byte
@@ -632,6 +691,10 @@ func (a *JWTAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginR
}
func (a *JWTAuthenticator) Logout(ctx context.Context, req LogoutRequest) error {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.JWTLogout) {
return a.jwtLogoutDirect(ctx, req)
}
var success bool
var errorMsg sql.NullString
@@ -681,14 +744,23 @@ func (a *JWTAuthenticator) Authenticate(r *http.Request) (*UserContext, error) {
// All database operations go through stored procedures
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
type DatabaseColumnSecurityProvider struct {
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
sqlNames *SQLNames
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
sqlNames *SQLNames
queryMode QueryMode
capability *dbCapability
}
func NewDatabaseColumnSecurityProvider(db *sql.DB, names ...*SQLNames) *DatabaseColumnSecurityProvider {
return &DatabaseColumnSecurityProvider{db: db, sqlNames: resolveSQLNames(names...)}
return &DatabaseColumnSecurityProvider{db: db, sqlNames: resolveSQLNames(names...), capability: newDBCapability()}
}
// WithQueryMode selects stored-procedure vs Direct-mode SQL (default ModeAuto).
// Direct mode is unsupported for column security (see ErrDirectModeUnsupported).
func (p *DatabaseColumnSecurityProvider) WithQueryMode(mode QueryMode) *DatabaseColumnSecurityProvider {
p.queryMode = mode
return p
}
func (p *DatabaseColumnSecurityProvider) WithDBFactory(factory func() (*sql.DB, error)) *DatabaseColumnSecurityProvider {
@@ -713,10 +785,17 @@ func (p *DatabaseColumnSecurityProvider) reconnectDB() error {
p.dbMu.Lock()
p.db = newDB
p.dbMu.Unlock()
if p.capability != nil {
p.capability.reset()
}
return nil
}
func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error) {
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.ColumnSecurity) {
return nil, ErrDirectModeUnsupported
}
var rules []ColumnSecurity
var success bool
@@ -781,14 +860,23 @@ func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context,
// All database operations go through stored procedures
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
type DatabaseRowSecurityProvider struct {
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
sqlNames *SQLNames
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
sqlNames *SQLNames
queryMode QueryMode
capability *dbCapability
}
func NewDatabaseRowSecurityProvider(db *sql.DB, names ...*SQLNames) *DatabaseRowSecurityProvider {
return &DatabaseRowSecurityProvider{db: db, sqlNames: resolveSQLNames(names...)}
return &DatabaseRowSecurityProvider{db: db, sqlNames: resolveSQLNames(names...), capability: newDBCapability()}
}
// WithQueryMode selects stored-procedure vs Direct-mode SQL (default ModeAuto).
// Direct mode is unsupported for row security (see ErrDirectModeUnsupported).
func (p *DatabaseRowSecurityProvider) WithQueryMode(mode QueryMode) *DatabaseRowSecurityProvider {
p.queryMode = mode
return p
}
func (p *DatabaseRowSecurityProvider) WithDBFactory(factory func() (*sql.DB, error)) *DatabaseRowSecurityProvider {
@@ -813,16 +901,23 @@ func (p *DatabaseRowSecurityProvider) reconnectDB() error {
p.dbMu.Lock()
p.db = newDB
p.dbMu.Unlock()
if p.capability != nil {
p.capability.reset()
}
return nil
}
func (p *DatabaseRowSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
func (p *DatabaseRowSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.RowSecurity) {
return RowSecurity{}, ErrDirectModeUnsupported
}
var template string
var hasBlock bool
runQuery := func() error {
query := fmt.Sprintf(`SELECT p_template, p_block FROM %s($1, $2, $3)`, p.sqlNames.RowSecurity)
return p.getDB().QueryRowContext(ctx, query, schema, table, userID).Scan(&template, &hasBlock)
return p.getDB().QueryRowContext(ctx, query, schema, table, userRef).Scan(&template, &hasBlock)
}
err := runQuery()
if isDBClosed(err) {
@@ -837,7 +932,7 @@ func (p *DatabaseRowSecurityProvider) GetRowSecurity(ctx context.Context, userID
return RowSecurity{
Schema: schema,
Tablename: table,
UserID: userID,
UserID: userRef,
Template: template,
HasBlock: hasBlock,
}, nil
@@ -874,14 +969,14 @@ func NewConfigRowSecurityProvider(templates map[string]string, blocked map[strin
}
}
func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (RowSecurity, error) {
key := fmt.Sprintf("%s.%s", schema, table)
if p.blocked[key] {
return RowSecurity{
Schema: schema,
Tablename: table,
UserID: userID,
UserID: userRef,
HasBlock: true,
}, nil
}
@@ -890,7 +985,7 @@ func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userID i
return RowSecurity{
Schema: schema,
Tablename: table,
UserID: userID,
UserID: userRef,
Template: template,
HasBlock: false,
}, nil
@@ -950,6 +1045,9 @@ func generateRandomString(length int) string {
// RequestPasswordReset implements PasswordResettable. It calls the stored procedure
// resolvespec_password_reset_request and returns the reset token and expiry.
func (a *DatabaseAuthenticator) RequestPasswordReset(ctx context.Context, req PasswordResetRequest) (*PasswordResetResponse, error) {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.PasswordResetRequest) {
return a.requestPasswordResetDirect(ctx, req)
}
reqJSON, err := json.Marshal(req)
if err != nil {
return nil, fmt.Errorf("failed to marshal password reset request: %w", err)
@@ -987,6 +1085,9 @@ func (a *DatabaseAuthenticator) RequestPasswordReset(ctx context.Context, req Pa
// CompletePasswordReset implements PasswordResettable. It validates the token and
// updates the user's password via resolvespec_password_reset.
func (a *DatabaseAuthenticator) CompletePasswordReset(ctx context.Context, req PasswordResetCompleteRequest) error {
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.PasswordResetComplete) {
return a.completePasswordResetDirect(ctx, req)
}
reqJSON, err := json.Marshal(req)
if err != nil {
return fmt.Errorf("failed to marshal password reset complete request: %w", err)
+473
View File
@@ -0,0 +1,473 @@
package security
import (
"context"
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/hex"
"errors"
"fmt"
"strings"
"time"
)
// Direct-mode implementations for DatabaseAuthenticator and JWTAuthenticator.
// These mirror the plpgsql bodies in database_schema.sql using plain
// parameterized SQL against the configured TableNames, so they work on
// SQLite, MySQL, or Postgres without the resolvespec_* functions installed.
//
// Password verification is intentionally not implemented here: the stored
// procedures never verify the password hash either (see the TODOs in
// database_schema.sql), so Direct mode matches that behavior exactly rather
// than introducing a mismatch between modes.
var (
errUsernameExists = errors.New("username already exists")
errEmailExists = errors.New("email already exists")
)
func (a *DatabaseAuthenticator) loginDirect(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
var userID int
var email, roles, programUserTable sql.NullString
var userLevel, programUserID sql.NullInt64
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`SELECT id, email, user_level, roles, program_user_id, program_user_table FROM %s WHERE username = ? AND is_active = ?`,
a.tableNames.Users))
return db.QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles, &programUserID, &programUserTable)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("invalid credentials")
}
return nil, fmt.Errorf("login query failed: %w", err)
}
sessionToken, err := generateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
now := time.Now()
expiresAt := now.Add(24 * time.Hour)
ipAddress, userAgent := claimStrings(req.Claims)
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
insertQuery := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (session_token, user_id, expires_at, ip_address, user_agent, last_activity_at, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`,
a.tableNames.UserSessions))
if _, err := db.ExecContext(ctx, insertQuery, sessionToken, userID, expiresAt, ipAddress, userAgent, now, now); err != nil {
return err
}
updateQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET last_login_at = ? WHERE id = ?`, a.tableNames.Users))
_, err := db.ExecContext(ctx, updateQuery, now, userID)
return err
})
if err != nil {
return nil, fmt.Errorf("login query failed: %w", err)
}
userCtx := &UserContext{
UserID: userID,
UserName: req.Username,
Email: email.String,
UserLevel: int(userLevel.Int64),
Roles: parseRoles(roles.String),
SessionID: sessionToken,
ProgramUserID: int(programUserID.Int64),
ProgramUserTable: programUserTable.String,
}
return &LoginResponse{
Token: sessionToken,
User: userCtx,
ExpiresIn: 86400,
}, nil
}
func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req RegisterRequest) (*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")
}
rolesStr := strings.Join(req.Roles, ",")
now := time.Now()
ipAddress, userAgent := claimStrings(req.Claims)
var userID int64
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
var count int
checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT COUNT(*) FROM %s WHERE username = ?`, a.tableNames.Users))
if err := db.QueryRowContext(ctx, checkQuery, req.Username).Scan(&count); err != nil {
return err
}
if count > 0 {
return errUsernameExists
}
checkQuery2 := rewritePlaceholders(db, fmt.Sprintf(`SELECT COUNT(*) FROM %s WHERE email = ?`, a.tableNames.Users))
if err := db.QueryRowContext(ctx, checkQuery2, req.Email).Scan(&count); err != nil {
return err
}
if count > 0 {
return errEmailExists
}
insertQuery := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (username, email, password, user_level, roles, is_active, created_at, updated_at, program_user_id, program_user_table) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
a.tableNames.Users))
res, err := db.ExecContext(ctx, insertQuery, req.Username, req.Email, req.Password, req.UserLevel, rolesStr, true, now, now, 0, "")
if err != nil {
return err
}
userID, err = res.LastInsertId()
return err
})
if err != nil {
if errors.Is(err, errUsernameExists) {
return nil, errUsernameExists
}
if errors.Is(err, errEmailExists) {
return nil, errEmailExists
}
return nil, fmt.Errorf("register query failed: %w", err)
}
sessionToken, err := generateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
expiresAt := now.Add(24 * time.Hour)
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
insertSession := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (session_token, user_id, expires_at, ip_address, user_agent, last_activity_at, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`,
a.tableNames.UserSessions))
if _, err := db.ExecContext(ctx, insertSession, sessionToken, userID, expiresAt, ipAddress, userAgent, now, now); err != nil {
return err
}
updUser := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET last_login_at = ? WHERE id = ?`, a.tableNames.Users))
_, err := db.ExecContext(ctx, updUser, now, userID)
return err
})
if err != nil {
return nil, fmt.Errorf("register query failed: %w", err)
}
userCtx := &UserContext{
UserID: int(userID),
UserName: req.Username,
Email: req.Email,
UserLevel: req.UserLevel,
Roles: parseRoles(rolesStr),
SessionID: sessionToken,
ProgramUserID: 0,
ProgramUserTable: "",
}
return &LoginResponse{
Token: sessionToken,
User: userCtx,
ExpiresIn: 86400,
}, nil
}
func (a *DatabaseAuthenticator) logoutDirect(ctx context.Context, req LogoutRequest) error {
token := req.Token
token = strings.TrimPrefix(token, "Bearer ")
token = strings.TrimPrefix(token, "bearer ")
var rows int64
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE session_token = ? AND user_id = ?`, a.tableNames.UserSessions))
res, err := db.ExecContext(ctx, query, token, req.UserID)
if err != nil {
return err
}
rows, err = res.RowsAffected()
return err
})
if err != nil {
return fmt.Errorf("logout query failed: %w", err)
}
if rows == 0 {
return fmt.Errorf("session not found")
}
if req.Token != "" {
cacheKey := fmt.Sprintf("auth:session:%s", req.Token)
_ = a.cache.Delete(ctx, cacheKey)
}
return nil
}
func (a *DatabaseAuthenticator) sessionDirect(ctx context.Context, token string) (*UserContext, error) {
var userID int
var username, email, roles, programUserTable sql.NullString
var userLevel, programUserID sql.NullInt64
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`SELECT s.user_id, u.username, u.email, u.user_level, u.roles, u.program_user_id, u.program_user_table
FROM %s s JOIN %s u ON s.user_id = u.id
WHERE s.session_token = ? AND s.expires_at > ? AND u.is_active = ?`,
a.tableNames.UserSessions, a.tableNames.Users))
return db.QueryRowContext(ctx, query, token, time.Now(), true).Scan(&userID, &username, &email, &userLevel, &roles, &programUserID, &programUserTable)
})
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 &UserContext{
UserID: userID,
UserName: username.String,
Email: email.String,
UserLevel: int(userLevel.Int64),
SessionID: token,
Roles: parseRoles(roles.String),
ProgramUserID: int(programUserID.Int64),
ProgramUserTable: programUserTable.String,
}, nil
}
func (a *DatabaseAuthenticator) updateSessionActivityDirect(ctx context.Context, sessionToken string) error {
return a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET last_activity_at = ? WHERE session_token = ? AND expires_at > ?`, a.tableNames.UserSessions))
_, err := db.ExecContext(ctx, query, time.Now(), sessionToken, time.Now())
return err
})
}
func (a *DatabaseAuthenticator) refreshTokenDirect(ctx context.Context, oldToken string) (*LoginResponse, error) {
var userID int
var username, email, roles, ipAddress, userAgent, programUserTable sql.NullString
var userLevel, programUserID sql.NullInt64
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(
`SELECT s.user_id, u.username, u.email, u.user_level, u.roles, s.ip_address, s.user_agent, u.program_user_id, u.program_user_table
FROM %s s JOIN %s u ON s.user_id = u.id
WHERE s.session_token = ? AND s.expires_at > ? AND u.is_active = ?`,
a.tableNames.UserSessions, a.tableNames.Users))
return db.QueryRowContext(ctx, query, oldToken, time.Now(), true).Scan(&userID, &username, &email, &userLevel, &roles, &ipAddress, &userAgent, &programUserID, &programUserTable)
})
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)
}
newToken, err := generateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
now := time.Now()
expiresAt := now.Add(24 * time.Hour)
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
insertQuery := rewritePlaceholders(db, fmt.Sprintf(
`INSERT INTO %s (session_token, user_id, expires_at, ip_address, user_agent, last_activity_at, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`,
a.tableNames.UserSessions))
if _, err := db.ExecContext(ctx, insertQuery, newToken, userID, expiresAt, ipAddress.String, userAgent.String, now, now); err != nil {
return err
}
delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE session_token = ?`, a.tableNames.UserSessions))
_, err := db.ExecContext(ctx, delQuery, oldToken)
return err
})
if err != nil {
return nil, fmt.Errorf("refresh token generation failed: %w", err)
}
userCtx := &UserContext{
UserID: userID,
UserName: username.String,
Email: email.String,
UserLevel: int(userLevel.Int64),
SessionID: newToken,
Roles: parseRoles(roles.String),
ProgramUserID: int(programUserID.Int64),
ProgramUserTable: programUserTable.String,
}
return &LoginResponse{
Token: newToken,
User: userCtx,
ExpiresIn: int64(24 * time.Hour.Seconds()),
}, nil
}
func (a *DatabaseAuthenticator) requestPasswordResetDirect(ctx context.Context, req PasswordResetRequest) (*PasswordResetResponse, error) {
if req.Email == "" && req.Username == "" {
return nil, fmt.Errorf("email or username is required")
}
var userID int
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
var query string
var arg string
if req.Email != "" {
query = rewritePlaceholders(db, fmt.Sprintf(`SELECT id FROM %s WHERE email = ? AND is_active = ?`, a.tableNames.Users))
arg = req.Email
} else {
query = rewritePlaceholders(db, fmt.Sprintf(`SELECT id FROM %s WHERE username = ? AND is_active = ?`, a.tableNames.Users))
arg = req.Username
}
return db.QueryRowContext(ctx, query, arg, true).Scan(&userID)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
// Return generic success even when user not found to avoid user enumeration.
return &PasswordResetResponse{Token: "", ExpiresIn: 0}, nil
}
return nil, fmt.Errorf("password reset request query failed: %w", err)
}
rawBytes := make([]byte, 32)
if _, err := rand.Read(rawBytes); err != nil {
return nil, fmt.Errorf("failed to generate reset token: %w", err)
}
rawToken := hex.EncodeToString(rawBytes)
hash := sha256.Sum256([]byte(rawToken))
tokenHash := hex.EncodeToString(hash[:])
now := time.Now()
expiresAt := now.Add(1 * time.Hour)
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ? AND used = ?`, a.tableNames.UserPasswordResets))
if _, err := db.ExecContext(ctx, delQuery, userID, false); err != nil {
return err
}
insQuery := rewritePlaceholders(db, fmt.Sprintf(`INSERT INTO %s (user_id, token_hash, expires_at, created_at, used) VALUES (?, ?, ?, ?, ?)`, a.tableNames.UserPasswordResets))
_, err := db.ExecContext(ctx, insQuery, userID, tokenHash, expiresAt, now, false)
return err
})
if err != nil {
return nil, fmt.Errorf("password reset request query failed: %w", err)
}
return &PasswordResetResponse{Token: rawToken, ExpiresIn: 3600}, nil
}
func (a *DatabaseAuthenticator) completePasswordResetDirect(ctx context.Context, req PasswordResetCompleteRequest) error {
if req.Token == "" {
return fmt.Errorf("token is required")
}
if req.NewPassword == "" {
return fmt.Errorf("new_password is required")
}
hash := sha256.Sum256([]byte(req.Token))
tokenHash := hex.EncodeToString(hash[:])
var resetID, userID int
var expiresAt time.Time
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT id, user_id, expires_at FROM %s WHERE token_hash = ? AND used = ?`, a.tableNames.UserPasswordResets))
return db.QueryRowContext(ctx, query, tokenHash, false).Scan(&resetID, &userID, &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(time.Now()) {
return fmt.Errorf("invalid or expired token")
}
now := time.Now()
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
updUser := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET password = ?, updated_at = ? WHERE id = ?`, a.tableNames.Users))
if _, err := db.ExecContext(ctx, updUser, req.NewPassword, now, userID); err != nil {
return err
}
delSessions := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, a.tableNames.UserSessions))
if _, err := db.ExecContext(ctx, delSessions, userID); err != nil {
return err
}
updReset := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET used = ?, used_at = ? WHERE id = ?`, a.tableNames.UserPasswordResets))
_, err := db.ExecContext(ctx, updReset, true, now, resetID)
return err
})
if err != nil {
return fmt.Errorf("password reset complete query failed: %w", err)
}
return nil
}
// claimStrings extracts ip_address/user_agent from a request's Claims map, mirroring
// p_request->'claims'->>'ip_address' / 'user_agent' in the plpgsql procedures.
func claimStrings(claims map[string]any) (ipAddress, userAgent string) {
if claims == nil {
return "", ""
}
if v, ok := claims["ip_address"].(string); ok {
ipAddress = v
}
if v, ok := claims["user_agent"].(string); ok {
userAgent = v
}
return ipAddress, userAgent
}
// jwtLoginDirect mirrors resolvespec_jwt_login.
func (a *JWTAuthenticator) jwtLoginDirect(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
var userID int
var email, roles sql.NullString
var userLevel sql.NullInt64
runQuery := func() error {
query := rewritePlaceholders(a.getDB(), fmt.Sprintf(`SELECT id, email, user_level, roles FROM %s WHERE username = ? AND is_active = ?`, a.tableNames.Users))
return a.getDB().QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles)
}
err := runQuery()
if isDBClosed(err) {
if reconnErr := a.reconnectDB(); reconnErr == nil {
err = runQuery()
}
}
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("invalid credentials")
}
return nil, fmt.Errorf("login query failed: %w", err)
}
expiresAt := time.Now().Add(24 * time.Hour)
tokenString := fmt.Sprintf("token_%d_%d", userID, expiresAt.Unix())
return &LoginResponse{
Token: tokenString,
User: &UserContext{
UserID: userID,
UserName: req.Username,
Email: email.String,
UserLevel: int(userLevel.Int64),
Roles: parseRoles(roles.String),
},
ExpiresIn: int64(24 * time.Hour.Seconds()),
}, nil
}
// jwtLogoutDirect mirrors resolvespec_jwt_logout (adds token to the blacklist table).
func (a *JWTAuthenticator) jwtLogoutDirect(ctx context.Context, req LogoutRequest) error {
db := a.getDB()
expiresAt := time.Now().Add(24 * time.Hour)
query := rewritePlaceholders(db, fmt.Sprintf(`INSERT INTO %s (token, user_id, expires_at, created_at) VALUES (?, ?, ?, ?)`, a.tableNames.TokenBlacklist))
_, err := db.ExecContext(ctx, query, req.Token, req.UserID, expiresAt, time.Now())
if err != nil {
return fmt.Errorf("logout query failed: %w", err)
}
return nil
}
+169
View File
@@ -0,0 +1,169 @@
package security
import (
"context"
"crypto/rand"
"database/sql"
"encoding/hex"
"fmt"
"strconv"
"strings"
"sync"
"time"
)
// QueryMode selects how a provider talks to the database: via the configured
// resolvespec_* stored procedure (ModeProcedure), via portable Go/SQL logic
// (ModeDirect), or auto-detected per connection (ModeAuto, the default).
type QueryMode int
const (
// ModeAuto probes whether the configured stored procedure exists on a
// Postgres connection and uses it if so, otherwise falls back to Direct mode.
ModeAuto QueryMode = iota
// ModeProcedure always calls the configured resolvespec_* stored procedure.
ModeProcedure
// ModeDirect always uses the portable Go/SQL implementation, never the
// stored procedure.
ModeDirect
)
// dbCapability probes and caches whether a given *sql.DB is Postgres and
// whether specific stored procedures exist on it. One instance is shared by
// a provider (DatabaseAuthenticator, DatabaseTwoFactorProvider, etc.) across
// all of its operations.
type dbCapability struct {
funcExists sync.Map // procName (string) -> exists (bool)
}
// newDBCapability creates a new, empty capability cache.
func newDBCapability() *dbCapability {
return &dbCapability{}
}
// reset clears all cached probe results. Call after reconnecting to a
// (possibly different) database.
func (c *dbCapability) reset() {
c.funcExists.Range(func(key, _ any) bool {
c.funcExists.Delete(key)
return true
})
}
// probeFunctionExists checks, via a Postgres-specific system catalog query,
// whether a function named procName exists. Any error (wrong dialect,
// placeholder syntax rejected, relation missing, etc.) is treated as "does
// not exist" rather than propagated - the probe must never be able to panic
// or block resolution of the query mode.
func probeFunctionExists(ctx context.Context, db *sql.DB, procName string) bool {
if db == nil {
return false
}
var exists bool
defer func() {
// Guard against any unexpected panic from a misbehaving driver.
_ = recover()
}()
row := db.QueryRowContext(ctx, `SELECT EXISTS (SELECT 1 FROM pg_proc WHERE proname = $1 LIMIT 1)`, procName)
if err := row.Scan(&exists); err != nil {
return false
}
return exists
}
// ShouldUseProcedure decides whether the stored-procedure code path should be
// used for procName given mode. ModeProcedure/ModeDirect are unconditional.
//
// ModeAuto resolves as follows:
// - A recognized portable-only driver (SQLite, MySQL) always uses Direct
// mode; no query is issued.
// - A recognized Postgres driver (lib/pq, pgx) probes pg_proc for procName
// and uses the stored procedure only if it actually exists there.
// - Any other/unrecognized driver (including test doubles such as
// sqlmock) cannot be safely dialect-probed without risking an
// unexpected query against a strictly-ordered mock, so it defaults to
// the stored-procedure path, preserving pre-existing behavior for
// callers that configure Postgres-flavored access without an
// identifiable driver type.
func (c *dbCapability) ShouldUseProcedure(ctx context.Context, mode QueryMode, db *sql.DB, procName string) bool {
switch mode {
case ModeProcedure:
return true
case ModeDirect:
return false
default:
if cached, ok := c.funcExists.Load(procName); ok {
return cached.(bool)
}
var exists bool
switch {
case driverIsPortableOnly(db):
exists = false
case driverIsPostgres(db):
exists = probeFunctionExists(ctx, db, procName)
default:
exists = true
}
c.funcExists.Store(procName, exists)
return exists
}
}
// driverIsPostgres reports whether db's underlying driver looks like a
// Postgres driver (lib/pq or pgx), based on the driver's Go type name.
func driverIsPostgres(db *sql.DB) bool {
if db == nil {
return false
}
t := strings.ToLower(fmt.Sprintf("%T", db.Driver()))
return strings.Contains(t, "pq.") || strings.Contains(t, "pgx") || strings.Contains(t, "postgres")
}
// driverIsPortableOnly reports whether db's underlying driver is a dialect
// that never has the resolvespec_* Postgres functions available (SQLite,
// MySQL), so ModeAuto can skip probing entirely and go straight to Direct.
func driverIsPortableOnly(db *sql.DB) bool {
if db == nil {
return false
}
t := strings.ToLower(fmt.Sprintf("%T", db.Driver()))
return strings.Contains(t, "sqlite") || strings.Contains(t, "mysql")
}
// rewritePlaceholders converts a query written with "?" placeholders
// (SQLite/MySQL style) to Postgres "$1", "$2", ... style when db's driver is
// Postgres. All Direct-mode SQL in this package is written with "?" and
// passed through this helper before execution so the same query source works
// against SQLite, MySQL, and (in the rare fallback case) Postgres without the
// resolvespec_* functions installed.
func rewritePlaceholders(db *sql.DB, query string) string {
if !driverIsPostgres(db) {
return query
}
var b strings.Builder
n := 0
for _, r := range query {
if r == '?' {
n++
b.WriteString("$")
b.WriteString(strconv.Itoa(n))
} else {
b.WriteRune(r)
}
}
return b.String()
}
// generateSessionToken produces a session token in the same shape the
// plpgsql stored procedures generate: "sess_" + hex(32 random bytes) + "_" +
// unix timestamp, so downstream code that parses/displays tokens is
// unaffected by which mode created them.
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
}
+165
View File
@@ -0,0 +1,165 @@
package security
import (
"context"
"database/sql"
"testing"
"github.com/DATA-DOG/go-sqlmock"
_ "github.com/mattn/go-sqlite3"
)
func openTestSQLite(t *testing.T) *sql.DB {
t.Helper()
db, err := sql.Open("sqlite3", ":memory:")
if err != nil {
t.Fatalf("failed to open sqlite db: %v", err)
}
t.Cleanup(func() { _ = db.Close() })
return db
}
func TestShouldUseProcedure_ModeProcedure_AlwaysTrue(t *testing.T) {
db := openTestSQLite(t)
c := newDBCapability()
if !c.ShouldUseProcedure(context.Background(), ModeProcedure, db, "resolvespec_login") {
t.Error("ModeProcedure should always resolve to true")
}
}
func TestShouldUseProcedure_ModeDirect_AlwaysFalse(t *testing.T) {
db := openTestSQLite(t)
c := newDBCapability()
if c.ShouldUseProcedure(context.Background(), ModeDirect, db, "resolvespec_login") {
t.Error("ModeDirect should always resolve to false")
}
}
func TestShouldUseProcedure_Auto_SQLite_ResolvesDirect(t *testing.T) {
db := openTestSQLite(t)
c := newDBCapability()
if c.ShouldUseProcedure(context.Background(), ModeAuto, db, "resolvespec_login") {
t.Error("ModeAuto against a SQLite connection should resolve to Direct (false)")
}
}
func TestShouldUseProcedure_Auto_UnknownDriver_DefaultsToProcedure(t *testing.T) {
db, mock, err := sqlmock.New()
if err != nil {
t.Fatalf("failed to create mock db: %v", err)
}
defer db.Close()
c := newDBCapability()
if !c.ShouldUseProcedure(context.Background(), ModeAuto, db, "resolvespec_login") {
t.Error("ModeAuto against an unrecognized driver (e.g. a test double) should default to Procedure (true) to preserve existing behavior")
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Errorf("unfulfilled expectations: %v", err)
}
}
func TestShouldUseProcedure_CachesResult(t *testing.T) {
db := openTestSQLite(t)
c := newDBCapability()
first := c.ShouldUseProcedure(context.Background(), ModeAuto, db, "resolvespec_login")
if _, ok := c.funcExists.Load("resolvespec_login"); !ok {
t.Error("expected result to be cached")
}
second := c.ShouldUseProcedure(context.Background(), ModeAuto, db, "resolvespec_login")
if first != second {
t.Errorf("cached result changed: first=%v second=%v", first, second)
}
}
func TestDBCapability_Reset_ClearsCache(t *testing.T) {
c := newDBCapability()
c.funcExists.Store("resolvespec_login", true)
c.reset()
if _, ok := c.funcExists.Load("resolvespec_login"); ok {
t.Error("reset should clear all cached entries")
}
}
func TestProbeFunctionExists_Postgres_FunctionExists(t *testing.T) {
db, mock, err := sqlmock.New()
if err != nil {
t.Fatalf("failed to create mock db: %v", err)
}
defer db.Close()
rows := sqlmock.NewRows([]string{"exists"}).AddRow(true)
mock.ExpectQuery(`SELECT EXISTS \(SELECT 1 FROM pg_proc WHERE proname = \$1 LIMIT 1\)`).
WithArgs("resolvespec_login").
WillReturnRows(rows)
if !probeFunctionExists(context.Background(), db, "resolvespec_login") {
t.Error("expected probeFunctionExists to return true")
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Errorf("unfulfilled expectations: %v", err)
}
}
func TestProbeFunctionExists_Postgres_FunctionMissing(t *testing.T) {
db, mock, err := sqlmock.New()
if err != nil {
t.Fatalf("failed to create mock db: %v", err)
}
defer db.Close()
rows := sqlmock.NewRows([]string{"exists"}).AddRow(false)
mock.ExpectQuery(`SELECT EXISTS \(SELECT 1 FROM pg_proc WHERE proname = \$1 LIMIT 1\)`).
WithArgs("resolvespec_login").
WillReturnRows(rows)
if probeFunctionExists(context.Background(), db, "resolvespec_login") {
t.Error("expected probeFunctionExists to return false")
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Errorf("unfulfilled expectations: %v", err)
}
}
func TestProbeFunctionExists_QueryError_ReturnsFalse(t *testing.T) {
db, mock, err := sqlmock.New()
if err != nil {
t.Fatalf("failed to create mock db: %v", err)
}
defer db.Close()
mock.ExpectQuery(`SELECT EXISTS \(SELECT 1 FROM pg_proc WHERE proname = \$1 LIMIT 1\)`).
WithArgs("resolvespec_login").
WillReturnError(sql.ErrConnDone)
if probeFunctionExists(context.Background(), db, "resolvespec_login") {
t.Error("expected probeFunctionExists to return false on query error")
}
}
func TestProbeFunctionExists_NilDB(t *testing.T) {
if probeFunctionExists(context.Background(), nil, "resolvespec_login") {
t.Error("expected probeFunctionExists(nil) to return false")
}
}
func TestRewritePlaceholders_NonPostgres_NoOp(t *testing.T) {
db := openTestSQLite(t)
q := "SELECT * FROM users WHERE id = ? AND name = ?"
if got := rewritePlaceholders(db, q); got != q {
t.Errorf("rewritePlaceholders on non-Postgres = %q, want unchanged %q", got, q)
}
}
func TestGenerateSessionToken_Format(t *testing.T) {
token, err := generateSessionToken()
if err != nil {
t.Fatalf("generateSessionToken() error = %v", err)
}
if len(token) < len("sess_")+64+2 {
t.Errorf("generateSessionToken() = %q, unexpected length", token)
}
if token[:5] != "sess_" {
t.Errorf("generateSessionToken() = %q, want prefix 'sess_'", token)
}
}
+101
View File
@@ -0,0 +1,101 @@
package security
import (
"errors"
"fmt"
"reflect"
)
// ErrDirectModeUnsupported is returned by Direct-mode operations that have no
// portable equivalent because they depend on an external schema this package
// does not own (e.g. core.secaccess / core.hub_link for column/row security).
// Callers needing that functionality must run against Postgres with the
// resolvespec_column_security / resolvespec_row_security stored procedures
// installed (ModeProcedure or ModeAuto with the procedures present).
var ErrDirectModeUnsupported = errors.New("direct mode does not support column/row security; requires the resolvespec_column_security/resolvespec_row_security stored procedures")
// TableNames defines all configurable table names used by Direct-mode SQL
// in the security package. Override individual fields to remap to custom
// table names. Use DefaultTableNames() for baseline defaults, and
// MergeTableNames() to apply partial overrides.
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"
}
// DefaultTableNames returns a TableNames with all default table names.
func DefaultTableNames() *TableNames {
return &TableNames{
Users: "users",
UserSessions: "user_sessions",
TokenBlacklist: "token_blacklist",
UserTOTPBackupCodes: "user_totp_backup_codes",
UserPasskeyCredentials: "user_passkey_credentials",
UserPasswordResets: "user_password_resets",
OAuthClients: "oauth_clients",
OAuthCodes: "oauth_codes",
}
}
// MergeTableNames returns a copy of base with any non-empty fields from override applied.
// If override is nil, a copy of base is returned.
func MergeTableNames(base, override *TableNames) *TableNames {
if override == nil {
copied := *base
return &copied
}
merged := *base
if override.Users != "" {
merged.Users = override.Users
}
if override.UserSessions != "" {
merged.UserSessions = override.UserSessions
}
if override.TokenBlacklist != "" {
merged.TokenBlacklist = override.TokenBlacklist
}
if override.UserTOTPBackupCodes != "" {
merged.UserTOTPBackupCodes = override.UserTOTPBackupCodes
}
if override.UserPasskeyCredentials != "" {
merged.UserPasskeyCredentials = override.UserPasskeyCredentials
}
if override.UserPasswordResets != "" {
merged.UserPasswordResets = override.UserPasswordResets
}
if override.OAuthClients != "" {
merged.OAuthClients = override.OAuthClients
}
if override.OAuthCodes != "" {
merged.OAuthCodes = override.OAuthCodes
}
return &merged
}
// ValidateTableNames checks that all non-empty fields in names are valid SQL identifiers.
func ValidateTableNames(names *TableNames) error {
v := reflect.ValueOf(names).Elem()
typ := v.Type()
for i := 0; i < v.NumField(); i++ {
field := v.Field(i)
if field.Kind() != reflect.String {
continue
}
val := field.String()
if val != "" && !validSQLIdentifier.MatchString(val) {
return fmt.Errorf("TableNames.%s contains invalid characters: %q", typ.Field(i).Name, val)
}
}
return nil
}
// resolveTableNames merges an optional override with defaults.
func resolveTableNames(override *TableNames) *TableNames {
return MergeTableNames(DefaultTableNames(), override)
}
+134
View File
@@ -0,0 +1,134 @@
package security
import (
"reflect"
"testing"
)
func TestDefaultTableNames_AllFieldsNonEmpty(t *testing.T) {
names := DefaultTableNames()
v := reflect.ValueOf(names).Elem()
typ := v.Type()
for i := 0; i < v.NumField(); i++ {
field := v.Field(i)
if field.Kind() != reflect.String {
continue
}
if field.String() == "" {
t.Errorf("DefaultTableNames().%s is empty", typ.Field(i).Name)
}
}
}
func TestMergeTableNames_PartialOverride(t *testing.T) {
base := DefaultTableNames()
override := &TableNames{Users: "custom_users", OAuthCodes: "custom_oauth_codes"}
merged := MergeTableNames(base, override)
if merged.Users != "custom_users" {
t.Errorf("MergeTableNames().Users = %q, want %q", merged.Users, "custom_users")
}
if merged.OAuthCodes != "custom_oauth_codes" {
t.Errorf("MergeTableNames().OAuthCodes = %q, want %q", merged.OAuthCodes, "custom_oauth_codes")
}
if merged.UserSessions != "user_sessions" {
t.Errorf("MergeTableNames().UserSessions = %q, want default", merged.UserSessions)
}
}
func TestMergeTableNames_NilOverride(t *testing.T) {
base := DefaultTableNames()
merged := MergeTableNames(base, nil)
if merged == base {
t.Error("MergeTableNames with nil override should return a copy, not the same pointer")
}
if *merged != *base {
t.Errorf("MergeTableNames(base, nil) = %+v, want %+v", merged, base)
}
}
func TestMergeTableNames_DoesNotMutateBase(t *testing.T) {
base := DefaultTableNames()
original := base.Users
override := &TableNames{Users: "custom_users"}
_ = MergeTableNames(base, override)
if base.Users != original {
t.Errorf("MergeTableNames mutated base: Users = %q, want %q", base.Users, original)
}
}
func TestValidateTableNames_Valid(t *testing.T) {
if err := ValidateTableNames(DefaultTableNames()); err != nil {
t.Errorf("ValidateTableNames(defaults) error = %v", err)
}
}
func TestValidateTableNames_Invalid(t *testing.T) {
names := DefaultTableNames()
names.Users = "users; DROP TABLE users; --"
if err := ValidateTableNames(names); err == nil {
t.Error("ValidateTableNames should reject names with invalid characters")
}
}
func TestResolveTableNames_NoOverride(t *testing.T) {
names := resolveTableNames(nil)
if names.Users != "users" {
t.Errorf("resolveTableNames(nil).Users = %q, want default", names.Users)
}
}
func TestResolveTableNames_WithOverride(t *testing.T) {
names := resolveTableNames(&TableNames{Users: "custom_users"})
if names.Users != "custom_users" {
t.Errorf("resolveTableNames().Users = %q, want %q", names.Users, "custom_users")
}
if names.UserSessions != "user_sessions" {
t.Errorf("resolveTableNames().UserSessions = %q, want default", names.UserSessions)
}
}
func TestDefaultKeyStoreTableNames(t *testing.T) {
names := DefaultKeyStoreTableNames()
if names.UserKeys != "user_keys" {
t.Errorf("DefaultKeyStoreTableNames().UserKeys = %q, want %q", names.UserKeys, "user_keys")
}
}
func TestMergeKeyStoreTableNames_PartialOverride(t *testing.T) {
base := DefaultKeyStoreTableNames()
merged := MergeKeyStoreTableNames(base, &KeyStoreTableNames{UserKeys: "custom_keys"})
if merged.UserKeys != "custom_keys" {
t.Errorf("MergeKeyStoreTableNames().UserKeys = %q, want %q", merged.UserKeys, "custom_keys")
}
}
func TestMergeKeyStoreTableNames_NilOverride(t *testing.T) {
base := DefaultKeyStoreTableNames()
merged := MergeKeyStoreTableNames(base, nil)
if merged == base {
t.Error("MergeKeyStoreTableNames with nil override should return a copy, not the same pointer")
}
if *merged != *base {
t.Errorf("MergeKeyStoreTableNames(base, nil) = %+v, want %+v", merged, base)
}
}
func TestValidateKeyStoreTableNames_Invalid(t *testing.T) {
names := &KeyStoreTableNames{UserKeys: "bad name!"}
if err := ValidateKeyStoreTableNames(names); err == nil {
t.Error("ValidateKeyStoreTableNames should reject names with invalid characters")
}
}
func TestValidateKeyStoreTableNames_Valid(t *testing.T) {
if err := ValidateKeyStoreTableNames(DefaultKeyStoreTableNames()); err != nil {
t.Errorf("ValidateKeyStoreTableNames(defaults) error = %v", err)
}
}
+109 -12
View File
@@ -1,20 +1,27 @@
package security
import (
"context"
"crypto/sha256"
"database/sql"
"encoding/hex"
"encoding/json"
"fmt"
"sync"
)
// DatabaseTwoFactorProvider implements TwoFactorAuthProvider using PostgreSQL stored procedures
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
// See totp_database_schema.sql for procedure definitions
type DatabaseTwoFactorProvider struct {
db *sql.DB
totpGen *TOTPGenerator
sqlNames *SQLNames
db *sql.DB
dbMu sync.RWMutex
dbFactory func() (*sql.DB, error)
totpGen *TOTPGenerator
sqlNames *SQLNames
tableNames *TableNames
queryMode QueryMode
capability *dbCapability
}
// NewDatabaseTwoFactorProvider creates a new database-backed 2FA provider
@@ -23,12 +30,69 @@ func NewDatabaseTwoFactorProvider(db *sql.DB, config *TwoFactorConfig, names ...
config = DefaultTwoFactorConfig()
}
return &DatabaseTwoFactorProvider{
db: db,
totpGen: NewTOTPGenerator(config),
sqlNames: resolveSQLNames(names...),
db: db,
totpGen: NewTOTPGenerator(config),
sqlNames: resolveSQLNames(names...),
tableNames: DefaultTableNames(),
capability: newDBCapability(),
}
}
// WithDBFactory configures a factory used to reopen the database connection if it is closed.
func (p *DatabaseTwoFactorProvider) WithDBFactory(factory func() (*sql.DB, error)) *DatabaseTwoFactorProvider {
p.dbFactory = factory
return p
}
// WithTableNames configures Direct-mode table names. If names is nil, defaults are used.
func (p *DatabaseTwoFactorProvider) WithTableNames(names *TableNames) *DatabaseTwoFactorProvider {
p.tableNames = resolveTableNames(names)
return p
}
// WithQueryMode selects stored-procedure vs Direct-mode SQL (default ModeAuto).
func (p *DatabaseTwoFactorProvider) WithQueryMode(mode QueryMode) *DatabaseTwoFactorProvider {
p.queryMode = mode
return p
}
func (p *DatabaseTwoFactorProvider) getDB() *sql.DB {
p.dbMu.RLock()
defer p.dbMu.RUnlock()
return p.db
}
func (p *DatabaseTwoFactorProvider) reconnectDB() error {
if p.dbFactory == nil {
return fmt.Errorf("no db factory configured for reconnect")
}
newDB, err := p.dbFactory()
if err != nil {
return err
}
p.dbMu.Lock()
p.db = newDB
p.dbMu.Unlock()
if p.capability != nil {
p.capability.reset()
}
return nil
}
func (p *DatabaseTwoFactorProvider) runDBOpWithReconnect(run func(*sql.DB) error) error {
db := p.getDB()
if db == nil {
return fmt.Errorf("database connection is nil")
}
err := run(db)
if isDBClosed(err) {
if reconnErr := p.reconnectDB(); reconnErr == nil {
err = run(p.getDB())
}
}
return err
}
// Generate2FASecret creates a new secret for a user
func (p *DatabaseTwoFactorProvider) Generate2FASecret(userID int, issuer, accountName string) (*TwoFactorSecret, error) {
secret, err := p.totpGen.GenerateSecret()
@@ -72,12 +136,17 @@ func (p *DatabaseTwoFactorProvider) Enable2FA(userID int, secret string, backupC
return fmt.Errorf("failed to marshal backup codes: %w", err)
}
ctx := context.Background()
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPEnable) {
return p.enable2FADirect(ctx, userID, secret, hashedCodes)
}
// Call stored procedure
var success bool
var errorMsg sql.NullString
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3::jsonb)`, p.sqlNames.TOTPEnable)
err = p.db.QueryRow(query, userID, secret, string(codesJSON)).Scan(&success, &errorMsg)
err = p.getDB().QueryRow(query, userID, secret, string(codesJSON)).Scan(&success, &errorMsg)
if err != nil {
return fmt.Errorf("enable 2FA query failed: %w", err)
}
@@ -94,11 +163,16 @@ func (p *DatabaseTwoFactorProvider) Enable2FA(userID int, secret string, backupC
// Disable2FA deactivates 2FA for a user
func (p *DatabaseTwoFactorProvider) Disable2FA(userID int) error {
ctx := context.Background()
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPDisable) {
return p.disable2FADirect(ctx, userID)
}
var success bool
var errorMsg sql.NullString
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1)`, p.sqlNames.TOTPDisable)
err := p.db.QueryRow(query, userID).Scan(&success, &errorMsg)
err := p.getDB().QueryRow(query, userID).Scan(&success, &errorMsg)
if err != nil {
return fmt.Errorf("disable 2FA query failed: %w", err)
}
@@ -115,12 +189,17 @@ func (p *DatabaseTwoFactorProvider) Disable2FA(userID int) error {
// Get2FAStatus checks if user has 2FA enabled
func (p *DatabaseTwoFactorProvider) Get2FAStatus(userID int) (bool, error) {
ctx := context.Background()
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPGetStatus) {
return p.get2FAStatusDirect(ctx, userID)
}
var success bool
var errorMsg sql.NullString
var enabled bool
query := fmt.Sprintf(`SELECT p_success, p_error, p_enabled FROM %s($1)`, p.sqlNames.TOTPGetStatus)
err := p.db.QueryRow(query, userID).Scan(&success, &errorMsg, &enabled)
err := p.getDB().QueryRow(query, userID).Scan(&success, &errorMsg, &enabled)
if err != nil {
return false, fmt.Errorf("get 2FA status query failed: %w", err)
}
@@ -137,12 +216,17 @@ func (p *DatabaseTwoFactorProvider) Get2FAStatus(userID int) (bool, error) {
// Get2FASecret retrieves the user's 2FA secret
func (p *DatabaseTwoFactorProvider) Get2FASecret(userID int) (string, error) {
ctx := context.Background()
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPGetSecret) {
return p.get2FASecretDirect(ctx, userID)
}
var success bool
var errorMsg sql.NullString
var secret sql.NullString
query := fmt.Sprintf(`SELECT p_success, p_error, p_secret FROM %s($1)`, p.sqlNames.TOTPGetSecret)
err := p.db.QueryRow(query, userID).Scan(&success, &errorMsg, &secret)
err := p.getDB().QueryRow(query, userID).Scan(&success, &errorMsg, &secret)
if err != nil {
return "", fmt.Errorf("get 2FA secret query failed: %w", err)
}
@@ -175,6 +259,14 @@ func (p *DatabaseTwoFactorProvider) GenerateBackupCodes(userID int, count int) (
hashedCodes[i] = hex.EncodeToString(hash[:])
}
ctx := context.Background()
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPRegenerateBackup) {
if err := p.regenerateBackupCodesDirect(ctx, userID, hashedCodes); err != nil {
return nil, err
}
return codes, nil
}
// Convert to JSON array
codesJSON, err := json.Marshal(hashedCodes)
if err != nil {
@@ -186,7 +278,7 @@ func (p *DatabaseTwoFactorProvider) GenerateBackupCodes(userID int, count int) (
var errorMsg sql.NullString
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2::jsonb)`, p.sqlNames.TOTPRegenerateBackup)
err = p.db.QueryRow(query, userID, string(codesJSON)).Scan(&success, &errorMsg)
err = p.getDB().QueryRow(query, userID, string(codesJSON)).Scan(&success, &errorMsg)
if err != nil {
return nil, fmt.Errorf("regenerate backup codes query failed: %w", err)
}
@@ -208,12 +300,17 @@ func (p *DatabaseTwoFactorProvider) ValidateBackupCode(userID int, code string)
hash := sha256.Sum256([]byte(code))
codeHash := hex.EncodeToString(hash[:])
ctx := context.Background()
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.TOTPValidateBackupCode) {
return p.validateBackupCodeDirect(ctx, userID, codeHash)
}
var success bool
var errorMsg sql.NullString
var valid bool
query := fmt.Sprintf(`SELECT p_success, p_error, p_valid FROM %s($1, $2)`, p.sqlNames.TOTPValidateBackupCode)
err := p.db.QueryRow(query, userID, codeHash).Scan(&success, &errorMsg, &valid)
err := p.getDB().QueryRow(query, userID, codeHash).Scan(&success, &errorMsg, &valid)
if err != nil {
return false, fmt.Errorf("validate backup code query failed: %w", err)
}
@@ -0,0 +1,151 @@
package security
import (
"context"
"database/sql"
"errors"
"fmt"
"time"
)
// Direct-mode implementations mirroring the resolvespec_totp_* stored
// procedures in database_schema.sql using plain SQL against TableNames.
func (p *DatabaseTwoFactorProvider) enable2FADirect(ctx context.Context, userID int, secret string, hashedCodes []string) error {
return p.runDBOpWithReconnect(func(db *sql.DB) error {
updQuery := rewritePlaceholders(db, fmt.Sprintf(
`UPDATE %s SET totp_secret = ?, totp_enabled = ?, totp_enabled_at = ? WHERE id = ?`, p.tableNames.Users))
res, err := db.ExecContext(ctx, updQuery, secret, true, time.Now(), userID)
if err != nil {
return err
}
if rows, err := res.RowsAffected(); err != nil {
return err
} else if rows == 0 {
return fmt.Errorf("user not found")
}
delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, p.tableNames.UserTOTPBackupCodes))
if _, err := db.ExecContext(ctx, delQuery, userID); err != nil {
return err
}
insQuery := rewritePlaceholders(db, fmt.Sprintf(`INSERT INTO %s (user_id, code_hash) VALUES (?, ?)`, p.tableNames.UserTOTPBackupCodes))
for _, hash := range hashedCodes {
if _, err := db.ExecContext(ctx, insQuery, userID, hash); err != nil {
return err
}
}
return nil
})
}
func (p *DatabaseTwoFactorProvider) disable2FADirect(ctx context.Context, userID int) error {
return p.runDBOpWithReconnect(func(db *sql.DB) error {
updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET totp_secret = NULL, totp_enabled = ? WHERE id = ?`, p.tableNames.Users))
res, err := db.ExecContext(ctx, updQuery, false, userID)
if err != nil {
return err
}
if rows, err := res.RowsAffected(); err != nil {
return err
} else if rows == 0 {
return fmt.Errorf("user not found")
}
delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, p.tableNames.UserTOTPBackupCodes))
_, err = db.ExecContext(ctx, delQuery, userID)
return err
})
}
func (p *DatabaseTwoFactorProvider) get2FAStatusDirect(ctx context.Context, userID int) (bool, error) {
var enabled sql.NullBool
err := p.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT totp_enabled FROM %s WHERE id = ?`, p.tableNames.Users))
return db.QueryRowContext(ctx, query, userID).Scan(&enabled)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return false, fmt.Errorf("user not found")
}
return false, fmt.Errorf("get 2FA status query failed: %w", err)
}
return enabled.Bool, nil
}
func (p *DatabaseTwoFactorProvider) get2FASecretDirect(ctx context.Context, userID int) (string, error) {
var secret sql.NullString
var enabled sql.NullBool
err := p.runDBOpWithReconnect(func(db *sql.DB) error {
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT totp_secret, totp_enabled FROM %s WHERE id = ?`, p.tableNames.Users))
return db.QueryRowContext(ctx, query, userID).Scan(&secret, &enabled)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return "", fmt.Errorf("user not found")
}
return "", fmt.Errorf("get 2FA secret query failed: %w", err)
}
if !enabled.Bool {
return "", fmt.Errorf("TOTP not enabled for user")
}
return secret.String, nil
}
func (p *DatabaseTwoFactorProvider) regenerateBackupCodesDirect(ctx context.Context, userID int, hashedCodes []string) error {
return p.runDBOpWithReconnect(func(db *sql.DB) error {
var count int
checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT COUNT(*) FROM %s WHERE id = ? AND totp_enabled = ?`, p.tableNames.Users))
if err := db.QueryRowContext(ctx, checkQuery, userID, true).Scan(&count); err != nil {
return err
}
if count == 0 {
return fmt.Errorf("user not found or TOTP not enabled")
}
delQuery := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, p.tableNames.UserTOTPBackupCodes))
if _, err := db.ExecContext(ctx, delQuery, userID); err != nil {
return err
}
insQuery := rewritePlaceholders(db, fmt.Sprintf(`INSERT INTO %s (user_id, code_hash) VALUES (?, ?)`, p.tableNames.UserTOTPBackupCodes))
for _, hash := range hashedCodes {
if _, err := db.ExecContext(ctx, insQuery, userID, hash); err != nil {
return err
}
}
return nil
})
}
func (p *DatabaseTwoFactorProvider) validateBackupCodeDirect(ctx context.Context, userID int, codeHash string) (bool, error) {
var valid bool
err := p.runDBOpWithReconnect(func(db *sql.DB) error {
var codeID int
var used bool
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT id, used FROM %s WHERE user_id = ? AND code_hash = ?`, p.tableNames.UserTOTPBackupCodes))
err := db.QueryRowContext(ctx, query, userID, codeHash).Scan(&codeID, &used)
if errors.Is(err, sql.ErrNoRows) {
valid = false
return nil
}
if err != nil {
return err
}
if used {
return fmt.Errorf("backup code already used")
}
updQuery := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET used = ?, used_at = ? WHERE id = ?`, p.tableNames.UserTOTPBackupCodes))
if _, err := db.ExecContext(ctx, updQuery, true, time.Now(), codeID); err != nil {
return err
}
valid = true
return nil
})
if err != nil {
return false, err
}
return valid, nil
}
+28 -4
View File
@@ -221,7 +221,7 @@ func (h *Handler) handleRequest(conn *Connection, msg *Message) {
// handleRead processes a read operation
func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContext) {
// Execute before hook
if err := h.hooks.Execute(BeforeRead, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil {
logger.Error("[WebSocketSpec] BeforeRead hook failed: %v", err)
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
_ = conn.SendJSON(errResp)
@@ -273,7 +273,7 @@ func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContex
// handleCreate processes a create operation
func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookContext) {
// Execute before hook
if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
logger.Error("[WebSocketSpec] BeforeCreate hook failed: %v", err)
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
_ = conn.SendJSON(errResp)
@@ -311,7 +311,7 @@ func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookCont
// handleUpdate processes an update operation
func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookContext) {
// Execute before hook
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil {
logger.Error("[WebSocketSpec] BeforeUpdate hook failed: %v", err)
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
_ = conn.SendJSON(errResp)
@@ -349,7 +349,7 @@ func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookCont
// handleDelete processes a delete operation
func (h *Handler) handleDelete(conn *Connection, msg *Message, hookCtx *HookContext) {
// Execute before hook
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil {
logger.Error("[WebSocketSpec] BeforeDelete hook failed: %v", err)
errResp := NewErrorResponse(msg.ID, "hook_error", err.Error())
_ = conn.SendJSON(errResp)
@@ -574,6 +574,18 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) {
}
}
// Execute BeforeScan hooks - pass query so hooks (e.g. row-level security) can modify it
if hookCtx.Metadata == nil {
hookCtx.Metadata = make(map[string]interface{})
}
hookCtx.Metadata["query"] = query
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
return nil, fmt.Errorf("BeforeScan hook failed: %w", err)
}
if modifiedQuery, ok := hookCtx.Metadata["query"].(common.SelectQuery); ok {
query = modifiedQuery
}
// Execute query
if err := query.ScanModel(hookCtx.Context); err != nil {
return nil, fmt.Errorf("failed to read record: %w", err)
@@ -624,6 +636,18 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
}
}
// Execute BeforeScan hooks - pass query so hooks (e.g. row-level security) can modify it
if hookCtx.Metadata == nil {
hookCtx.Metadata = make(map[string]interface{})
}
hookCtx.Metadata["query"] = query
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
return nil, nil, fmt.Errorf("BeforeScan hook failed: %w", err)
}
if modifiedQuery, ok := hookCtx.Metadata["query"].(common.SelectQuery); ok {
query = modifiedQuery
}
// Execute query
if err := query.ScanModel(hookCtx.Context); err != nil {
return nil, nil, fmt.Errorf("failed to read records: %w", err)
+20
View File
@@ -35,6 +35,11 @@ const (
// AfterDelete is called after a delete operation
AfterDelete HookType = "after_delete"
// BeforeScan is called right before a read query is executed against the database,
// after all filters/sort/pagination have been applied. Use this for row-level
// security that needs to modify the query (stored in HookContext.Metadata["query"]).
BeforeScan HookType = "before_scan"
// BeforeSubscribe is called before creating a subscription
BeforeSubscribe HookType = "before_subscribe"
// AfterSubscribe is called after creating a subscription
@@ -54,6 +59,11 @@ const (
BeforeDisconnect HookType = "before_disconnect"
// AfterDisconnect is called after a connection is closed
AfterDisconnect HookType = "after_disconnect"
// BeforeOp fires immediately before every SQL operation (read, create, update, delete).
// Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each
// individual SQL-operation hook point, so it runs once per statement executed.
BeforeOp HookType = "before_op"
)
// HookContext contains context information for hook execution
@@ -197,6 +207,16 @@ func (hr *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
return nil
}
// ExecuteBeforeOp executes the BeforeOp hook followed by the given SQL-operation hook
// (BeforeRead, BeforeCreate, BeforeUpdate, or BeforeDelete). BeforeOp always runs first
// so it can observe/veto every SQL operation regardless of type.
func (hr *HookRegistry) ExecuteBeforeOp(hookType HookType, ctx *HookContext) error {
if err := hr.Execute(BeforeOp, ctx); err != nil {
return err
}
return hr.Execute(hookType, ctx)
}
// HasHooks checks if any hooks are registered for a hook type
func (hr *HookRegistry) HasHooks(hookType HookType) bool {
hooks, exists := hr.hooks[hookType]
+21 -4
View File
@@ -27,25 +27,31 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
return security.LoadSecurityRules(secCtx, securityList)
})
// Hook 2: AfterRead - Apply column-level security (masking)
// Hook 2: BeforeScan - Apply row-level security filters
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
return security.ApplyRowSecurity(secCtx, securityList)
})
// Hook 3: AfterRead - Apply column-level security (masking)
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
return security.ApplyColumnSecurity(secCtx, securityList)
})
// Hook 3 (Optional): Audit logging
// Hook 4 (Optional): Audit logging
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
return security.LogDataAccess(secCtx)
})
// Hook 4: BeforeUpdate - enforce CanUpdate rule from context/registry
// Hook 5: BeforeUpdate - enforce CanUpdate rule from context/registry
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
return security.CheckModelUpdateAllowed(secCtx)
})
// Hook 5: BeforeDelete - enforce CanDelete rule from context/registry
// Hook 6: BeforeDelete - enforce CanDelete rule from context/registry
handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
return security.CheckModelDeleteAllowed(secCtx)
@@ -71,6 +77,17 @@ func (s *securityContext) GetUserID() (int, bool) {
return security.GetUserID(s.ctx.Context)
}
// GetUserRef returns an opaque user identifier for row security lookups.
// It prefers the full *security.UserContext (so providers can read JWT claims,
// e.g. a UUID subject) and falls back to the int user ID.
func (s *securityContext) GetUserRef() (any, bool) {
if userCtx, ok := security.GetUserContext(s.ctx.Context); ok {
return userCtx, true
}
userID, ok := security.GetUserID(s.ctx.Context)
return userID, ok
}
func (s *securityContext) GetSchema() string {
return s.ctx.Schema
}