feat(security): full OAuth 2.1 / OpenID Connect server and OIDC relying-party client

Authorization server: consent and scopes, OIDC (nonce, auth_time, acr, sid,
at_hash, signed userinfo, RP-initiated and back-channel logout), managed
refresh tokens with rotation and reuse detection, RFC 9068 JWT access tokens,
DPoP, PAR, device grant, token exchange, private_key_jwt, RFC 7591/7592
registration, RFC 9207 iss, signing keyring with rotation.

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

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

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

Docs: OAUTH2_SERVER.md, oauth2_full_example.go, breaking_changes.md step 8.
This commit is contained in:
Hein
2026-10-01 14:42:12 +02:00
parent f54b707040
commit 640faeeeaf
51 changed files with 8682 additions and 1125 deletions
+2
View File
@@ -646,6 +646,8 @@ Authentication and authorization framework with hooks integration. Database-back
For documentation, see [pkg/security/README.md](pkg/security/README.md) (see "Database access (lookup)" for the SQLite/portable-SQL path).
It includes a standards-based OAuth 2.1 / OpenID Connect authorization server (consent, rotating refresh tokens, JWT access tokens, DPoP, PAR, device grant, token exchange, logout) and an OIDC relying-party client; see [pkg/security/OAUTH2_SERVER.md](pkg/security/OAUTH2_SERVER.md).
#### Middleware
HTTP middleware collection for common tasks (CORS, logging, metrics, rate limiting, etc.).
+2
View File
@@ -179,6 +179,8 @@ It can operate as:
- **An OAuth2 federation layer** — delegates to external providers (Google, GitHub, Microsoft, etc.)
- **Both simultaneously**
> The underlying `security.OAuthServer` also supports consent, OpenID Connect, rotating refresh tokens, JWT access tokens, DPoP, PAR, the device grant and token exchange; they are opt-in `OAuthServerConfig` options described in [pkg/security/OAUTH2_SERVER.md](../security/OAUTH2_SERVER.md). The options of `resolvemcp.OAuth2Config` are unchanged.
### Standard endpoints served
| Path | Spec | Purpose |
+3
View File
@@ -4,6 +4,8 @@
The security package provides OAuth2 authentication support for any OAuth2-compliant provider including Google, GitHub, Microsoft, Facebook, and custom providers.
> **Full OAuth 2.1 / OpenID Connect**: this guide covers the plain OAuth2 client login. For the OIDC relying party (discovery, PKCE, nonce, id_token validation, logout) and the complete authorization server (consent, refresh rotation, DPoP, PAR, device grant, token exchange) see [OAUTH2_SERVER.md](OAUTH2_SERVER.md).
## Features
- **Universal OAuth2 Support**: Works with any OAuth2 provider
@@ -14,6 +16,7 @@ The security package provides OAuth2 authentication support for any OAuth2-compl
- **Token Refresh**: Automatic token refresh support
- **State Validation**: Built-in CSRF protection
- **User Auto-Creation**: Automatically creates users on first login
- **OpenID Connect** (opt-in): `WithOIDC` discovery, PKCE, nonce and id_token validation, RP-initiated logout
- **Unified Authentication**: OAuth2 and traditional auth share same session storage
## Quick Start
@@ -1,5 +1,7 @@
# OAuth2 Refresh Token - Quick Reference
> This covers refreshing tokens of an upstream provider with `OAuth2RefreshToken`. For refresh tokens issued by `OAuthServer` (rotation, reuse detection, downscoping) see [OAUTH2_SERVER.md](OAUTH2_SERVER.md#refresh-token-rotation).
## Quick Setup (3 Steps)
### 1. Initialize Authenticator
+248
View File
@@ -0,0 +1,248 @@
# OAuth 2.1 / OpenID Connect Authorization Server
`OAuthServer` turns a `DatabaseAuthenticator` into a standards-based identity provider, and `OIDCConfig` / `WithOIDC` make the same package a relying party for any OpenID Connect provider. This guide covers both. For the older "log in with Google/GitHub" client flow see [OAUTH2.md](OAUTH2.md); a runnable end-to-end wiring is in [`oauth2_full_example.go`](oauth2_full_example.go) (`ExampleOAuth2FullServer`, `ExampleOAuth2FullClient`).
Every feature beyond the original authorization-code flow is **opt-in**: the zero value of each `OAuthServerConfig` option keeps the previous behaviour.
## Contents
1. [Architecture](#architecture)
2. [Quick start](#quick-start)
3. [Endpoints](#endpoints)
4. [Configuration reference](#configuration-reference)
5. [Flows](#flows)
6. [Resource servers](#resource-servers)
7. [Client side: logging in with an OpenID Connect provider](#client-side-relying-party)
8. [Database setup](#database-setup)
9. [Security checklist](#security-checklist)
10. [Not supported](#not-supported)
## Architecture
```
browser / app ──► OAuthServer.HTTPHandler() ──► DatabaseAuthenticator (users, sessions, login)
│
└────────────────────► lookup.Provider ──► procedure | direct backend
oauth_clients, oauth_codes, oauth_consents,
oauth_refresh_tokens, oauth_device_codes,
oauth_par_requests, oauth_jti
```
- **State is in the database** (clients, codes, consents, rotating refresh tokens, device codes, pushed requests, the replay cache), so any number of instances can serve one issuer.
- **Stateless pieces** (the SSO cookie, the login/consent form state, the device and logout state) are HMAC-sealed with `CookieSecret`. Instances must share it, or share the first signing key it is derived from.
- **Signing keys**: RSA or ECDSA (P-256/P-384); `kid` is the RFC 7638 thumbprint. All keys are published in the JWKS.
- Access tokens are either the opaque session token (default) or RFC 9068 JWTs. Either way an *access-grant record* stores scope, client, audience and DPoP binding, so introspection, revocation and logout work for both.
## Quick start
```go
auth := security.NewDatabaseAuthenticatorWithOptions(db, security.DatabaseAuthenticatorOptions{})
srv := security.NewOAuthServer(security.OAuthServerConfig{
Issuer: "https://auth.example.com",
PersistClients: true,
PersistCodes: true,
RequireConsent: true,
ManagedRefreshTokens: true,
}, auth)
defer srv.Close()
mux := http.NewServeMux()
mux.Handle("/", srv.HTTPHandler())
```
Register a first-party client from code (no consent screen), or let clients register themselves with `POST /oauth/register`:
```go
client, secret, err := srv.RegisterTrustedClient(ctx, security.OAuthServerClient{
ClientName: "Admin console",
RedirectURIs: []string{"https://console.example.com/callback"},
})
```
## Endpoints
| Method | Path | Spec | Notes |
|---|---|---|---|
| GET | `/.well-known/oauth-authorization-server` | RFC 8414 | Also `/{path}` variants for issuers with a path |
| GET | `/.well-known/openid-configuration` | OIDC Discovery | Same document, plus OIDC fields |
| GET | `/.well-known/oauth-protected-resource` | RFC 9728 | |
| POST | `/oauth/register` | RFC 7591 | Needs `InitialAccessToken` when configured |
| GET PUT DELETE | `/oauth/register/{client_id}` | RFC 7592 | `registration_access_token` as Bearer |
| POST | `/oauth/register/{client_id}/rotate-secret` | RFC 7592 | |
| GET POST | `/oauth/authorize` | RFC 6749, OIDC Core | PKCE S256 required; `response_mode` `query` or `form_post` |
| POST | `/oauth/token` | RFC 6749 | `authorization_code`, `refresh_token`, `client_credentials`, device code, token exchange |
| POST | `/oauth/par` | RFC 9126 | `EnablePAR` / `RequirePAR` |
| POST | `/oauth/device_authorization` | RFC 8628 | `EnableDeviceFlow` |
| GET POST | `/oauth/device` | RFC 8628 | Verification page (login, consent) |
| POST | `/oauth/revoke` | RFC 7009 | Client authentication required |
| POST | `/oauth/introspect` | RFC 7662 | Client authentication required |
| GET POST | `/oauth/userinfo` | OIDC Core | Scope-filtered; signed response if the client asks |
| GET | `/oauth/jwks.json` | RFC 7517 | |
| GET POST | `/oauth/logout` | OIDC RP-Initiated + Back-Channel Logout | |
| GET | `ProviderCallbackPath` | | Callback of a federated upstream provider |
Client authentication methods: `client_secret_basic`, `client_secret_post`, `private_key_jwt` (`jwks` or `jwks_uri`), `none` (public clients).
## Configuration reference
| Option | Default | Purpose |
|---|---|---|
| `Issuer` | required | Public base URL; `iss` of every token. May contain a path |
| `SigningKeys` / `SigningKey` | generated RSA-2048 | Persistent keys for multi-instance and restarts; first signs |
| `CookieSecret` | derived from first key | HMAC key of cookie and form state |
| `SSOCookie` | `resolvespec_sso`, 8h, Lax, Secure when issuer is https | Enables `prompt=none`, `max_age`, single sign-on, logout |
| `PersistClients`, `PersistCodes` | false | Store clients/codes in the DB (needed for several instances) |
| `RequireConsent`, `ConsentTTL` | false, 90 days | Consent screen for non-first-party clients; remembered per user and client |
| `ScopeDescriptions` | built-in for the OIDC scopes | Text on the consent screen |
| `ManagedRefreshTokens`, `RefreshTokenTTL` | false, 30 days | Server-issued rotating refresh tokens with reuse detection |
| `JWTAccessTokens`, `AccessTokenAudience` | false, `ResourceIdentifier` | RFC 9068 access tokens |
| `AccessTokenTTL`, `AuthCodeTTL` | 24h, 2 min | |
| `EnableDPoP` | false | RFC 9449 sender-constrained tokens |
| `EnablePAR`, `RequirePAR`, `PARTTL` | false, false, 90s | |
| `EnableDeviceFlow`, `DeviceCodeTTL`, `DevicePollSeconds` | false, 10 min, 5 | |
| `EnableTokenExchange` | false | RFC 8693 |
| `ClaimsProvider` | `sub`, `preferred_username`, `email` | Source of profile/email/address/phone/custom claims |
| `SupportedACR` | none | Advertised and accepted `acr_values` |
| `DisableLogout` | false | Do not serve `/oauth/logout` |
| `InitialAccessToken` | none | Bearer secret required to register clients |
| `AllowAnonymousIntrospection` | false | Skip client authentication at introspect/revoke |
| `RateLimiter` | none | `func(r, endpoint) bool`; false answers 429 |
| `AllowPrivateNetworkFetch` | false | Allow `jwks_uri` fetches to private addresses (SSRF guard) |
| `LoginTemplate`, `ConsentTemplate` | built-in | `html/template` overrides (`OAuthLoginPage`, `OAuthConsentPage`) |
## Flows
### Authorization code with PKCE
```
GET /oauth/authorize?response_type=code&client_id=ID&redirect_uri=https://app/cb
&scope=openid%20profile&state=S&nonce=N
&code_challenge=BASE64URL(SHA256(V))&code_challenge_method=S256
```
The user signs in (a cookie keeps the session), approves the consent screen if required, and is redirected to `redirect_uri?code=…&state=S&iss=<issuer>` (RFC 9207; compare `iss`). Exchange it:
```
curl -X POST https://auth.example.com/oauth/token \
-d grant_type=authorization_code -d code=CODE -d redirect_uri=https://app/cb \
-d client_id=ID -d code_verifier=V
```
Request parameters: `prompt` (`none`, `login`, `consent`, `select_account`), `max_age`, `id_token_hint`, `login_hint`, `acr_values`, `claims`, `resource` (RFC 8707), `response_mode=form_post`. Once the `redirect_uri` is validated, errors are returned to the client as redirects (`error`, `state`, `iss`); before that they are shown to the user. `prompt=none` without a session answers `login_required`.
The id_token carries `iss`, `sub`, `aud`, `exp`, `iat`, `nonce`, `auth_time`, `acr`, `amr`, `sid`, `at_hash`, `azp` and the claims the granted scopes entitle the client to.
### Consent and scopes
Requested scopes are intersected with the client's `AllowedScopes`; an empty result is `invalid_scope`. With `RequireConsent` (or per client `require_consent`), a third-party client sees a consent screen unless a stored consent already covers the scopes. A client marked `first_party` (`RegisterTrustedClient`) never does. Approval can be remembered for `ConsentTTL`; a denial redirects with `access_denied`.
### Refresh token rotation
With `ManagedRefreshTokens`, every refresh returns a **new** refresh token and invalidates the old one:
```
curl -X POST …/oauth/token -d grant_type=refresh_token -d refresh_token=R1 -d client_id=ID
```
Presenting a rotated token again (a stolen copy) answers `invalid_grant` and **revokes the whole family**, so the legitimate holder has to sign in again. A refresh may downscope (`scope=`) but never widen. Confidential clients must authenticate. Request `offline_access` or allow the `refresh_token` grant to receive one. Without `ManagedRefreshTokens` the refresh token of the underlying authenticator is passed through as in earlier versions.
### JWT access tokens (RFC 9068)
`JWTAccessTokens` issues `at+jwt` tokens with `iss sub aud exp iat jti client_id scope` (and `cnf` for DPoP). See [Resource servers](#resource-servers).
### DPoP (RFC 9449)
With `EnableDPoP` a client sends a `DPoP` proof header (typ `dpop+jwt`, `htm`, `htu`, `iat`, `jti`, public `jwk`) to the token endpoint. The access and refresh tokens are then bound to the proof key; the response has `token_type: DPoP`. Use them as `Authorization: DPoP <token>` plus a proof carrying `ath = base64url(SHA256(token))`. Proof `jti`s are single use (replay cache in `oauth_jti`). A DPoP-bound token is refused as a Bearer token. Set the client's `dpop_bound_access_tokens` to require proofs.
### Pushed authorization requests (RFC 9126)
```
curl -X POST …/oauth/par -d client_id=ID -d response_type=code … -d code_challenge=… # → {"request_uri": "urn:ietf:params:oauth:request_uri:…", "expires_in": 90}
GET /oauth/authorize?client_id=ID&request_uri=urn:ietf:params:oauth:request_uri:…
```
Confidential clients authenticate at the PAR endpoint. A `request_uri` is single use. `RequirePAR` rejects plain authorization requests.
### Device grant (RFC 8628)
```
curl -X POST …/oauth/device_authorization -d client_id=ID -d scope=openid
# → device_code, user_code, verification_uri, verification_uri_complete, interval
curl -X POST …/oauth/token -d grant_type=urn:ietf:params:oauth:grant-type:device_code -d device_code=… -d client_id=ID
```
The user opens `verification_uri` on another device, signs in, enters the code and approves. The device polls; answers are `authorization_pending`, `slow_down` (polled faster than `interval`), `access_denied`, `expired_token`. A code is consumed by the first successful poll.
### Token exchange (RFC 8693)
A confidential client holding the `urn:ietf:params:oauth:grant-type:token-exchange` grant swaps a user's access token for a narrower one (`scope`, `audience`/`resource`). The scope can only shrink; DPoP-bound subject tokens and `actor_token` are not accepted.
### Client credentials
Unchanged: `grant_type=client_credentials` with client authentication returns a token for the client itself.
### Dynamic registration (RFC 7591/7592)
`POST /oauth/register` accepts `redirect_uris` (https, loopback http, or a custom scheme; no fragments), `grant_types`, `response_types`, `token_endpoint_auth_method`, `scope`, `jwks` / `jwks_uri`, `client_name`, `client_uri`, `logo_uri`, `contacts`, `post_logout_redirect_uris`, `backchannel_logout_uri`, `id_token_signed_response_alg`, `userinfo_signed_response_alg`, `dpop_bound_access_tokens`. The response includes `registration_access_token` and `registration_client_uri`; use them to read, update or delete the client and to rotate its secret. Secrets are stored hashed and shown once. Loopback redirect URIs ignore the port (RFC 8252).
### Logout
`/oauth/logout?id_token_hint=…&post_logout_redirect_uri=…&state=…` ends the SSO session, revokes the session's tokens and refresh families, and redirects to a `post_logout_redirect_uri` the client registered. Without a hint the user is asked to confirm. Clients with `backchannel_logout_uri` receive a signed `logout_token` (best effort, in the background).
### Federation
`srv.RegisterExternalProvider(auth, "google")` lets users sign in through an upstream provider; the server remains the issuer for your clients and your users get the same tokens, consent and logout behaviour. `login_hint` pre-fills the built-in login form.
## Resource servers
```go
claims, err := srv.VerifyAccessToken(ctx, token, security.VerifyAccessTokenOptions{
Audience: "https://api.example.com",
Scopes: []string{"orders:read"},
})
```
JWT tokens are verified locally against the key set; opaque tokens are looked up. `claims` holds `Subject`, `UserID`, `ClientID`, `Scopes`, `Audience`, `JTI` and the DPoP key thumbprint. For DPoP-bound tokens also check the proof (`verifyDPoP` runs on the server's own endpoints; an external API validates the `DPoP` header and compares `DPoPKey`). A ready-made middleware is in `ExampleOAuth2FullServer`. Services in another process can call `/oauth/introspect` with their client credentials, or verify the JWT with `GET /oauth/jwks.json`.
## Client side (relying party)
`WithOIDC` discovers the endpoints from `{issuer}/.well-known/openid-configuration` and registers a provider with these protections switched on:
- PKCE (S256) and a `nonce`, both kept with the `state` and used once.
- The `id_token` is verified: signature against the provider's JWKS (refetched once when a `kid` is unknown), algorithm allow-list (`RS256 PS256 ES256 ES384` by default), `iss`, `aud`/`azp`, `exp` (1 minute skew), `nonce`, `at_hash`.
- The userinfo `sub` must equal the id_token `sub`; the RFC 9207 `iss` parameter must match.
```go
auth.WithOIDC(ctx, security.OIDCConfig{Issuer: "https://auth.example.com", ClientID: id, ClientSecret: secret,
RedirectURL: "https://app/auth/callback", ProviderName: "company"})
url, _ := auth.OAuth2GetAuthURLWithOptions("company", state, security.OAuth2AuthOptions{Prompt: "login"})
login, err := auth.OAuth2HandleCallbackRequest(ctx, "company", r) // r is the callback request
logout, _ := auth.OAuth2LogoutURL(ctx, "company", login.Meta["id_token"].(string), "https://app/", "")
```
`OAuth2HandleCallback(ctx, provider, code, state)` still works. The raw `id_token` is returned in `LoginResponse.Meta["id_token"]` (keep it for the logout hint); `OAuth2RefreshToken` re-validates a new id_token when the provider returns one. The Google preset validates id_tokens as well; for other providers set `Issuer` and `JWKSURL` on `OAuth2Config`, or use `WithOIDC`.
## Database setup
Apply the schema for your backend (see the lookup section of the [README](README.md)):
- Postgres procedures: `lookup/database_schema.sql`
- Table-only (any dialect, direct mode): `lookup/ddl/{postgres,sqlite,mysql,mssql}.sql`
The OAuth additions are the `oauth_clients.metadata` and `oauth_codes.extra` JSON columns plus the tables `oauth_consents`, `oauth_refresh_tokens`, `oauth_device_codes`, `oauth_par_requests` and `oauth_jti`. Existing installs: see [breaking_changes.md](breaking_changes.md#step-8-full-oauth2--openid-connect) for the ALTER statements. Purge expired rows of the new tables periodically (`expires_at < now`).
## Security checklist
- Serve the issuer over HTTPS only; the SSO cookie is `Secure` automatically then.
- Persist `SigningKeys` and share `CookieSecret` across instances.
- Set `InitialAccessToken` unless open dynamic registration is intended.
- Use `ManagedRefreshTokens` for public clients; keep `RequireConsent` on for third-party clients.
- Put a `RateLimiter` in front of `token`, `authorize`, `device` and `register`.
- Only PKCE S256 is accepted; redirect URIs match exactly (except the loopback port).
- Keep `AllowPrivateNetworkFetch` off; `jwks_uri` fetches are SSRF-guarded.
- Authenticate callers of `introspect` and `revoke` (the default).
## Not supported
`client_secret_jwt`, signed request objects (`request` / `request_uri` to a remote JWT), the DPoP server nonce, `c_hash`, encrypted id_tokens, and actor tokens in token exchange. Discovery does not advertise them.
+7 -4
View File
@@ -13,7 +13,7 @@ Type-safe, composable security system for ResolveSpec with support for authentic
- ✅ **Extensible** - Implement custom providers for your needs
- ✅ **Stored Procedures** - Database operations use PostgreSQL stored procedures where available, for security and maintainability
- ✅ **Direct Mode** - Portable Go/SQL fallback for SQLite, MySQL, or Postgres without the stored procedures installed — no code changes required
- ✅ **OAuth2 Authorization Server** - Built-in OAuth 2.1 + PKCE server (RFC 8414, 7591, 7009, 7662) with login form and external provider federation
- ✅ **OAuth2 / OpenID Connect** - Built-in OAuth 2.1 + PKCE authorization server and OIDC provider (RFC 8414, 7591/7592, 7009, 7662, 9068, 9126, 9207, 9449, 8628, 8693): consent, rotating refresh tokens, JWT access tokens, logout, federation; plus an OIDC relying-party client. See [OAUTH2_SERVER.md](OAUTH2_SERVER.md)
- ✅ **Password Reset** - Self-service password reset with secure token generation and session invalidation
## Stored Procedure Architecture
@@ -1068,6 +1068,8 @@ lookup.ProcNames{
## OAuth2 Authorization Server
> The complete guide (consent, OIDC, refresh rotation, DPoP, PAR, device grant, token exchange, logout, relying-party client) is in [OAUTH2_SERVER.md](OAUTH2_SERVER.md). The table below lists the original endpoints.
`OAuthServer` is a generic OAuth 2.1 + PKCE authorization server. It is not tied to any spec — `pkg/resolvemcp` uses it, but it can be used standalone with any `http.ServeMux`.
### Endpoints
@@ -1229,12 +1231,13 @@ The main changes:
|------|-------------|
| **QUICK_REFERENCE.md** | Quick reference guide with examples |
| **KEYSTORE.md** | Per-user auth keys and key stores |
| **OAUTH2.md** | OAuth2 client login and the authorization server |
| **OAUTH2.md** | OAuth2 client login |
| **OAUTH2_SERVER.md** | OAuth 2.1 / OpenID Connect server and relying-party client (full guide) |
| **OAUTH2_REFRESH_QUICK_REFERENCE.md** / **OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md** | OAuth2 refresh tokens |
| **PASSKEY_QUICK_REFERENCE.md** | WebAuthn passkeys |
| **SECURITY_FEATURES.md** | Security feature overview |
| **breaking_changes.md** | Migration notes for the `lookup` refactor |
| **examples.go**, **examples_funcspec.go**, **oauth2_examples.go**, **passkey_examples.go** | Working provider implementations |
| **breaking_changes.md** | Migration notes (`lookup` refactor, full OAuth2/OIDC schema changes) |
| **examples.go**, **examples_funcspec.go**, **oauth2_examples.go**, **oauth2_full_example.go**, **passkey_examples.go** | Working provider implementations |
## API Reference
+34
View File
@@ -123,3 +123,37 @@ Behaviour changes:
- Removed the unexported `password.go` from `pkg/security` (bcrypt helpers live in `lookup/direct`).
- `README.md`, `KEYSTORE.md` and the root README describe `lookup.Config` instead of `QueryMode`,
`SQLNames` and `TableNames`.
## Step 8: full OAuth2 / OpenID Connect
Full guide: [OAUTH2_SERVER.md](OAUTH2_SERVER.md). New features are opt-in; the items below are what existing installs must do or notice.
### Schema (existing installs)
Fresh installs use `lookup/database_schema.sql` or `lookup/ddl/<dialect>.sql`. Existing databases need:
- `ALTER TABLE oauth_clients ADD COLUMN metadata <json>` (client metadata: logout URIs, jwks, require_consent, first_party, dpop_bound, signing algs, ...)
- `ALTER TABLE oauth_codes ADD COLUMN extra <json>` (nonce, auth_time, acr, amr, claims, user_id, dpop_jkt, resource)
- New tables `oauth_consents`, `oauth_refresh_tokens`, `oauth_device_codes`, `oauth_par_requests`, `oauth_jti` (copy them from the schema files). Access-grant records are stored in `oauth_refresh_tokens`.
- Postgres procedure mode: reapply `lookup/database_schema.sql` (new `resolvespec_oauth_*` functions, listed in `lookup/procs.go`).
`<json>` is `jsonb` on Postgres, `TEXT` on SQLite, `JSON` on MySQL and `NVARCHAR(MAX)` on SQL Server. New lookup operations and `lookup.OAuthGrantStore` (`Provider.OAuthGrant`) are added to the procedure and direct backends and to the conformance suite; custom `lookup.Config.Procs` overrides gain the new names.
### New API (no action needed)
`OAuthServerConfig` options (see the guide), `OAuthSigningKey`, `OAuthServer.RegisterTrustedClient`, `VerifyAccessToken`, `OAuthClaimsProvider`; `OIDCConfig`, `DatabaseAuthenticator.WithOIDC`, `OAuth2GetAuthURLWithOptions`, `OAuth2HandleCallbackRequest`, `OAuth2LogoutURL`; `OAuth2Config` gains `Issuer`, `JWKSURL`, `EndSessionURL`, `UsePKCE`, `AllowedAlgs`, `AuthStyle`, `HTTPClient`, `ClockSkew`. `DatabaseAuthenticator` gains `OAuthUpdateClient`, `OAuthDeleteClient`, `OAuthGetUser`, `OAuthGrants`.
### Behaviour changes
- `/oauth/introspect` and `/oauth/revoke` require client authentication. Set `AllowAnonymousIntrospection` for the old behaviour.
- Once the `redirect_uri` is validated, authorization errors are redirected to the client (`error`, `state`, `iss`) instead of being returned as JSON. Authorization responses carry `iss` (RFC 9207).
- Only PKCE `S256` is accepted.
- The login form is an `html/template` page with a signed state field; direct form POSTs of earlier versions are still accepted.
- Default grant types of a dynamically registered client include `refresh_token`.
- Authorization-code grants mint a fresh session for the grant. Tokens saved directly with `OAuthSaveCode(SessionToken: ...)` keep working.
- `OAuth2Provider` keeps its PKCE verifier and nonce with the `state`; `Google` preset now validates id_tokens and uses the OpenID Connect endpoints.
- Unauthenticated `userinfo` and discovery routes are unchanged; `userinfo` also answers POST and releases only the claims the granted scopes allow.
### Not supported
`client_secret_jwt`, signed request objects, the DPoP server nonce, `c_hash`, encrypted id_tokens and `actor_token`.
+4
View File
@@ -71,6 +71,9 @@ func New(db *sql.DB, cfg lookup.Config, opts Options) (*lookup.Provider, error)
OAuthUser: &oauthUserRouter{c: c,
proc: procedure.NewOAuthUsers(run, res.Procs),
direct: direct.NewOAuthUsers(base)},
OAuthGrant: &oauthGrantRouter{c: c,
proc: procedure.NewOAuthGrants(run, res.Procs),
direct: direct.NewOAuthGrants(base)},
Passkey: &passkeyRouter{c: c, proc: p, direct: direct.NewPasskey(base)},
TOTP: &totpRouter{c: c,
proc: procedure.NewTOTP(run, res.Procs),
@@ -90,6 +93,7 @@ func Failed(err error) *lookup.Provider {
Keys: &keysRouter{c: c},
OAuthClient: &oauthClientRouter{c: c},
OAuthUser: &oauthUserRouter{c: c},
OAuthGrant: &oauthGrantRouter{c: c},
Passkey: &passkeyRouter{c: c},
TOTP: &totpRouter{c: c},
Policy: &policyRouter{c: c},
@@ -113,6 +113,11 @@ func cleanup(t *testing.T, db *sql.DB, d dialect.Dialect, prefix string) {
like := prefix + "%"
for _, q := range []struct{ table, col string }{
{"oauth_codes", "code"},
{"oauth_consents", "client_id"},
{"oauth_refresh_tokens", "client_id"},
{"oauth_device_codes", "client_id"},
{"oauth_par_requests", "client_id"},
{"oauth_jti", "jti_key"},
{"oauth_clients", "client_id"},
{"token_blacklist", "token"},
{"sec_column_rules", "schema_name"},
@@ -154,3 +154,27 @@ func TestConformanceMSSQLContainer(t *testing.T) {
}
runOnServer(t, "sqlserver", dsn("cf_direct"), "mssql", lookup.Config{}, true)
}
// TestContainerLifecycle checks the start/stop plumbing the container tests rely on: the
// container comes up and accepts connections, and after stop it is gone (it runs with --rm).
func TestContainerLifecycle(t *testing.T) {
rt := containerRuntime(t)
port := startContainer(t, rt, "docker.io/library/postgres:16-alpine", "5432", map[string]string{"POSTGRES_PASSWORD": containerPassword})
waitReady(t, "pgx", fmt.Sprintf("postgres://postgres:%s@127.0.0.1:%s/postgres?sslmode=disable", containerPassword, port), 90*time.Second)
listed := func() string {
return run(t, 30*time.Second, rt, "ps", "-q", "--filter", "ancestor=docker.io/library/postgres:16-alpine")
}
id := listed()
if id == "" {
t.Fatal("container is not running after start")
}
run(t, time.Minute, rt, "stop", "-t", "2", id)
deadline := time.Now().Add(30 * time.Second)
for listed() != "" {
if time.Now().After(deadline) {
t.Fatal("container still present after stop")
}
time.Sleep(500 * time.Millisecond)
}
}
+147
View File
@@ -199,6 +199,22 @@ func (r *oauthClientRouter) Revoke(ctx context.Context, token string) error {
return st.Revoke(ctx, token)
}
func (r *oauthClientRouter) UpdateClient(ctx context.Context, client *sectypes.OAuthServerClient) error {
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthUpdateClient, r.c.procs.OAuthUpdateClient, r.proc, r.direct)
if err != nil {
return err
}
return st.UpdateClient(ctx, client)
}
func (r *oauthClientRouter) DeleteClient(ctx context.Context, clientID string) error {
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthDeleteClient, r.c.procs.OAuthDeleteClient, r.proc, r.direct)
if err != nil {
return err
}
return st.DeleteClient(ctx, clientID)
}
type oauthUserRouter struct {
c *chooser
proc, direct lookup.OAuthUserStore
@@ -394,3 +410,134 @@ func (r *policyRouter) RowSecurity(ctx context.Context, userRef any, schema, tab
}
return st.RowSecurity(ctx, userRef, schema, table)
}
type oauthGrantRouter struct {
c *chooser
proc, direct lookup.OAuthGrantStore
}
var _ lookup.OAuthGrantStore = (*oauthGrantRouter)(nil)
func (r *oauthGrantRouter) pick(ctx context.Context, op lookup.Op, proc string) (lookup.OAuthGrantStore, error) {
return pick[lookup.OAuthGrantStore](r.c, ctx, op, proc, r.proc, r.direct)
}
func (r *oauthGrantRouter) SaveConsent(ctx context.Context, c lookup.Consent) error {
st, err := r.pick(ctx, lookup.OpOAuthSaveConsent, r.c.procs.OAuthSaveConsent)
if err != nil {
return err
}
return st.SaveConsent(ctx, c)
}
func (r *oauthGrantRouter) GetConsent(ctx context.Context, userID int, clientID string) (*lookup.Consent, error) {
st, err := r.pick(ctx, lookup.OpOAuthGetConsent, r.c.procs.OAuthGetConsent)
if err != nil {
return nil, err
}
return st.GetConsent(ctx, userID, clientID)
}
func (r *oauthGrantRouter) RevokeConsent(ctx context.Context, userID int, clientID string) error {
st, err := r.pick(ctx, lookup.OpOAuthRevokeConsent, r.c.procs.OAuthRevokeConsent)
if err != nil {
return err
}
return st.RevokeConsent(ctx, userID, clientID)
}
func (r *oauthGrantRouter) SaveRefresh(ctx context.Context, t lookup.RefreshToken) error {
st, err := r.pick(ctx, lookup.OpOAuthSaveRefresh, r.c.procs.OAuthSaveRefresh)
if err != nil {
return err
}
return st.SaveRefresh(ctx, t)
}
func (r *oauthGrantRouter) RotateRefresh(ctx context.Context, oldHash string, next lookup.RefreshToken) (*lookup.RefreshToken, error) {
st, err := r.pick(ctx, lookup.OpOAuthRotateRefresh, r.c.procs.OAuthRotateRefresh)
if err != nil {
return nil, err
}
return st.RotateRefresh(ctx, oldHash, next)
}
func (r *oauthGrantRouter) PeekRefresh(ctx context.Context, hash string) (*lookup.RefreshToken, error) {
st, err := r.pick(ctx, lookup.OpOAuthPeekRefresh, r.c.procs.OAuthPeekRefresh)
if err != nil {
return nil, err
}
return st.PeekRefresh(ctx, hash)
}
func (r *oauthGrantRouter) RevokeRefreshFamily(ctx context.Context, familyID string) error {
st, err := r.pick(ctx, lookup.OpOAuthRevokeRefreshFamily, r.c.procs.OAuthRevokeRefreshFamily)
if err != nil {
return err
}
return st.RevokeRefreshFamily(ctx, familyID)
}
func (r *oauthGrantRouter) RevokeRefreshBySession(ctx context.Context, sessionToken string) error {
st, err := r.pick(ctx, lookup.OpOAuthRevokeRefreshByUser, r.c.procs.OAuthRevokeRefreshByUser)
if err != nil {
return err
}
return st.RevokeRefreshBySession(ctx, sessionToken)
}
func (r *oauthGrantRouter) CreateDevice(ctx context.Context, d lookup.DeviceCode) error {
st, err := r.pick(ctx, lookup.OpOAuthCreateDevice, r.c.procs.OAuthCreateDevice)
if err != nil {
return err
}
return st.CreateDevice(ctx, d)
}
func (r *oauthGrantRouter) DeviceByUserCode(ctx context.Context, userCode string) (*lookup.DeviceCode, error) {
st, err := r.pick(ctx, lookup.OpOAuthDeviceByUserCode, r.c.procs.OAuthDeviceByUserCode)
if err != nil {
return nil, err
}
return st.DeviceByUserCode(ctx, userCode)
}
func (r *oauthGrantRouter) DeviceDecide(ctx context.Context, userCode string, approve bool, userID int, sessionToken string) error {
st, err := r.pick(ctx, lookup.OpOAuthDeviceDecide, r.c.procs.OAuthDeviceDecide)
if err != nil {
return err
}
return st.DeviceDecide(ctx, userCode, approve, userID, sessionToken)
}
func (r *oauthGrantRouter) DevicePoll(ctx context.Context, deviceHash string) (*lookup.DeviceCode, error) {
st, err := r.pick(ctx, lookup.OpOAuthDevicePoll, r.c.procs.OAuthDevicePoll)
if err != nil {
return nil, err
}
return st.DevicePoll(ctx, deviceHash)
}
func (r *oauthGrantRouter) SavePushedRequest(ctx context.Context, req lookup.PushedRequest) error {
st, err := r.pick(ctx, lookup.OpOAuthSavePAR, r.c.procs.OAuthSavePAR)
if err != nil {
return err
}
return st.SavePushedRequest(ctx, req)
}
func (r *oauthGrantRouter) ConsumePushedRequest(ctx context.Context, requestURI string) (*lookup.PushedRequest, error) {
st, err := r.pick(ctx, lookup.OpOAuthConsumePAR, r.c.procs.OAuthConsumePAR)
if err != nil {
return nil, err
}
return st.ConsumePushedRequest(ctx, requestURI)
}
func (r *oauthGrantRouter) SeenJTI(ctx context.Context, key string, expires time.Time) (bool, error) {
st, err := r.pick(ctx, lookup.OpOAuthSeenJTI, r.c.procs.OAuthSeenJTI)
if err != nil {
return false, err
}
return st.SeenJTI(ctx, key, expires)
}
@@ -52,6 +52,12 @@ func Run(t *testing.T, env Env) {
t.Run("OAuthClientAndCodes", s.oauthClientAndCodes)
t.Run("OAuthIntrospectRevoke", s.oauthIntrospectRevoke)
t.Run("OAuthUsers", s.oauthUsers)
t.Run("OAuthClientMetadata", s.oauthClientMetadata)
t.Run("OAuthGrantConsent", s.oauthGrantConsent)
t.Run("OAuthGrantRefresh", s.oauthGrantRefresh)
t.Run("OAuthGrantDevice", s.oauthGrantDevice)
t.Run("OAuthGrantPAR", s.oauthGrantPAR)
t.Run("OAuthGrantJTI", s.oauthGrantJTI)
t.Run("Passkey", s.passkey)
t.Run("TOTP", s.totp)
t.Run("Policy", s.policy)
@@ -562,3 +568,305 @@ func (s *suite) policy(t *testing.T) {
}
func ptr[T any](v T) *T { return &v }
// --- OAuth server grant state -------------------------------------------------------------
func (s *suite) oauthClientMetadata(t *testing.T) {
st := s.Provider.OAuthClient
id := s.name("meta-client")
reg, err := st.RegisterClient(ctx, &sectypes.OAuthServerClient{
ClientID: id, RedirectURIs: []string{"https://app.example/cb"}, ClientName: "Meta",
PostLogoutRedirectURIs: []string{"https://app.example/bye"}, RequireConsent: true, FirstParty: false,
IDTokenSignedResponseAlg: "RS256", Contacts: []string{"ops@example.test"}, DPoPBoundAccessTokens: true,
})
if err != nil || reg.ClientID != id {
t.Fatalf("register: %+v %v", reg, err)
}
got, err := st.GetClient(ctx, id)
if err != nil {
t.Fatalf("get: %v", err)
}
if !got.RequireConsent || !got.DPoPBoundAccessTokens || got.IDTokenSignedResponseAlg != "RS256" ||
len(got.PostLogoutRedirectURIs) != 1 || got.PostLogoutRedirectURIs[0] != "https://app.example/bye" ||
len(got.Contacts) != 1 {
t.Fatalf("metadata lost: %+v", got)
}
got.ClientName = "Renamed"
got.RequireConsent = false
got.RedirectURIs = []string{"https://app.example/cb", "https://app.example/cb2"}
if err := st.UpdateClient(ctx, got); err != nil {
t.Fatalf("update: %v", err)
}
again, err := st.GetClient(ctx, id)
if err != nil || again.ClientName != "Renamed" || again.RequireConsent || len(again.RedirectURIs) != 2 || !again.DPoPBoundAccessTokens {
t.Fatalf("after update: %+v %v", again, err)
}
if err := st.DeleteClient(ctx, id); err != nil {
t.Fatalf("delete: %v", err)
}
_, err = st.GetClient(ctx, id)
rejected(t, "deleted client", err)
// Code extras round-trip.
code := s.name("meta-code")
err = st.SaveCode(ctx, &sectypes.OAuthCode{
Code: code, ClientID: id, RedirectURI: "https://app.example/cb", CodeChallenge: "chal", CodeChallengeMethod: "S256",
SessionToken: "sess", Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Minute),
Nonce: "n-0S6", AuthTime: 1700000000, ACR: "urn:acr:1", AMR: []string{"pwd"}, UserID: 7,
Claims: map[string]any{"id_token": map[string]any{"email": nil}}, Resource: []string{"https://api.example"}, DPoPJKT: "jkt",
})
if err != nil {
t.Fatalf("save code: %v", err)
}
c, err := st.ExchangeCode(ctx, code)
if err != nil {
t.Fatalf("exchange: %v", err)
}
if c.Nonce != "n-0S6" || c.AuthTime != 1700000000 || c.ACR != "urn:acr:1" || c.UserID != 7 || c.DPoPJKT != "jkt" ||
len(c.AMR) != 1 || len(c.Resource) != 1 || c.Claims["id_token"] == nil {
t.Fatalf("code extra lost: %+v", c)
}
}
func (s *suite) oauthGrantConsent(t *testing.T) {
g := s.Provider.OAuthGrant
client := s.name("consent-client")
if _, err := g.GetConsent(ctx, 1, client); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("missing consent: %v", err)
}
if err := g.SaveConsent(ctx, lookup.Consent{UserID: 1, ClientID: client, Scopes: []string{"openid", "email"}, ExpiresAt: time.Now().Add(time.Hour)}); err != nil {
t.Fatalf("save: %v", err)
}
c, err := g.GetConsent(ctx, 1, client)
if err != nil || len(c.Scopes) != 2 {
t.Fatalf("get: %+v %v", c, err)
}
// Saving again replaces.
if err := g.SaveConsent(ctx, lookup.Consent{UserID: 1, ClientID: client, Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Hour)}); err != nil {
t.Fatalf("resave: %v", err)
}
if c, err = g.GetConsent(ctx, 1, client); err != nil || len(c.Scopes) != 1 {
t.Fatalf("replaced: %+v %v", c, err)
}
// Another user is separate.
if _, err := g.GetConsent(ctx, 2, client); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("other user: %v", err)
}
// Expired consents are not returned.
if err := g.SaveConsent(ctx, lookup.Consent{UserID: 3, ClientID: client, Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(-time.Minute)}); err != nil {
t.Fatalf("save expired: %v", err)
}
if _, err := g.GetConsent(ctx, 3, client); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("expired consent: %v", err)
}
if err := g.RevokeConsent(ctx, 1, client); err != nil {
t.Fatalf("revoke: %v", err)
}
if _, err := g.GetConsent(ctx, 1, client); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("revoked consent: %v", err)
}
}
func (s *suite) oauthGrantRefresh(t *testing.T) {
g := s.Provider.OAuthGrant
client := s.name("refresh-client")
mk := func(n string, session string) lookup.RefreshToken {
return lookup.RefreshToken{TokenHash: s.name(n), FamilyID: s.name("fam-" + n), ClientID: client, UserID: 5,
SessionToken: session, Scopes: []string{"openid", "offline_access"},
Extra: map[string]any{"nonce": "abc"}, ExpiresAt: time.Now().Add(time.Hour)}
}
first := mk("r1", s.name("sess1"))
if err := g.SaveRefresh(ctx, first); err != nil {
t.Fatalf("save: %v", err)
}
peek, err := g.PeekRefresh(ctx, first.TokenHash)
if err != nil || peek.UserID != 5 || peek.ClientID != client || len(peek.Scopes) != 2 || peek.Extra["nonce"] != "abc" {
t.Fatalf("peek: %+v %v", peek, err)
}
if _, err := g.PeekRefresh(ctx, s.name("nope")); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("peek unknown: %v", err)
}
next := lookup.RefreshToken{TokenHash: s.name("r2"), ExpiresAt: time.Now().Add(time.Hour), Scopes: []string{"openid"}}
old, err := g.RotateRefresh(ctx, first.TokenHash, next)
if err != nil || old.FamilyID != first.FamilyID || old.SessionToken != first.SessionToken {
t.Fatalf("rotate: %+v %v", old, err)
}
// The new token belongs to the same family, client and user.
n, err := g.PeekRefresh(ctx, next.TokenHash)
if err != nil || n.FamilyID != first.FamilyID || n.ClientID != client || n.UserID != 5 || n.SessionToken != first.SessionToken {
t.Fatalf("next: %+v %v", n, err)
}
// The consumed token is still visible to Peek, so that presenting it reaches RotateRefresh.
if _, err := g.PeekRefresh(ctx, first.TokenHash); err != nil {
t.Fatalf("peek consumed: %v", err)
}
// Presenting the consumed token again is reuse: the family (including the new token) dies.
third := lookup.RefreshToken{TokenHash: s.name("r3"), ExpiresAt: time.Now().Add(time.Hour)}
reused, err := g.RotateRefresh(ctx, first.TokenHash, third)
if !errors.Is(err, lookup.ErrRefreshReused) || reused == nil || reused.FamilyID != first.FamilyID {
t.Fatalf("reuse: %+v %v", reused, err)
}
if _, err := g.PeekRefresh(ctx, next.TokenHash); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("family survived reuse: %v", err)
}
if _, err := g.RotateRefresh(ctx, next.TokenHash, lookup.RefreshToken{TokenHash: s.name("r4"), ExpiresAt: time.Now().Add(time.Hour)}); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("rotate revoked: %v", err)
}
if _, err := g.PeekRefresh(ctx, third.TokenHash); err == nil {
t.Fatal("a rejected rotation must not store the next token")
}
// Expired tokens cannot rotate.
exp := mk("rexp", s.name("sess2"))
exp.ExpiresAt = time.Now().Add(-time.Minute)
if err := g.SaveRefresh(ctx, exp); err != nil {
t.Fatalf("save expired: %v", err)
}
if _, err := g.RotateRefresh(ctx, exp.TokenHash, lookup.RefreshToken{TokenHash: s.name("rexp2"), ExpiresAt: time.Now().Add(time.Hour)}); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("rotate expired: %v", err)
}
// Revoking by family and by session.
fam := mk("rfam", s.name("sess3"))
if err := g.SaveRefresh(ctx, fam); err != nil {
t.Fatal(err)
}
if err := g.RevokeRefreshFamily(ctx, fam.FamilyID); err != nil {
t.Fatalf("revoke family: %v", err)
}
if _, err := g.PeekRefresh(ctx, fam.TokenHash); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("family revoked: %v", err)
}
bs := mk("rsess", s.name("sess4"))
if err := g.SaveRefresh(ctx, bs); err != nil {
t.Fatal(err)
}
if err := g.RevokeRefreshBySession(ctx, bs.SessionToken); err != nil {
t.Fatalf("revoke session: %v", err)
}
if _, err := g.PeekRefresh(ctx, bs.TokenHash); !errors.Is(err, lookup.ErrRefreshInvalid) {
t.Fatalf("session revoked: %v", err)
}
}
func (s *suite) oauthGrantDevice(t *testing.T) {
g := s.Provider.OAuthGrant
client := s.name("device-client")
mk := func(n string) lookup.DeviceCode {
return lookup.DeviceCode{DeviceHash: s.name("dh-" + n), UserCode: strings.ToUpper(s.name("uc-" + n)), ClientID: client,
Scopes: []string{"openid"}, Interval: 1, ExpiresAt: time.Now().Add(time.Minute)}
}
// pending -> approved
d := mk("a")
if err := g.CreateDevice(ctx, d); err != nil {
t.Fatalf("create: %v", err)
}
got, err := g.DeviceByUserCode(ctx, strings.ToLower(d.UserCode)) // user codes are case-insensitive
if err != nil || got.ClientID != client || got.DeviceHash != d.DeviceHash || len(got.Scopes) != 1 {
t.Fatalf("by user code: %+v %v", got, err)
}
if _, err := g.DevicePoll(ctx, d.DeviceHash); !errors.Is(err, lookup.ErrDevicePending) {
t.Fatalf("first poll: %v", err)
}
if _, err := g.DevicePoll(ctx, d.DeviceHash); !errors.Is(err, lookup.ErrDeviceSlowDown) {
t.Fatalf("immediate re-poll: %v", err)
}
if err := g.DeviceDecide(ctx, d.UserCode, true, 9, s.name("dsess")); err != nil {
t.Fatalf("approve: %v", err)
}
if _, err := g.DeviceByUserCode(ctx, d.UserCode); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("decided code still pending: %v", err)
}
if err := g.DeviceDecide(ctx, d.UserCode, true, 9, "x"); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("second decision: %v", err)
}
time.Sleep(1100 * time.Millisecond)
done, err := g.DevicePoll(ctx, d.DeviceHash)
if err != nil || done.UserID != 9 || done.SessionToken != s.name("dsess") || done.ClientID != client {
t.Fatalf("approved poll: %+v %v", done, err)
}
if _, err := g.DevicePoll(ctx, d.DeviceHash); !errors.Is(err, lookup.ErrDeviceExpired) {
t.Fatalf("consumed code: %v", err)
}
// denied
dd := mk("d")
if err := g.CreateDevice(ctx, dd); err != nil {
t.Fatal(err)
}
if err := g.DeviceDecide(ctx, dd.UserCode, false, 0, ""); err != nil {
t.Fatalf("deny: %v", err)
}
if _, err := g.DevicePoll(ctx, dd.DeviceHash); !errors.Is(err, lookup.ErrDeviceDenied) {
t.Fatalf("denied poll: %v", err)
}
// expired
de := mk("e")
de.ExpiresAt = time.Now().Add(-time.Second)
if err := g.CreateDevice(ctx, de); err != nil {
t.Fatal(err)
}
if _, err := g.DevicePoll(ctx, de.DeviceHash); !errors.Is(err, lookup.ErrDeviceExpired) {
t.Fatalf("expired poll: %v", err)
}
if _, err := g.DeviceByUserCode(ctx, de.UserCode); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("expired by user code: %v", err)
}
if _, err := g.DevicePoll(ctx, s.name("unknown")); !errors.Is(err, lookup.ErrDeviceExpired) {
t.Fatalf("unknown poll: %v", err)
}
}
func (s *suite) oauthGrantPAR(t *testing.T) {
g := s.Provider.OAuthGrant
uri := "urn:ietf:params:oauth:request_uri:" + s.name("par")
if len(uri) > 255 {
t.Fatal("test request_uri too long")
}
req := lookup.PushedRequest{RequestURI: uri, ClientID: s.name("par-client"),
Params: map[string]string{"redirect_uri": "https://app.example/cb", "scope": "openid"}, ExpiresAt: time.Now().Add(time.Minute)}
if err := g.SavePushedRequest(ctx, req); err != nil {
t.Fatalf("save: %v", err)
}
got, err := g.ConsumePushedRequest(ctx, uri)
if err != nil || got.ClientID != req.ClientID || got.Params["scope"] != "openid" {
t.Fatalf("consume: %+v %v", got, err)
}
if _, err := g.ConsumePushedRequest(ctx, uri); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("second consume: %v", err)
}
exp := lookup.PushedRequest{RequestURI: uri + "-exp", ClientID: s.name("par-client"), Params: map[string]string{"a": "b"}, ExpiresAt: time.Now().Add(-time.Minute)}
if err := g.SavePushedRequest(ctx, exp); err != nil {
t.Fatal(err)
}
if _, err := g.ConsumePushedRequest(ctx, exp.RequestURI); !errors.Is(err, lookup.ErrNotFound) {
t.Fatalf("expired consume: %v", err)
}
}
func (s *suite) oauthGrantJTI(t *testing.T) {
g := s.Provider.OAuthGrant
key := s.name("jti")
seen, err := g.SeenJTI(ctx, key, time.Now().Add(time.Minute))
if err != nil || seen {
t.Fatalf("first: %v %v", seen, err)
}
seen, err = g.SeenJTI(ctx, key, time.Now().Add(time.Minute))
if err != nil || !seen {
t.Fatalf("replay: %v %v", seen, err)
}
// An expired entry is forgotten.
old := s.name("jti-old")
if _, err := g.SeenJTI(ctx, old, time.Now().Add(-time.Minute)); err != nil {
t.Fatal(err)
}
seen, err = g.SeenJTI(ctx, old, time.Now().Add(time.Minute))
if err != nil || seen {
t.Fatalf("after expiry: %v %v", seen, err)
}
}
+384 -7
View File
@@ -1765,8 +1765,10 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
client_secret_hash TEXT, -- sha256 hex of the confidential-client secret; NULL for public clients
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
is_active BOOLEAN DEFAULT true,
metadata jsonb, -- every other RFC 7591 field (see sectypes.OAuthServerClient)
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
ALTER TABLE oauth_clients ADD COLUMN IF NOT EXISTS metadata jsonb;
-- oauth_codes: short-lived authorization codes (for multi-instance deployments)
-- Note: client_id is stored without a foreign key so codes can be persisted even
@@ -1783,8 +1785,10 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
refresh_token TEXT,
scopes TEXT[],
expires_at TIMESTAMP NOT NULL,
extra jsonb, -- nonce, auth_time, acr, claims, user_id, dpop_jkt ... (see sectypes.OAuthCode)
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
ALTER TABLE oauth_codes ADD COLUMN IF NOT EXISTS extra jsonb;
CREATE INDEX IF NOT EXISTS idx_oauth_codes_code ON oauth_codes(code);
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
@@ -1801,7 +1805,7 @@ DECLARE
BEGIN
v_client_id := p_request->>'client_id';
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method)
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method, metadata)
VALUES (
v_client_id,
ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)),
@@ -1809,9 +1813,10 @@ BEGIN
CASE WHEN jsonb_typeof(p_request->'grant_types') = 'array' AND jsonb_array_length(p_request->'grant_types') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'grant_types')) ELSE ARRAY['authorization_code'] END,
CASE WHEN jsonb_typeof(p_request->'allowed_scopes') = 'array' AND jsonb_array_length(p_request->'allowed_scopes') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'allowed_scopes')) ELSE ARRAY['openid','profile','email'] END,
NULLIF(p_request->>'client_secret_hash', ''),
COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), 'none')
COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), 'none'),
NULLIF(p_request - ARRAY['client_id','redirect_uris','client_name','grant_types','allowed_scopes','client_secret_hash','token_endpoint_auth_method'], '{}'::jsonb)
)
RETURNING to_jsonb(oauth_clients.*) INTO v_row;
RETURNING (to_jsonb(oauth_clients.*) - 'metadata') || COALESCE(metadata, '{}'::jsonb) INTO v_row;
RETURN QUERY SELECT true, null::text, v_row;
EXCEPTION WHEN OTHERS THEN
@@ -1825,7 +1830,7 @@ LANGUAGE plpgsql AS $$
DECLARE
v_row jsonb;
BEGIN
SELECT to_jsonb(oauth_clients.*)
SELECT (to_jsonb(oauth_clients.*) - 'metadata') || COALESCE(metadata, '{}'::jsonb)
INTO v_row
FROM oauth_clients
WHERE client_id = p_client_id AND is_active = true;
@@ -1842,7 +1847,7 @@ CREATE OR REPLACE FUNCTION resolvespec_oauth_save_code(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text)
LANGUAGE plpgsql AS $$
BEGIN
INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at)
INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at, extra)
VALUES (
p_request->>'code',
p_request->>'client_id',
@@ -1853,7 +1858,8 @@ BEGIN
p_request->>'session_token',
p_request->>'refresh_token',
ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'scopes') = 'array' THEN p_request->'scopes' ELSE '[]'::jsonb END)),
(p_request->>'expires_at')::timestamptz::timestamp
(p_request->>'expires_at')::timestamptz::timestamp,
NULLIF(p_request - ARRAY['code','client_id','redirect_uri','client_state','code_challenge','code_challenge_method','session_token','refresh_token','scopes','expires_at'], '{}'::jsonb)
);
RETURN QUERY SELECT true, null::text;
@@ -1879,7 +1885,7 @@ BEGIN
'session_token', session_token,
'refresh_token', refresh_token,
'scopes', to_jsonb(scopes)
) INTO v_row;
) || COALESCE(extra, '{}'::jsonb) INTO v_row;
IF v_row IS NULL THEN
RETURN QUERY SELECT false, 'invalid or expired code'::text, null::jsonb;
@@ -1930,3 +1936,374 @@ BEGIN
RETURN QUERY SELECT true, null::text;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_update_client(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text)
LANGUAGE plpgsql AS $$
DECLARE
v_rows int;
BEGIN
UPDATE oauth_clients SET
redirect_uris = ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)),
client_name = p_request->>'client_name',
grant_types = CASE WHEN jsonb_typeof(p_request->'grant_types') = 'array' AND jsonb_array_length(p_request->'grant_types') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'grant_types')) ELSE grant_types END,
allowed_scopes = CASE WHEN jsonb_typeof(p_request->'allowed_scopes') = 'array' AND jsonb_array_length(p_request->'allowed_scopes') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'allowed_scopes')) ELSE allowed_scopes END,
client_secret_hash = NULLIF(p_request->>'client_secret_hash', ''),
token_endpoint_auth_method = COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), token_endpoint_auth_method),
metadata = NULLIF(p_request - ARRAY['client_id','redirect_uris','client_name','grant_types','allowed_scopes','client_secret_hash','token_endpoint_auth_method'], '{}'::jsonb)
WHERE client_id = p_request->>'client_id' AND is_active = true;
GET DIAGNOSTICS v_rows = ROW_COUNT;
IF v_rows = 0 THEN
RETURN QUERY SELECT false, 'client not found'::text;
ELSE
RETURN QUERY SELECT true, null::text;
END IF;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_delete_client(p_client_id text)
RETURNS TABLE(p_success bool, p_error text)
LANGUAGE plpgsql AS $$
BEGIN
UPDATE oauth_clients SET is_active = false WHERE client_id = p_client_id;
RETURN QUERY SELECT true, null::text;
END;
$$;
-- ============================================
-- OAuth2 Server grant state (consents, refresh tokens, device codes, PAR, replay cache)
-- ============================================
-- Procedure-backend tables use jsonb for scopes/extra/params. Every procedure takes one jsonb
-- request and returns (p_success, p_error, p_data). p_error carries a stable code for the
-- failures the Go side maps to errors: not_found, refresh_invalid, refresh_reused,
-- device_pending, device_slowdown, device_denied, device_expired.
CREATE TABLE IF NOT EXISTS oauth_consents (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL,
client_id VARCHAR(255) NOT NULL,
scopes jsonb,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
id SERIAL PRIMARY KEY,
token_hash VARCHAR(64) NOT NULL UNIQUE, -- sha256 hex of the raw refresh token
family_id VARCHAR(64) NOT NULL,
client_id VARCHAR(255) NOT NULL,
user_id INTEGER NOT NULL,
session_token VARCHAR(255),
scopes jsonb,
extra jsonb,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
used_at TIMESTAMP,
revoked_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
CREATE TABLE IF NOT EXISTS oauth_device_codes (
id SERIAL PRIMARY KEY,
device_hash VARCHAR(64) NOT NULL UNIQUE,
user_code VARCHAR(32) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
scopes jsonb,
status VARCHAR(16) NOT NULL DEFAULT 'pending',
user_id INTEGER,
session_token VARCHAR(255),
poll_interval INTEGER NOT NULL DEFAULT 5,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
last_polled_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_device_expires ON oauth_device_codes(expires_at);
CREATE TABLE IF NOT EXISTS oauth_par_requests (
id SERIAL PRIMARY KEY,
request_uri VARCHAR(255) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
params jsonb,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_par_expires ON oauth_par_requests(expires_at);
CREATE TABLE IF NOT EXISTS oauth_jti (
id SERIAL PRIMARY KEY,
jti_key VARCHAR(255) NOT NULL UNIQUE,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_jti_expires ON oauth_jti(expires_at);
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_consent(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
DELETE FROM oauth_consents
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id';
INSERT INTO oauth_consents (user_id, client_id, scopes, expires_at)
VALUES ((p_request->>'user_id')::int, p_request->>'client_id', COALESCE(p_request->'scopes', '[]'::jsonb),
(p_request->>'expires_at')::timestamptz::timestamp);
RETURN QUERY SELECT true, null::text, null::jsonb;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_get_consent(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_row jsonb;
BEGIN
SELECT jsonb_build_object('user_id', user_id, 'client_id', client_id, 'scopes', COALESCE(scopes, '[]'::jsonb), 'expires_at', expires_at)
INTO v_row
FROM oauth_consents
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id' AND expires_at > now();
IF v_row IS NULL THEN
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
ELSE
RETURN QUERY SELECT true, null::text, v_row;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_consent(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
DELETE FROM oauth_consents
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id';
RETURN QUERY SELECT true, null::text, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_refresh(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
INSERT INTO oauth_refresh_tokens (token_hash, family_id, client_id, user_id, session_token, scopes, extra, expires_at)
VALUES (p_request->>'token_hash', p_request->>'family_id', p_request->>'client_id', (p_request->>'user_id')::int,
p_request->>'session_token', p_request->'scopes', p_request->'extra',
(p_request->>'expires_at')::timestamptz::timestamp);
RETURN QUERY SELECT true, null::text, null::jsonb;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_rotate_refresh(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
r oauth_refresh_tokens%ROWTYPE;
v_next jsonb := p_request->'next';
v_old jsonb;
BEGIN
SELECT * INTO r FROM oauth_refresh_tokens WHERE token_hash = p_request->>'old_hash' FOR UPDATE;
IF NOT FOUND OR r.revoked_at IS NOT NULL OR r.expires_at <= now() THEN
RETURN QUERY SELECT false, 'refresh_invalid'::text, null::jsonb;
RETURN;
END IF;
v_old := jsonb_build_object('token_hash', r.token_hash, 'family_id', r.family_id, 'client_id', r.client_id,
'user_id', r.user_id, 'session_token', r.session_token,
'scopes', COALESCE(r.scopes, '[]'::jsonb), 'extra', COALESCE(r.extra, '{}'::jsonb),
'expires_at', r.expires_at);
IF r.used_at IS NOT NULL THEN
-- A rotated token came back: revoke the whole family. Returning (not raising) keeps the revoke.
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE family_id = r.family_id AND revoked_at IS NULL;
RETURN QUERY SELECT false, 'refresh_reused'::text, v_old;
RETURN;
END IF;
UPDATE oauth_refresh_tokens SET used_at = now() WHERE id = r.id;
INSERT INTO oauth_refresh_tokens (token_hash, family_id, client_id, user_id, session_token, scopes, extra, expires_at)
VALUES (v_next->>'token_hash', r.family_id, r.client_id, r.user_id,
COALESCE(NULLIF(v_next->>'session_token', ''), r.session_token),
v_next->'scopes', v_next->'extra', (v_next->>'expires_at')::timestamptz::timestamp);
RETURN QUERY SELECT true, null::text, v_old;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_peek_refresh(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_row jsonb;
BEGIN
SELECT jsonb_build_object('token_hash', token_hash, 'family_id', family_id, 'client_id', client_id,
'user_id', user_id, 'session_token', session_token,
'scopes', COALESCE(scopes, '[]'::jsonb), 'extra', COALESCE(extra, '{}'::jsonb),
'expires_at', expires_at)
INTO v_row
FROM oauth_refresh_tokens
WHERE token_hash = p_request->>'token_hash' AND revoked_at IS NULL AND expires_at > now();
IF v_row IS NULL THEN
RETURN QUERY SELECT false, 'refresh_invalid'::text, null::jsonb;
ELSE
RETURN QUERY SELECT true, null::text, v_row;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_refresh_family(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE family_id = p_request->>'family_id' AND revoked_at IS NULL;
RETURN QUERY SELECT true, null::text, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_refresh_session(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE session_token = p_request->>'session_token' AND revoked_at IS NULL;
RETURN QUERY SELECT true, null::text, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_create_device(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
INSERT INTO oauth_device_codes (device_hash, user_code, client_id, scopes, status, poll_interval, expires_at)
VALUES (p_request->>'device_hash', upper(p_request->>'user_code'), p_request->>'client_id', p_request->'scopes',
COALESCE(NULLIF(p_request->>'status', ''), 'pending'), COALESCE((p_request->>'interval')::int, 5),
(p_request->>'expires_at')::timestamptz::timestamp);
RETURN QUERY SELECT true, null::text, null::jsonb;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_by_user_code(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_row jsonb;
BEGIN
SELECT jsonb_build_object('device_hash', device_hash, 'user_code', user_code, 'client_id', client_id,
'scopes', COALESCE(scopes, '[]'::jsonb), 'status', status, 'interval', poll_interval,
'expires_at', expires_at)
INTO v_row
FROM oauth_device_codes
WHERE user_code = upper(p_request->>'user_code') AND status = 'pending' AND expires_at > now();
IF v_row IS NULL THEN
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
ELSE
RETURN QUERY SELECT true, null::text, v_row;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_decide(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_rows int;
v_approve boolean := COALESCE((p_request->>'approve')::boolean, false);
BEGIN
UPDATE oauth_device_codes
SET status = CASE WHEN v_approve THEN 'approved' ELSE 'denied' END,
user_id = CASE WHEN v_approve THEN (p_request->>'user_id')::int ELSE user_id END,
session_token = CASE WHEN v_approve THEN p_request->>'session_token' ELSE session_token END
WHERE user_code = upper(p_request->>'user_code') AND status = 'pending' AND expires_at > now();
GET DIAGNOSTICS v_rows = ROW_COUNT;
IF v_rows = 0 THEN
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
ELSE
RETURN QUERY SELECT true, null::text, null::jsonb;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_poll(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
d oauth_device_codes%ROWTYPE;
v_slow boolean;
BEGIN
SELECT * INTO d FROM oauth_device_codes WHERE device_hash = p_request->>'device_hash' FOR UPDATE;
IF NOT FOUND THEN
RETURN QUERY SELECT false, 'device_expired'::text, null::jsonb;
RETURN;
END IF;
IF d.expires_at <= now() THEN
DELETE FROM oauth_device_codes WHERE id = d.id;
RETURN QUERY SELECT false, 'device_expired'::text, null::jsonb;
RETURN;
END IF;
v_slow := d.last_polled_at IS NOT NULL AND (now() - d.last_polled_at) < make_interval(secs => d.poll_interval);
UPDATE oauth_device_codes SET last_polled_at = now() WHERE id = d.id;
IF v_slow THEN
RETURN QUERY SELECT false, 'device_slowdown'::text, null::jsonb;
ELSIF d.status = 'denied' THEN
DELETE FROM oauth_device_codes WHERE id = d.id;
RETURN QUERY SELECT false, 'device_denied'::text, null::jsonb;
ELSIF d.status = 'approved' THEN
DELETE FROM oauth_device_codes WHERE id = d.id;
RETURN QUERY SELECT true, null::text, jsonb_build_object('device_hash', d.device_hash, 'user_code', d.user_code,
'client_id', d.client_id, 'scopes', COALESCE(d.scopes, '[]'::jsonb), 'status', d.status,
'user_id', d.user_id, 'session_token', d.session_token, 'interval', d.poll_interval, 'expires_at', d.expires_at);
ELSE
RETURN QUERY SELECT false, 'device_pending'::text, null::jsonb;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_par(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
BEGIN
INSERT INTO oauth_par_requests (request_uri, client_id, params, expires_at)
VALUES (p_request->>'request_uri', p_request->>'client_id', p_request->'params',
(p_request->>'expires_at')::timestamptz::timestamp);
RETURN QUERY SELECT true, null::text, null::jsonb;
EXCEPTION WHEN OTHERS THEN
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_consume_par(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_row jsonb;
BEGIN
DELETE FROM oauth_par_requests
WHERE request_uri = p_request->>'request_uri' AND expires_at > now()
RETURNING jsonb_build_object('request_uri', request_uri, 'client_id', client_id,
'params', COALESCE(params, '{}'::jsonb), 'expires_at', expires_at)
INTO v_row;
IF v_row IS NULL THEN
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
ELSE
RETURN QUERY SELECT true, null::text, v_row;
END IF;
END;
$$;
CREATE OR REPLACE FUNCTION resolvespec_oauth_seen_jti(p_request jsonb)
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
LANGUAGE plpgsql AS $$
DECLARE
v_rows int;
BEGIN
DELETE FROM oauth_jti WHERE expires_at < now();
INSERT INTO oauth_jti (jti_key, expires_at)
VALUES (p_request->>'key', (p_request->>'expires_at')::timestamptz::timestamp)
ON CONFLICT (jti_key) DO NOTHING;
GET DIAGNOSTICS v_rows = ROW_COUNT;
RETURN QUERY SELECT true, null::text, jsonb_build_object('seen', v_rows = 0);
END;
$$;
+85
View File
@@ -147,6 +147,7 @@ CREATE TABLE oauth_clients (
client_secret_hash NVARCHAR(MAX),
token_endpoint_auth_method NVARCHAR(30) DEFAULT 'none',
is_active BIT DEFAULT 1,
metadata NVARCHAR(MAX),
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
);
@@ -166,6 +167,7 @@ CREATE TABLE oauth_codes (
refresh_token NVARCHAR(MAX),
scopes NVARCHAR(MAX),
expires_at DATETIME2 NOT NULL,
extra NVARCHAR(MAX),
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
);
@@ -173,6 +175,89 @@ IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_codes_expires'
CREATE INDEX idx_oauth_codes_expires ON oauth_codes(expires_at);
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
IF OBJECT_ID(N'oauth_consents', N'U') IS NULL
CREATE TABLE oauth_consents (
id INT IDENTITY(1,1) PRIMARY KEY,
user_id INT NOT NULL,
client_id NVARCHAR(255) NOT NULL,
scopes NVARCHAR(MAX),
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
expires_at DATETIME2 NOT NULL
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_consents_user_client' AND object_id = OBJECT_ID(N'oauth_consents'))
CREATE INDEX idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
IF OBJECT_ID(N'oauth_refresh_tokens', N'U') IS NULL
CREATE TABLE oauth_refresh_tokens (
id INT IDENTITY(1,1) PRIMARY KEY,
token_hash NVARCHAR(64) NOT NULL UNIQUE,
family_id NVARCHAR(64) NOT NULL,
client_id NVARCHAR(255) NOT NULL,
user_id INT NOT NULL,
session_token NVARCHAR(255),
scopes NVARCHAR(MAX),
extra NVARCHAR(MAX),
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
expires_at DATETIME2 NOT NULL,
used_at DATETIME2,
revoked_at DATETIME2
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_refresh_family' AND object_id = OBJECT_ID(N'oauth_refresh_tokens'))
CREATE INDEX idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_refresh_session' AND object_id = OBJECT_ID(N'oauth_refresh_tokens'))
CREATE INDEX idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_refresh_expires' AND object_id = OBJECT_ID(N'oauth_refresh_tokens'))
CREATE INDEX idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
IF OBJECT_ID(N'oauth_device_codes', N'U') IS NULL
CREATE TABLE oauth_device_codes (
id INT IDENTITY(1,1) PRIMARY KEY,
device_hash NVARCHAR(64) NOT NULL UNIQUE,
user_code NVARCHAR(32) NOT NULL UNIQUE,
client_id NVARCHAR(255) NOT NULL,
scopes NVARCHAR(MAX),
status NVARCHAR(16) NOT NULL DEFAULT 'pending',
user_id INT,
session_token NVARCHAR(255),
poll_interval INT NOT NULL DEFAULT 5,
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
expires_at DATETIME2 NOT NULL,
last_polled_at DATETIME2
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_device_expires' AND object_id = OBJECT_ID(N'oauth_device_codes'))
CREATE INDEX idx_oauth_device_expires ON oauth_device_codes(expires_at);
IF OBJECT_ID(N'oauth_par_requests', N'U') IS NULL
CREATE TABLE oauth_par_requests (
id INT IDENTITY(1,1) PRIMARY KEY,
request_uri NVARCHAR(255) NOT NULL UNIQUE,
client_id NVARCHAR(255) NOT NULL,
params NVARCHAR(MAX),
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
expires_at DATETIME2 NOT NULL
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_par_expires' AND object_id = OBJECT_ID(N'oauth_par_requests'))
CREATE INDEX idx_oauth_par_expires ON oauth_par_requests(expires_at);
IF OBJECT_ID(N'oauth_jti', N'U') IS NULL
CREATE TABLE oauth_jti (
id INT IDENTITY(1,1) PRIMARY KEY,
jti_key NVARCHAR(255) NOT NULL UNIQUE,
expires_at DATETIME2 NOT NULL
);
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_jti_expires' AND object_id = OBJECT_ID(N'oauth_jti'))
CREATE INDEX idx_oauth_jti_expires ON oauth_jti(expires_at);
-- key_hash: SHA-256 hex
-- scopes: JSON-encoded array
-- meta: JSON-encoded object
+66
View File
@@ -126,6 +126,7 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
client_secret_hash TEXT,
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
is_active TINYINT(1) DEFAULT 1,
metadata TEXT,
created_at DATETIME NULL
);
@@ -144,11 +145,76 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
refresh_token TEXT,
scopes TEXT,
expires_at DATETIME NOT NULL,
extra TEXT,
created_at DATETIME NULL,
INDEX idx_oauth_codes_expires (expires_at)
);
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
CREATE TABLE IF NOT EXISTS oauth_consents (
id INT AUTO_INCREMENT PRIMARY KEY,
user_id INT NOT NULL,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
created_at DATETIME NULL,
expires_at DATETIME NOT NULL,
INDEX idx_oauth_consents_user_client (user_id, client_id)
);
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
id INT AUTO_INCREMENT PRIMARY KEY,
token_hash VARCHAR(64) NOT NULL UNIQUE,
family_id VARCHAR(64) NOT NULL,
client_id VARCHAR(255) NOT NULL,
user_id INT NOT NULL,
session_token VARCHAR(255),
scopes TEXT,
extra TEXT,
created_at DATETIME NULL,
expires_at DATETIME NOT NULL,
used_at DATETIME,
revoked_at DATETIME,
INDEX idx_oauth_refresh_family (family_id),
INDEX idx_oauth_refresh_session (session_token),
INDEX idx_oauth_refresh_expires (expires_at)
);
CREATE TABLE IF NOT EXISTS oauth_device_codes (
id INT AUTO_INCREMENT PRIMARY KEY,
device_hash VARCHAR(64) NOT NULL UNIQUE,
user_code VARCHAR(32) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
status VARCHAR(16) NOT NULL DEFAULT 'pending',
user_id INT,
session_token VARCHAR(255),
poll_interval INT NOT NULL DEFAULT 5,
created_at DATETIME NULL,
expires_at DATETIME NOT NULL,
last_polled_at DATETIME,
INDEX idx_oauth_device_expires (expires_at)
);
CREATE TABLE IF NOT EXISTS oauth_par_requests (
id INT AUTO_INCREMENT PRIMARY KEY,
request_uri VARCHAR(255) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
params TEXT,
created_at DATETIME NULL,
expires_at DATETIME NOT NULL,
INDEX idx_oauth_par_expires (expires_at)
);
CREATE TABLE IF NOT EXISTS oauth_jti (
id INT AUTO_INCREMENT PRIMARY KEY,
jti_key VARCHAR(255) NOT NULL UNIQUE,
expires_at DATETIME NOT NULL,
INDEX idx_oauth_jti_expires (expires_at)
);
-- key_hash: SHA-256 hex
-- scopes: JSON-encoded array
-- meta: JSON-encoded object
+73
View File
@@ -137,6 +137,7 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
client_secret_hash TEXT,
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
is_active BOOLEAN DEFAULT true,
metadata TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
@@ -155,12 +156,84 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
refresh_token TEXT,
scopes TEXT,
expires_at TIMESTAMP NOT NULL,
extra TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
CREATE TABLE IF NOT EXISTS oauth_consents (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
id SERIAL PRIMARY KEY,
token_hash VARCHAR(64) NOT NULL UNIQUE,
family_id VARCHAR(64) NOT NULL,
client_id VARCHAR(255) NOT NULL,
user_id INTEGER NOT NULL,
session_token VARCHAR(255),
scopes TEXT,
extra TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
used_at TIMESTAMP,
revoked_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
CREATE TABLE IF NOT EXISTS oauth_device_codes (
id SERIAL PRIMARY KEY,
device_hash VARCHAR(64) NOT NULL UNIQUE,
user_code VARCHAR(32) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
status VARCHAR(16) NOT NULL DEFAULT 'pending',
user_id INTEGER,
session_token VARCHAR(255),
poll_interval INTEGER NOT NULL DEFAULT 5,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
last_polled_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_device_expires ON oauth_device_codes(expires_at);
CREATE TABLE IF NOT EXISTS oauth_par_requests (
id SERIAL PRIMARY KEY,
request_uri VARCHAR(255) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
params TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_par_expires ON oauth_par_requests(expires_at);
CREATE TABLE IF NOT EXISTS oauth_jti (
id SERIAL PRIMARY KEY,
jti_key VARCHAR(255) NOT NULL UNIQUE,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_jti_expires ON oauth_jti(expires_at);
-- key_hash: SHA-256 hex
-- scopes: JSON-encoded array
-- meta: JSON-encoded object
+73
View File
@@ -130,6 +130,7 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
client_secret_hash TEXT,
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
is_active BOOLEAN DEFAULT 1,
metadata TEXT,
created_at TIMESTAMP
);
@@ -148,12 +149,84 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
refresh_token TEXT,
scopes TEXT,
expires_at TIMESTAMP NOT NULL,
extra TEXT,
created_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
CREATE TABLE IF NOT EXISTS oauth_consents (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
created_at TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
id INTEGER PRIMARY KEY AUTOINCREMENT,
token_hash VARCHAR(64) NOT NULL UNIQUE,
family_id VARCHAR(64) NOT NULL,
client_id VARCHAR(255) NOT NULL,
user_id INTEGER NOT NULL,
session_token VARCHAR(255),
scopes TEXT,
extra TEXT,
created_at TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
used_at TIMESTAMP,
revoked_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
CREATE TABLE IF NOT EXISTS oauth_device_codes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
device_hash VARCHAR(64) NOT NULL UNIQUE,
user_code VARCHAR(32) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
scopes TEXT,
status VARCHAR(16) NOT NULL DEFAULT 'pending',
user_id INTEGER,
session_token VARCHAR(255),
poll_interval INTEGER NOT NULL DEFAULT 5,
created_at TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
last_polled_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_oauth_device_expires ON oauth_device_codes(expires_at);
CREATE TABLE IF NOT EXISTS oauth_par_requests (
id INTEGER PRIMARY KEY AUTOINCREMENT,
request_uri VARCHAR(255) NOT NULL UNIQUE,
client_id VARCHAR(255) NOT NULL,
params TEXT,
created_at TIMESTAMP,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_par_expires ON oauth_par_requests(expires_at);
CREATE TABLE IF NOT EXISTS oauth_jti (
id INTEGER PRIMARY KEY AUTOINCREMENT,
jti_key VARCHAR(255) NOT NULL UNIQUE,
expires_at TIMESTAMP NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_oauth_jti_expires ON oauth_jti(expires_at);
-- key_hash: SHA-256 hex
-- scopes: JSON-encoded array
-- meta: JSON-encoded object
+5
View File
@@ -197,6 +197,11 @@ func Gt(c lookup.Column, v any) Cond {
return func(bl *builder) string { return bl.col(c) + " > " + bl.ph(v) }
}
// Lt is `col < value`.
func Lt(c lookup.Column, v any) Cond {
return func(bl *builder) string { return bl.col(c) + " < " + bl.ph(v) }
}
// IsNull is `col IS NULL`.
func IsNull(c lookup.Column) Cond {
return func(bl *builder) string { return bl.col(c) + " IS NULL" }
+92 -19
View File
@@ -60,6 +60,10 @@ func (o *OAuthClients) RegisterClient(ctx context.Context, client *sectypes.OAut
return nil, fmt.Errorf("failed to marshal allowed_scopes: %w", err)
}
meta, err := client.ClientMetadataJSON()
if err != nil {
return nil, fmt.Errorf("failed to marshal client metadata: %w", err)
}
err = o.do(func(q Querier) error {
return o.Insert(lookup.EntityOAuthClients).Set(
Set(lookup.OAuthClientsClientID, client.ClientID),
@@ -70,33 +74,86 @@ func (o *OAuthClients) RegisterClient(ctx context.Context, client *sectypes.OAut
Set(lookup.OAuthClientsClientSecretHash, nullIfEmpty(client.ClientSecretHash)),
Set(lookup.OAuthClientsTokenEndpointAuthMethod, authMethod),
Set(lookup.OAuthClientsIsActive, true),
Set(lookup.OAuthClientsMetadata, nullIfEmpty(meta)),
Set(lookup.OAuthClientsCreatedAt, o.Now()),
).Exec(ctx, q)
})
if err != nil {
return nil, fmt.Errorf("failed to register client: %w", err)
}
return &sectypes.OAuthServerClient{
ClientID: client.ClientID,
RedirectURIs: client.RedirectURIs,
ClientName: client.ClientName,
GrantTypes: grantTypes,
AllowedScopes: allowedScopes,
ClientSecretHash: client.ClientSecretHash,
TokenEndpointAuthMethod: authMethod,
}, nil
res := *client
res.GrantTypes = grantTypes
res.AllowedScopes = allowedScopes
res.TokenEndpointAuthMethod = authMethod
return &res, nil
}
// UpdateClient implements lookup.OAuthClientStore: it rewrites the mutable registration
// fields of an existing client (RFC 7592 management).
func (o *OAuthClients) UpdateClient(ctx context.Context, client *sectypes.OAuthServerClient) error {
redirects, err := o.d.EncodeJSON(client.RedirectURIs)
if err != nil {
return err
}
if redirects == nil {
redirects = "[]"
}
grants, err := o.d.EncodeJSON(client.GrantTypes)
if err != nil {
return err
}
scopes, err := o.d.EncodeJSON(client.AllowedScopes)
if err != nil {
return err
}
meta, err := client.ClientMetadataJSON()
if err != nil {
return err
}
var n int64
err = o.do(func(q Querier) error {
var err error
n, err = o.Update(lookup.EntityOAuthClients).Set(
Set(lookup.OAuthClientsRedirectURIs, redirects),
Set(lookup.OAuthClientsClientName, client.ClientName),
Set(lookup.OAuthClientsGrantTypes, grants),
Set(lookup.OAuthClientsAllowedScopes, scopes),
Set(lookup.OAuthClientsClientSecretHash, nullIfEmpty(client.ClientSecretHash)),
Set(lookup.OAuthClientsTokenEndpointAuthMethod, client.TokenEndpointAuthMethod),
Set(lookup.OAuthClientsMetadata, nullIfEmpty(meta)),
).Where(Eq(lookup.OAuthClientsClientID, client.ClientID), Eq(lookup.OAuthClientsIsActive, true)).Exec(ctx, q)
return err
})
if err != nil {
return fmt.Errorf("failed to update client: %w", err)
}
if n == 0 {
return fmt.Errorf("client not found")
}
return nil
}
// DeleteClient implements lookup.OAuthClientStore: the client is deactivated.
func (o *OAuthClients) DeleteClient(ctx context.Context, clientID string) error {
return o.do(func(q Querier) error {
_, err := o.Update(lookup.EntityOAuthClients).Set(Set(lookup.OAuthClientsIsActive, false)).
Where(Eq(lookup.OAuthClientsClientID, clientID)).Exec(ctx, q)
return err
})
}
// GetClient implements lookup.OAuthClientStore.
func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
var redirects, grants, scopes any
var name, secret, method sql.NullString
var meta any
err := o.do(func(q Querier) error {
return o.From(lookup.EntityOAuthClients).
Cols(lookup.OAuthClientsRedirectURIs, lookup.OAuthClientsClientName, lookup.OAuthClientsGrantTypes,
lookup.OAuthClientsAllowedScopes, lookup.OAuthClientsClientSecretHash, lookup.OAuthClientsTokenEndpointAuthMethod).
lookup.OAuthClientsAllowedScopes, lookup.OAuthClientsClientSecretHash, lookup.OAuthClientsTokenEndpointAuthMethod,
lookup.OAuthClientsMetadata).
Where(Eq(lookup.OAuthClientsClientID, clientID), Eq(lookup.OAuthClientsIsActive, true)).
QueryRow(ctx, q, &redirects, &name, &grants, &scopes, &secret, &method)
QueryRow(ctx, q, &redirects, &name, &grants, &scopes, &secret, &method, &meta)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
@@ -104,12 +161,17 @@ func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectype
}
return nil, fmt.Errorf("failed to get client: %w", err)
}
res := &sectypes.OAuthServerClient{
ClientID: clientID,
ClientName: name.String,
ClientSecretHash: secret.String,
TokenEndpointAuthMethod: method.String,
res := &sectypes.OAuthServerClient{}
switch v := meta.(type) {
case []byte:
_ = res.ApplyClientMetadata(string(v))
case string:
_ = res.ApplyClientMetadata(v)
}
res.ClientID = clientID
res.ClientName = name.String
res.ClientSecretHash = secret.String
res.TokenEndpointAuthMethod = method.String
_ = o.d.DecodeJSON(redirects, &res.RedirectURIs)
_ = o.d.DecodeJSON(grants, &res.GrantTypes)
_ = o.d.DecodeJSON(scopes, &res.AllowedScopes)
@@ -126,6 +188,10 @@ func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) e
if method == "" {
method = "S256"
}
extra, err := code.CodeExtraJSON()
if err != nil {
return fmt.Errorf("failed to marshal code extra: %w", err)
}
return o.do(func(q Querier) error {
return o.Insert(lookup.EntityOAuthCodes).Set(
Set(lookup.OAuthCodesCode, code.Code),
@@ -138,6 +204,7 @@ func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) e
Set(lookup.OAuthCodesRefreshToken, code.RefreshToken),
Set(lookup.OAuthCodesScopes, scopes),
Set(lookup.OAuthCodesExpiresAt, code.ExpiresAt),
Set(lookup.OAuthCodesExtra, nullIfEmpty(extra)),
Set(lookup.OAuthCodesCreatedAt, o.Now()),
).Exec(ctx, q)
})
@@ -148,15 +215,15 @@ func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) e
func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) {
var res sectypes.OAuthCode
var state, refresh sql.NullString
var scopes any
var scopes, extra any
err := o.tx(ctx, func(q Querier) error {
err := o.From(lookup.EntityOAuthCodes).
Cols(lookup.OAuthCodesClientID, lookup.OAuthCodesRedirectURI, lookup.OAuthCodesClientState,
lookup.OAuthCodesCodeChallenge, lookup.OAuthCodesCodeChallengeMethod, lookup.OAuthCodesSessionToken,
lookup.OAuthCodesRefreshToken, lookup.OAuthCodesScopes).
lookup.OAuthCodesRefreshToken, lookup.OAuthCodesScopes, lookup.OAuthCodesExtra).
Where(Eq(lookup.OAuthCodesCode, code), Gt(lookup.OAuthCodesExpiresAt, o.Now())).
QueryRow(ctx, q, &res.ClientID, &res.RedirectURI, &state, &res.CodeChallenge, &res.CodeChallengeMethod,
&res.SessionToken, &refresh, &scopes)
&res.SessionToken, &refresh, &scopes, &extra)
if err != nil {
return err
}
@@ -179,6 +246,12 @@ func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes
res.ClientState = state.String
res.RefreshToken = refresh.String
_ = o.d.DecodeJSON(scopes, &res.Scopes)
switch v := extra.(type) {
case []byte:
_ = res.ApplyCodeExtra(string(v))
case string:
_ = res.ApplyCodeExtra(v)
}
return &res, nil
}
+475
View File
@@ -0,0 +1,475 @@
package direct
import (
"context"
"database/sql"
"errors"
"fmt"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// OAuthGrants implements lookup.OAuthGrantStore on tables. Every multi-step operation runs in
// one transaction, and single-use records (refresh rotation, device codes, pushed requests)
// are consumed with a conditional write so concurrent callers cannot both succeed.
type OAuthGrants struct{ *Base }
var _ lookup.OAuthGrantStore = (*OAuthGrants)(nil)
// NewOAuthGrants creates the direct OAuthGrantStore.
func NewOAuthGrants(b *Base) *OAuthGrants { return &OAuthGrants{Base: b} }
// optTime reads a nullable time column scanned into an `any`.
func (o *OAuthGrants) optTime(src any) (time.Time, bool) {
if src == nil {
return time.Time{}, false
}
t, err := o.d.ScanTime(src)
if err != nil || t.IsZero() {
return time.Time{}, false
}
return t, true
}
// jsonArg encodes v for a JSON/TEXT column; an empty value is NULL.
func (o *OAuthGrants) jsonArg(v any) (any, error) { return o.d.EncodeJSON(v) }
// SaveConsent implements lookup.OAuthGrantStore.
func (o *OAuthGrants) SaveConsent(ctx context.Context, c lookup.Consent) error {
scopes, err := o.jsonArg(c.Scopes)
if err != nil {
return err
}
return o.tx(ctx, func(q Querier) error {
if _, err := o.Delete(lookup.EntityOAuthConsents).
Where(Eq(lookup.OAuthConsentsUserID, c.UserID), Eq(lookup.OAuthConsentsClientID, c.ClientID)).Exec(ctx, q); err != nil {
return err
}
return o.Insert(lookup.EntityOAuthConsents).Set(
Set(lookup.OAuthConsentsUserID, c.UserID),
Set(lookup.OAuthConsentsClientID, c.ClientID),
Set(lookup.OAuthConsentsScopes, scopes),
Set(lookup.OAuthConsentsCreatedAt, o.Now()),
Set(lookup.OAuthConsentsExpiresAt, c.ExpiresAt),
).Exec(ctx, q)
})
}
// GetConsent implements lookup.OAuthGrantStore.
func (o *OAuthGrants) GetConsent(ctx context.Context, userID int, clientID string) (*lookup.Consent, error) {
var scopes any
var exp time.Time
err := o.do(func(q Querier) error {
return o.From(lookup.EntityOAuthConsents).
Cols(lookup.OAuthConsentsScopes, lookup.OAuthConsentsExpiresAt).
Where(Eq(lookup.OAuthConsentsUserID, userID), Eq(lookup.OAuthConsentsClientID, clientID),
Gt(lookup.OAuthConsentsExpiresAt, o.Now())).
QueryRow(ctx, q, &scopes, o.timeDest(&exp))
})
if errors.Is(err, sql.ErrNoRows) {
return nil, lookup.ErrNotFound
}
if err != nil {
return nil, fmt.Errorf("failed to get consent: %w", err)
}
c := &lookup.Consent{UserID: userID, ClientID: clientID, ExpiresAt: exp}
_ = o.d.DecodeJSON(scopes, &c.Scopes)
return c, nil
}
// RevokeConsent implements lookup.OAuthGrantStore.
func (o *OAuthGrants) RevokeConsent(ctx context.Context, userID int, clientID string) error {
return o.do(func(q Querier) error {
_, err := o.Delete(lookup.EntityOAuthConsents).
Where(Eq(lookup.OAuthConsentsUserID, userID), Eq(lookup.OAuthConsentsClientID, clientID)).Exec(ctx, q)
return err
})
}
func (o *OAuthGrants) insertRefresh(ctx context.Context, q Querier, t lookup.RefreshToken) error {
scopes, err := o.jsonArg(t.Scopes)
if err != nil {
return err
}
extra, err := o.jsonArg(t.Extra)
if err != nil {
return err
}
return o.Insert(lookup.EntityOAuthRefreshTokens).Set(
Set(lookup.OAuthRefreshTokenHash, t.TokenHash),
Set(lookup.OAuthRefreshFamilyID, t.FamilyID),
Set(lookup.OAuthRefreshClientID, t.ClientID),
Set(lookup.OAuthRefreshUserID, t.UserID),
Set(lookup.OAuthRefreshSessionToken, t.SessionToken),
Set(lookup.OAuthRefreshScopes, scopes),
Set(lookup.OAuthRefreshExtra, extra),
Set(lookup.OAuthRefreshCreatedAt, o.Now()),
Set(lookup.OAuthRefreshExpiresAt, t.ExpiresAt),
).Exec(ctx, q)
}
// SaveRefresh implements lookup.OAuthGrantStore.
func (o *OAuthGrants) SaveRefresh(ctx context.Context, t lookup.RefreshToken) error {
return o.do(func(q Querier) error { return o.insertRefresh(ctx, q, t) })
}
type refreshRow struct {
lookup.RefreshToken
used, revoked bool
expired bool
}
func (o *OAuthGrants) loadRefresh(ctx context.Context, q Querier, hash string) (*refreshRow, error) {
var scopes, extra, usedAt, revokedAt any
var exp time.Time
var session sql.NullString
r := &refreshRow{}
r.TokenHash = hash
err := o.From(lookup.EntityOAuthRefreshTokens).
Cols(lookup.OAuthRefreshFamilyID, lookup.OAuthRefreshClientID, lookup.OAuthRefreshUserID,
lookup.OAuthRefreshSessionToken, lookup.OAuthRefreshScopes, lookup.OAuthRefreshExtra,
lookup.OAuthRefreshExpiresAt, lookup.OAuthRefreshUsedAt, lookup.OAuthRefreshRevokedAt).
Where(Eq(lookup.OAuthRefreshTokenHash, hash)).
QueryRow(ctx, q, &r.FamilyID, &r.ClientID, &r.UserID, &session, &scopes, &extra, o.timeDest(&exp), &usedAt, &revokedAt)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, lookup.ErrRefreshInvalid
}
return nil, err
}
r.SessionToken = session.String
r.ExpiresAt = exp
_ = o.d.DecodeJSON(scopes, &r.Scopes)
_ = o.d.DecodeJSON(extra, &r.Extra)
_, r.used = o.optTime(usedAt)
_, r.revoked = o.optTime(revokedAt)
r.expired = !exp.After(o.Now())
return r, nil
}
func (o *OAuthGrants) revokeFamilyTx(ctx context.Context, q Querier, family string) error {
_, err := o.Update(lookup.EntityOAuthRefreshTokens).
Set(Set(lookup.OAuthRefreshRevokedAt, o.Now())).
Where(Eq(lookup.OAuthRefreshFamilyID, family), IsNull(lookup.OAuthRefreshRevokedAt)).Exec(ctx, q)
return err
}
// RotateRefresh implements lookup.OAuthGrantStore.
func (o *OAuthGrants) RotateRefresh(ctx context.Context, oldHash string, next lookup.RefreshToken) (*lookup.RefreshToken, error) {
var old *refreshRow
var reused bool
err := o.tx(ctx, func(q Querier) error {
r, err := o.loadRefresh(ctx, q, oldHash)
if err != nil {
return err
}
if r.revoked || r.expired {
return lookup.ErrRefreshInvalid
}
old = r
if r.used {
// A rotated token came back: the family is compromised. The revoke must commit, so
// the reuse is reported after the transaction instead of rolling it back.
reused = true
return o.revokeFamilyTx(ctx, q, r.FamilyID)
}
n, err := o.Update(lookup.EntityOAuthRefreshTokens).
Set(Set(lookup.OAuthRefreshUsedAt, o.Now())).
Where(Eq(lookup.OAuthRefreshTokenHash, oldHash), IsNull(lookup.OAuthRefreshUsedAt)).Exec(ctx, q)
if err != nil {
return err
}
if n == 0 { // lost a race with a concurrent rotation of the same token
reused = true
return o.revokeFamilyTx(ctx, q, r.FamilyID)
}
next.FamilyID = r.FamilyID
next.ClientID = r.ClientID
next.UserID = r.UserID
if next.SessionToken == "" {
next.SessionToken = r.SessionToken
}
return o.insertRefresh(ctx, q, next)
})
if err != nil {
if errors.Is(err, lookup.ErrRefreshInvalid) {
return nil, err
}
return nil, fmt.Errorf("failed to rotate refresh token: %w", err)
}
tok := old.RefreshToken
if reused {
return &tok, lookup.ErrRefreshReused
}
return &tok, nil
}
// PeekRefresh implements lookup.OAuthGrantStore.
func (o *OAuthGrants) PeekRefresh(ctx context.Context, hash string) (*lookup.RefreshToken, error) {
var r *refreshRow
err := o.do(func(q Querier) error {
var err error
r, err = o.loadRefresh(ctx, q, hash)
return err
})
if err != nil {
if errors.Is(err, lookup.ErrRefreshInvalid) {
return nil, err
}
return nil, fmt.Errorf("failed to read refresh token: %w", err)
}
if r.revoked || r.expired { // a rotated token is still returned so its reuse is detected by RotateRefresh
return nil, lookup.ErrRefreshInvalid
}
t := r.RefreshToken
return &t, nil
}
// RevokeRefreshFamily implements lookup.OAuthGrantStore.
func (o *OAuthGrants) RevokeRefreshFamily(ctx context.Context, familyID string) error {
return o.do(func(q Querier) error { return o.revokeFamilyTx(ctx, q, familyID) })
}
// RevokeRefreshBySession implements lookup.OAuthGrantStore.
func (o *OAuthGrants) RevokeRefreshBySession(ctx context.Context, sessionToken string) error {
return o.do(func(q Querier) error {
_, err := o.Update(lookup.EntityOAuthRefreshTokens).
Set(Set(lookup.OAuthRefreshRevokedAt, o.Now())).
Where(Eq(lookup.OAuthRefreshSessionToken, sessionToken), IsNull(lookup.OAuthRefreshRevokedAt)).Exec(ctx, q)
return err
})
}
// CreateDevice implements lookup.OAuthGrantStore.
func (o *OAuthGrants) CreateDevice(ctx context.Context, d lookup.DeviceCode) error {
scopes, err := o.jsonArg(d.Scopes)
if err != nil {
return err
}
status := d.Status
if status == "" {
status = lookup.DevicePending
}
return o.do(func(q Querier) error {
return o.Insert(lookup.EntityOAuthDeviceCodes).Set(
Set(lookup.OAuthDeviceHash, d.DeviceHash),
Set(lookup.OAuthDeviceUserCode, strings.ToUpper(d.UserCode)),
Set(lookup.OAuthDeviceClientID, d.ClientID),
Set(lookup.OAuthDeviceScopes, scopes),
Set(lookup.OAuthDeviceStatus, string(status)),
Set(lookup.OAuthDeviceInterval, d.Interval),
Set(lookup.OAuthDeviceCreatedAt, o.Now()),
Set(lookup.OAuthDeviceExpiresAt, d.ExpiresAt),
).Exec(ctx, q)
})
}
// DeviceByUserCode implements lookup.OAuthGrantStore.
func (o *OAuthGrants) DeviceByUserCode(ctx context.Context, userCode string) (*lookup.DeviceCode, error) {
var d lookup.DeviceCode
var scopes any
var status string
var exp time.Time
userCode = strings.ToUpper(userCode)
err := o.do(func(q Querier) error {
return o.From(lookup.EntityOAuthDeviceCodes).
Cols(lookup.OAuthDeviceHash, lookup.OAuthDeviceClientID, lookup.OAuthDeviceScopes, lookup.OAuthDeviceStatus,
lookup.OAuthDeviceInterval, lookup.OAuthDeviceExpiresAt).
Where(Eq(lookup.OAuthDeviceUserCode, userCode), Eq(lookup.OAuthDeviceStatus, string(lookup.DevicePending)),
Gt(lookup.OAuthDeviceExpiresAt, o.Now())).
QueryRow(ctx, q, &d.DeviceHash, &d.ClientID, &scopes, &status, &d.Interval, o.timeDest(&exp))
})
if errors.Is(err, sql.ErrNoRows) {
return nil, lookup.ErrNotFound
}
if err != nil {
return nil, fmt.Errorf("failed to read device code: %w", err)
}
d.UserCode = userCode
d.Status = lookup.DeviceStatus(status)
d.ExpiresAt = exp
_ = o.d.DecodeJSON(scopes, &d.Scopes)
return &d, nil
}
// DeviceDecide implements lookup.OAuthGrantStore.
func (o *OAuthGrants) DeviceDecide(ctx context.Context, userCode string, approve bool, userID int, sessionToken string) error {
status := lookup.DeviceDenied
sets := []Assignment{}
if approve {
status = lookup.DeviceApproved
sets = append(sets, Set(lookup.OAuthDeviceUserID, userID), Set(lookup.OAuthDeviceSessionToken, sessionToken))
}
sets = append(sets, Set(lookup.OAuthDeviceStatus, string(status)))
var n int64
err := o.do(func(q Querier) error {
var err error
n, err = o.Update(lookup.EntityOAuthDeviceCodes).Set(sets...).
Where(Eq(lookup.OAuthDeviceUserCode, strings.ToUpper(userCode)),
Eq(lookup.OAuthDeviceStatus, string(lookup.DevicePending)),
Gt(lookup.OAuthDeviceExpiresAt, o.Now())).Exec(ctx, q)
return err
})
if err != nil {
return fmt.Errorf("failed to decide device code: %w", err)
}
if n == 0 {
return lookup.ErrNotFound
}
return nil
}
// DevicePoll implements lookup.OAuthGrantStore.
func (o *OAuthGrants) DevicePoll(ctx context.Context, deviceHash string) (*lookup.DeviceCode, error) {
var out *lookup.DeviceCode
var outErr error
err := o.tx(ctx, func(q Querier) error {
var d lookup.DeviceCode
var scopes, polled any
var status string
var userID sql.NullInt64
var session sql.NullString
var exp time.Time
err := o.From(lookup.EntityOAuthDeviceCodes).
Cols(lookup.OAuthDeviceUserCode, lookup.OAuthDeviceClientID, lookup.OAuthDeviceScopes, lookup.OAuthDeviceStatus,
lookup.OAuthDeviceUserID, lookup.OAuthDeviceSessionToken, lookup.OAuthDeviceInterval,
lookup.OAuthDeviceExpiresAt, lookup.OAuthDeviceLastPolledAt).
Where(Eq(lookup.OAuthDeviceHash, deviceHash)).
QueryRow(ctx, q, &d.UserCode, &d.ClientID, &scopes, &status, &userID, &session, &d.Interval, o.timeDest(&exp), &polled)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
outErr = lookup.ErrDeviceExpired
return nil
}
return err
}
now := o.Now()
del := func() error {
_, err := o.Delete(lookup.EntityOAuthDeviceCodes).Where(Eq(lookup.OAuthDeviceHash, deviceHash)).Exec(ctx, q)
return err
}
if !exp.After(now) {
outErr = lookup.ErrDeviceExpired
return del()
}
if last, ok := o.optTime(polled); ok && now.Sub(last) < time.Duration(d.Interval)*time.Second {
outErr = lookup.ErrDeviceSlowDown
}
if _, err := o.Update(lookup.EntityOAuthDeviceCodes).Set(Set(lookup.OAuthDeviceLastPolledAt, now)).
Where(Eq(lookup.OAuthDeviceHash, deviceHash)).Exec(ctx, q); err != nil {
return err
}
if outErr != nil {
return nil
}
switch lookup.DeviceStatus(status) {
case lookup.DeviceDenied:
outErr = lookup.ErrDeviceDenied
return del()
case lookup.DeviceApproved:
d.DeviceHash = deviceHash
d.Status = lookup.DeviceApproved
d.UserID = int(userID.Int64)
d.SessionToken = session.String
d.ExpiresAt = exp
_ = o.d.DecodeJSON(scopes, &d.Scopes)
out = &d
return del()
}
outErr = lookup.ErrDevicePending
return nil
})
if err != nil {
return nil, fmt.Errorf("failed to poll device code: %w", err)
}
return out, outErr
}
// SavePushedRequest implements lookup.OAuthGrantStore.
func (o *OAuthGrants) SavePushedRequest(ctx context.Context, r lookup.PushedRequest) error {
params, err := o.jsonArg(r.Params)
if err != nil {
return err
}
return o.do(func(q Querier) error {
return o.Insert(lookup.EntityOAuthPARRequests).Set(
Set(lookup.OAuthPARRequestURI, r.RequestURI),
Set(lookup.OAuthPARClientID, r.ClientID),
Set(lookup.OAuthPARParams, params),
Set(lookup.OAuthPARCreatedAt, o.Now()),
Set(lookup.OAuthPARExpiresAt, r.ExpiresAt),
).Exec(ctx, q)
})
}
// ConsumePushedRequest implements lookup.OAuthGrantStore.
func (o *OAuthGrants) ConsumePushedRequest(ctx context.Context, requestURI string) (*lookup.PushedRequest, error) {
r := &lookup.PushedRequest{RequestURI: requestURI}
var params any
var exp time.Time
err := o.tx(ctx, func(q Querier) error {
if err := o.From(lookup.EntityOAuthPARRequests).
Cols(lookup.OAuthPARClientID, lookup.OAuthPARParams, lookup.OAuthPARExpiresAt).
Where(Eq(lookup.OAuthPARRequestURI, requestURI), Gt(lookup.OAuthPARExpiresAt, o.Now())).
QueryRow(ctx, q, &r.ClientID, &params, o.timeDest(&exp)); err != nil {
return err
}
n, err := o.Delete(lookup.EntityOAuthPARRequests).Where(Eq(lookup.OAuthPARRequestURI, requestURI)).Exec(ctx, q)
if err != nil {
return err
}
if n == 0 {
return sql.ErrNoRows
}
return nil
})
if errors.Is(err, sql.ErrNoRows) {
return nil, lookup.ErrNotFound
}
if err != nil {
return nil, fmt.Errorf("failed to consume pushed request: %w", err)
}
r.ExpiresAt = exp
_ = o.d.DecodeJSON(params, &r.Params)
return r, nil
}
// SeenJTI implements lookup.OAuthGrantStore.
func (o *OAuthGrants) SeenJTI(ctx context.Context, key string, expires time.Time) (bool, error) {
seen := false
err := o.tx(ctx, func(q Querier) error {
if _, err := o.Delete(lookup.EntityOAuthJTI).Where(Lt(lookup.OAuthJTIExpiresAt, o.Now())).Exec(ctx, q); err != nil {
return err
}
exists, err := o.From(lookup.EntityOAuthJTI).Cols(lookup.OAuthJTIKey).Where(Eq(lookup.OAuthJTIKey, key)).Exists(ctx, q)
if err != nil {
return err
}
if exists {
seen = true
return nil
}
return o.Insert(lookup.EntityOAuthJTI).Set(
Set(lookup.OAuthJTIKey, key), Set(lookup.OAuthJTIExpiresAt, expires)).Exec(ctx, q)
})
if err != nil {
// A concurrent insert of the same key violates the unique index: that is a replay.
if ok, qerr := o.keyExists(ctx, key); qerr == nil && ok {
return true, nil
}
return false, fmt.Errorf("failed to record jti: %w", err)
}
return seen, nil
}
func (o *OAuthGrants) keyExists(ctx context.Context, key string) (bool, error) {
var ok bool
err := o.do(func(q Querier) error {
var err error
ok, err = o.From(lookup.EntityOAuthJTI).Cols(lookup.OAuthJTIKey).Where(Eq(lookup.OAuthJTIKey, key)).Exists(ctx, q)
return err
})
return ok, err
}
+117
View File
@@ -0,0 +1,117 @@
package lookup
import (
"context"
"errors"
"time"
)
// Errors returned by OAuthGrantStore. Callers compare with errors.Is.
var (
// ErrRefreshInvalid: the refresh token is unknown, expired or revoked.
ErrRefreshInvalid = errors.New("invalid refresh token")
// ErrRefreshReused: a refresh token that was already rotated was presented again. The
// store has revoked the whole token family.
ErrRefreshReused = errors.New("refresh token reuse detected")
// ErrDevicePending, ErrDeviceSlowDown, ErrDeviceDenied and ErrDeviceExpired are the RFC 8628
// polling outcomes other than success.
ErrDevicePending = errors.New("authorization pending")
ErrDeviceSlowDown = errors.New("slow down")
ErrDeviceDenied = errors.New("access denied")
ErrDeviceExpired = errors.New("device code expired")
// ErrNotFound: the requested record does not exist or has expired.
ErrNotFound = errors.New("not found")
)
// Consent records that a user allowed a client to act with Scopes.
type Consent struct {
UserID int `json:"user_id"`
ClientID string `json:"client_id"`
Scopes []string `json:"scopes"`
ExpiresAt time.Time `json:"expires_at"`
}
// RefreshToken is a server-managed refresh token. Only the SHA-256 hash of the raw token
// is stored.
type RefreshToken struct {
TokenHash string `json:"token_hash"`
FamilyID string `json:"family_id"`
ClientID string `json:"client_id"`
UserID int `json:"user_id"`
SessionToken string `json:"session_token"`
Scopes []string `json:"scopes,omitempty"`
Extra map[string]any `json:"extra,omitempty"` // nonce, auth_time, acr, sid, dpop_jkt, resource
ExpiresAt time.Time `json:"expires_at"`
}
// DeviceStatus is the state of an RFC 8628 device authorization.
type DeviceStatus string
const (
DevicePending DeviceStatus = "pending"
DeviceApproved DeviceStatus = "approved"
DeviceDenied DeviceStatus = "denied"
)
// DeviceCode is a pending RFC 8628 device authorization. DeviceHash is the SHA-256 hash of the
// device_code returned to the device; UserCode is stored as typed by the user (normalised).
type DeviceCode struct {
DeviceHash string `json:"device_hash"`
UserCode string `json:"user_code"`
ClientID string `json:"client_id"`
Scopes []string `json:"scopes,omitempty"`
Status DeviceStatus `json:"status"`
UserID int `json:"user_id,omitempty"`
SessionToken string `json:"session_token,omitempty"`
Interval int `json:"interval"` // minimum seconds between polls
ExpiresAt time.Time `json:"expires_at"`
}
// PushedRequest is an RFC 9126 pushed authorization request.
type PushedRequest struct {
RequestURI string `json:"request_uri"`
ClientID string `json:"client_id"`
Params map[string]string `json:"params"`
ExpiresAt time.Time `json:"expires_at"`
}
// OAuthGrantStore persists the OAuth2 authorization server state that is not a client, a code
// or a session: consents, refresh tokens, device codes, pushed requests and the replay cache.
type OAuthGrantStore interface {
// SaveConsent replaces the consent of (UserID, ClientID).
SaveConsent(ctx context.Context, c Consent) error
// GetConsent returns the unexpired consent or ErrNotFound.
GetConsent(ctx context.Context, userID int, clientID string) (*Consent, error)
RevokeConsent(ctx context.Context, userID int, clientID string) error
SaveRefresh(ctx context.Context, t RefreshToken) error
// RotateRefresh atomically consumes the token with hash oldHash and stores next in the same
// family. It returns the consumed token. An unknown, expired or revoked token is
// ErrRefreshInvalid. A token that was already consumed revokes its family and returns the
// token together with ErrRefreshReused so the caller can end the session.
RotateRefresh(ctx context.Context, oldHash string, next RefreshToken) (*RefreshToken, error)
// PeekRefresh returns the token without consuming it. Unknown, expired and revoked tokens are
// ErrRefreshInvalid; an already rotated token is returned so RotateRefresh can report its reuse.
PeekRefresh(ctx context.Context, hash string) (*RefreshToken, error)
RevokeRefreshFamily(ctx context.Context, familyID string) error
// RevokeRefreshBySession revokes every refresh token bound to a session token.
RevokeRefreshBySession(ctx context.Context, sessionToken string) error
CreateDevice(ctx context.Context, d DeviceCode) error
// DeviceByUserCode returns the unexpired pending device authorization or ErrNotFound.
DeviceByUserCode(ctx context.Context, userCode string) (*DeviceCode, error)
// DeviceDecide approves or denies the device authorization of userCode.
DeviceDecide(ctx context.Context, userCode string, approve bool, userID int, sessionToken string) error
// DevicePoll implements the token endpoint side: it enforces the poll interval and returns
// one of ErrDevicePending, ErrDeviceSlowDown, ErrDeviceDenied, ErrDeviceExpired or, once
// approved, the record (consumed: it cannot be polled again).
DevicePoll(ctx context.Context, deviceHash string) (*DeviceCode, error)
SavePushedRequest(ctx context.Context, r PushedRequest) error
// ConsumePushedRequest returns and deletes the request or ErrNotFound.
ConsumePushedRequest(ctx context.Context, requestURI string) (*PushedRequest, error)
// SeenJTI records key until expires and reports whether it was already recorded. It is the
// replay cache for DPoP proofs and client assertions.
SeenJTI(ctx context.Context, key string, expires time.Time) (bool, error)
}
+5
View File
@@ -62,6 +62,10 @@ type OAuthClientStore interface {
ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error)
Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error)
Revoke(ctx context.Context, token string) error
// UpdateClient rewrites the registration fields of an existing client (RFC 7592).
UpdateClient(ctx context.Context, client *sectypes.OAuthServerClient) error
// DeleteClient deactivates a client.
DeleteClient(ctx context.Context, clientID string) error
}
// OAuthSession is the session row written after an OAuth2 client login.
@@ -147,6 +151,7 @@ type Provider struct {
Keys KeyStore
OAuthClient OAuthClientStore
OAuthUser OAuthUserStore
OAuthGrant OAuthGrantStore
Passkey PasskeyStore
TOTP TOTPStore
Policy PolicyStore
+35
View File
@@ -61,6 +61,8 @@ const (
OpOAuthExchangeCode Op = "oauth_exchange_code"
OpOAuthIntrospect Op = "oauth_introspect"
OpOAuthRevoke Op = "oauth_revoke"
OpOAuthUpdateClient Op = "oauth_update_client"
OpOAuthDeleteClient Op = "oauth_delete_client"
OpOAuthGetOrCreateUser Op = "oauth_get_or_create_user"
OpOAuthCreateSession Op = "oauth_create_session"
@@ -68,6 +70,22 @@ const (
OpOAuthUpdateRefreshToken Op = "oauth_update_refresh_token" //nolint:gosec // operation name, not a credential
OpOAuthGetUser Op = "oauth_get_user"
OpOAuthSaveConsent Op = "oauth_save_consent"
OpOAuthGetConsent Op = "oauth_get_consent"
OpOAuthRevokeConsent Op = "oauth_revoke_consent"
OpOAuthSaveRefresh Op = "oauth_save_refresh" //nolint:gosec // operation name, not a credential
OpOAuthRotateRefresh Op = "oauth_rotate_refresh" //nolint:gosec // operation name, not a credential
OpOAuthPeekRefresh Op = "oauth_peek_refresh" //nolint:gosec // operation name, not a credential
OpOAuthRevokeRefreshFamily Op = "oauth_revoke_refresh_family" //nolint:gosec // operation name, not a credential
OpOAuthRevokeRefreshByUser Op = "oauth_revoke_refresh_session" //nolint:gosec // operation name, not a credential
OpOAuthCreateDevice Op = "oauth_create_device"
OpOAuthDeviceByUserCode Op = "oauth_device_by_user_code"
OpOAuthDeviceDecide Op = "oauth_device_decide"
OpOAuthDevicePoll Op = "oauth_device_poll"
OpOAuthSavePAR Op = "oauth_save_par"
OpOAuthConsumePAR Op = "oauth_consume_par"
OpOAuthSeenJTI Op = "oauth_seen_jti"
OpPasskeyStore Op = "passkey_store"
OpPasskeyGet Op = "passkey_get"
OpPasskeyUpdateCounter Op = "passkey_update_counter"
@@ -144,11 +162,28 @@ func AllOps() []Op {
OpOAuthExchangeCode,
OpOAuthIntrospect,
OpOAuthRevoke,
OpOAuthUpdateClient,
OpOAuthDeleteClient,
OpOAuthGetOrCreateUser,
OpOAuthCreateSession,
OpOAuthGetRefreshToken,
OpOAuthUpdateRefreshToken,
OpOAuthGetUser,
OpOAuthSaveConsent,
OpOAuthGetConsent,
OpOAuthRevokeConsent,
OpOAuthSaveRefresh,
OpOAuthRotateRefresh,
OpOAuthPeekRefresh,
OpOAuthRevokeRefreshFamily,
OpOAuthRevokeRefreshByUser,
OpOAuthCreateDevice,
OpOAuthDeviceByUserCode,
OpOAuthDeviceDecide,
OpOAuthDevicePoll,
OpOAuthSavePAR,
OpOAuthConsumePAR,
OpOAuthSeenJTI,
OpPasskeyStore,
OpPasskeyGet,
OpPasskeyUpdateCounter,
+35
View File
@@ -313,3 +313,38 @@ func (o *OAuthClients) Revoke(ctx context.Context, token string) error {
}
return nil
}
// UpdateClient implements lookup.OAuthClientStore.
func (o *OAuthClients) UpdateClient(ctx context.Context, client *sectypes.OAuthServerClient) error {
input, err := json.Marshal(client)
if err != nil {
return fmt.Errorf("failed to marshal client: %w", err)
}
var success bool
var errMsg sql.NullString
err = o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error
FROM %s($1::jsonb)
`, o.procs.OAuthUpdateClient), input).Scan(&success, &errMsg)
})
if err != nil {
return fmt.Errorf("failed to update client: %w", err)
}
if !success {
return failure(errMsg, "failed to update client")
}
return nil
}
// DeleteClient implements lookup.OAuthClientStore.
func (o *OAuthClients) DeleteClient(ctx context.Context, clientID string) error {
ok, errMsg, err := o.callNoData(ctx, o.procs.OAuthDeleteClient, clientID)
if err != nil {
return fmt.Errorf("failed to delete client: %w", err)
}
if !ok {
return failure(errMsg, "failed to delete client")
}
return nil
}
@@ -0,0 +1,217 @@
package procedure
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// OAuthGrants implements lookup.OAuthGrantStore with the resolvespec_oauth_* grant procedures.
// Every procedure takes one jsonb request and returns (p_success, p_error, p_data). A failure
// that maps to a lookup sentinel carries a stable code in p_error (see the grantErrors table).
type OAuthGrants struct {
run Runner
procs lookup.ProcNames
}
var _ lookup.OAuthGrantStore = (*OAuthGrants)(nil)
// NewOAuthGrants creates the procedure-backed OAuthGrantStore.
func NewOAuthGrants(run Runner, procs lookup.ProcNames) *OAuthGrants {
return &OAuthGrants{run: run, procs: procs}
}
// grantErrors maps the codes a grant procedure puts in p_error to the lookup sentinels.
var grantErrors = map[string]error{
"not_found": lookup.ErrNotFound,
"refresh_invalid": lookup.ErrRefreshInvalid,
"refresh_reused": lookup.ErrRefreshReused,
"device_pending": lookup.ErrDevicePending,
"device_slowdown": lookup.ErrDeviceSlowDown,
"device_denied": lookup.ErrDeviceDenied,
"device_expired": lookup.ErrDeviceExpired,
}
// call runs proc with the JSON-encoded request. The returned data is the p_data of the
// procedure, also when it reports a failure (rotate returns the reused token that way).
func (o *OAuthGrants) call(ctx context.Context, proc string, req any) (data []byte, err error) {
input, err := json.Marshal(req)
if err != nil {
return nil, fmt.Errorf("failed to marshal request: %w", err)
}
var success bool
var errMsg sql.NullString
err = o.run.Run(func(db *sql.DB) error {
return db.QueryRowContext(ctx, fmt.Sprintf(`
SELECT p_success, p_error, p_data::text
FROM %s($1::jsonb)
`, proc), input).Scan(&success, &errMsg, &data)
})
if err != nil {
return nil, fmt.Errorf("%s: %w", proc, err)
}
if success {
return data, nil
}
if e, ok := grantErrors[errMsg.String]; ok {
return data, e
}
return data, failure(errMsg, proc+" failed")
}
// SaveConsent implements lookup.OAuthGrantStore.
func (o *OAuthGrants) SaveConsent(ctx context.Context, c lookup.Consent) error {
_, err := o.call(ctx, o.procs.OAuthSaveConsent, c)
return err
}
// GetConsent implements lookup.OAuthGrantStore.
func (o *OAuthGrants) GetConsent(ctx context.Context, userID int, clientID string) (*lookup.Consent, error) {
data, err := o.call(ctx, o.procs.OAuthGetConsent, map[string]any{"user_id": userID, "client_id": clientID})
if err != nil {
return nil, err
}
var c lookup.Consent
if err := json.Unmarshal(normalizeTimes(data), &c); err != nil {
return nil, fmt.Errorf("failed to parse consent: %w", err)
}
return &c, nil
}
// RevokeConsent implements lookup.OAuthGrantStore.
func (o *OAuthGrants) RevokeConsent(ctx context.Context, userID int, clientID string) error {
_, err := o.call(ctx, o.procs.OAuthRevokeConsent, map[string]any{"user_id": userID, "client_id": clientID})
return err
}
// SaveRefresh implements lookup.OAuthGrantStore.
func (o *OAuthGrants) SaveRefresh(ctx context.Context, t lookup.RefreshToken) error {
_, err := o.call(ctx, o.procs.OAuthSaveRefresh, t)
return err
}
func parseRefresh(data []byte) (*lookup.RefreshToken, error) {
if len(data) == 0 {
return nil, nil
}
var t lookup.RefreshToken
if err := json.Unmarshal(normalizeTimes(data), &t); err != nil {
return nil, fmt.Errorf("failed to parse refresh token: %w", err)
}
return &t, nil
}
// RotateRefresh implements lookup.OAuthGrantStore.
func (o *OAuthGrants) RotateRefresh(ctx context.Context, oldHash string, next lookup.RefreshToken) (*lookup.RefreshToken, error) {
data, err := o.call(ctx, o.procs.OAuthRotateRefresh, map[string]any{"old_hash": oldHash, "next": next})
if err != nil && err != lookup.ErrRefreshReused { //nolint:errorlint // sentinel returned unwrapped by call
return nil, err
}
t, perr := parseRefresh(data)
if perr != nil {
return nil, perr
}
return t, err
}
// PeekRefresh implements lookup.OAuthGrantStore.
func (o *OAuthGrants) PeekRefresh(ctx context.Context, hash string) (*lookup.RefreshToken, error) {
data, err := o.call(ctx, o.procs.OAuthPeekRefresh, map[string]any{"token_hash": hash})
if err != nil {
return nil, err
}
return parseRefresh(data)
}
// RevokeRefreshFamily implements lookup.OAuthGrantStore.
func (o *OAuthGrants) RevokeRefreshFamily(ctx context.Context, familyID string) error {
_, err := o.call(ctx, o.procs.OAuthRevokeRefreshFamily, map[string]any{"family_id": familyID})
return err
}
// RevokeRefreshBySession implements lookup.OAuthGrantStore.
func (o *OAuthGrants) RevokeRefreshBySession(ctx context.Context, sessionToken string) error {
_, err := o.call(ctx, o.procs.OAuthRevokeRefreshByUser, map[string]any{"session_token": sessionToken})
return err
}
// CreateDevice implements lookup.OAuthGrantStore.
func (o *OAuthGrants) CreateDevice(ctx context.Context, d lookup.DeviceCode) error {
if d.Status == "" {
d.Status = lookup.DevicePending
}
_, err := o.call(ctx, o.procs.OAuthCreateDevice, d)
return err
}
func parseDevice(data []byte) (*lookup.DeviceCode, error) {
var d lookup.DeviceCode
if err := json.Unmarshal(normalizeTimes(data), &d); err != nil {
return nil, fmt.Errorf("failed to parse device code: %w", err)
}
return &d, nil
}
// DeviceByUserCode implements lookup.OAuthGrantStore.
func (o *OAuthGrants) DeviceByUserCode(ctx context.Context, userCode string) (*lookup.DeviceCode, error) {
data, err := o.call(ctx, o.procs.OAuthDeviceByUserCode, map[string]any{"user_code": userCode})
if err != nil {
return nil, err
}
return parseDevice(data)
}
// DeviceDecide implements lookup.OAuthGrantStore.
func (o *OAuthGrants) DeviceDecide(ctx context.Context, userCode string, approve bool, userID int, sessionToken string) error {
_, err := o.call(ctx, o.procs.OAuthDeviceDecide, map[string]any{
"user_code": userCode, "approve": approve, "user_id": userID, "session_token": sessionToken,
})
return err
}
// DevicePoll implements lookup.OAuthGrantStore.
func (o *OAuthGrants) DevicePoll(ctx context.Context, deviceHash string) (*lookup.DeviceCode, error) {
data, err := o.call(ctx, o.procs.OAuthDevicePoll, map[string]any{"device_hash": deviceHash})
if err != nil {
return nil, err
}
return parseDevice(data)
}
// SavePushedRequest implements lookup.OAuthGrantStore.
func (o *OAuthGrants) SavePushedRequest(ctx context.Context, r lookup.PushedRequest) error {
_, err := o.call(ctx, o.procs.OAuthSavePAR, r)
return err
}
// ConsumePushedRequest implements lookup.OAuthGrantStore.
func (o *OAuthGrants) ConsumePushedRequest(ctx context.Context, requestURI string) (*lookup.PushedRequest, error) {
data, err := o.call(ctx, o.procs.OAuthConsumePAR, map[string]any{"request_uri": requestURI})
if err != nil {
return nil, err
}
var r lookup.PushedRequest
if err := json.Unmarshal(normalizeTimes(data), &r); err != nil {
return nil, fmt.Errorf("failed to parse pushed request: %w", err)
}
return &r, nil
}
// SeenJTI implements lookup.OAuthGrantStore.
func (o *OAuthGrants) SeenJTI(ctx context.Context, key string, expires time.Time) (bool, error) {
data, err := o.call(ctx, o.procs.OAuthSeenJTI, map[string]any{"key": key, "expires_at": expires})
if err != nil {
return false, err
}
var out struct {
Seen bool `json:"seen"`
}
if err := json.Unmarshal(data, &out); err != nil {
return false, fmt.Errorf("failed to parse jti result: %w", err)
}
return out.Seen, nil
}
+37
View File
@@ -63,6 +63,26 @@ type ProcNames struct {
OAuthExchangeCode string // default: "resolvespec_oauth_exchange_code"
OAuthIntrospect string // default: "resolvespec_oauth_introspect"
OAuthRevoke string // default: "resolvespec_oauth_revoke"
OAuthUpdateClient string // default: "resolvespec_oauth_update_client"
OAuthDeleteClient string // default: "resolvespec_oauth_delete_client"
// OAuth2 server grant procedures (consents, refresh tokens, device codes, PAR, replay cache).
// Each takes a jsonb request and returns (p_success, p_error, p_data).
OAuthSaveConsent string // default: "resolvespec_oauth_save_consent"
OAuthGetConsent string // default: "resolvespec_oauth_get_consent"
OAuthRevokeConsent string // default: "resolvespec_oauth_revoke_consent"
OAuthSaveRefresh string // default: "resolvespec_oauth_save_refresh"
OAuthRotateRefresh string // default: "resolvespec_oauth_rotate_refresh"
OAuthPeekRefresh string // default: "resolvespec_oauth_peek_refresh"
OAuthRevokeRefreshFamily string // default: "resolvespec_oauth_revoke_refresh_family"
OAuthRevokeRefreshByUser string // default: "resolvespec_oauth_revoke_refresh_session"
OAuthCreateDevice string // default: "resolvespec_oauth_create_device"
OAuthDeviceByUserCode string // default: "resolvespec_oauth_device_by_user_code"
OAuthDeviceDecide string // default: "resolvespec_oauth_device_decide"
OAuthDevicePoll string // default: "resolvespec_oauth_device_poll"
OAuthSavePAR string // default: "resolvespec_oauth_save_par"
OAuthConsumePAR string // default: "resolvespec_oauth_consume_par"
OAuthSeenJTI string // default: "resolvespec_oauth_seen_jti"
// Keystore procedures (KeyStore)
KeystoreGetUserKeys string // default: "resolvespec_keystore_get_user_keys"
@@ -112,6 +132,23 @@ func DefaultProcNames() ProcNames {
OAuthExchangeCode: "resolvespec_oauth_exchange_code",
OAuthIntrospect: "resolvespec_oauth_introspect",
OAuthRevoke: "resolvespec_oauth_revoke",
OAuthUpdateClient: "resolvespec_oauth_update_client",
OAuthDeleteClient: "resolvespec_oauth_delete_client",
OAuthSaveConsent: "resolvespec_oauth_save_consent",
OAuthGetConsent: "resolvespec_oauth_get_consent",
OAuthRevokeConsent: "resolvespec_oauth_revoke_consent",
OAuthSaveRefresh: "resolvespec_oauth_save_refresh",
OAuthRotateRefresh: "resolvespec_oauth_rotate_refresh",
OAuthPeekRefresh: "resolvespec_oauth_peek_refresh",
OAuthRevokeRefreshFamily: "resolvespec_oauth_revoke_refresh_family",
OAuthRevokeRefreshByUser: "resolvespec_oauth_revoke_refresh_session",
OAuthCreateDevice: "resolvespec_oauth_create_device",
OAuthDeviceByUserCode: "resolvespec_oauth_device_by_user_code",
OAuthDeviceDecide: "resolvespec_oauth_device_decide",
OAuthDevicePoll: "resolvespec_oauth_device_poll",
OAuthSavePAR: "resolvespec_oauth_save_par",
OAuthConsumePAR: "resolvespec_oauth_consume_par",
OAuthSeenJTI: "resolvespec_oauth_seen_jti",
KeystoreGetUserKeys: "resolvespec_keystore_get_user_keys",
KeystoreCreateKey: "resolvespec_keystore_create_key",
KeystoreDeleteKey: "resolvespec_keystore_delete_key",
+63 -2
View File
@@ -26,6 +26,11 @@ const (
EntityUserPasswordResets Entity = "user_password_resets"
EntityOAuthClients Entity = "oauth_clients"
EntityOAuthCodes Entity = "oauth_codes"
EntityOAuthConsents Entity = "oauth_consents"
EntityOAuthRefreshTokens Entity = "oauth_refresh_tokens" //nolint:gosec // table name, not a credential
EntityOAuthDeviceCodes Entity = "oauth_device_codes"
EntityOAuthPARRequests Entity = "oauth_par_requests"
EntityOAuthJTI Entity = "oauth_jti"
EntityUserKeys Entity = "user_keys"
EntitySecGroupMembers Entity = "sec_group_members"
EntitySecColumnRules Entity = "sec_column_rules"
@@ -121,6 +126,7 @@ var (
OAuthClientsClientSecretHash = col(EntityOAuthClients, "client_secret_hash")
OAuthClientsTokenEndpointAuthMethod = col(EntityOAuthClients, "token_endpoint_auth_method")
OAuthClientsIsActive = col(EntityOAuthClients, "is_active")
OAuthClientsMetadata = col(EntityOAuthClients, "metadata")
OAuthClientsCreatedAt = col(EntityOAuthClients, "created_at")
OAuthCodesID = col(EntityOAuthCodes, "id")
@@ -135,6 +141,51 @@ var (
OAuthCodesScopes = col(EntityOAuthCodes, "scopes")
OAuthCodesExpiresAt = col(EntityOAuthCodes, "expires_at")
OAuthCodesCreatedAt = col(EntityOAuthCodes, "created_at")
OAuthCodesExtra = col(EntityOAuthCodes, "extra")
OAuthConsentsID = col(EntityOAuthConsents, "id")
OAuthConsentsUserID = col(EntityOAuthConsents, "user_id")
OAuthConsentsClientID = col(EntityOAuthConsents, "client_id")
OAuthConsentsScopes = col(EntityOAuthConsents, "scopes")
OAuthConsentsCreatedAt = col(EntityOAuthConsents, "created_at")
OAuthConsentsExpiresAt = col(EntityOAuthConsents, "expires_at")
OAuthRefreshID = col(EntityOAuthRefreshTokens, "id")
OAuthRefreshTokenHash = col(EntityOAuthRefreshTokens, "token_hash")
OAuthRefreshFamilyID = col(EntityOAuthRefreshTokens, "family_id")
OAuthRefreshClientID = col(EntityOAuthRefreshTokens, "client_id")
OAuthRefreshUserID = col(EntityOAuthRefreshTokens, "user_id")
OAuthRefreshSessionToken = col(EntityOAuthRefreshTokens, "session_token")
OAuthRefreshScopes = col(EntityOAuthRefreshTokens, "scopes")
OAuthRefreshExtra = col(EntityOAuthRefreshTokens, "extra")
OAuthRefreshCreatedAt = col(EntityOAuthRefreshTokens, "created_at")
OAuthRefreshExpiresAt = col(EntityOAuthRefreshTokens, "expires_at")
OAuthRefreshUsedAt = col(EntityOAuthRefreshTokens, "used_at")
OAuthRefreshRevokedAt = col(EntityOAuthRefreshTokens, "revoked_at")
OAuthDeviceID = col(EntityOAuthDeviceCodes, "id")
OAuthDeviceHash = col(EntityOAuthDeviceCodes, "device_hash")
OAuthDeviceUserCode = col(EntityOAuthDeviceCodes, "user_code")
OAuthDeviceClientID = col(EntityOAuthDeviceCodes, "client_id")
OAuthDeviceScopes = col(EntityOAuthDeviceCodes, "scopes")
OAuthDeviceStatus = col(EntityOAuthDeviceCodes, "status")
OAuthDeviceUserID = col(EntityOAuthDeviceCodes, "user_id")
OAuthDeviceSessionToken = col(EntityOAuthDeviceCodes, "session_token")
OAuthDeviceInterval = col(EntityOAuthDeviceCodes, "poll_interval")
OAuthDeviceCreatedAt = col(EntityOAuthDeviceCodes, "created_at")
OAuthDeviceExpiresAt = col(EntityOAuthDeviceCodes, "expires_at")
OAuthDeviceLastPolledAt = col(EntityOAuthDeviceCodes, "last_polled_at")
OAuthPARID = col(EntityOAuthPARRequests, "id")
OAuthPARRequestURI = col(EntityOAuthPARRequests, "request_uri")
OAuthPARClientID = col(EntityOAuthPARRequests, "client_id")
OAuthPARParams = col(EntityOAuthPARRequests, "params")
OAuthPARCreatedAt = col(EntityOAuthPARRequests, "created_at")
OAuthPARExpiresAt = col(EntityOAuthPARRequests, "expires_at")
OAuthJTIID = col(EntityOAuthJTI, "id")
OAuthJTIKey = col(EntityOAuthJTI, "jti_key")
OAuthJTIExpiresAt = col(EntityOAuthJTI, "expires_at")
KeysID = col(EntityUserKeys, "id")
KeysUserID = col(EntityUserKeys, "user_id")
@@ -190,10 +241,20 @@ var allColumns = []Column{
ResetsID, ResetsUserID, ResetsTokenHash, ResetsExpiresAt, ResetsCreatedAt, ResetsUsed, ResetsUsedAt,
OAuthClientsID, OAuthClientsClientID, OAuthClientsRedirectURIs, OAuthClientsClientName, OAuthClientsGrantTypes,
OAuthClientsAllowedScopes, OAuthClientsClientSecretHash, OAuthClientsTokenEndpointAuthMethod,
OAuthClientsIsActive, OAuthClientsCreatedAt,
OAuthClientsIsActive, OAuthClientsCreatedAt, OAuthClientsMetadata,
OAuthCodesID, OAuthCodesCode, OAuthCodesClientID, OAuthCodesRedirectURI, OAuthCodesClientState,
OAuthCodesCodeChallenge, OAuthCodesCodeChallengeMethod, OAuthCodesSessionToken, OAuthCodesRefreshToken,
OAuthCodesScopes, OAuthCodesExpiresAt, OAuthCodesCreatedAt,
OAuthCodesScopes, OAuthCodesExpiresAt, OAuthCodesCreatedAt, OAuthCodesExtra,
OAuthConsentsID, OAuthConsentsUserID, OAuthConsentsClientID, OAuthConsentsScopes, OAuthConsentsCreatedAt,
OAuthConsentsExpiresAt,
OAuthRefreshID, OAuthRefreshTokenHash, OAuthRefreshFamilyID, OAuthRefreshClientID, OAuthRefreshUserID,
OAuthRefreshSessionToken, OAuthRefreshScopes, OAuthRefreshExtra, OAuthRefreshCreatedAt, OAuthRefreshExpiresAt,
OAuthRefreshUsedAt, OAuthRefreshRevokedAt,
OAuthDeviceID, OAuthDeviceHash, OAuthDeviceUserCode, OAuthDeviceClientID, OAuthDeviceScopes, OAuthDeviceStatus,
OAuthDeviceUserID, OAuthDeviceSessionToken, OAuthDeviceInterval, OAuthDeviceCreatedAt, OAuthDeviceExpiresAt,
OAuthDeviceLastPolledAt,
OAuthPARID, OAuthPARRequestURI, OAuthPARClientID, OAuthPARParams, OAuthPARCreatedAt, OAuthPARExpiresAt,
OAuthJTIID, OAuthJTIKey, OAuthJTIExpiresAt,
KeysID, KeysUserID, KeysKeyType, KeysKeyHash, KeysName, KeysScopes, KeysMeta, KeysExpiresAt,
KeysCreatedAt, KeysLastUsedAt, KeysIsActive,
GroupMembersGroupID, GroupMembersUserID,
+174
View File
@@ -0,0 +1,174 @@
package security
import (
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"database/sql"
"fmt"
"log"
"net/http"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// ExampleOAuth2FullServer runs a complete OAuth 2.1 / OpenID Connect provider: login, consent,
// rotating refresh tokens, JWT access tokens, DPoP, PAR, the device grant and token exchange.
// OAUTH2_SERVER.md walks through every endpoint.
func ExampleOAuth2FullServer() {
db, _ := sql.Open("postgres", "postgres://user:pass@localhost/app?sslmode=disable")
// 1. The authenticator holds users and sessions. The OAuth state (clients, codes, consents,
// refresh tokens, device codes, PAR requests, replay cache) lives in the same database:
// apply lookup/database_schema.sql (Postgres procedures) or lookup/ddl/<dialect>.sql.
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{
Lookup: lookup.Config{},
})
// 2. Signing keys are persistent so every instance publishes and accepts the same keys.
// The first key signs; add the next key here first when rotating.
key, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) // load from your secret store instead
srv := NewOAuthServer(OAuthServerConfig{ //nolint:gosec // example secrets
Issuer: "https://auth.example.com",
SigningKeys: []OAuthSigningKey{{Key: key, Alg: "ES256"}},
// Share this secret between instances so the SSO cookie works on all of them.
CookieSecret: []byte("a-32-byte-or-longer-random-secret"),
PersistClients: true, // clients survive restarts
PersistCodes: true, // codes work across instances
RequireConsent: true, // ask the user before releasing scopes to a third-party client
ManagedRefreshTokens: true, // rotate refresh tokens, revoke the family on reuse
JWTAccessTokens: true, // RFC 9068: resource servers verify locally
AccessTokenAudience: "https://api.example.com",
EnableDPoP: true, // RFC 9449 sender-constrained tokens
EnablePAR: true, // RFC 9126
EnableDeviceFlow: true, // RFC 8628 for TVs and CLIs
EnableTokenExchange: true, // RFC 8693 downscoping for service calls
InitialAccessToken: "registration-secret", // only trusted callers may register clients
ScopeDescriptions: map[string]string{"orders:read": "Read your orders"},
RateLimiter: func(r *http.Request, endpoint string) bool {
return true // plug in your limiter; false answers 429
},
}, auth)
// 3. First-party applications skip the consent screen. The secret is shown once.
app, secret, err := srv.RegisterTrustedClient(context.Background(), OAuthServerClient{
ClientName: "Admin console",
RedirectURIs: []string{"https://console.example.com/callback"},
GrantTypes: []string{"authorization_code", "refresh_token"},
AllowedScopes: []string{"openid", "profile", "email", "offline_access"},
})
if err != nil {
srv.Close()
log.Fatal(err)
}
fmt.Println(app.ClientID, secret)
mux := http.NewServeMux()
mux.Handle("/", srv.HTTPHandler())
// 4. A protected API verifies the JWT access token locally (no database call).
mux.Handle("/api/orders", requireScope(srv, "orders:read", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
claims, _ := r.Context().Value(accessClaimsKey{}).(*AccessTokenClaims)
fmt.Fprintf(w, "orders of %s", claims.Subject)
})))
httpSrv := &http.Server{Addr: ":8443", Handler: mux, ReadHeaderTimeout: 10 * time.Second}
err = httpSrv.ListenAndServe()
srv.Close()
log.Fatal(err)
}
type accessClaimsKey struct{}
// requireScope is a resource-server middleware around OAuthServer.VerifyAccessToken.
func requireScope(srv *OAuthServer, scope string, next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
token := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ")
claims, err := srv.VerifyAccessToken(r.Context(), token, VerifyAccessTokenOptions{
Audience: "https://api.example.com",
Scopes: []string{scope},
})
if err != nil {
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`)
http.Error(w, "unauthorized", http.StatusUnauthorized)
return
}
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), accessClaimsKey{}, claims)))
})
}
// ExampleOAuth2FullClient is the relying-party side: log users in with any OpenID Connect
// provider (including the server above). Discovery, PKCE, nonce and id_token validation are
// automatic.
func ExampleOAuth2FullClient() {
db, _ := sql.Open("postgres", "postgres://user:pass@localhost/app?sslmode=disable")
auth := NewDatabaseAuthenticator(db)
if _, err := auth.WithOIDC(context.Background(), OIDCConfig{
Issuer: "https://auth.example.com",
ClientID: "my-client-id",
ClientSecret: "my-client-secret", // empty for a public client
RedirectURL: "https://app.example.com/auth/callback",
ProviderName: "company",
Scopes: []string{"openid", "profile", "email", "offline_access"},
}); err != nil {
log.Fatal(err)
}
mux := http.NewServeMux()
mux.HandleFunc("/auth/login", func(w http.ResponseWriter, r *http.Request) {
state, _ := auth.OAuth2GenerateState()
// Keep state in a cookie so the callback can be tied to this browser.
http.SetCookie(w, &http.Cookie{Name: "oauth_state", Value: state, Path: "/", HttpOnly: true, Secure: true, SameSite: http.SameSiteLaxMode, MaxAge: 600})
maxAge := 3600
url, err := auth.OAuth2GetAuthURLWithOptions("company", state, OAuth2AuthOptions{MaxAge: &maxAge})
if err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
http.Redirect(w, r, url, http.StatusFound)
})
mux.HandleFunc("/auth/callback", func(w http.ResponseWriter, r *http.Request) {
if c, err := r.Cookie("oauth_state"); err != nil || c.Value != r.URL.Query().Get("state") {
http.Error(w, "state mismatch", http.StatusBadRequest)
return
}
// Checks state, PKCE, the RFC 9207 iss parameter, the id_token (signature, iss, aud, exp,
// nonce, at_hash) and the userinfo subject; then creates the local user and session.
login, err := auth.OAuth2HandleCallbackRequest(r.Context(), "company", r)
if err != nil {
http.Error(w, "login failed", http.StatusUnauthorized)
return
}
idToken, _ := login.Meta["id_token"].(string) // keep it for logout
http.SetCookie(w, &http.Cookie{Name: "session", Value: login.Token, Path: "/", HttpOnly: true, Secure: true, SameSite: http.SameSiteLaxMode})
http.SetCookie(w, &http.Cookie{Name: "id_token", Value: idToken, Path: "/", HttpOnly: true, Secure: true, SameSite: http.SameSiteLaxMode})
http.Redirect(w, r, "/", http.StatusFound)
})
mux.HandleFunc("/auth/logout", func(w http.ResponseWriter, r *http.Request) {
hint := ""
if c, err := r.Cookie("id_token"); err == nil {
hint = c.Value
}
// Ends the provider session too (RP-initiated logout).
url, err := auth.OAuth2LogoutURL(r.Context(), "company", hint, "https://app.example.com/", "bye")
if err != nil {
http.Redirect(w, r, "/", http.StatusFound)
return
}
http.Redirect(w, r, url, http.StatusFound)
})
srv := &http.Server{Addr: ":8080", Handler: mux, ReadHeaderTimeout: 10 * time.Second}
log.Fatal(srv.ListenAndServe())
}
+240 -43
View File
@@ -8,6 +8,8 @@ import (
"encoding/json"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"sync"
"time"
@@ -32,6 +34,45 @@ type OAuth2Config struct {
// Optional: Custom user info parser
// If not provided, will use standard claims (sub, email, name)
UserInfoParser func(userInfo map[string]any) (*UserContext, error)
// --- OpenID Connect (see oidc_client.go) ---
// Issuer turns the provider into an OpenID Connect provider: PKCE and a nonce are used and
// the id_token returned by the token endpoint is validated (signature, iss, aud, exp, nonce,
// at_hash). WithOIDC fills the endpoints in by discovery; with WithOAuth2 set JWKSURL
// as well. UserInfoURL stays optional: the id_token claims are used when it is empty.
Issuer string
// JWKSURL is the provider's key set. Only needed with WithOAuth2; WithOIDC discovers it.
JWKSURL string
// EndSessionURL is the provider's RP-initiated logout endpoint (discovered by WithOIDC).
EndSessionURL string
// UsePKCE sends a PKCE S256 challenge for a provider that is not OIDC. It is always on in OIDC mode.
UsePKCE bool
// AllowedAlgs lists the id_token signature algorithms to accept. Default: RS256, PS256, ES256, ES384.
AllowedAlgs []string
// AuthStyle selects how the client authenticates at the token endpoint: "basic", "post" or ""
// (try basic, fall back to post).
AuthStyle string
// HTTPClient is used for discovery, JWKS, token and userinfo requests.
HTTPClient *http.Client
// ClockSkew tolerates clock differences when validating the id_token. Default 1 minute.
ClockSkew time.Duration
}
// OAuth2AuthOptions are optional OpenID Connect authentication request parameters.
type OAuth2AuthOptions struct {
LoginHint string
Prompt string // none, login, consent, select_account
MaxAge *int
ACRValues string
Extra map[string]string
}
// oauth2State is what the login redirect remembers until the callback.
type oauth2State struct {
expiry time.Time
verifier string // PKCE code_verifier
nonce string
}
// OAuth2Provider holds configuration and state for a single OAuth2 provider
@@ -40,7 +81,10 @@ type OAuth2Provider struct {
userInfoURL string
userInfoParser func(userInfo map[string]any) (*UserContext, error)
providerName string
states map[string]time.Time // state -> expiry time
states map[string]*oauth2State
oidc *oidcProvider // nil for plain OAuth2
usePKCE bool
httpClient *http.Client
statesMutex sync.RWMutex
stopCh chan struct{} // closed to stop cleanupStates
stopOnce sync.Once
@@ -58,6 +102,13 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
cfg.UserInfoParser = defaultOAuth2UserInfoParser
}
authStyle := oauth2.AuthStyleAutoDetect
switch cfg.AuthStyle {
case "basic":
authStyle = oauth2.AuthStyleInHeader
case "post":
authStyle = oauth2.AuthStyleInParams
}
provider := &OAuth2Provider{
config: &oauth2.Config{
ClientID: cfg.ClientID,
@@ -65,15 +116,22 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
RedirectURL: cfg.RedirectURL,
Scopes: cfg.Scopes,
Endpoint: oauth2.Endpoint{
AuthURL: cfg.AuthURL,
TokenURL: cfg.TokenURL,
AuthURL: cfg.AuthURL,
TokenURL: cfg.TokenURL,
AuthStyle: authStyle,
},
},
userInfoURL: cfg.UserInfoURL,
userInfoParser: cfg.UserInfoParser,
providerName: cfg.ProviderName,
states: make(map[string]time.Time),
states: make(map[string]*oauth2State),
stopCh: make(chan struct{}),
usePKCE: cfg.UsePKCE,
httpClient: cfg.HTTPClient,
}
if cfg.Issuer != "" {
provider.oidc = newOIDCProvider(&cfg)
provider.usePKCE = true
}
// Initialize providers map if needed
@@ -97,17 +155,56 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
// OAuth2GetAuthURL returns the OAuth2 authorization URL for redirecting users
func (a *DatabaseAuthenticator) OAuth2GetAuthURL(providerName, state string) (string, error) {
return a.OAuth2GetAuthURLWithOptions(providerName, state, OAuth2AuthOptions{})
}
// OAuth2GetAuthURLWithOptions is OAuth2GetAuthURL with OpenID Connect request parameters. For an
// OIDC provider (and with UsePKCE) it also creates the PKCE verifier and the nonce, which are
// kept with the state until the callback.
func (a *DatabaseAuthenticator) OAuth2GetAuthURLWithOptions(providerName, state string, opts OAuth2AuthOptions) (string, error) {
provider, err := a.getOAuth2Provider(providerName)
if err != nil {
return "", err
}
if provider.oidc != nil {
dctx, cancel := context.WithTimeout(provider.withHTTPClient(context.Background()), 15*time.Second)
defer cancel()
if err := provider.oidc.ensureEndpoints(dctx, provider); err != nil {
return "", err
}
}
st := &oauth2State{expiry: time.Now().Add(10 * time.Minute)}
var params []oauth2.AuthCodeOption
if provider.usePKCE {
st.verifier = oauth2.GenerateVerifier()
params = append(params, oauth2.S256ChallengeOption(st.verifier))
}
if provider.oidc != nil {
if st.nonce, err = randomOAuthToken(); err != nil {
return "", err
}
params = append(params, oauth2.SetAuthURLParam("nonce", st.nonce))
}
set := func(k, v string) {
if v != "" {
params = append(params, oauth2.SetAuthURLParam(k, v))
}
}
set("login_hint", opts.LoginHint)
set("prompt", opts.Prompt)
set("acr_values", opts.ACRValues)
if opts.MaxAge != nil {
set("max_age", strconv.Itoa(*opts.MaxAge))
}
for k, v := range opts.Extra {
set(k, v)
}
// Store state for validation
provider.statesMutex.Lock()
provider.states[state] = time.Now().Add(10 * time.Minute)
provider.states[state] = st
provider.statesMutex.Unlock()
return provider.config.AuthCodeURL(state), nil
return provider.config.AuthCodeURL(state, params...), nil
}
// OAuth2GenerateState generates a random state string for CSRF protection
@@ -121,42 +218,97 @@ func (a *DatabaseAuthenticator) OAuth2GenerateState() (string, error) {
// OAuth2HandleCallback handles the OAuth2 callback and exchanges code for token
func (a *DatabaseAuthenticator) OAuth2HandleCallback(ctx context.Context, providerName, code, state string) (*LoginResponse, error) {
return a.oauth2Callback(ctx, providerName, code, state, "")
}
// OAuth2HandleCallbackRequest is OAuth2HandleCallback for the redirect request itself. Besides
// code and state it honours the error parameters and the RFC 9207 "iss" parameter, which
// protects against mix-up attacks when several providers are in use.
func (a *DatabaseAuthenticator) OAuth2HandleCallbackRequest(ctx context.Context, providerName string, r *http.Request) (*LoginResponse, error) {
q := r.URL.Query()
if e := q.Get("error"); e != "" {
return nil, fmt.Errorf("provider returned an error: %s %s", e, q.Get("error_description"))
}
return a.oauth2Callback(ctx, providerName, q.Get("code"), q.Get("state"), q.Get("iss"))
}
func (a *DatabaseAuthenticator) oauth2Callback(ctx context.Context, providerName, code, state, iss string) (*LoginResponse, error) {
provider, err := a.getOAuth2Provider(providerName)
if err != nil {
return nil, err
}
// Validate state
if !provider.validateState(state) {
st, ok := provider.validateState(state)
if !ok {
return nil, fmt.Errorf("invalid state parameter")
}
if code == "" {
return nil, fmt.Errorf("missing authorization code")
}
if provider.oidc != nil && iss != "" && iss != provider.oidc.issuer {
return nil, fmt.Errorf("authorization response issuer mismatch")
}
if ctx = provider.withHTTPClient(ctx); ctx == nil {
return nil, fmt.Errorf("no context")
}
// Exchange code for token
token, err := provider.config.Exchange(ctx, code)
var exchange []oauth2.AuthCodeOption
if st.verifier != "" {
exchange = append(exchange, oauth2.VerifierOption(st.verifier))
}
if provider.oidc != nil {
if err := provider.oidc.ensureEndpoints(ctx, provider); err != nil {
return nil, err
}
}
token, err := provider.config.Exchange(ctx, code, exchange...)
if err != nil {
return nil, fmt.Errorf("failed to exchange code: %w", err)
}
// OpenID Connect: validate the id_token.
var rawIDToken string
var idClaims map[string]any
if provider.oidc != nil {
rawIDToken, _ = token.Extra("id_token").(string)
if rawIDToken == "" && oauthSliceContains(provider.config.Scopes, "openid") {
return nil, fmt.Errorf("token response contains no id_token")
}
if rawIDToken != "" {
if idClaims, err = provider.oidc.validateIDToken(ctx, rawIDToken, st.nonce, token.AccessToken); err != nil {
return nil, fmt.Errorf("invalid id_token: %w", err)
}
}
}
// Fetch user info
client := provider.config.Client(ctx, token)
resp, err := client.Get(provider.userInfoURL)
if err != nil {
return nil, fmt.Errorf("failed to fetch user info: %w", err)
userInfo := map[string]any{}
if provider.userInfoURL != "" {
fetched, err := provider.fetchUserInfo(ctx, token)
switch {
case err == nil:
if sub, _ := idClaims["sub"].(string); sub != "" {
if us, _ := fetched["sub"].(string); us != "" && us != sub {
return nil, fmt.Errorf("userinfo subject does not match the id_token")
}
}
userInfo = fetched
case provider.oidc == nil || idClaims == nil:
return nil, err
}
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("failed to read user info: %w", err)
claims := map[string]any{}
for k, v := range idClaims {
claims[k] = v
}
var userInfo map[string]any
if err := json.Unmarshal(body, &userInfo); err != nil {
return nil, fmt.Errorf("failed to parse user info: %w", err)
for k, v := range userInfo {
claims[k] = v
}
// Parse user info
userCtx, err := provider.userInfoParser(userInfo)
userCtx, err := provider.userInfoParser(claims)
if err != nil {
return nil, fmt.Errorf("failed to parse user context: %w", err)
}
@@ -187,12 +339,47 @@ func (a *DatabaseAuthenticator) OAuth2HandleCallback(ctx context.Context, provid
userCtx.SessionID = sessionToken
return &LoginResponse{
resp := &LoginResponse{
Token: sessionToken,
RefreshToken: token.RefreshToken,
User: userCtx,
ExpiresIn: int64(time.Until(expiresAt).Seconds()),
}, nil
}
if rawIDToken != "" {
// Keep the id_token: it is the id_token_hint of OAuth2LogoutURL.
resp.Meta = map[string]any{"id_token": rawIDToken}
}
return resp, nil
}
// fetchUserInfo calls the provider's userinfo endpoint with the access token.
func (p *OAuth2Provider) fetchUserInfo(ctx context.Context, token *oauth2.Token) (map[string]any, error) {
resp, err := p.config.Client(ctx, token).Get(p.userInfoURL)
if err != nil {
return nil, fmt.Errorf("failed to fetch user info: %w", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
if err != nil {
return nil, fmt.Errorf("failed to read user info: %w", err)
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("user info request failed: status %d", resp.StatusCode)
}
var userInfo map[string]any
if err := json.Unmarshal(body, &userInfo); err != nil {
return nil, fmt.Errorf("failed to parse user info: %w", err)
}
return userInfo, nil
}
// withHTTPClient makes oauth2 use the provider's HTTP client.
func (p *OAuth2Provider) withHTTPClient(ctx context.Context) context.Context {
if p.httpClient == nil {
return ctx
}
return context.WithValue(ctx, oauth2.HTTPClient, p.httpClient)
}
// OAuth2GetProviders returns list of configured OAuth2 provider names
@@ -251,23 +438,20 @@ func (a *DatabaseAuthenticator) oauth2CreateSession(ctx context.Context, session
})
}
// validateState validates state using in-memory storage
func (p *OAuth2Provider) validateState(state string) bool {
// validateState validates state using in-memory storage and returns what was remembered with it.
func (p *OAuth2Provider) validateState(state string) (*oauth2State, bool) {
p.statesMutex.Lock()
defer p.statesMutex.Unlock()
expiry, ok := p.states[state]
st, ok := p.states[state]
if !ok {
return false
return nil, false
}
if time.Now().After(expiry) {
delete(p.states, state)
return false
}
delete(p.states, state) // One-time use
return true
if time.Now().After(st.expiry) {
return nil, false
}
return st, true
}
// cleanupStates removes expired states periodically
@@ -284,8 +468,8 @@ func (p *OAuth2Provider) cleanupStates() {
}
p.statesMutex.Lock()
now := time.Now()
for state, expiry := range p.states {
if now.After(expiry) {
for state, st := range p.states {
if now.After(st.expiry) {
delete(p.states, state)
}
}
@@ -363,7 +547,7 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
}
// Use OAuth2 provider to refresh the token
tokenSource := provider.config.TokenSource(ctx, oldToken)
tokenSource := provider.config.TokenSource(provider.withHTTPClient(ctx), oldToken)
newToken, err := tokenSource.Token()
if err != nil {
return nil, fmt.Errorf("failed to refresh token with provider: %w", err)
@@ -388,12 +572,21 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
userCtx.SessionID = newSessionToken
return &LoginResponse{
resp := &LoginResponse{
Token: newSessionToken,
RefreshToken: newToken.RefreshToken,
User: userCtx,
ExpiresIn: int64(time.Until(newToken.Expiry).Seconds()),
}, nil
}
if provider.oidc != nil {
if raw, _ := newToken.Extra("id_token").(string); raw != "" {
if _, err := provider.oidc.validateIDToken(provider.withHTTPClient(ctx), raw, "", newToken.AccessToken); err != nil {
return nil, fmt.Errorf("invalid id_token in refresh response: %w", err)
}
resp.Meta = map[string]any{"id_token": raw}
}
}
return resp, nil
}
// Pre-configured OAuth2 factory methods
@@ -406,10 +599,14 @@ func NewGoogleAuthenticator(clientID, clientSecret, redirectURL string, db *sql.
ClientSecret: clientSecret,
RedirectURL: redirectURL,
Scopes: []string{"openid", "profile", "email"},
AuthURL: "https://accounts.google.com/o/oauth2/auth",
AuthURL: "https://accounts.google.com/o/oauth2/v2/auth",
TokenURL: "https://oauth2.googleapis.com/token",
UserInfoURL: "https://www.googleapis.com/oauth2/v2/userinfo",
UserInfoURL: "https://openidconnect.googleapis.com/v1/userinfo",
ProviderName: "google",
// OpenID Connect: PKCE, nonce and id_token validation against Google's published keys.
Issuer: "https://accounts.google.com",
JWKSURL: "https://www.googleapis.com/oauth2/v3/certs",
EndSessionURL: "",
})
}
+619
View File
@@ -0,0 +1,619 @@
package security
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/url"
"regexp"
"strconv"
"strings"
"time"
)
// authzRequest is a validated authorization request. It is carried through the login and consent
// forms as a sealed (HMAC-protected) value, so the server stays stateless between the steps.
type authzRequest struct {
ClientID string `json:"cid"`
RedirectURI string `json:"ru"`
State string `json:"st,omitempty"`
Nonce string `json:"n,omitempty"`
CodeChallenge string `json:"cc"`
ResponseMode string `json:"rm,omitempty"`
Scopes []string `json:"sc,omitempty"`
Prompt []string `json:"pr,omitempty"`
MaxAge int `json:"ma"` // -1 = not requested
IDTokenHint string `json:"ith,omitempty"`
LoginHint string `json:"lh,omitempty"`
ACRValues []string `json:"acr,omitempty"`
Claims map[string]any `json:"cl,omitempty"`
Resource []string `json:"res,omitempty"`
Provider string `json:"pv,omitempty"`
DPoPJKT string `json:"dj,omitempty"`
ViaPAR bool `json:"par,omitempty"`
LoginDone bool `json:"ld,omitempty"`
ConsentDone bool `json:"cd,omitempty"`
Bind string `json:"b,omitempty"` // sid the consent form is bound to
Tx string `json:"tx,omitempty"` // browser binding (login CSRF)
Sess *ssoSession `json:"ss,omitempty"` // only when the SSO cookie is disabled
}
func (r *authzRequest) hasPrompt(p string) bool { return oauthSliceContains(r.Prompt, p) }
var pkceChallengeRE = regexp.MustCompile(`^[A-Za-z0-9_-]{43}$`)
const txCookie = "resolvespec_oauth_tx"
// --------------------------------------------------------------------------
// Authorization endpoint — GET + POST /oauth/authorize
// --------------------------------------------------------------------------
func (s *OAuthServer) authorizeHandler(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodGet:
s.authorizeGet(w, r)
case http.MethodPost:
s.authorizePost(w, r)
default:
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
}
}
// authzDirectError answers an error that must not be redirected (the client or redirect_uri is
// not trusted): JSON for API callers, a page for browsers.
func (s *OAuthServer) authzDirectError(w http.ResponseWriter, r *http.Request, code, desc string) {
if strings.Contains(r.Header.Get("Accept"), "text/html") {
s.renderMessage(w, http.StatusBadRequest, "Authorization error", desc, true)
return
}
writeOAuthError(w, code, desc, http.StatusBadRequest)
}
func (s *OAuthServer) authorizeGet(w http.ResponseWriter, r *http.Request) {
q := r.URL.Query()
viaPAR := false
if uri := q.Get("request_uri"); uri != "" {
pushed, e := s.resolvePushedRequest(r.Context(), q.Get("client_id"), uri)
if e != nil {
s.authzDirectError(w, r, e.Code, e.Desc)
return
}
q, viaPAR = pushed, true
}
req, fail := s.parseAuthz(r.Context(), q, viaPAR)
if fail != nil {
fail.respond(s, w, r)
return
}
s.continueAuthorize(w, r, req)
}
// authorizePost handles the login form, the consent form and, for compatibility, the legacy
// login form that repeated the request parameters as hidden fields.
func (s *OAuthServer) authorizePost(w http.ResponseWriter, r *http.Request) {
if err := r.ParseForm(); err != nil {
http.Error(w, "invalid form", http.StatusBadRequest)
return
}
var req *authzRequest
if blob := r.PostFormValue("req"); blob != "" {
req = &authzRequest{}
if err := s.open("authz", blob, req); err != nil {
s.renderMessage(w, http.StatusBadRequest, "Request expired", "This sign-in request has expired. Please go back to the application and start again.", true)
return
}
if !s.txMatches(r, req) {
s.renderMessage(w, http.StatusBadRequest, "Request rejected", "The browser session of this request does not match. Please start again.", true)
return
}
} else {
form := url.Values{}
for k, v := range r.PostForm {
form[k] = v
}
if form.Get("state") == "" {
form.Set("state", form.Get("client_state"))
}
if form.Get("response_type") == "" {
form.Set("response_type", "code")
}
var fail *authzFailure
if req, fail = s.parseAuthz(r.Context(), form, false); fail != nil {
fail.respond(s, w, r)
return
}
}
if r.PostFormValue("decision") != "" {
s.consentSubmit(w, r, req)
return
}
s.loginSubmit(w, r, req)
}
// --------------------------------------------------------------------------
// Parsing and validation
// --------------------------------------------------------------------------
// authzFailure is an invalid authorization request. When req is set the client and redirect_uri
// are trusted and the error goes back to the client; otherwise it is shown directly.
type authzFailure struct {
req *authzRequest
code string
desc string
}
func (f *authzFailure) respond(s *OAuthServer, w http.ResponseWriter, r *http.Request) {
if f.req == nil {
s.authzDirectError(w, r, f.code, f.desc)
return
}
s.authzRedirectError(w, r, f.req, f.code, f.desc)
}
// redirectURIMatches compares exactly, except that a registered http loopback URI matches any port
// (RFC 8252 §7.3).
func redirectURIMatches(registered []string, uri string) bool {
if oauthSliceContains(registered, uri) {
return true
}
u, err := url.Parse(uri)
if err != nil || u.Scheme != "http" || !isLoopbackHost(u.Hostname()) {
return false
}
for _, reg := range registered {
ru, err := url.Parse(reg)
if err != nil || ru.Scheme != "http" || !isLoopbackHost(ru.Hostname()) {
continue
}
if ru.Hostname() == u.Hostname() && ru.Path == u.Path && ru.RawQuery == u.RawQuery {
return true
}
}
return false
}
func isLoopbackHost(h string) bool { return h == "localhost" || h == "127.0.0.1" || h == "::1" }
func (s *OAuthServer) parseAuthz(ctx context.Context, q url.Values, viaPAR bool) (*authzRequest, *authzFailure) {
client, ok := s.lookupOrFetchClient(ctx, q.Get("client_id"))
if !ok {
return nil, &authzFailure{code: "invalid_client", desc: "unknown client_id"}
}
redirectURI := q.Get("redirect_uri")
if redirectURI == "" && len(client.RedirectURIs) == 1 {
redirectURI = client.RedirectURIs[0]
}
if !redirectURIMatches(client.RedirectURIs, redirectURI) {
return nil, &authzFailure{code: "invalid_request", desc: "redirect_uri not registered"}
}
req := &authzRequest{
ClientID: client.ClientID, RedirectURI: redirectURI, State: q.Get("state"), Nonce: q.Get("nonce"),
MaxAge: -1, IDTokenHint: q.Get("id_token_hint"), LoginHint: q.Get("login_hint"),
Provider: q.Get("provider"), DPoPJKT: q.Get("dpop_jkt"), ViaPAR: viaPAR,
ResponseMode: q.Get("response_mode"), CodeChallenge: q.Get("code_challenge"),
}
bad := func(code, desc string) (*authzRequest, *authzFailure) {
return nil, &authzFailure{req: req, code: code, desc: desc}
}
if req.ResponseMode != "" && req.ResponseMode != "query" && req.ResponseMode != "form_post" {
req.ResponseMode = ""
return bad("invalid_request", "unsupported response_mode")
}
if q.Get("response_type") != "code" {
return bad("unsupported_response_type", "only 'code' is supported")
}
if q.Get("request") != "" {
return bad("request_not_supported", "request objects are not supported")
}
if q.Get("request_uri") != "" && !viaPAR {
return bad("request_uri_not_supported", "use the pushed authorization request endpoint")
}
if (s.cfg.RequirePAR || client.RequirePAR) && !viaPAR {
return bad("invalid_request", "this client must use pushed authorization requests")
}
if req.CodeChallenge == "" {
return bad("invalid_request", "code_challenge required (PKCE S256)")
}
if m := q.Get("code_challenge_method"); m != "" && m != "S256" {
return bad("invalid_request", "only S256 code_challenge_method is supported")
}
if !pkceChallengeRE.MatchString(req.CodeChallenge) {
return bad("invalid_request", "code_challenge must be a base64url SHA-256 value")
}
requested := strings.Fields(q.Get("scope"))
req.Scopes = requested
if len(client.AllowedScopes) > 0 && len(requested) > 0 {
req.Scopes = nil
for _, sc := range requested {
if oauthSliceContains(client.AllowedScopes, sc) {
req.Scopes = append(req.Scopes, sc)
}
}
if len(req.Scopes) == 0 {
return bad("invalid_scope", "none of the requested scopes is allowed for this client")
}
}
req.Prompt = strings.Fields(q.Get("prompt"))
for _, p := range req.Prompt {
if !oauthSliceContains([]string{"none", "login", "consent", "select_account"}, p) {
return bad("invalid_request", "unsupported prompt value")
}
}
if req.hasPrompt("none") && len(req.Prompt) > 1 {
return bad("invalid_request", "prompt=none cannot be combined with other values")
}
if v := q.Get("max_age"); v != "" {
n, err := strconv.Atoi(v)
if err != nil || n < 0 {
return bad("invalid_request", "max_age must be a non-negative integer")
}
req.MaxAge = n
}
if v := q.Get("claims"); v != "" {
if err := json.Unmarshal([]byte(v), &req.Claims); err != nil {
return bad("invalid_request", "claims must be a JSON object")
}
}
req.ACRValues = strings.Fields(q.Get("acr_values"))
for _, res := range q["resource"] {
u, err := url.Parse(res)
if err != nil || !u.IsAbs() || u.Fragment != "" {
return bad("invalid_target", "resource must be an absolute URI without fragment")
}
req.Resource = append(req.Resource, res)
}
return req, nil
}
// --------------------------------------------------------------------------
// Flow
// --------------------------------------------------------------------------
func userIDOfLogin(ctx context.Context, a *DatabaseAuthenticator, resp *LoginResponse) int {
if resp.User != nil && resp.User.UserID != 0 {
return resp.User.UserID
}
if info, err := a.OAuthIntrospectToken(ctx, resp.Token); err == nil && info.Active {
id, _ := strconv.Atoi(info.Sub)
return id
}
return 0
}
// continueAuthorize runs the authorization pipeline for the browser's current SSO session.
func (s *OAuthServer) continueAuthorize(w http.ResponseWriter, r *http.Request, req *authzRequest) {
s.continueWith(w, r, req, s.ssoFromRequest(r))
}
func (s *OAuthServer) continueWith(w http.ResponseWriter, r *http.Request, req *authzRequest, sso *ssoSession) {
client, ok := s.lookupOrFetchClient(r.Context(), req.ClientID)
if !ok {
s.authzDirectError(w, r, "invalid_client", "unknown client_id")
return
}
needLogin := sso == nil
if sso != nil && !req.LoginDone {
switch {
case req.hasPrompt("login"):
needLogin = true
case req.MaxAge >= 0 && time.Now().Unix()-sso.AuthTime > int64(req.MaxAge):
needLogin = true
case req.IDTokenHint != "":
if sub, _ := s.hintSubject(req.IDTokenHint); sub != "" && sub != strconv.Itoa(sso.UserID) {
needLogin = true
}
}
}
if needLogin {
if req.hasPrompt("none") {
s.authzRedirectError(w, r, req, "login_required", "no authenticated session")
return
}
switch {
case s.hasProviders():
s.redirectToExternalProvider(w, r, req)
case s.auth != nil:
s.renderLogin(w, r, req, client, "")
default:
http.Error(w, "no authentication provider configured", http.StatusInternalServerError)
}
return
}
need, err := s.consentRequired(r.Context(), req, client, sso.UserID)
if err != nil {
s.authzRedirectError(w, r, req, "server_error", "could not evaluate consent")
return
}
if need {
if req.hasPrompt("none") {
s.authzRedirectError(w, r, req, "consent_required", "user consent is required")
return
}
s.renderConsent(w, r, req, client, sso)
return
}
s.issueCode(w, r, req, sso)
}
// txBinding sets the browser binding cookie and returns its value.
func (s *OAuthServer) txBinding(w http.ResponseWriter, r *http.Request) string {
if c, err := r.Cookie(txCookie); err == nil && len(c.Value) >= 16 {
return c.Value
}
v, err := randomOAuthToken()
if err != nil {
return ""
}
http.SetCookie(w, &http.Cookie{ //nolint:gosec // Secure follows the issuer scheme (cookieSecure)
Name: txCookie, Value: v, Path: "/", MaxAge: 900, HttpOnly: true,
Secure: s.cookieSecure(), SameSite: http.SameSiteLaxMode})
return v
}
func (s *OAuthServer) txMatches(r *http.Request, req *authzRequest) bool {
if req.Tx == "" {
return true
}
c, err := r.Cookie(txCookie)
return err == nil && c.Value == req.Tx
}
func (s *OAuthServer) sealRequest(w http.ResponseWriter, r *http.Request, req *authzRequest) (string, error) {
req.Tx = s.txBinding(w, r)
return s.seal("authz", req, 15*time.Minute)
}
func (s *OAuthServer) renderLogin(w http.ResponseWriter, r *http.Request, req *authzRequest, client *OAuthServerClient, errMsg string) {
state, err := s.sealRequest(w, r, req)
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
s.renderHTML(w, http.StatusOK, s.tmpl.login, "login", OAuthLoginPage{
Title: s.cfg.LoginTitle, Error: errMsg, Action: "authorize", State: state,
ClientName: client.ClientName, LoginHint: req.LoginHint,
})
}
// loginSubmit verifies the posted credentials and continues the pipeline.
func (s *OAuthServer) loginSubmit(w http.ResponseWriter, r *http.Request, req *authzRequest) {
client, ok := s.lookupOrFetchClient(r.Context(), req.ClientID)
if !ok {
s.authzDirectError(w, r, "invalid_client", "unknown client_id")
return
}
if s.auth == nil {
http.Error(w, "no authentication provider configured", http.StatusInternalServerError)
return
}
loginResp, err := s.auth.Login(r.Context(), LoginRequest{
Username: r.PostFormValue("username"),
Password: r.PostFormValue("password"),
})
if err != nil || loginResp == nil || loginResp.Token == "" {
msg := "Invalid username or password"
if loginResp != nil && loginResp.Requires2FA {
msg = "Two-factor authentication is not supported on this sign-in form"
}
s.renderLogin(w, r, req, client, msg)
return
}
userID := userIDOfLogin(r.Context(), s.auth, loginResp)
if userID == 0 {
s.renderLogin(w, r, req, client, "Invalid username or password")
return
}
sso, err := s.newSSO(loginResp.Token, userID, "")
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
s.setSSO(w, sso)
req.LoginDone = true
if s.cfg.SSOCookie.Disable {
req.Sess = sso
}
s.continueWith(w, r, req, sso)
}
// redirectToExternalProvider stores the request and sends the browser to the configured provider.
func (s *OAuthServer) redirectToExternalProvider(w http.ResponseWriter, r *http.Request, req *authzRequest) {
var provider *externalProvider
if req.Provider != "" {
if provider = s.providerByName(req.Provider); provider == nil {
http.Error(w, fmt.Sprintf("provider %q not found", req.Provider), http.StatusBadRequest)
return
}
} else {
s.mu.RLock()
provider = &s.providers[0]
s.mu.RUnlock()
}
providerState, err := randomOAuthToken()
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
s.mu.Lock()
s.pending[providerState] = &pendingAuth{Req: req, Provider: provider.providerName, ExpiresAt: time.Now().Add(10 * time.Minute)}
s.mu.Unlock()
authURL, err := provider.auth.OAuth2GetAuthURL(provider.providerName, providerState)
if err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
http.Redirect(w, r, authURL, http.StatusFound)
}
// --------------------------------------------------------------------------
// External provider callback — GET {ProviderCallbackPath}
// --------------------------------------------------------------------------
func (s *OAuthServer) providerCallbackHandler(w http.ResponseWriter, r *http.Request) {
code := r.URL.Query().Get("code")
providerState := r.URL.Query().Get("state")
s.mu.Lock()
pending, ok := s.pending[providerState]
if ok {
delete(s.pending, providerState)
}
s.mu.Unlock()
if !ok || time.Now().After(pending.ExpiresAt) {
http.Error(w, "invalid or expired state", http.StatusBadRequest)
return
}
if e := r.URL.Query().Get("error"); e != "" {
s.authzRedirectError(w, r, pending.Req, "access_denied", "the identity provider refused the request")
return
}
if code == "" {
http.Error(w, "missing code", http.StatusBadRequest)
return
}
provider := s.providerByName(pending.Provider)
if provider == nil {
http.Error(w, fmt.Sprintf("provider %q not found", pending.Provider), http.StatusInternalServerError)
return
}
loginResp, err := provider.auth.OAuth2HandleCallback(r.Context(), pending.Provider, code, providerState)
if err != nil {
http.Error(w, err.Error(), http.StatusUnauthorized)
return
}
userID := userIDOfLogin(r.Context(), provider.auth, loginResp)
if userID == 0 {
http.Error(w, "could not resolve the user", http.StatusInternalServerError)
return
}
sso, err := s.newSSO(loginResp.Token, userID, pending.Provider)
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
s.setSSO(w, sso)
req := pending.Req
req.LoginDone = true
if s.cfg.SSOCookie.Disable {
req.Sess = sso
}
s.continueWith(w, r, req, sso)
}
// --------------------------------------------------------------------------
// Code issuance and responses
// --------------------------------------------------------------------------
func (s *OAuthServer) saveCode(ctx context.Context, c *OAuthCode) error {
if s.cfg.PersistCodes && s.auth != nil {
return s.auth.OAuthSaveCode(ctx, c)
}
s.mu.Lock()
s.codes[c.Code] = c
s.mu.Unlock()
return nil
}
// takeCode returns and invalidates a code (single use).
func (s *OAuthServer) takeCode(ctx context.Context, code string) (*OAuthCode, bool) {
if s.cfg.PersistCodes && s.auth != nil {
c, err := s.auth.OAuthExchangeCode(ctx, code)
return c, err == nil
}
s.mu.Lock()
c, ok := s.codes[code]
if ok {
delete(s.codes, code)
}
s.mu.Unlock()
if !ok || time.Now().After(c.ExpiresAt) {
return nil, false
}
return c, true
}
func (s *OAuthServer) issueCode(w http.ResponseWriter, r *http.Request, req *authzRequest, sso *ssoSession) {
authCode, err := randomOAuthToken()
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
acr := ""
if len(s.cfg.SupportedACR) > 0 {
acr = s.cfg.SupportedACR[0]
}
amr := []string{"pwd"}
if sso.Provider != "" {
amr = []string{"fed"}
}
c := &OAuthCode{
Code: authCode, ClientID: req.ClientID, RedirectURI: req.RedirectURI, ClientState: req.State,
CodeChallenge: req.CodeChallenge, CodeChallengeMethod: "S256",
SessionToken: sso.SID, Scopes: req.Scopes, ExpiresAt: time.Now().Add(s.cfg.AuthCodeTTL),
UserID: sso.UserID, Nonce: req.Nonce, AuthTime: sso.AuthTime, ACR: acr, AMR: amr, SessionID: sso.SID,
Claims: req.Claims, Resource: req.Resource, DPoPJKT: req.DPoPJKT, ResponseType: "code",
}
if err := s.saveCode(r.Context(), c); err != nil {
s.authzRedirectError(w, r, req, "server_error", "could not issue the authorization code")
return
}
s.noteClient(w, sso, req.ClientID)
s.authzRespond(w, r, req, map[string]string{"code": authCode})
}
func (s *OAuthServer) authzRedirectError(w http.ResponseWriter, r *http.Request, req *authzRequest, code, desc string) {
p := map[string]string{"error": code}
if desc != "" {
p["error_description"] = desc
}
s.authzRespond(w, r, req, p)
}
// authzRespond sends params to the client's redirect_uri (query or form_post) with state and the
// RFC 9207 iss parameter.
func (s *OAuthServer) authzRespond(w http.ResponseWriter, r *http.Request, req *authzRequest, params map[string]string) {
if req.State != "" {
params["state"] = req.State
}
params["iss"] = s.cfg.Issuer
if req.ResponseMode == "form_post" {
h := w.Header()
h.Set("Content-Type", "text/html; charset=utf-8")
h.Set("Cache-Control", "no-store")
h.Set("Content-Security-Policy", "script-src 'unsafe-inline'; frame-ancestors 'none'")
var b strings.Builder
b.WriteString(`<!DOCTYPE html><html><body onload="document.forms[0].submit()"><form method="post" action="`)
b.WriteString(htmlEscape(req.RedirectURI))
b.WriteString(`">`)
for k, v := range params {
b.WriteString(`<input type="hidden" name="` + htmlEscape(k) + `" value="` + htmlEscape(v) + `">`)
}
b.WriteString(`<noscript><button type="submit">Continue</button></noscript></form></body></html>`)
w.Write([]byte(b.String())) //nolint:errcheck,gosec // G104: best-effort write, G705: values are HTML-escaped
return
}
u, err := url.Parse(req.RedirectURI)
if err != nil {
http.Error(w, "invalid redirect_uri", http.StatusInternalServerError)
return
}
qp := u.Query()
for k, v := range params {
qp.Set(k, v)
}
u.RawQuery = qp.Encode()
w.Header().Set("Cache-Control", "no-store")
http.Redirect(w, r, u.String(), http.StatusFound) //nolint:gosec // u is built from the registered redirect URI
}
+189
View File
@@ -0,0 +1,189 @@
package security
import (
"context"
"crypto/subtle"
"net/http"
"net/url"
"time"
"github.com/golang-jwt/jwt/v5"
)
const clientAssertionTypeJWT = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer"
// authedClient is a client that has identified itself at an endpoint.
type authedClient struct {
Client *OAuthServerClient
Method string // none, client_secret_basic, client_secret_post, private_key_jwt
}
func invalidClient(desc string, basic bool) *oauthError {
e := oerr("invalid_client", desc, http.StatusUnauthorized)
if basic {
e.WWWAuth = `Basic realm="oauth"`
}
return e
}
// authenticateClient identifies the client of a request from client_secret_basic,
// client_secret_post, private_key_jwt or (for public clients) client_id alone. It returns
// (nil, nil) when the request carries nothing that identifies a client.
func (s *OAuthServer) authenticateClient(r *http.Request) (*authedClient, *oauthError) {
if assertion := r.FormValue("client_assertion"); assertion != "" {
return s.authenticateAssertion(r, assertion)
}
id, secret, basic := r.BasicAuth()
method := "client_secret_basic"
if basic {
// RFC 6749 §2.3.1: the id and secret are form-urlencoded before Base64 encoding.
if v, err := url.QueryUnescape(id); err == nil {
id = v
}
if v, err := url.QueryUnescape(secret); err == nil {
secret = v
}
} else {
id, secret = r.FormValue("client_id"), r.FormValue("client_secret")
method = "client_secret_post"
}
if id == "" {
return nil, nil
}
client, ok := s.lookupOrFetchClient(r.Context(), id)
if !ok {
return nil, invalidClient("invalid client credentials", basic)
}
if secret == "" {
if basic || needsClientAuth(client) {
return nil, invalidClient("client authentication required", basic)
}
return &authedClient{Client: client, Method: "none"}, nil
}
if client.ClientSecretHash == "" ||
subtle.ConstantTimeCompare([]byte(hashClientSecret(secret)), []byte(client.ClientSecretHash)) != 1 {
return nil, invalidClient("invalid client credentials", basic)
}
if client.ClientSecretExpiresAt != 0 && time.Now().Unix() > client.ClientSecretExpiresAt {
return nil, invalidClient("client secret expired", basic)
}
return &authedClient{Client: client, Method: method}, nil
}
// requireClient is authenticateClient for endpoints that need an identified client.
func (s *OAuthServer) requireClient(r *http.Request) (*authedClient, *oauthError) {
ac, e := s.authenticateClient(r)
if e != nil {
return nil, e
}
if ac == nil {
return nil, invalidClient("client authentication required", false)
}
return ac, nil
}
func (s *OAuthServer) authenticateAssertion(r *http.Request, assertion string) (*authedClient, *oauthError) {
if r.FormValue("client_assertion_type") != clientAssertionTypeJWT {
return nil, invalidClient("unsupported client_assertion_type", false)
}
unverified, _, err := jwt.NewParser().ParseUnverified(assertion, &jwt.RegisteredClaims{})
if err != nil {
return nil, invalidClient("malformed client_assertion", false)
}
rc, _ := unverified.Claims.(*jwt.RegisteredClaims)
clientID := rc.Subject
if formID := r.FormValue("client_id"); formID != "" && formID != clientID {
return nil, invalidClient("client_id does not match the assertion", false)
}
client, ok := s.lookupOrFetchClient(r.Context(), clientID)
if !ok || client.TokenEndpointAuthMethod != "private_key_jwt" {
return nil, invalidClient("invalid client credentials", false)
}
set, err := s.clientKeySet(r.Context(), client, false)
if err != nil {
return nil, invalidClient("client keys are unavailable", false)
}
alg := []string{"RS256", "PS256", "ES256", "ES384"}
if client.TokenEndpointAuthSigningAlg != "" {
alg = []string{client.TokenEndpointAuthSigningAlg}
}
claims := &jwt.RegisteredClaims{}
verify := func(set *jwkSet) error {
_, err := verifyJWTWithSet(assertion, set, alg, claims,
jwt.WithIssuer(clientID), jwt.WithSubject(clientID), jwt.WithExpirationRequired(), jwt.WithLeeway(30*time.Second))
return err
}
if err := verify(set); err != nil && client.JWKSURI != "" {
// The client may have rotated its keys: refetch once and try again.
if set, ferr := s.clientKeySet(r.Context(), client, true); ferr == nil {
err = verify(set)
}
if err != nil {
return nil, invalidClient("client_assertion rejected", false)
}
} else if err != nil {
return nil, invalidClient("client_assertion rejected", false)
}
if !s.assertionAudienceOK(claims.Audience, r) {
return nil, invalidClient("client_assertion audience mismatch", false)
}
exp := claims.ExpiresAt.Time
if exp.After(time.Now().Add(10 * time.Minute)) {
return nil, invalidClient("client_assertion lifetime too long", false)
}
if claims.ID == "" {
return nil, invalidClient("client_assertion needs a jti", false)
}
if seen, err := s.replayed(r.Context(), "cla:"+clientID+":"+claims.ID, exp.Add(time.Minute)); err != nil {
return nil, serverErr()
} else if seen {
return nil, invalidClient("client_assertion replayed", false)
}
return &authedClient{Client: client, Method: "private_key_jwt"}, nil
}
func (s *OAuthServer) assertionAudienceOK(aud jwt.ClaimStrings, r *http.Request) bool {
for _, a := range aud {
if a == s.cfg.Issuer || a == s.endpoint("/oauth/token") || a == s.requestURL(r) {
return true
}
}
return false
}
// clientKeySet returns the verification keys of a private_key_jwt client.
func (s *OAuthServer) clientKeySet(ctx context.Context, c *OAuthServerClient, refresh bool) (*jwkSet, error) {
if len(c.JWKS) > 0 {
return parseJWKS(c.JWKS)
}
return s.jwks.get(ctx, c.JWKSURI, refresh)
}
// replayed records key until expires and reports whether it was seen before.
func (s *OAuthServer) replayed(ctx context.Context, key string, expires time.Time) (bool, error) {
g := s.grants()
if g == nil {
return false, nil
}
if len(key) > 250 { // keys are bounded by the jti_key column
key = hashToken(key)
}
return g.SeenJTI(ctx, key, expires)
}
// requestURL is the public URL of the request, built from the issuer so it survives reverse proxies.
func (s *OAuthServer) requestURL(r *http.Request) string {
path := r.URL.Path
prefix := ""
if p := s.issuerURL.Path; p != "" && p != "/" && !hasPathPrefix(path, p) {
prefix = p
}
return s.issuerURL.Scheme + "://" + s.issuerURL.Host + prefix + path
}
func hasPathPrefix(path, prefix string) bool {
return len(path) >= len(prefix) && path[:len(prefix)] == prefix
}
+141
View File
@@ -0,0 +1,141 @@
package security
import (
"context"
"errors"
"net/http"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
var builtinScopeDescriptions = map[string]string{
"openid": "Verify your identity",
"profile": "Read your username and profile",
"email": "Read your email address",
"offline_access": "Stay signed in and refresh access without asking again",
}
func (s *OAuthServer) scopeInfos(scopes []string) []OAuthScopeInfo {
out := make([]OAuthScopeInfo, 0, len(scopes))
for _, sc := range scopes {
d := s.cfg.ScopeDescriptions[sc]
if d == "" {
d = builtinScopeDescriptions[sc]
}
out = append(out, OAuthScopeInfo{Name: sc, Description: d})
}
return out
}
func scopesCovered(have, want []string) bool {
for _, w := range want {
if !oauthSliceContains(have, w) {
return false
}
}
return true
}
// consentRequired reports whether the user has to approve the request on the consent screen.
func (s *OAuthServer) consentRequired(ctx context.Context, req *authzRequest, client *OAuthServerClient, userID int) (bool, error) {
if client.FirstParty || req.ConsentDone {
return false, nil
}
if !s.cfg.RequireConsent && !client.RequireConsent {
return false, nil
}
if req.hasPrompt("consent") {
return true, nil
}
g := s.grants()
if g == nil {
return true, nil
}
c, err := g.GetConsent(ctx, userID, client.ClientID)
if errors.Is(err, lookup.ErrNotFound) {
return true, nil
}
if err != nil {
return false, err
}
return !scopesCovered(c.Scopes, req.Scopes), nil
}
func (s *OAuthServer) renderConsent(w http.ResponseWriter, r *http.Request, req *authzRequest, client *OAuthServerClient, sso *ssoSession) {
req.Bind = sso.SID
req.LoginDone = true
if s.cfg.SSOCookie.Disable {
req.Sess = sso
}
state, err := s.sealRequest(w, r, req)
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
name := client.ClientName
if name == "" {
name = client.ClientID
}
user := ""
if a := s.anyAuth(); a != nil {
if info, err := a.OAuthIntrospectToken(r.Context(), sso.Token); err == nil && info.Active {
user = info.Username
if user == "" {
user = info.Email
}
}
}
s.renderHTML(w, http.StatusOK, s.tmpl.consent, "consent", OAuthConsentPage{
Title: "Authorize " + name, Action: "authorize", State: state, ClientName: name,
ClientURI: safeWebURL(client.ClientURI), LogoURI: safeWebURL(client.LogoURI),
Scopes: s.scopeInfos(req.Scopes), User: user,
})
}
// safeWebURL returns u when it is an http(s) URL, so client-supplied links cannot be javascript: URLs.
func safeWebURL(u string) string {
if strings.HasPrefix(u, "https://") || strings.HasPrefix(u, "http://") {
return u
}
return ""
}
// consentSubmit handles the consent form.
func (s *OAuthServer) consentSubmit(w http.ResponseWriter, r *http.Request, req *authzRequest) {
sso := s.ssoFromRequest(r)
if sso == nil && req.Sess != nil {
if a := s.anyAuth(); a != nil {
if info, err := a.OAuthIntrospectToken(r.Context(), req.Sess.Token); err == nil && info.Active {
sso = req.Sess
}
}
}
if sso == nil || req.Bind == "" || req.Bind != sso.SID {
s.renderMessage(w, http.StatusBadRequest, "Session expired", "Your session has expired. Please start again from the application.", true)
return
}
if r.PostFormValue("decision") != "allow" {
s.authzRedirectError(w, r, req, "access_denied", "the user denied the request")
return
}
if g := s.grants(); g != nil && r.PostFormValue("remember") == "1" {
scopes := req.Scopes
if old, err := g.GetConsent(r.Context(), sso.UserID, req.ClientID); err == nil {
for _, sc := range old.Scopes {
if !oauthSliceContains(scopes, sc) {
scopes = append(scopes, sc)
}
}
}
if err := g.SaveConsent(r.Context(), lookup.Consent{
UserID: sso.UserID, ClientID: req.ClientID, Scopes: scopes, ExpiresAt: time.Now().Add(s.cfg.ConsentTTL),
}); err != nil {
s.authzRedirectError(w, r, req, "server_error", "could not store the consent")
return
}
}
req.ConsentDone = true
s.issueCode(w, r, req, sso)
}
+331
View File
@@ -0,0 +1,331 @@
package security
import (
"crypto/rand"
"errors"
"net/http"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// userCodeAlphabet avoids vowels and look-alike characters (RFC 8628 §6.1).
const userCodeAlphabet = "BCDFGHJKLMNPQRSTVWXZ"
func newUserCode() (string, error) {
b := make([]byte, 8)
if _, err := rand.Read(b); err != nil {
return "", err
}
out := make([]byte, 8)
for i, v := range b {
out[i] = userCodeAlphabet[int(v)%len(userCodeAlphabet)]
}
return string(out), nil
}
func normalizeUserCode(c string) string {
c = strings.ToUpper(c)
c = strings.ReplaceAll(c, "-", "")
return strings.ReplaceAll(c, " ", "")
}
func formatUserCode(c string) string {
if len(c) == 8 {
return c[:4] + "-" + c[4:]
}
return c
}
// filterScopes intersects the requested scopes with the client's allowed scopes.
func filterScopes(c *OAuthServerClient, requested []string) ([]string, bool) {
if len(c.AllowedScopes) == 0 || len(requested) == 0 {
return requested, true
}
var out []string
for _, sc := range requested {
if oauthSliceContains(c.AllowedScopes, sc) {
out = append(out, sc)
}
}
return out, len(out) > 0
}
// --------------------------------------------------------------------------
// RFC 8628 — device authorization: POST /oauth/device_authorization
// --------------------------------------------------------------------------
func (s *OAuthServer) deviceAuthorizationHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
if err := r.ParseForm(); err != nil {
writeOAuthError(w, "invalid_request", "cannot parse form", http.StatusBadRequest)
return
}
ac, e := s.requireClient(r)
if e != nil {
e.write(w)
return
}
client := ac.Client
if !grantAllowed(client, grantDeviceCode) {
writeOAuthError(w, "unauthorized_client", "client may not use the device grant", http.StatusBadRequest)
return
}
scopes, ok := filterScopes(client, strings.Fields(r.PostFormValue("scope")))
if !ok {
writeOAuthError(w, "invalid_scope", "none of the requested scopes is allowed for this client", http.StatusBadRequest)
return
}
gs := s.grants()
if gs == nil {
writeOAuthError(w, "server_error", "", http.StatusInternalServerError)
return
}
deviceCode, err := randomOAuthToken()
if err != nil {
writeOAuthError(w, "server_error", "", http.StatusInternalServerError)
return
}
var userCode string
for attempt := 0; ; attempt++ {
if userCode, err = newUserCode(); err == nil {
err = gs.CreateDevice(r.Context(), lookup.DeviceCode{
DeviceHash: hashToken(deviceCode), UserCode: userCode, ClientID: client.ClientID, Scopes: scopes,
Interval: s.cfg.DevicePollSeconds, ExpiresAt: time.Now().Add(s.cfg.DeviceCodeTTL),
})
}
if err == nil {
break
}
if attempt >= 3 { // a user-code collision is the only expected failure; give up after a few tries
writeOAuthError(w, "server_error", "", http.StatusInternalServerError)
return
}
}
verification := s.endpoint("/oauth/device")
writeJSON(w, http.StatusOK, map[string]any{
"device_code": deviceCode,
"user_code": formatUserCode(userCode),
"verification_uri": verification,
"verification_uri_complete": verification + "?user_code=" + formatUserCode(userCode),
"expires_in": int(s.cfg.DeviceCodeTTL.Seconds()),
"interval": s.cfg.DevicePollSeconds,
})
}
// --------------------------------------------------------------------------
// Verification page: GET/POST /oauth/device
// --------------------------------------------------------------------------
type deviceState struct {
UserCode string `json:"uc"`
Bind string `json:"b,omitempty"`
Tx string `json:"tx,omitempty"`
}
func (s *OAuthServer) deviceVerificationHandler(w http.ResponseWriter, r *http.Request) {
if s.cfg.SSOCookie.Disable {
s.renderMessage(w, http.StatusNotImplemented, "Not available", "The device flow needs the SSO cookie, which is disabled on this server.", true)
return
}
if err := r.ParseForm(); err != nil {
http.Error(w, "invalid form", http.StatusBadRequest)
return
}
if r.Method == http.MethodGet {
if code := r.FormValue("user_code"); code != "" {
s.deviceNext(w, r, normalizeUserCode(code))
return
}
s.renderHTML(w, http.StatusOK, nil, "device", oauthDevicePage{Title: "Connect a device", Action: "device"})
return
}
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
switch r.PostFormValue("step") {
case "code":
s.deviceNext(w, r, normalizeUserCode(r.PostFormValue("user_code")))
case "login":
st, ok := s.openDeviceState(w, r, "device-login")
if !ok {
return
}
if s.auth == nil {
http.Error(w, "no authentication provider configured", http.StatusInternalServerError)
return
}
resp, err := s.auth.Login(r.Context(), LoginRequest{Username: r.PostFormValue("username"), Password: r.PostFormValue("password")})
uid := 0
if err == nil && resp != nil && resp.Token != "" {
uid = userIDOfLogin(r.Context(), s.auth, resp)
}
if uid == 0 {
s.renderDeviceLogin(w, r, st.UserCode, "Invalid username or password")
return
}
sso, err := s.newSSO(resp.Token, uid, "")
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
s.setSSO(w, sso)
s.deviceConsent(w, r, st.UserCode, sso)
case "decide":
st, ok := s.openDeviceState(w, r, "device-consent")
if !ok {
return
}
sso := s.ssoFromRequest(r)
if sso == nil || st.Bind != sso.SID {
s.renderMessage(w, http.StatusBadRequest, "Session expired", "Your session has expired. Start again from the device.", true)
return
}
err := s.grants().DeviceDecide(r.Context(), st.UserCode, r.PostFormValue("decision") == "allow", sso.UserID, sso.SID)
if err != nil {
s.renderMessage(w, http.StatusBadRequest, "Code not valid", "This code is unknown, expired or already used.", true)
return
}
if r.PostFormValue("decision") == "allow" {
s.renderMessage(w, http.StatusOK, "Device connected", "You can return to your device now.", false)
} else {
s.renderMessage(w, http.StatusOK, "Request denied", "The device was not given access.", false)
}
default:
http.Error(w, "invalid request", http.StatusBadRequest)
}
}
func (s *OAuthServer) openDeviceState(w http.ResponseWriter, r *http.Request, kind string) (*deviceState, bool) {
var st deviceState
if err := s.open(kind, r.PostFormValue("req"), &st); err != nil {
s.renderMessage(w, http.StatusBadRequest, "Request expired", "This request has expired. Enter the code again.", true)
return nil, false
}
if c, err := r.Cookie(txCookie); st.Tx != "" && (err != nil || c.Value != st.Tx) {
s.renderMessage(w, http.StatusBadRequest, "Request rejected", "The browser session of this request does not match.", true)
return nil, false
}
return &st, true
}
// deviceNext continues with the entered code: login when needed, then the approval screen.
func (s *OAuthServer) deviceNext(w http.ResponseWriter, r *http.Request, userCode string) {
if _, err := s.grants().DeviceByUserCode(r.Context(), userCode); err != nil {
code := http.StatusBadRequest
if !errors.Is(err, lookup.ErrNotFound) {
code = http.StatusInternalServerError
}
s.renderHTML(w, code, nil, "device", oauthDevicePage{Title: "Connect a device", Action: "device",
Error: "This code is unknown or has expired."})
return
}
if sso := s.ssoFromRequest(r); sso != nil {
s.deviceConsent(w, r, userCode, sso)
return
}
if s.auth == nil {
s.renderMessage(w, http.StatusNotImplemented, "Sign-in unavailable",
"Device sign-in needs the server's own login form, which is not configured.", true)
return
}
s.renderDeviceLogin(w, r, userCode, "")
}
func (s *OAuthServer) renderDeviceLogin(w http.ResponseWriter, r *http.Request, userCode, errMsg string) {
state, err := s.seal("device-login", deviceState{UserCode: userCode, Tx: s.txBinding(w, r)}, 15*time.Minute)
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
s.renderHTML(w, http.StatusOK, s.tmpl.login, "login", OAuthLoginPage{
Title: s.cfg.LoginTitle, Error: errMsg, Action: "device", State: state, Hidden: map[string]string{"step": "login"},
})
}
func (s *OAuthServer) deviceConsent(w http.ResponseWriter, r *http.Request, userCode string, sso *ssoSession) {
dc, err := s.grants().DeviceByUserCode(r.Context(), userCode)
if err != nil {
s.renderMessage(w, http.StatusBadRequest, "Code not valid", "This code is unknown, expired or already used.", true)
return
}
client, ok := s.lookupOrFetchClient(r.Context(), dc.ClientID)
if !ok {
s.renderMessage(w, http.StatusBadRequest, "Unknown application", "The application that requested this code no longer exists.", true)
return
}
state, err := s.seal("device-consent", deviceState{UserCode: userCode, Bind: sso.SID, Tx: s.txBinding(w, r)}, 15*time.Minute)
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
name := client.ClientName
if name == "" {
name = client.ClientID
}
user := ""
if a := s.anyAuth(); a != nil {
if info, err := a.OAuthIntrospectToken(r.Context(), sso.Token); err == nil && info.Active {
user = info.Username
}
}
s.renderHTML(w, http.StatusOK, s.tmpl.consent, "consent", OAuthConsentPage{
Title: "Connect " + name, Action: "device", State: state, ClientName: name,
ClientURI: safeWebURL(client.ClientURI), LogoURI: safeWebURL(client.LogoURI),
Scopes: s.scopeInfos(dc.Scopes), User: user, Hidden: map[string]string{"step": "decide"},
})
}
// --------------------------------------------------------------------------
// Token endpoint side
// --------------------------------------------------------------------------
func (s *OAuthServer) handleDeviceGrant(r *http.Request) (map[string]any, *oauthError) {
ac, e := s.requireClient(r)
if e != nil {
return nil, e
}
client := ac.Client
if !grantAllowed(client, grantDeviceCode) {
return nil, oerr("unauthorized_client", "client may not use the device grant", http.StatusBadRequest)
}
code := r.PostFormValue("device_code")
if code == "" {
return nil, oerr("invalid_request", "device_code required", http.StatusBadRequest)
}
gs := s.grants()
if gs == nil {
return nil, serverErr()
}
dc, err := gs.DevicePoll(r.Context(), hashToken(code))
switch {
case errors.Is(err, lookup.ErrDevicePending):
return nil, oerr("authorization_pending", "", http.StatusBadRequest)
case errors.Is(err, lookup.ErrDeviceSlowDown):
return nil, oerr("slow_down", "", http.StatusBadRequest)
case errors.Is(err, lookup.ErrDeviceDenied):
return nil, oerr("access_denied", "", http.StatusBadRequest)
case errors.Is(err, lookup.ErrDeviceExpired):
return nil, oerr("expired_token", "", http.StatusBadRequest)
case err != nil:
return nil, serverErr()
}
if dc.ClientID != client.ClientID {
return nil, oerr("invalid_grant", "device_code was issued to another client", http.StatusBadRequest)
}
jkt, e := s.dpopForClient(r, client)
if e != nil {
return nil, e
}
return s.mintTokens(r.Context(), &tokenGrant{
Client: client, UserID: dc.UserID, Scopes: dc.Scopes, SID: dc.SessionToken, AuthTime: time.Now().Unix(),
AMR: []string{"pwd"}, DPoPJKT: jkt, IDToken: true, IssueRefresh: refreshAllowed(client, dc.Scopes),
})
}
+106
View File
@@ -0,0 +1,106 @@
package security
import (
"crypto/sha256"
"crypto/subtle"
"net/http"
"net/url"
"strings"
"time"
"github.com/golang-jwt/jwt/v5"
)
// dpopProof is a validated RFC 9449 proof.
type dpopProof struct {
JKT string // base64url SHA-256 thumbprint of the proof key
}
// verifyDPoP validates the DPoP header of r. It returns (nil, nil) when the request has none.
// accessToken, when non-empty, is the token the proof must be bound to (the "ath" claim).
func (s *OAuthServer) verifyDPoP(r *http.Request, accessToken string) (*dpopProof, *oauthError) {
vals := r.Header.Values("DPoP")
if len(vals) == 0 {
return nil, nil
}
bad := func(desc string) (*dpopProof, *oauthError) {
return nil, oerr("invalid_dpop_proof", desc, http.StatusBadRequest)
}
if !s.cfg.EnableDPoP {
return bad("DPoP is not enabled")
}
if len(vals) != 1 {
return bad("exactly one DPoP header is required")
}
var jkt string
claims := &struct {
jwt.RegisteredClaims
HTM string `json:"htm"`
HTU string `json:"htu"`
ATH string `json:"ath"`
}{}
tok, err := jwt.NewParser(jwt.WithValidMethods([]string{"ES256", "ES384", "RS256", "PS256"})).
ParseWithClaims(vals[0], claims, func(t *jwt.Token) (any, error) {
if t.Header["typ"] != "dpop+jwt" {
return nil, jwt.ErrTokenUnverifiable
}
m, ok := t.Header["jwk"].(map[string]any)
if !ok {
return nil, jwt.ErrTokenUnverifiable
}
if _, private := m["d"]; private {
return nil, jwt.ErrTokenUnverifiable
}
pub, err := publicFromJWK(m)
if err != nil {
return nil, err
}
if jkt, err = jwkThumbprint(pub); err != nil {
return nil, err
}
return pub, nil
})
if err != nil || !tok.Valid {
return bad("proof signature or header invalid")
}
if claims.HTM != r.Method {
return bad("htm does not match the request method")
}
if !sameURLNoQuery(claims.HTU, s.requestURL(r)) {
return bad("htu does not match the request URL")
}
if claims.IssuedAt == nil || claims.ID == "" {
return bad("iat and jti are required")
}
iat := claims.IssuedAt.Time
if d := time.Since(iat); d > 2*time.Minute || d < -time.Minute {
return bad("proof is not fresh")
}
if accessToken != "" {
sum := sha256.Sum256([]byte(accessToken))
if subtle.ConstantTimeCompare([]byte(b64u(sum[:])), []byte(claims.ATH)) != 1 {
return bad("ath does not match the access token")
}
}
seen, err := s.replayed(r.Context(), "dpop:"+jkt+":"+claims.ID, iat.Add(3*time.Minute))
if err != nil {
return nil, serverErr()
}
if seen {
return bad("proof replayed")
}
return &dpopProof{JKT: jkt}, nil
}
func sameURLNoQuery(a, b string) bool {
ua, err1 := url.Parse(a)
ub, err2 := url.Parse(b)
if err1 != nil || err2 != nil {
return false
}
norm := func(u *url.URL) string {
return strings.ToLower(u.Scheme) + "://" + strings.ToLower(u.Host) + u.EscapedPath()
}
return norm(ua) == norm(ub)
}
+76
View File
@@ -0,0 +1,76 @@
package security
import (
"net/http"
"net/url"
"strings"
)
const (
tokenTypeAccess = "urn:ietf:params:oauth:token-type:access_token" //nolint:gosec // RFC 8693 URN, not a credential
tokenTypeJWT = "urn:ietf:params:oauth:token-type:jwt" //nolint:gosec // RFC 8693 URN, not a credential
)
// handleTokenExchange implements RFC 8693 for access tokens issued by this server: the caller
// trades a subject token for one with a narrower scope and/or another audience. Delegation with an
// actor_token is not supported.
func (s *OAuthServer) handleTokenExchange(r *http.Request) (map[string]any, *oauthError) {
ac, e := s.requireClient(r)
if e != nil {
return nil, e
}
client := ac.Client
if ac.Method == "none" {
return nil, invalidClient("token exchange requires a confidential client", true)
}
if !grantAllowed(client, grantTokenExchange) {
return nil, oerr("unauthorized_client", "client may not use token exchange", http.StatusBadRequest)
}
if r.PostFormValue("actor_token") != "" {
return nil, oerr("invalid_request", "actor_token is not supported", http.StatusBadRequest)
}
if t := r.PostFormValue("subject_token_type"); t != "" && t != tokenTypeAccess && t != tokenTypeJWT {
return nil, oerr("invalid_request", "unsupported subject_token_type", http.StatusBadRequest)
}
if t := r.PostFormValue("requested_token_type"); t != "" && t != tokenTypeAccess {
return nil, oerr("invalid_request", "only access tokens can be issued", http.StatusBadRequest)
}
subject := r.PostFormValue("subject_token")
if subject == "" {
return nil, oerr("invalid_request", "subject_token required", http.StatusBadRequest)
}
ti := s.resolveAccessToken(r.Context(), subject)
if ti == nil {
return nil, oerr("invalid_grant", "subject_token is invalid or inactive", http.StatusBadRequest)
}
if ti.JKT != "" {
return nil, oerr("invalid_request", "DPoP-bound tokens cannot be exchanged", http.StatusBadRequest)
}
scopes := ti.Scopes
if req := strings.Fields(r.PostFormValue("scope")); len(req) > 0 {
if !scopesCovered(ti.Scopes, req) || (len(client.AllowedScopes) > 0 && !scopesCovered(client.AllowedScopes, req)) {
return nil, oerr("invalid_scope", "scope exceeds the subject token or the client", http.StatusBadRequest)
}
scopes = req
}
var resource []string
for _, v := range append(r.PostForm["audience"], r.PostForm["resource"]...) {
if v == "" {
continue
}
if u, err := url.Parse(v); r.PostForm["resource"] != nil && (err != nil || !u.IsAbs() || u.Fragment != "") {
return nil, oerr("invalid_target", "resource must be an absolute URI", http.StatusBadRequest)
}
resource = append(resource, v)
}
resp, e := s.mintTokens(r.Context(), &tokenGrant{
Client: client, UserID: ti.UserID, Scopes: scopes, Resource: resource, SID: ti.SID,
})
if e != nil {
return nil, e
}
resp["issued_token_type"] = tokenTypeAccess
return resp, nil
}
+563
View File
@@ -0,0 +1,563 @@
package security
import (
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"net/http"
"net/http/cookiejar"
"net/http/httptest"
"net/url"
"strings"
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
)
const flowRedirect = "https://rp.example.com/callback"
type flowEnv struct {
t *testing.T
ts *httptest.Server
srv *OAuthServer
auth *DatabaseAuthenticator
browser *http.Client
clientID string
}
// newFlowEnv starts a real HTTP server with the given config, one user (olivia/pw) and one
// registered public client.
func newFlowEnv(t *testing.T, cfg OAuthServerConfig) *flowEnv {
t.Helper()
var handler http.Handler
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { handler.ServeHTTP(w, r) }))
t.Cleanup(ts.Close)
auth := NewDatabaseAuthenticatorWithOptions(newDirectTestDB(t), DatabaseAuthenticatorOptions{Lookup: directConfig})
cfg.Issuer = ts.URL
cfg.PersistCodes = true
srv := NewOAuthServer(cfg, auth)
t.Cleanup(srv.Close)
handler = srv.HTTPHandler()
if _, err := auth.Register(context.Background(), RegisterRequest{Username: "olivia", Password: "pw", Email: "olivia@example.com"}); err != nil {
t.Fatal(err)
}
jar, _ := cookiejar.New(nil)
e := &flowEnv{t: t, ts: ts, srv: srv, auth: auth,
browser: &http.Client{Jar: jar, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }}}
_, reg := doJSON(t, handler, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{flowRedirect},
"grant_types": []string{"authorization_code", "refresh_token"},
})
e.clientID, _ = reg["client_id"].(string)
if e.clientID == "" {
t.Fatalf("registration failed: %v", reg)
}
return e
}
func (e *flowEnv) authURL(extra url.Values) (string, string) {
verifier := "verifier-0123456789-0123456789-0123456789-abcdef"
q := url.Values{
"response_type": {"code"}, "client_id": {e.clientID}, "redirect_uri": {flowRedirect},
"code_challenge": {s256Challenge(verifier)}, "code_challenge_method": {"S256"},
"scope": {"openid profile"}, "state": {"st"},
}
for k, v := range extra {
q[k] = v
}
return e.ts.URL + "/oauth/authorize?" + q.Encode(), verifier
}
// login drives the browser through authorize until it is redirected to the RP and returns that URL.
func (e *flowEnv) login(authURL string, decision string) *url.URL {
e.t.Helper()
resp, err := e.browser.Get(authURL)
if err != nil {
e.t.Fatal(err)
}
for i := 0; i < 4; i++ {
if loc := resp.Header.Get("Location"); loc != "" {
bodyOf(e.t, resp)
u, _ := url.Parse(loc)
if !u.IsAbs() {
u = resp.Request.URL.ResolveReference(u)
}
if u.Host == "rp.example.com" {
return u
}
if resp, err = e.browser.Get(u.String()); err != nil {
e.t.Fatal(err)
}
continue
}
page := bodyOf(e.t, resp)
switch {
case strings.Contains(page, `name="password"`):
resp = browserSubmit(e.t, e.browser, authURL, page, url.Values{"username": {"olivia"}, "password": {"pw"}})
case strings.Contains(page, `name="decision"`):
resp = browserSubmit(e.t, e.browser, authURL, page, url.Values{"decision": {decision}})
default:
e.t.Fatalf("unexpected page (status %d): %s", resp.StatusCode, page)
}
}
e.t.Fatal("too many steps")
return nil
}
func (e *flowEnv) post(path string, form url.Values) (*httptest.ResponseRecorder, map[string]interface{}) {
return doForm(e.t, e.srv.HTTPHandler(), path, form, "", "")
}
// tokens runs the full code flow and returns the token response.
func (e *flowEnv) tokens(extra url.Values) map[string]interface{} {
e.t.Helper()
authURL, verifier := e.authURL(extra)
cb := e.login(authURL, "allow")
code := cb.Query().Get("code")
if code == "" {
e.t.Fatalf("no code in %s", cb)
}
rec, tok := e.post("/oauth/token", url.Values{
"grant_type": {"authorization_code"}, "code": {code}, "redirect_uri": {flowRedirect},
"client_id": {e.clientID}, "code_verifier": {verifier},
})
if rec.Code != http.StatusOK {
e.t.Fatalf("token: %d %s", rec.Code, rec.Body.String())
}
return tok
}
func TestOAuthFlow_ConsentDenyAndScopeFilteredUserinfo(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{RequireConsent: true})
authURL, _ := e.authURL(nil)
cb := e.login(authURL, "deny")
if cb.Query().Get("error") != "access_denied" || cb.Query().Get("state") != "st" || cb.Query().Get("iss") != e.ts.URL {
t.Errorf("deny redirect = %s", cb)
}
tok := e.tokens(url.Values{"scope": {"openid"}})
rec, info := doGet(e.srv.HTTPHandler(), "/oauth/userinfo", tok["access_token"].(string))
if rec.Code != http.StatusOK {
t.Fatalf("userinfo %d %s", rec.Code, rec.Body.String())
}
if info["sub"] == nil {
t.Error("sub missing")
}
if _, ok := info["email"]; ok {
t.Errorf("email must not be released without the email scope: %v", info)
}
}
func TestOAuthFlow_PromptNoneWithoutSession(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{})
authURL, _ := e.authURL(url.Values{"prompt": {"none"}})
cb := e.login(authURL, "allow")
if cb.Query().Get("error") != "login_required" {
t.Errorf("want login_required, got %s", cb)
}
}
func TestOAuthFlow_RefreshRotationAndReuseDetection(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{ManagedRefreshTokens: true})
tok := e.tokens(nil)
r1, _ := tok["refresh_token"].(string)
if r1 == "" {
t.Fatalf("no refresh token: %v", tok)
}
refresh := func(rt string) (int, map[string]interface{}) {
rec, body := e.post("/oauth/token", url.Values{"grant_type": {"refresh_token"}, "refresh_token": {rt}, "client_id": {e.clientID}})
return rec.Code, body
}
code, t2 := refresh(r1)
if code != http.StatusOK {
t.Fatalf("refresh: %d %v", code, t2)
}
r2, _ := t2["refresh_token"].(string)
if r2 == "" || r2 == r1 {
t.Fatalf("refresh token must rotate: %q -> %q", r1, r2)
}
// Presenting the used token again is theft: it fails and kills the family.
if code, body := refresh(r1); code != http.StatusBadRequest || body["error"] != "invalid_grant" {
t.Fatalf("reuse must fail with invalid_grant: %d %v", code, body)
}
if code, _ := refresh(r2); code == http.StatusOK {
t.Fatal("the family must be revoked after reuse")
}
}
func TestOAuthFlow_JWTAccessTokenVerify(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{JWTAccessTokens: true, AccessTokenAudience: "https://api.example.com"})
tok := e.tokens(nil)
at := tok["access_token"].(string)
if strings.Count(at, ".") != 2 {
t.Fatalf("access token is not a JWT: %s", at)
}
claims, err := e.srv.VerifyAccessToken(context.Background(), at, VerifyAccessTokenOptions{Audience: "https://api.example.com", Scopes: []string{"openid"}})
if err != nil {
t.Fatal(err)
}
if !claims.JWT || claims.ClientID != e.clientID {
t.Errorf("claims = %+v", claims)
}
if _, err := e.srv.VerifyAccessToken(context.Background(), at, VerifyAccessTokenOptions{Audience: "https://other"}); err == nil {
t.Error("wrong audience must fail")
}
if _, err := e.srv.VerifyAccessToken(context.Background(), at, VerifyAccessTokenOptions{Scopes: []string{"admin"}}); err == nil {
t.Error("missing scope must fail")
}
}
func TestOAuthFlow_IntrospectionRequiresClientAuth(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{})
tok := e.tokens(nil)
rec, _ := e.post("/oauth/introspect", url.Values{"token": {tok["access_token"].(string)}})
if rec.Code != http.StatusUnauthorized && rec.Code != http.StatusBadRequest {
t.Errorf("anonymous introspection must be refused, got %d", rec.Code)
}
}
func TestOAuthFlow_PAR(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{EnablePAR: true})
_, verifier := e.authURL(nil)
rec, par := e.post("/oauth/par", url.Values{
"response_type": {"code"}, "client_id": {e.clientID}, "redirect_uri": {flowRedirect},
"code_challenge": {s256Challenge(verifier)}, "code_challenge_method": {"S256"}, "scope": {"openid"}, "state": {"par-state"},
})
if rec.Code != http.StatusCreated {
t.Fatalf("par: %d %s", rec.Code, rec.Body.String())
}
uri, _ := par["request_uri"].(string)
if !strings.HasPrefix(uri, "urn:ietf:params:oauth:request_uri:") {
t.Fatalf("request_uri = %q", uri)
}
cb := e.login(e.ts.URL+"/oauth/authorize?"+url.Values{"client_id": {e.clientID}, "request_uri": {uri}}.Encode(), "allow")
if cb.Query().Get("code") == "" || cb.Query().Get("state") != "par-state" {
t.Fatalf("callback = %s", cb)
}
// A request_uri is single use.
resp, _ := e.browser.Get(e.ts.URL + "/oauth/authorize?" + url.Values{"client_id": {e.clientID}, "request_uri": {uri}}.Encode())
if loc := resp.Header.Get("Location"); strings.Contains(loc, "code=") {
t.Errorf("request_uri was reusable: %s", loc)
}
bodyOf(t, resp)
}
func TestOAuthFlow_DeviceGrant(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{EnableDeviceFlow: true, DevicePollSeconds: 1})
// the device client
_, reg := doJSON(t, e.srv.HTTPHandler(), http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{flowRedirect},
"grant_types": []string{"urn:ietf:params:oauth:grant-type:device_code"},
"token_endpoint_auth_method": "none",
})
cid, _ := reg["client_id"].(string)
rec, dev := e.post("/oauth/device_authorization", url.Values{"client_id": {cid}, "scope": {"openid"}})
if rec.Code != http.StatusOK {
t.Fatalf("device_authorization: %d %s", rec.Code, rec.Body.String())
}
poll := func() (int, map[string]interface{}) {
r, b := e.post("/oauth/token", url.Values{
"grant_type": {"urn:ietf:params:oauth:grant-type:device_code"}, "device_code": {dev["device_code"].(string)}, "client_id": {cid},
})
return r.Code, b
}
if code, b := poll(); code != http.StatusBadRequest || b["error"] != "authorization_pending" {
t.Fatalf("first poll: %d %v", code, b)
}
if _, b := poll(); b["error"] != "slow_down" {
t.Errorf("polling too fast must give slow_down, got %v", b)
}
// The user approves in the browser.
resp, _ := e.browser.Get(e.ts.URL + "/oauth/device?user_code=" + url.QueryEscape(dev["user_code"].(string)))
page := bodyOf(t, resp)
for i := 0; i < 4 && !strings.Contains(page, "Approve") && !strings.Contains(page, `value="allow"`); i++ {
switch {
case strings.Contains(page, `name="password"`):
resp = browserSubmit(t, e.browser, e.ts.URL+"/oauth/device", page, url.Values{"username": {"olivia"}, "password": {"pw"}})
case strings.Contains(page, `name="user_code"`):
resp = browserSubmit(t, e.browser, e.ts.URL+"/oauth/device", page, url.Values{"user_code": {dev["user_code"].(string)}})
default:
t.Fatalf("unexpected device page: %s", page)
}
page = bodyOf(t, resp)
}
if strings.Contains(page, `value="allow"`) {
resp = browserSubmit(t, e.browser, e.ts.URL+"/oauth/device", page, url.Values{"decision": {"allow"}})
bodyOf(t, resp)
}
time.Sleep(1100 * time.Millisecond) // past the polling interval
if code, b := poll(); code != http.StatusOK || b["access_token"] == nil {
t.Fatalf("poll after approval: %d %v", code, b)
}
if code, b := poll(); code == http.StatusOK {
t.Errorf("device code must be single use: %v", b)
}
}
func TestOAuthFlow_Discovery(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{EnablePAR: true, EnableDeviceFlow: true, EnableDPoP: true})
rec, doc := doGet(e.srv.HTTPHandler(), "/.well-known/openid-configuration", "")
if rec.Code != http.StatusOK {
t.Fatal(rec.Code)
}
for _, k := range []string{"end_session_endpoint", "device_authorization_endpoint", "pushed_authorization_request_endpoint",
"dpop_signing_alg_values_supported", "authorization_response_iss_parameter_supported", "prompt_values_supported", "response_modes_supported"} {
if doc[k] == nil {
t.Errorf("discovery lacks %s", k)
}
}
if methods, _ := doc["code_challenge_methods_supported"].([]interface{}); len(methods) != 1 || methods[0] != "S256" {
t.Errorf("code_challenge_methods_supported = %v", methods)
}
// The issuer has no path, so a path-inserted form names another issuer.
if rec, _ := doGet(e.srv.HTTPHandler(), "/.well-known/oauth-authorization-server/tenant1", ""); rec.Code != http.StatusNotFound {
t.Errorf("metadata for a foreign issuer path = %d, want 404", rec.Code)
}
}
func TestOAuthServer_DiscoveryPathInsertion(t *testing.T) {
auth := NewDatabaseAuthenticatorWithOptions(newDirectTestDB(t), DatabaseAuthenticatorOptions{Lookup: directConfig})
srv := NewOAuthServer(OAuthServerConfig{Issuer: "https://auth.example.com/tenant1"}, auth)
defer srv.Close()
h := srv.HTTPHandler()
for _, p := range []string{
"/.well-known/oauth-authorization-server/tenant1", // RFC 8414 path insertion
"/.well-known/openid-configuration/tenant1",
"/tenant1/.well-known/openid-configuration", // OIDC Discovery appending
} {
rec, doc := doGet(h, p, "")
if rec.Code != http.StatusOK || doc["issuer"] != "https://auth.example.com/tenant1" {
t.Errorf("%s: %d %v", p, rec.Code, doc["issuer"])
}
}
}
func TestOAuthFlow_RegistrationManagement(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{})
h := e.srv.HTTPHandler()
_, reg := doJSON(t, h, http.MethodPost, "/oauth/register", map[string]interface{}{"redirect_uris": []string{flowRedirect}, "client_name": "App"})
id, _ := reg["client_id"].(string)
rat, _ := reg["registration_access_token"].(string)
if rat == "" {
t.Fatalf("no registration_access_token: %v", reg)
}
req := httptest.NewRequest(http.MethodGet, "/oauth/register/"+id, nil)
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusUnauthorized {
t.Errorf("management without token = %d", rec.Code)
}
req = httptest.NewRequest(http.MethodGet, "/oauth/register/"+id, nil)
req.Header.Set("Authorization", "Bearer "+rat)
rec = httptest.NewRecorder()
h.ServeHTTP(rec, req)
var got map[string]interface{}
_ = json.Unmarshal(rec.Body.Bytes(), &got)
if rec.Code != http.StatusOK || got["client_name"] != "App" {
t.Fatalf("management read: %d %s", rec.Code, rec.Body.String())
}
req = httptest.NewRequest(http.MethodDelete, "/oauth/register/"+id, nil)
req.Header.Set("Authorization", "Bearer "+rat)
rec = httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusNoContent {
t.Errorf("delete = %d", rec.Code)
}
}
func TestOAuthFlow_Logout(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{})
tok := e.tokens(nil)
idt, _ := tok["id_token"].(string)
if idt == "" {
t.Fatal("no id_token")
}
resp, err := e.browser.Get(e.ts.URL + "/oauth/logout?" + url.Values{"id_token_hint": {idt}, "client_id": {e.clientID}}.Encode())
if err != nil {
t.Fatal(err)
}
page := bodyOf(t, resp)
if resp.StatusCode == http.StatusOK && strings.Contains(page, `name="confirm"`) {
resp = browserSubmit(t, e.browser, e.ts.URL+"/oauth/logout", page, url.Values{"confirm": {"yes"}})
bodyOf(t, resp)
}
// SSO is gone: prompt=none now reports login_required.
authURL, _ := e.authURL(url.Values{"prompt": {"none"}})
if cb := e.login(authURL, "allow"); cb.Query().Get("error") != "login_required" {
t.Errorf("after logout want login_required, got %s", cb)
}
}
func makeDPoP(t *testing.T, key *ecdsa.PrivateKey, method, target, accessToken, jti string) string {
t.Helper()
jwk, err := jwkFromPublic(&key.PublicKey, "", "ES256")
if err != nil {
t.Fatal(err)
}
claims := jwt.MapClaims{"htm": method, "htu": target, "iat": time.Now().Unix(), "jti": jti}
if accessToken != "" {
h := sha256.Sum256([]byte(accessToken))
claims["ath"] = base64.RawURLEncoding.EncodeToString(h[:])
}
tok := jwt.NewWithClaims(jwt.SigningMethodES256, claims)
tok.Header["typ"] = "dpop+jwt"
tok.Header["jwk"] = jwk
s, err := tok.SignedString(key)
if err != nil {
t.Fatal(err)
}
return s
}
func TestOAuthFlow_DPoP(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{EnableDPoP: true})
h := e.srv.HTTPHandler()
key, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
authURL, verifier := e.authURL(nil)
code := e.login(authURL, "allow").Query().Get("code")
tokenURL := e.ts.URL + "/oauth/token"
form := url.Values{"grant_type": {"authorization_code"}, "code": {code}, "redirect_uri": {flowRedirect},
"client_id": {e.clientID}, "code_verifier": {verifier}}
// A proof for another URL is refused (and does not burn the code).
req := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("DPoP", makeDPoP(t, key, "POST", e.ts.URL+"/other", "", "j1"))
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest || !strings.Contains(rec.Body.String(), "invalid_dpop_proof") {
t.Fatalf("wrong htu: %d %s", rec.Code, rec.Body.String())
}
req = httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("DPoP", makeDPoP(t, key, "POST", tokenURL, "", "j2"))
rec = httptest.NewRecorder()
h.ServeHTTP(rec, req)
var tok map[string]interface{}
_ = json.Unmarshal(rec.Body.Bytes(), &tok)
if rec.Code != http.StatusOK || tok["token_type"] != "DPoP" {
t.Fatalf("token: %d %s", rec.Code, rec.Body.String())
}
at := tok["access_token"].(string)
userinfo := func(scheme, proof string) int {
req := httptest.NewRequest(http.MethodGet, "/oauth/userinfo", nil)
req.Header.Set("Authorization", scheme+" "+at)
if proof != "" {
req.Header.Set("DPoP", proof)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
return rec.Code
}
if c := userinfo("Bearer", ""); c != http.StatusUnauthorized {
t.Errorf("a DPoP-bound token must not work as a Bearer token, got %d", c)
}
if c := userinfo("DPoP", makeDPoP(t, key, "GET", e.ts.URL+"/oauth/userinfo", at, "j3")); c != http.StatusOK {
t.Errorf("valid DPoP userinfo = %d", c)
}
if c := userinfo("DPoP", makeDPoP(t, key, "GET", e.ts.URL+"/oauth/userinfo", at, "j3")); c == http.StatusOK {
t.Error("a replayed proof (same jti) must be refused")
}
other, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if c := userinfo("DPoP", makeDPoP(t, other, "GET", e.ts.URL+"/oauth/userinfo", at, "j4")); c == http.StatusOK {
t.Error("a proof from another key must be refused")
}
}
func TestOAuthFlow_TokenExchange(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{EnableTokenExchange: true})
h := e.srv.HTTPHandler()
_, reg := doJSON(t, h, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{flowRedirect},
"grant_types": []string{"urn:ietf:params:oauth:grant-type:token-exchange"},
"token_endpoint_auth_method": "client_secret_basic",
})
cid, secret := reg["client_id"].(string), reg["client_secret"].(string)
user := e.tokens(url.Values{"scope": {"openid profile"}})
rec, out := doForm(t, h, "/oauth/token", url.Values{
"grant_type": {"urn:ietf:params:oauth:grant-type:token-exchange"},
"subject_token": {user["access_token"].(string)}, "subject_token_type": {"urn:ietf:params:oauth:token-type:access_token"},
"scope": {"openid"},
}, cid, secret)
if rec.Code != http.StatusOK || out["access_token"] == nil || out["issued_token_type"] != "urn:ietf:params:oauth:token-type:access_token" {
t.Fatalf("exchange: %d %s", rec.Code, rec.Body.String())
}
if out["scope"] != "openid" {
t.Errorf("scope must be downscoped, got %v", out["scope"])
}
// Widening the scope is refused.
rec, _ = doForm(t, h, "/oauth/token", url.Values{
"grant_type": {"urn:ietf:params:oauth:grant-type:token-exchange"},
"subject_token": {out["access_token"].(string)}, "scope": {"openid email"},
}, cid, secret)
if rec.Code != http.StatusBadRequest {
t.Errorf("scope widening = %d", rec.Code)
}
// The public client of the user cannot exchange.
rec, _ = e.post("/oauth/token", url.Values{
"grant_type": {"urn:ietf:params:oauth:grant-type:token-exchange"},
"subject_token": {user["access_token"].(string)}, "client_id": {e.clientID},
})
if rec.Code == http.StatusOK {
t.Error("public clients must not use token exchange")
}
}
func TestOAuthFlow_PrivateKeyJWTClientAuth(t *testing.T) {
e := newFlowEnv(t, OAuthServerConfig{})
h := e.srv.HTTPHandler()
key, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
jwk, err := jwkFromPublic(&key.PublicKey, "k1", "ES256")
if err != nil {
t.Fatal(err)
}
rec, reg := doJSON(t, h, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{flowRedirect},
"grant_types": []string{"client_credentials"},
"token_endpoint_auth_method": "private_key_jwt",
"jwks": map[string]interface{}{"keys": []interface{}{jwk}},
})
if rec.Code != http.StatusCreated {
t.Fatalf("register: %d %s", rec.Code, rec.Body.String())
}
cid := reg["client_id"].(string)
assertion := func(jti, aud string) string {
tok := jwt.NewWithClaims(jwt.SigningMethodES256, jwt.MapClaims{
"iss": cid, "sub": cid, "aud": aud, "jti": jti, "exp": time.Now().Add(time.Minute).Unix(),
})
tok.Header["kid"] = "k1"
s, err := tok.SignedString(key)
if err != nil {
t.Fatal(err)
}
return s
}
call := func(a string) int {
rec, _ := doForm(t, h, "/oauth/token", url.Values{
"grant_type": {"client_credentials"}, "client_id": {cid},
"client_assertion_type": {"urn:ietf:params:oauth:client-assertion-type:jwt-bearer"}, "client_assertion": {a},
}, "", "")
return rec.Code
}
good := assertion("a1", e.ts.URL+"/oauth/token")
if c := call(good); c != http.StatusOK {
t.Fatalf("valid assertion = %d", c)
}
if c := call(good); c == http.StatusOK {
t.Error("a replayed assertion must be refused")
}
if c := call(assertion("a2", "https://elsewhere.example/token")); c == http.StatusOK {
t.Error("an assertion for another audience must be refused")
}
}
+124
View File
@@ -0,0 +1,124 @@
package security
import (
"errors"
"net/http"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// introspectionCaller authenticates the caller of the revocation / introspection endpoints.
func (s *OAuthServer) introspectionCaller(r *http.Request) (*authedClient, *oauthError) {
ac, e := s.authenticateClient(r)
if e != nil {
return nil, e
}
if (ac == nil || ac.Method == "none") && !s.cfg.AllowAnonymousIntrospection {
return nil, invalidClient("client authentication required", true)
}
return ac, nil
}
// --------------------------------------------------------------------------
// RFC 7662 — Token introspection
// --------------------------------------------------------------------------
func (s *OAuthServer) introspectHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
if err := r.ParseForm(); err != nil {
writeOAuthError(w, "invalid_request", "cannot parse form", http.StatusBadRequest)
return
}
if _, e := s.introspectionCaller(r); e != nil {
e.write(w)
return
}
token := r.PostFormValue("token")
inactive := map[string]any{"active": false}
if token == "" {
writeJSON(w, http.StatusOK, inactive)
return
}
if r.PostFormValue("token_type_hint") != "access_token" {
if gs := s.grants(); gs != nil {
if rt, err := gs.PeekRefresh(r.Context(), hashToken(token)); err == nil && !isAccessRecord(rt) {
writeJSON(w, http.StatusOK, map[string]any{
"active": true, "sub": itoa(rt.UserID), "client_id": rt.ClientID, "scope": joinScopes(rt.Scopes),
"exp": rt.ExpiresAt.Unix(), "iss": s.cfg.Issuer, "token_type": "refresh_token",
})
return
}
}
}
ti := s.resolveAccessToken(r.Context(), token)
if ti == nil {
writeJSON(w, http.StatusOK, inactive)
return
}
writeJSON(w, http.StatusOK, s.introspectionInfo(ti))
}
// --------------------------------------------------------------------------
// RFC 7009 — Token revocation
// --------------------------------------------------------------------------
func (s *OAuthServer) revokeHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
if err := r.ParseForm(); err != nil {
writeOAuthError(w, "invalid_request", "cannot parse form", http.StatusBadRequest)
return
}
ac, e := s.introspectionCaller(r)
if e != nil {
e.write(w)
return
}
token := r.PostFormValue("token")
if token == "" {
w.WriteHeader(http.StatusOK)
return
}
owns := func(clientID string) bool { return ac == nil || clientID == "" || ac.Client.ClientID == clientID }
ctx := r.Context()
if gs := s.grants(); gs != nil {
if rt, err := gs.PeekRefresh(ctx, hashToken(token)); err == nil && !isAccessRecord(rt) {
if owns(rt.ClientID) {
_ = gs.RevokeRefreshFamily(ctx, rt.FamilyID)
}
w.WriteHeader(http.StatusOK)
return
} else if err != nil && !errors.Is(err, lookup.ErrRefreshInvalid) {
w.WriteHeader(http.StatusOK)
return
}
if ti := s.resolveAccessToken(ctx, token); ti != nil {
if owns(ti.ClientID) {
id := token
if ti.JWT {
id = ti.JTI
}
_ = gs.RevokeRefreshFamily(ctx, accessKey(id))
if !ti.JWT && ti.ClientID != "" {
if a := s.anyAuth(); a != nil {
_ = a.OAuthRevokeToken(ctx, token)
}
}
}
w.WriteHeader(http.StatusOK)
return
}
}
// Tokens issued by earlier versions (session tokens without a recorded grant, pass-through refresh tokens).
if a := s.anyAuth(); a != nil {
_ = a.OAuthRevokeToken(ctx, token)
}
w.WriteHeader(http.StatusOK)
}
+421
View File
@@ -0,0 +1,421 @@
package security
import (
"context"
"crypto"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"math/big"
"net"
"net/http"
"sort"
"sync"
"syscall"
"time"
"github.com/golang-jwt/jwt/v5"
)
// OAuthSigningKey is a key the authorization server signs tokens with. The first configured
// key is the default; the others are published in the JWKS so tokens signed before a rotation
// stay verifiable, and a client can ask for one by id_token_signed_response_alg.
type OAuthSigningKey struct {
// ID is the JWKS "kid". Derived from the public key (RFC 7638 thumbprint) when empty.
ID string
// Key is an *rsa.PrivateKey (RS256) or an *ecdsa.PrivateKey (ES256 for P-256, ES384 for P-384).
Key crypto.Signer
// Alg overrides the algorithm inferred from Key (RS256, PS256, ES256, ES384).
Alg string
}
type oauthKey struct {
id string
alg string
signer crypto.Signer
method jwt.SigningMethod
}
// oauthKeyring holds the server's signing keys.
type oauthKeyring struct {
keys []oauthKey
}
func newOAuthKeyring(cfg *OAuthServerConfig) (*oauthKeyring, error) {
in := cfg.SigningKeys
if len(in) == 0 && cfg.SigningKey != nil {
in = []OAuthSigningKey{{Key: cfg.SigningKey}}
}
if len(in) == 0 {
k, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
return nil, fmt.Errorf("generate signing key: %w", err)
}
in = []OAuthSigningKey{{Key: k}}
}
kr := &oauthKeyring{}
for _, k := range in {
if k.Key == nil {
return nil, fmt.Errorf("signing key without a private key")
}
alg := k.Alg
if alg == "" {
switch key := k.Key.(type) {
case *rsa.PrivateKey:
alg = "RS256"
case *ecdsa.PrivateKey:
switch key.Curve {
case elliptic.P256():
alg = "ES256"
case elliptic.P384():
alg = "ES384"
default:
return nil, fmt.Errorf("unsupported signing curve %s", key.Curve.Params().Name)
}
default:
return nil, fmt.Errorf("unsupported signing key type %T", k.Key)
}
}
method := jwt.GetSigningMethod(alg)
if method == nil {
return nil, fmt.Errorf("unsupported signing algorithm %q", alg)
}
id := k.ID
if id == "" {
tp, err := jwkThumbprint(k.Key.Public())
if err != nil {
return nil, err
}
id = tp[:16]
}
kr.keys = append(kr.keys, oauthKey{id: id, alg: alg, signer: k.Key, method: method})
}
return kr, nil
}
// forAlg returns the first key signing with alg, or the default key when alg is empty or unknown.
func (kr *oauthKeyring) forAlg(alg string) *oauthKey {
if alg != "" {
for i := range kr.keys {
if kr.keys[i].alg == alg {
return &kr.keys[i]
}
}
}
return &kr.keys[0]
}
func (kr *oauthKeyring) algs() []string {
var out []string
for _, k := range kr.keys {
if !oauthSliceContains(out, k.alg) {
out = append(out, k.alg)
}
}
return out
}
func (kr *oauthKeyring) jwks() []map[string]any {
out := make([]map[string]any, 0, len(kr.keys))
for _, k := range kr.keys {
if jwk, err := jwkFromPublic(k.signer.Public(), k.id, k.alg); err == nil {
out = append(out, jwk)
}
}
return out
}
// publicFor returns the public key with kid (or the default one when kid is empty).
func (kr *oauthKeyring) publicFor(kid string) (crypto.PublicKey, *oauthKey) {
for i := range kr.keys {
if kid == "" || kr.keys[i].id == kid {
return kr.keys[i].signer.Public(), &kr.keys[i]
}
}
return nil, nil
}
// sign signs claims with k, adding typ when non-empty.
func (k *oauthKey) sign(claims jwt.Claims, typ string) (string, error) {
t := jwt.NewWithClaims(k.method, claims)
t.Header["kid"] = k.id
if typ != "" {
t.Header["typ"] = typ
}
return t.SignedString(k.signer)
}
// --- JWK encoding ----------------------------------------------------------------------------
func b64u(b []byte) string { return base64.RawURLEncoding.EncodeToString(b) }
func jwkFromPublic(pub crypto.PublicKey, kid, alg string) (map[string]any, error) {
jwk := map[string]any{"use": "sig"}
if kid != "" {
jwk["kid"] = kid
}
if alg != "" {
jwk["alg"] = alg
}
switch k := pub.(type) {
case *rsa.PublicKey:
jwk["kty"] = "RSA"
jwk["n"] = b64u(k.N.Bytes())
jwk["e"] = b64u(big.NewInt(int64(k.E)).Bytes())
case *ecdsa.PublicKey:
size := (k.Curve.Params().BitSize + 7) / 8
jwk["kty"] = "EC"
jwk["crv"] = k.Curve.Params().Name
jwk["x"] = b64u(k.X.FillBytes(make([]byte, size)))
jwk["y"] = b64u(k.Y.FillBytes(make([]byte, size)))
default:
return nil, fmt.Errorf("unsupported key type %T", pub)
}
return jwk, nil
}
// jwkThumbprint is the RFC 7638 SHA-256 thumbprint (base64url).
func jwkThumbprint(pub crypto.PublicKey) (string, error) {
var canonical string
switch k := pub.(type) {
case *rsa.PublicKey:
canonical = fmt.Sprintf(`{"e":%q,"kty":"RSA","n":%q}`, b64u(big.NewInt(int64(k.E)).Bytes()), b64u(k.N.Bytes()))
case *ecdsa.PublicKey:
size := (k.Curve.Params().BitSize + 7) / 8
canonical = fmt.Sprintf(`{"crv":%q,"kty":"EC","x":%q,"y":%q}`, k.Curve.Params().Name,
b64u(k.X.FillBytes(make([]byte, size))), b64u(k.Y.FillBytes(make([]byte, size))))
default:
return "", fmt.Errorf("unsupported key type %T", pub)
}
sum := sha256.Sum256([]byte(canonical))
return b64u(sum[:]), nil
}
func publicFromJWK(m map[string]any) (crypto.PublicKey, error) {
str := func(k string) string { s, _ := m[k].(string); return s }
dec := func(k string) (*big.Int, error) {
raw, err := base64.RawURLEncoding.DecodeString(str(k))
if err != nil || len(raw) == 0 {
return nil, fmt.Errorf("invalid JWK member %q", k)
}
return new(big.Int).SetBytes(raw), nil
}
switch str("kty") {
case "RSA":
n, err := dec("n")
if err != nil {
return nil, err
}
e, err := dec("e")
if err != nil {
return nil, err
}
if n.BitLen() < 2048 || !e.IsInt64() || e.Int64() < 3 {
return nil, fmt.Errorf("RSA key too weak")
}
return &rsa.PublicKey{N: n, E: int(e.Int64())}, nil
case "EC":
var curve elliptic.Curve
switch str("crv") {
case "P-256":
curve = elliptic.P256()
case "P-384":
curve = elliptic.P384()
default:
return nil, fmt.Errorf("unsupported curve %q", str("crv"))
}
x, err := dec("x")
if err != nil {
return nil, err
}
y, err := dec("y")
if err != nil {
return nil, err
}
pub := &ecdsa.PublicKey{Curve: curve, X: x, Y: y}
if _, err := pub.ECDH(); err != nil { // rejects points that are not on the curve
return nil, fmt.Errorf("invalid EC point: %w", err)
}
return pub, nil
}
return nil, fmt.Errorf("unsupported key type %q", str("kty"))
}
// jwkEntry is one verification key of a JWK set.
type jwkEntry struct {
kid string
alg string
use string
pub crypto.PublicKey
}
// jwkSet is a parsed JWK set.
type jwkSet struct{ keys []jwkEntry }
func parseJWKS(data []byte) (*jwkSet, error) {
var doc struct {
Keys []map[string]any `json:"keys"`
}
if err := json.Unmarshal(data, &doc); err != nil {
return nil, fmt.Errorf("invalid JWKS: %w", err)
}
set := &jwkSet{}
for _, k := range doc.Keys {
pub, err := publicFromJWK(k)
if err != nil {
continue // skip keys of a type we cannot use
}
e := jwkEntry{pub: pub}
e.kid, _ = k["kid"].(string)
e.alg, _ = k["alg"].(string)
e.use, _ = k["use"].(string)
set.keys = append(set.keys, e)
}
if len(set.keys) == 0 {
return nil, fmt.Errorf("JWKS contains no usable keys")
}
return set, nil
}
// candidates returns the keys that may have produced a token with kid and alg.
func (s *jwkSet) candidates(kid, alg string) []crypto.PublicKey {
var out []crypto.PublicKey
for _, k := range s.keys {
if k.use != "" && k.use != "sig" {
continue
}
if kid != "" && k.kid != "" && k.kid != kid {
continue
}
if k.alg != "" && alg != "" && k.alg != alg {
continue
}
out = append(out, k.pub)
}
return out
}
// verifyJWTWithSet verifies token against the set. allowed lists the accepted algorithms.
func verifyJWTWithSet(token string, set *jwkSet, allowed []string, claims jwt.Claims, opts ...jwt.ParserOption) (*jwt.Token, error) {
opts = append([]jwt.ParserOption{jwt.WithValidMethods(allowed)}, opts...)
parser := jwt.NewParser(opts...)
lastErr := fmt.Errorf("no matching key")
unverified, _, err := parser.ParseUnverified(token, jwt.MapClaims{})
if err != nil {
return nil, err
}
kid, _ := unverified.Header["kid"].(string)
alg, _ := unverified.Header["alg"].(string)
for _, pub := range set.candidates(kid, alg) {
tok, err := parser.ParseWithClaims(token, claims, func(*jwt.Token) (any, error) { return pub, nil })
if err == nil {
return tok, nil
}
lastErr = err
}
return nil, lastErr
}
// --- remote JWKS fetching --------------------------------------------------------------------
// jwksCache fetches and caches remote JWK sets.
type jwksCache struct {
client *http.Client
ttl time.Duration
mu sync.Mutex
entries map[string]jwksCacheEntry
}
type jwksCacheEntry struct {
set *jwkSet
fetched time.Time
}
func newJWKSCache(client *http.Client) *jwksCache {
return &jwksCache{client: client, ttl: time.Hour, entries: map[string]jwksCacheEntry{}}
}
// get returns the set at uri. refresh forces a fetch (used when a kid is not in the cached set),
// rate limited to one fetch per 30 seconds per URI.
func (c *jwksCache) get(ctx context.Context, uri string, refresh bool) (*jwkSet, error) {
c.mu.Lock()
e, ok := c.entries[uri]
c.mu.Unlock()
age := time.Since(e.fetched)
if ok && age < c.ttl && (!refresh || age < 30*time.Second) {
return e.set, nil
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, uri, nil)
if err != nil {
return nil, err
}
resp, err := c.client.Do(req)
if err != nil {
if ok {
return e.set, nil // keep using the stale set while the endpoint is down
}
return nil, fmt.Errorf("fetch JWKS: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("fetch JWKS: status %d", resp.StatusCode)
}
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
if err != nil {
return nil, err
}
set, err := parseJWKS(body)
if err != nil {
return nil, err
}
c.mu.Lock()
c.entries[uri] = jwksCacheEntry{set: set, fetched: time.Now()}
c.mu.Unlock()
return set, nil
}
// publicHTTPClient returns an HTTP client for fetching URLs supplied by clients. It refuses to
// connect to loopback, private and link-local addresses (SSRF) unless allowPrivate is set.
func publicHTTPClient(allowPrivate bool) *http.Client {
dialer := &net.Dialer{Timeout: 5 * time.Second}
if !allowPrivate {
dialer.Control = func(_, address string, _ syscall.RawConn) error {
host, _, err := net.SplitHostPort(address)
if err != nil {
return err
}
ip := net.ParseIP(host)
if ip == nil || ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() ||
ip.IsUnspecified() || ip.IsMulticast() {
return fmt.Errorf("address %s is not allowed", host)
}
return nil
}
}
return &http.Client{
Timeout: 10 * time.Second,
Transport: &http.Transport{DialContext: dialer.DialContext, Proxy: nil},
CheckRedirect: func(_ *http.Request, via []*http.Request) error {
if len(via) >= 3 {
return fmt.Errorf("too many redirects")
}
return nil
},
}
}
func sortedKeys(m map[string]struct{}) []string {
out := make([]string, 0, len(m))
for k := range m {
out = append(out, k)
}
sort.Strings(out)
return out
}
+204
View File
@@ -0,0 +1,204 @@
package security
import (
"bytes"
"context"
"net/http"
"net/url"
"strconv"
"time"
"github.com/golang-jwt/jwt/v5"
)
const backchannelEvent = "http://schemas.openid.net/event/backchannel-logout"
// hintSubject returns the subject and client of an id_token_hint issued by this server ("" when
// the hint is not valid). The hint may be expired.
func (s *OAuthServer) hintSubject(hint string) (sub, clientID string) {
claims, err := s.parseOwnJWTOpts(hint, "", true)
if err != nil {
return "", ""
}
sub, _ = claims["sub"].(string)
if auds := audOf(claims["aud"]); len(auds) > 0 {
clientID = auds[0]
}
return sub, clientID
}
type logoutState struct {
ClientID string `json:"c,omitempty"`
RedirectURI string `json:"r,omitempty"`
State string `json:"s,omitempty"`
Sid string `json:"i,omitempty"`
Sub string `json:"u,omitempty"`
}
// --------------------------------------------------------------------------
// OIDC RP-Initiated Logout — GET/POST /oauth/logout
// --------------------------------------------------------------------------
func (s *OAuthServer) logoutHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
if err := r.ParseForm(); err != nil {
http.Error(w, "invalid form", http.StatusBadRequest)
return
}
// Confirmation answer to the page rendered below.
if blob := r.PostFormValue("lreq"); blob != "" {
var st logoutState
if err := s.open("logout", blob, &st); err != nil {
s.renderMessage(w, http.StatusBadRequest, "Request expired", "This sign-out request has expired.", true)
return
}
if r.PostFormValue("confirm") != "yes" {
s.finishLogout(w, r, &st, false)
return
}
s.performLogout(w, r, &st)
return
}
st := logoutState{State: r.FormValue("state")}
hint := r.FormValue("id_token_hint")
if hint != "" {
claims, err := s.parseOwnJWTOpts(hint, "", true)
if err != nil {
s.renderMessage(w, http.StatusBadRequest, "Invalid request", "The id_token_hint is not valid.", true)
return
}
st.Sub, _ = claims["sub"].(string)
st.Sid, _ = claims["sid"].(string)
if auds := audOf(claims["aud"]); len(auds) > 0 {
st.ClientID = auds[0]
}
if p := r.FormValue("client_id"); p != "" && p != st.ClientID {
s.renderMessage(w, http.StatusBadRequest, "Invalid request", "client_id does not match the id_token_hint.", true)
return
}
} else {
st.ClientID = r.FormValue("client_id")
}
if uri := r.FormValue("post_logout_redirect_uri"); uri != "" {
client, ok := s.lookupOrFetchClient(r.Context(), st.ClientID)
if !ok || !oauthSliceContains(client.PostLogoutRedirectURIs, uri) {
s.renderMessage(w, http.StatusBadRequest, "Invalid request", "post_logout_redirect_uri is not registered for this client.", true)
return
}
st.RedirectURI = uri
}
if hint != "" {
s.performLogout(w, r, &st)
return
}
blob, err := s.seal("logout", st, 15*time.Minute)
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
s.renderHTML(w, http.StatusOK, nil, "logout", oauthLogoutPage{
Title: "Sign out", Action: "logout", Hidden: map[string]string{"lreq": blob},
})
}
// performLogout ends the SSO session, revokes the refresh tokens bound to it and notifies the
// clients that hold tokens for it.
func (s *OAuthServer) performLogout(w http.ResponseWriter, r *http.Request, st *logoutState) {
ctx := r.Context()
sids := map[string]struct{}{}
if st.Sid != "" {
sids[st.Sid] = struct{}{}
}
clients := map[string]struct{}{}
if st.ClientID != "" {
clients[st.ClientID] = struct{}{}
}
sub, sid := st.Sub, st.Sid
if c, err := r.Cookie(s.cfg.SSOCookie.Name); err == nil {
var sess ssoSession
if s.open("sso", c.Value, &sess) == nil && (st.Sub == "" || st.Sub == strconv.Itoa(sess.UserID)) {
sids[sess.SID] = struct{}{}
for _, cl := range sess.Clients {
clients[cl] = struct{}{}
}
sub, sid = strconv.Itoa(sess.UserID), sess.SID
if a := s.anyAuth(); a != nil {
_ = a.Logout(ctx, LogoutRequest{Token: sess.Token})
}
}
}
if gs := s.grants(); gs != nil {
for id := range sids {
_ = gs.RevokeRefreshBySession(ctx, id)
}
}
s.clearSSO(w)
for id := range clients {
if client, ok := s.lookupOrFetchClient(ctx, id); ok && client.BackchannelLogoutURI != "" {
s.sendBackchannelLogout(client, sub, sid)
}
}
s.finishLogout(w, r, st, true)
}
func (s *OAuthServer) finishLogout(w http.ResponseWriter, r *http.Request, st *logoutState, done bool) {
if st.RedirectURI != "" {
u, err := url.Parse(st.RedirectURI)
if err == nil {
if st.State != "" {
q := u.Query()
q.Set("state", st.State)
u.RawQuery = q.Encode()
}
http.Redirect(w, r, u.String(), http.StatusFound)
return
}
}
if done {
s.renderMessage(w, http.StatusOK, "Signed out", "You have been signed out.", false)
return
}
s.renderMessage(w, http.StatusOK, "Still signed in", "You are still signed in.", false)
}
// sendBackchannelLogout posts a logout token to the client (OIDC Back-Channel Logout 1.0), best effort.
func (s *OAuthServer) sendBackchannelLogout(client *OAuthServerClient, sub, sid string) {
jti, err := randomOAuthToken()
if err != nil {
return
}
claims := jwt.MapClaims{
"iss": s.cfg.Issuer, "aud": client.ClientID, "iat": time.Now().Unix(), "jti": jti, "sub": sub,
"events": map[string]any{backchannelEvent: map[string]any{}},
}
if sid != "" {
claims["sid"] = sid
}
token, err := s.keys.forAlg(client.IDTokenSignedResponseAlg).sign(claims, "logout+jwt")
if err != nil {
return
}
target := client.BackchannelLogoutURI
s.bcWG.Add(1)
go func() {
defer s.bcWG.Done()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
body := url.Values{"logout_token": {token}}.Encode()
req, err := http.NewRequestWithContext(ctx, http.MethodPost, target, bytes.NewBufferString(body))
if err != nil {
return
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
resp, err := s.jwks.client.Do(req)
if err == nil {
_ = resp.Body.Close()
}
}()
}
+603
View File
@@ -0,0 +1,603 @@
package security
import (
"context"
"crypto/sha256"
"crypto/sha512"
"encoding/json"
"errors"
"hash"
"net/http"
"strconv"
"strings"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// standard claim sets released by scope (OIDC Core §5.4).
var scopeClaims = map[string][]string{
"profile": {"name", "family_name", "given_name", "middle_name", "nickname", "preferred_username", "profile",
"picture", "website", "gender", "birthdate", "zoneinfo", "locale", "updated_at"},
"email": {"email", "email_verified"},
"address": {"address"},
"phone": {"phone_number", "phone_number_verified"},
}
// --------------------------------------------------------------------------
// Discovery — RFC 8414 / OIDC Discovery / RFC 9728
// --------------------------------------------------------------------------
func (s *OAuthServer) grantTypesSupported() []string {
g := []string{"authorization_code", "refresh_token"}
if s.auth != nil {
g = append(g, "client_credentials")
}
if s.cfg.EnableDeviceFlow {
g = append(g, grantDeviceCode)
}
if s.cfg.EnableTokenExchange {
g = append(g, grantTokenExchange)
}
return g
}
func (s *OAuthServer) authMethodsSupported() []string {
return []string{"none", "client_secret_basic", "client_secret_post", "private_key_jwt"}
}
// serverMetadata builds the fields shared by RFC 8414 authorization-server
// metadata and OIDC discovery metadata.
func (s *OAuthServer) serverMetadata() map[string]any {
m := map[string]any{
"issuer": s.cfg.Issuer,
"authorization_endpoint": s.endpoint("/oauth/authorize"),
"token_endpoint": s.endpoint("/oauth/token"),
"registration_endpoint": s.endpoint("/oauth/register"),
"revocation_endpoint": s.endpoint("/oauth/revoke"),
"introspection_endpoint": s.endpoint("/oauth/introspect"),
"userinfo_endpoint": s.endpoint("/oauth/userinfo"),
"jwks_uri": s.endpoint("/oauth/jwks.json"),
"scopes_supported": s.scopesSupported(),
"response_types_supported": []string{"code"},
"response_modes_supported": []string{"query", "form_post"},
"grant_types_supported": s.grantTypesSupported(),
"code_challenge_methods_supported": []string{"S256"},
"token_endpoint_auth_methods_supported": s.authMethodsSupported(),
"token_endpoint_auth_signing_alg_values_supported": []string{"RS256", "PS256", "ES256", "ES384"},
"revocation_endpoint_auth_methods_supported": s.authMethodsSupported(),
"introspection_endpoint_auth_methods_supported": s.authMethodsSupported(),
"authorization_response_iss_parameter_supported": true,
"prompt_values_supported": []string{"none", "login", "consent"},
"request_parameter_supported": false,
"request_uri_parameter_supported": s.cfg.EnablePAR || s.cfg.RequirePAR,
"claims_parameter_supported": true,
"subject_types_supported": []string{"public"},
"id_token_signing_alg_values_supported": s.keys.algs(),
"userinfo_signing_alg_values_supported": append([]string{"none"}, s.keys.algs()...),
"claim_types_supported": []string{"normal"},
"claims_supported": s.claimsSupported(),
"service_documentation": "https://github.com/bitechdev/ResolveSpec/blob/main/pkg/security/OAUTH2_SERVER.md",
}
if len(s.cfg.SupportedACR) > 0 {
m["acr_values_supported"] = s.cfg.SupportedACR
}
if s.cfg.EnablePAR || s.cfg.RequirePAR {
m["pushed_authorization_request_endpoint"] = s.endpoint("/oauth/par")
m["require_pushed_authorization_requests"] = s.cfg.RequirePAR
}
if s.cfg.EnableDeviceFlow {
m["device_authorization_endpoint"] = s.endpoint("/oauth/device_authorization")
}
if s.cfg.EnableDPoP {
m["dpop_signing_alg_values_supported"] = []string{"ES256", "ES384", "RS256", "PS256"}
}
if !s.cfg.DisableLogout {
m["end_session_endpoint"] = s.endpoint("/oauth/logout")
m["backchannel_logout_supported"] = true
m["backchannel_logout_session_supported"] = true
}
return m
}
func (s *OAuthServer) scopesSupported() []string {
out := append([]string(nil), s.cfg.DefaultScopes...)
if s.cfg.ManagedRefreshTokens && !oauthSliceContains(out, "offline_access") {
out = append(out, "offline_access")
}
return out
}
func (s *OAuthServer) claimsSupported() []string {
set := map[string]struct{}{"sub": {}, "iss": {}, "aud": {}, "exp": {}, "iat": {}, "auth_time": {}, "nonce": {},
"acr": {}, "amr": {}, "azp": {}, "at_hash": {}, "sid": {}}
for _, sc := range s.scopesSupported() {
for _, c := range scopeClaims[sc] {
set[c] = struct{}{}
}
}
return sortedKeys(set)
}
// metadataHandler serves one of the well-known documents. The path-insertion forms
// (/.well-known/x/<issuer path>) are answered only for this server's issuer path.
func (s *OAuthServer) metadataHandler(name string) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if rest := r.PathValue("path"); rest != "" {
if "/"+strings.Trim(rest, "/") != strings.TrimSuffix(s.issuerURL.Path, "/") {
http.NotFound(w, r)
return
}
}
var doc map[string]any
switch name {
case "oauth-protected-resource":
doc = map[string]any{
"resource": s.cfg.ResourceIdentifier,
"authorization_servers": []string{s.cfg.Issuer},
"scopes_supported": s.cfg.DefaultScopes,
"bearer_methods_supported": []string{"header"},
}
if s.cfg.EnableDPoP {
doc["dpop_signing_alg_values_supported"] = []string{"ES256", "ES384", "RS256", "PS256"}
}
default:
doc = s.serverMetadata()
}
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "public, max-age=300")
json.NewEncoder(w).Encode(doc) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
}
func (s *OAuthServer) jwksHandler(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "public, max-age=300")
json.NewEncoder(w).Encode(map[string]any{"keys": s.keys.jwks()}) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored
}
// --------------------------------------------------------------------------
// id_token
// --------------------------------------------------------------------------
type idTokenParams struct {
Client *OAuthServerClient
UserID int
Scopes []string
Nonce string
AuthTime int64
ACR string
AMR []string
SID string
AccessToken string
Claims map[string]any // the OIDC "claims" request parameter
}
// halfHash is the left half of the hash of v (at_hash / c_hash), with the hash of the JWS alg.
func halfHash(alg, v string) string {
var h hash.Hash
switch {
case strings.HasSuffix(alg, "384"):
h = sha512.New384()
case strings.HasSuffix(alg, "512"):
h = sha512.New()
default:
h = sha256.New()
}
h.Write([]byte(v))
sum := h.Sum(nil)
return b64u(sum[:len(sum)/2])
}
func (s *OAuthServer) buildIDToken(ctx context.Context, p idTokenParams) (string, error) {
a := s.anyAuth()
if a == nil {
return "", errors.New("no authenticator configured")
}
user, err := a.OAuthGetUser(ctx, p.UserID)
if err != nil {
return "", err
}
key := s.keys.forAlg(p.Client.IDTokenSignedResponseAlg)
now := time.Now()
claims := jwt.MapClaims{
"iss": s.cfg.Issuer, "sub": strconv.Itoa(p.UserID), "aud": p.Client.ClientID, "azp": p.Client.ClientID,
"exp": now.Add(s.cfg.AccessTokenTTL).Unix(), "iat": now.Unix(),
}
if p.Nonce != "" {
claims["nonce"] = p.Nonce
}
if p.AuthTime != 0 {
claims["auth_time"] = p.AuthTime
}
if p.ACR != "" {
claims["acr"] = p.ACR
}
if len(p.AMR) > 0 {
claims["amr"] = p.AMR
}
if p.SID != "" {
claims["sid"] = p.SID
}
if p.AccessToken != "" {
claims["at_hash"] = halfHash(key.alg, p.AccessToken)
}
for k, v := range s.userClaims(ctx, user, p.UserID, p.Scopes, p.Claims, "id_token") {
if _, reserved := claims[k]; !reserved {
claims[k] = v
}
}
return key.sign(claims, "")
}
// userClaims returns the claims of a user released by scopes plus those requested one by one in the
// OIDC "claims" parameter for destination ("id_token" or "userinfo").
func (s *OAuthServer) userClaims(ctx context.Context, user *UserContext, userID int, scopes []string, claimsParam map[string]any, destination string) map[string]any {
base := map[string]any{}
if user != nil {
if user.UserName != "" {
base["preferred_username"] = user.UserName
}
if user.Email != "" {
base["email"] = user.Email
}
}
requested := map[string]struct{}{}
if m, ok := claimsParam[destination].(map[string]any); ok {
for k := range m {
requested[k] = struct{}{}
}
}
all := base
if s.cfg.ClaimsProvider != nil {
extra, err := s.cfg.ClaimsProvider(ctx, OAuthClaimsRequest{
UserID: userID, Sub: strconv.Itoa(userID), Scopes: scopes, Requested: sortedKeys(requested),
Destination: destination, Base: base,
})
if err == nil {
all = map[string]any{}
for k, v := range base {
all[k] = v
}
for k, v := range extra {
all[k] = v
}
}
}
out := map[string]any{}
allowed := map[string]struct{}{}
for _, sc := range scopes {
for _, c := range scopeClaims[sc] {
allowed[c] = struct{}{}
}
}
for k, v := range all {
_, byScope := allowed[k]
_, asked := requested[k]
if byScope || asked {
out[k] = v
}
}
return out
}
// --------------------------------------------------------------------------
// Access token resolution
// --------------------------------------------------------------------------
// tokenInfo is an active access token.
type tokenInfo struct {
UserID int
Sub string
Username string
Email string
Roles []string
UserLevel int
Scopes []string
ClientID string
Exp, Iat int64
JKT string
Aud []string
JTI string
SID string
Claims map[string]any
JWT bool
Legacy bool // a session token that was not issued through a grant (no recorded scope)
}
func (t *tokenInfo) tokenType() string {
if t.JKT != "" {
return "DPoP"
}
return "Bearer"
}
// resolveAccessToken returns the active access token, or nil. Tokens are our JWT access tokens or
// opaque session tokens.
func (s *OAuthServer) resolveAccessToken(ctx context.Context, token string) *tokenInfo {
if token == "" {
return nil
}
a := s.anyAuth()
if a == nil {
return nil
}
gs := s.grants()
if strings.Count(token, ".") == 2 {
claims, err := s.parseOwnJWT(token, "at+jwt")
if err != nil {
return nil
}
jti, _ := claims["jti"].(string)
sub, _ := claims["sub"].(string)
uid, _ := strconv.Atoi(sub)
if jti == "" || uid == 0 || gs == nil {
return nil
}
rec, err := gs.PeekRefresh(ctx, accessKey(jti))
if err != nil {
return nil // revoked or expired
}
user, err := a.OAuthGetUser(ctx, uid)
if err != nil {
return nil
}
ti := &tokenInfo{UserID: uid, Sub: sub, Username: user.UserName, Email: user.Email, Roles: user.Roles,
UserLevel: user.UserLevel, Scopes: rec.Scopes, ClientID: rec.ClientID, JTI: jti, JWT: true, SID: rec.SessionToken}
ti.Exp = int64(numberOf(claims["exp"]))
ti.Iat = int64(numberOf(claims["iat"]))
ti.JKT, _ = rec.Extra["jkt"].(string)
ti.Aud = stringsOf(rec.Extra["aud"])
ti.Claims, _ = rec.Extra["claims"].(map[string]any)
if len(ti.Aud) == 0 {
ti.Aud = audOf(claims["aud"])
}
return ti
}
info, err := a.OAuthIntrospectToken(ctx, token)
if err != nil || !info.Active {
return nil
}
uid, _ := strconv.Atoi(info.Sub)
ti := &tokenInfo{UserID: uid, Sub: info.Sub, Username: info.Username, Email: info.Email, Roles: info.Roles,
UserLevel: info.UserLevel, Exp: info.Exp, Iat: info.Iat, Legacy: true}
if gs != nil {
if rec, err := gs.PeekRefresh(ctx, accessKey(token)); err == nil {
ti.Legacy = false
ti.Scopes, ti.ClientID, ti.SID = rec.Scopes, rec.ClientID, rec.SessionToken
ti.JKT, _ = rec.Extra["jkt"].(string)
ti.Aud = stringsOf(rec.Extra["aud"])
ti.Claims, _ = rec.Extra["claims"].(map[string]any)
} else if !errors.Is(err, lookup.ErrRefreshInvalid) {
return nil
}
}
return ti
}
func audOf(v any) []string {
switch a := v.(type) {
case string:
return []string{a}
case []any:
return stringsOf(a)
}
return nil
}
// parseOwnJWT verifies a JWT signed by one of this server's keys and returns its claims.
func (s *OAuthServer) parseOwnJWT(token, typ string) (jwt.MapClaims, error) {
return s.parseOwnJWTOpts(token, typ, false)
}
// parseOwnJWTOpts is parseOwnJWT; allowExpired accepts tokens past their exp (id_token_hint).
func (s *OAuthServer) parseOwnJWTOpts(token, typ string, allowExpired bool) (jwt.MapClaims, error) {
claims := jwt.MapClaims{}
opts := []jwt.ParserOption{jwt.WithValidMethods(s.keys.algs()), jwt.WithIssuer(s.cfg.Issuer)}
if allowExpired {
opts = append(opts, jwt.WithoutClaimsValidation())
} else {
opts = append(opts, jwt.WithExpirationRequired())
}
tok, err := jwt.NewParser(opts...).
ParseWithClaims(token, claims, func(t *jwt.Token) (any, error) {
if typ != "" && t.Header["typ"] != typ {
return nil, errors.New("wrong token type")
}
kid, _ := t.Header["kid"].(string)
pub, k := s.keys.publicFor(kid)
if k == nil {
return nil, errors.New("unknown key")
}
return pub, nil
})
if err != nil || !tok.Valid {
return nil, errors.New("invalid token")
}
if allowExpired && claims["iss"] != s.cfg.Issuer {
return nil, errors.New("invalid token")
}
return claims, nil
}
// --------------------------------------------------------------------------
// UserInfo — GET/POST /oauth/userinfo
// --------------------------------------------------------------------------
// bearerFromRequest extracts an access token and its scheme (Bearer or DPoP).
func bearerFromRequest(r *http.Request) (token, scheme string) {
h := r.Header.Get("Authorization")
for _, sch := range []string{"Bearer", "DPoP"} {
if len(h) > len(sch)+1 && strings.EqualFold(h[:len(sch)], sch) && h[len(sch)] == ' ' {
return strings.TrimSpace(h[len(sch)+1:]), sch
}
}
if r.Method == http.MethodPost {
if err := r.ParseForm(); err == nil {
if t := r.PostFormValue("access_token"); t != "" {
return t, "Bearer"
}
}
}
return "", ""
}
func (s *OAuthServer) userinfoHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
challenge := func(code, desc string) {
h := `Bearer error="` + code + `"`
if s.cfg.EnableDPoP {
h += `, DPoP algs="ES256 ES384 RS256 PS256"`
}
w.Header().Set("WWW-Authenticate", h)
writeOAuthError(w, code, desc, http.StatusUnauthorized)
}
token, scheme := bearerFromRequest(r)
if token == "" {
challenge("invalid_token", "missing bearer token")
return
}
ti := s.resolveAccessToken(r.Context(), token)
if ti == nil {
challenge("invalid_token", "token is inactive or invalid")
return
}
if e := s.checkDPoPBinding(r, token, scheme, ti); e != nil {
e.write(w)
return
}
scopes := ti.Scopes
if ti.Legacy {
scopes = []string{"openid", "profile", "email"} // a session token without a recorded grant
} else if !oauthSliceContains(scopes, "openid") {
w.Header().Set("WWW-Authenticate", `Bearer error="insufficient_scope", scope="openid"`)
writeOAuthError(w, "insufficient_scope", "the openid scope is required", http.StatusForbidden)
return
}
user := &UserContext{UserName: ti.Username, Email: ti.Email}
out := s.userClaims(r.Context(), user, ti.UserID, scopes, ti.Claims, "userinfo")
out["sub"] = ti.Sub
var client *OAuthServerClient
if ti.ClientID != "" {
client, _ = s.lookupOrFetchClient(r.Context(), ti.ClientID)
}
if client != nil && client.UserinfoSignedResponseAlg != "" && client.UserinfoSignedResponseAlg != "none" {
key := s.keys.forAlg(client.UserinfoSignedResponseAlg)
claims := jwt.MapClaims{"iss": s.cfg.Issuer, "aud": client.ClientID, "iat": time.Now().Unix()}
for k, v := range out {
claims[k] = v
}
signed, err := key.sign(claims, "")
if err != nil {
writeOAuthError(w, "server_error", "could not sign the response", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "application/jwt")
w.Header().Set("Cache-Control", "no-store")
w.Write([]byte(signed)) //nolint:errcheck,gosec // G104: best-effort write
return
}
writeJSON(w, http.StatusOK, out)
}
// checkDPoPBinding enforces sender-constraining: a DPoP-bound token needs a matching proof and the
// DPoP scheme, and a Bearer-scheme request must not carry a bound token.
func (s *OAuthServer) checkDPoPBinding(r *http.Request, token, scheme string, ti *tokenInfo) *oauthError {
if ti.JKT == "" {
if scheme == "DPoP" {
return &oauthError{Code: "invalid_token", Desc: "token is not DPoP bound", Status: http.StatusUnauthorized,
WWWAuth: `Bearer error="invalid_token"`}
}
return nil
}
fail := func(code, desc string) *oauthError {
return &oauthError{Code: code, Desc: desc, Status: http.StatusUnauthorized,
WWWAuth: `DPoP error="` + code + `", algs="ES256 ES384 RS256 PS256"`}
}
if scheme != "DPoP" {
return fail("invalid_token", "DPoP-bound tokens must use the DPoP scheme")
}
proof, e := s.verifyDPoP(r, token)
if e != nil {
return fail("invalid_dpop_proof", e.Desc)
}
if proof == nil || proof.JKT != ti.JKT {
return fail("invalid_dpop_proof", "proof key does not match the token")
}
return nil
}
// --------------------------------------------------------------------------
// Resource server helpers
// --------------------------------------------------------------------------
// AccessTokenClaims describes a verified access token.
type AccessTokenClaims struct {
Subject string
UserID int
ClientID string
Scopes []string
Audience []string
JTI string
SessionID string
ExpiresAt time.Time
// DPoPKey is the thumbprint of the key the token is bound to, or "".
DPoPKey string
// JWT is true for RFC 9068 JWT access tokens.
JWT bool
}
// VerifyAccessTokenOptions tunes VerifyAccessToken.
type VerifyAccessTokenOptions struct {
// Audience, when set, must be one of the token's audiences.
Audience string
// Scopes that must all be granted.
Scopes []string
}
// VerifyAccessToken validates an access token issued by this server (JWT or opaque) against the
// store, so revoked tokens are rejected.
func (s *OAuthServer) VerifyAccessToken(ctx context.Context, token string, opts VerifyAccessTokenOptions) (*AccessTokenClaims, error) {
ti := s.resolveAccessToken(ctx, token)
if ti == nil {
return nil, errors.New("invalid or inactive access token")
}
if opts.Audience != "" && !oauthSliceContains(ti.Aud, opts.Audience) {
return nil, errors.New("access token audience mismatch")
}
if !scopesCovered(ti.Scopes, opts.Scopes) {
return nil, errors.New("insufficient scope")
}
return &AccessTokenClaims{
Subject: ti.Sub, UserID: ti.UserID, ClientID: ti.ClientID, Scopes: ti.Scopes, Audience: ti.Aud, JTI: ti.JTI,
SessionID: ti.SID, ExpiresAt: time.Unix(ti.Exp, 0), DPoPKey: ti.JKT, JWT: ti.JWT,
}, nil
}
// introspectionInfo converts a resolved token to the RFC 7662 response.
func (s *OAuthServer) introspectionInfo(ti *tokenInfo) *OAuthTokenInfo {
info := &OAuthTokenInfo{
Active: true, Sub: ti.Sub, Username: ti.Username, Email: ti.Email, UserLevel: ti.UserLevel, Roles: ti.Roles,
Exp: ti.Exp, Iat: ti.Iat, Scope: strings.Join(ti.Scopes, " "), ClientID: ti.ClientID,
TokenType: ti.tokenType(), Iss: s.cfg.Issuer, Aud: ti.Aud, Jti: ti.JTI,
}
if ti.JKT != "" {
info.Cnf = map[string]any{"jkt": ti.JKT}
}
return info
}
func itoa(i int) string { return strconv.Itoa(i) }
func joinScopes(s []string) string { return strings.Join(s, " ") }
+85
View File
@@ -0,0 +1,85 @@
package security
import (
"context"
"net/http"
"net/url"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
const parURIPrefix = "urn:ietf:params:oauth:request_uri:"
// --------------------------------------------------------------------------
// RFC 9126 — Pushed authorization requests: POST /oauth/par
// --------------------------------------------------------------------------
func (s *OAuthServer) parHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
if err := r.ParseForm(); err != nil {
writeOAuthError(w, "invalid_request", "cannot parse form", http.StatusBadRequest)
return
}
ac, e := s.requireClient(r)
if e != nil {
e.write(w)
return
}
if r.PostForm.Get("request_uri") != "" {
writeOAuthError(w, "invalid_request", "request_uri must not be pushed", http.StatusBadRequest)
return
}
form := url.Values{}
for k, v := range r.PostForm {
switch k {
case "client_secret", "client_assertion", "client_assertion_type":
continue
}
form[k] = v
}
form.Set("client_id", ac.Client.ClientID)
if _, fail := s.parseAuthz(r.Context(), form, true); fail != nil {
writeOAuthError(w, fail.code, fail.desc, http.StatusBadRequest)
return
}
id, err := randomOAuthToken()
gs := s.grants()
if err != nil || gs == nil {
writeOAuthError(w, "server_error", "", http.StatusInternalServerError)
return
}
uri := parURIPrefix + id
if err := gs.SavePushedRequest(r.Context(), lookup.PushedRequest{
RequestURI: uri, ClientID: ac.Client.ClientID, Params: map[string]string{"q": form.Encode()},
ExpiresAt: time.Now().Add(s.cfg.PARTTL),
}); err != nil {
writeOAuthError(w, "server_error", "", http.StatusInternalServerError)
return
}
writeJSON(w, http.StatusCreated, map[string]any{"request_uri": uri, "expires_in": int(s.cfg.PARTTL.Seconds())})
}
// resolvePushedRequest consumes the pushed request behind request_uri (single use).
func (s *OAuthServer) resolvePushedRequest(ctx context.Context, clientID, uri string) (url.Values, *oauthError) {
gs := s.grants()
if gs == nil || (!s.cfg.EnablePAR && !s.cfg.RequirePAR) {
return nil, oerr("request_uri_not_supported", "pushed authorization requests are not enabled", http.StatusBadRequest)
}
pr, err := gs.ConsumePushedRequest(ctx, uri)
if err != nil {
return nil, oerr("invalid_request_uri", "request_uri is unknown, expired or already used", http.StatusBadRequest)
}
if clientID != "" && clientID != pr.ClientID {
return nil, oerr("invalid_request", "client_id does not match the pushed request", http.StatusBadRequest)
}
q, err := url.ParseQuery(pr.Params["q"])
if err != nil {
return nil, oerr("invalid_request_uri", "stored request is unreadable", http.StatusBadRequest)
}
q.Set("client_id", pr.ClientID)
return q, nil
}
+213
View File
@@ -0,0 +1,213 @@
package security
import (
"context"
"errors"
"net/http"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// newRefreshToken stores a new managed refresh token (the first of a family) and returns it.
func (s *OAuthServer) newRefreshToken(ctx context.Context, g *tokenGrant, now time.Time) (string, *oauthError) {
gs := s.grants()
if gs == nil {
return "", serverErr()
}
rt, err := randomOAuthToken()
if err != nil {
return "", serverErr()
}
family := g.FamilyID
if family == "" {
if family, err = randomOAuthToken(); err != nil {
return "", serverErr()
}
family = family[:30]
}
extra := map[string]any{}
if g.AuthTime != 0 {
extra["auth_time"] = g.AuthTime
}
if g.ACR != "" {
extra["acr"] = g.ACR
}
if len(g.AMR) > 0 {
extra["amr"] = g.AMR
}
if len(g.Resource) > 0 {
extra["aud"] = g.Resource
}
if g.DPoPJKT != "" {
extra["dpop_jkt"] = g.DPoPJKT
}
if len(g.Claims) > 0 {
extra["claims"] = g.Claims
}
if err := gs.SaveRefresh(ctx, lookup.RefreshToken{
TokenHash: hashToken(rt), FamilyID: family, ClientID: g.Client.ClientID, UserID: g.UserID,
SessionToken: g.SID, Scopes: g.Scopes, Extra: extra, ExpiresAt: now.Add(s.cfg.RefreshTokenTTL),
}); err != nil {
return "", serverErr()
}
return rt, nil
}
func (s *OAuthServer) handleRefreshGrant(r *http.Request) (map[string]any, *oauthError) {
refreshToken := r.PostFormValue("refresh_token")
if refreshToken == "" {
return nil, oerr("invalid_request", "refresh_token required", http.StatusBadRequest)
}
if s.cfg.ManagedRefreshTokens {
resp, e, handled := s.handleManagedRefresh(r, refreshToken)
if handled {
return resp, e
}
// Not one of ours: it may be a pass-through token issued before ManagedRefreshTokens was enabled.
}
return s.handleLegacyRefresh(r, refreshToken)
}
// handleManagedRefresh rotates a server-issued refresh token. handled is false when the token is
// not known to the store.
func (s *OAuthServer) handleManagedRefresh(r *http.Request, refreshToken string) (map[string]any, *oauthError, bool) {
ctx := r.Context()
gs := s.grants()
if gs == nil {
return nil, nil, false
}
hash := hashToken(refreshToken)
peek, err := gs.PeekRefresh(ctx, hash)
if errors.Is(err, lookup.ErrRefreshInvalid) || (err == nil && isAccessRecord(peek)) {
return nil, nil, false
}
if err != nil {
return nil, serverErr(), true
}
ac, e := s.requireClient(r)
if e != nil {
return nil, e, true
}
client := ac.Client
if peek.ClientID != client.ClientID {
return nil, oerr("invalid_grant", "refresh token was issued to another client", http.StatusBadRequest), true
}
if !grantAllowed(client, "refresh_token") && !oauthSliceContains(peek.Scopes, "offline_access") {
return nil, oerr("unauthorized_client", "client may not use refresh_token", http.StatusBadRequest), true
}
scopes := peek.Scopes
if req := strings.Fields(r.PostFormValue("scope")); len(req) > 0 {
if !scopesCovered(peek.Scopes, req) {
return nil, oerr("invalid_scope", "scope exceeds the original grant", http.StatusBadRequest), true
}
scopes = req
}
bound, _ := peek.Extra["dpop_jkt"].(string)
proof, e := s.verifyDPoP(r, "")
if e != nil {
return nil, e, true
}
jkt := bound
switch {
case bound != "" && (proof == nil || proof.JKT != bound):
return nil, oerr("invalid_dpop_proof", "the refresh token is bound to another DPoP key", http.StatusBadRequest), true
case proof != nil:
jkt = proof.JKT
case client.DPoPBoundAccessTokens:
return nil, oerr("invalid_dpop_proof", "this client must present a DPoP proof", http.StatusBadRequest), true
}
next, err := randomOAuthToken()
if err != nil {
return nil, serverErr(), true
}
extra := map[string]any{}
for k, v := range peek.Extra {
extra[k] = v
}
if jkt != "" && bound == "" && proof != nil && ac.Method == "none" {
extra["dpop_jkt"] = jkt // a public client's first DPoP refresh binds the token
}
old, err := gs.RotateRefresh(ctx, hash, lookup.RefreshToken{
TokenHash: hashToken(next), Scopes: peek.Scopes, Extra: extra, ExpiresAt: peek.ExpiresAt,
})
switch {
case errors.Is(err, lookup.ErrRefreshReused):
return nil, oerr("invalid_grant", "refresh token reuse detected; the session was ended", http.StatusBadRequest), true
case errors.Is(err, lookup.ErrRefreshInvalid):
return nil, oerr("invalid_grant", "refresh token expired or revoked", http.StatusBadRequest), true
case err != nil:
return nil, serverErr(), true
}
g := &tokenGrant{
Client: client, UserID: old.UserID, Scopes: scopes, SID: old.SessionToken, DPoPJKT: jkt,
NextRefresh: next, IDToken: true,
}
g.AuthTime = int64(numberOf(old.Extra["auth_time"]))
g.ACR, _ = old.Extra["acr"].(string)
g.AMR = stringsOf(old.Extra["amr"])
g.Resource = stringsOf(old.Extra["aud"])
if cl, ok := old.Extra["claims"].(map[string]any); ok {
g.Claims = cl
}
resp, e := s.mintTokens(ctx, g)
return resp, e, true
}
func numberOf(v any) float64 {
f, _ := v.(float64)
return f
}
func stringsOf(v any) []string {
list, _ := v.([]any)
var out []string
for _, it := range list {
if str, ok := it.(string); ok {
out = append(out, str)
}
}
return out
}
// handleLegacyRefresh passes the token through to the authenticators (the behaviour of earlier versions).
func (s *OAuthServer) handleLegacyRefresh(r *http.Request, refreshToken string) (map[string]any, *oauthError) {
providerName := r.PostFormValue("provider")
clientID := r.PostFormValue("client_id")
ac, e := s.authenticateClient(r)
if e != nil {
return nil, e
}
var client *OAuthServerClient
if ac != nil {
client = ac.Client
} else if c, ok := s.lookupOrFetchClient(r.Context(), clientID); ok {
client = c
}
if client == nil {
client = &OAuthServerClient{ClientID: clientID}
}
var resp *LoginResponse
var err error
if provider := s.providerByName(providerName); provider != nil {
resp, err = provider.auth.OAuth2RefreshToken(r.Context(), refreshToken, providerName)
} else if s.auth != nil {
resp, err = s.auth.RefreshToken(r.Context(), refreshToken)
} else {
return nil, oerr("invalid_grant", "no provider available for refresh", http.StatusBadRequest)
}
if err != nil {
return nil, oerr("invalid_grant", err.Error(), http.StatusBadRequest)
}
return s.mintTokens(r.Context(), &tokenGrant{
Client: client, LegacyAccess: resp.Token, LegacyRefresh: resp.RefreshToken,
})
}
+352
View File
@@ -0,0 +1,352 @@
package security
import (
"crypto/subtle"
"encoding/json"
"net/http"
"net/url"
"strings"
"time"
)
var supportedGrantTypes = []string{"authorization_code", "refresh_token", "client_credentials", grantDeviceCode, grantTokenExchange}
// validateRedirectURI checks a redirect (or post-logout) URI at registration: absolute, no
// fragment, https, http only for loopback, or a private-use scheme (RFC 8252).
func validateRedirectURI(raw string) bool {
u, err := url.Parse(raw)
if err != nil || !u.IsAbs() || u.Fragment != "" || strings.Contains(raw, "#") {
return false
}
switch strings.ToLower(u.Scheme) {
case "https":
return u.Host != ""
case "http":
return isLoopbackHost(u.Hostname())
case "javascript", "data", "file", "vbscript", "about", "blob", "ftp":
return false
}
return true
}
func (s *OAuthServer) validateClientMetadata(c *OAuthServerClient) *oauthError {
meta := func(desc string) *oauthError { return oerr("invalid_client_metadata", desc, http.StatusBadRequest) }
for _, g := range c.GrantTypes {
if !oauthSliceContains(supportedGrantTypes, g) {
return meta("unsupported grant type " + g)
}
}
if oauthSliceContains(c.GrantTypes, grantDeviceCode) && !s.cfg.EnableDeviceFlow {
return meta("the device grant is not enabled on this server")
}
if oauthSliceContains(c.GrantTypes, grantTokenExchange) && !s.cfg.EnableTokenExchange {
return meta("token exchange is not enabled on this server")
}
for _, rt := range c.ResponseTypes {
if rt != "code" {
return meta("unsupported response type " + rt)
}
}
needsRedirect := false
for _, g := range c.GrantTypes {
if g == "authorization_code" {
needsRedirect = true
}
}
if needsRedirect && len(c.RedirectURIs) == 0 {
return oerr("invalid_redirect_uri", "redirect_uris required", http.StatusBadRequest)
}
for _, u := range c.RedirectURIs {
if !validateRedirectURI(u) {
return oerr("invalid_redirect_uri", "invalid redirect_uri "+u, http.StatusBadRequest)
}
}
for _, u := range c.PostLogoutRedirectURIs {
if !validateRedirectURI(u) {
return meta("invalid post_logout_redirect_uri " + u)
}
}
if !oauthSliceContains(s.authMethodsSupported(), c.TokenEndpointAuthMethod) {
return meta("unsupported token_endpoint_auth_method")
}
if c.TokenEndpointAuthMethod == "private_key_jwt" {
if (len(c.JWKS) == 0) == (c.JWKSURI == "") {
return meta("private_key_jwt needs exactly one of jwks and jwks_uri")
}
if len(c.JWKS) > 0 {
if _, err := parseJWKS(c.JWKS); err != nil {
return meta("jwks: " + err.Error())
}
}
}
if c.JWKSURI != "" && !validWebURL(c.JWKSURI, s.cfg.AllowPrivateNetworkFetch) {
return meta("jwks_uri must be an https URL")
}
if a := c.TokenEndpointAuthSigningAlg; a != "" && !oauthSliceContains([]string{"RS256", "PS256", "ES256", "ES384"}, a) {
return meta("unsupported token_endpoint_auth_signing_alg")
}
if a := c.IDTokenSignedResponseAlg; a != "" && !oauthSliceContains(s.keys.algs(), a) {
return meta("id_token_signed_response_alg is not offered by this server")
}
if a := c.UserinfoSignedResponseAlg; a != "" && a != "none" && !oauthSliceContains(s.keys.algs(), a) {
return meta("userinfo_signed_response_alg is not offered by this server")
}
if c.BackchannelLogoutURI != "" && !validWebURL(c.BackchannelLogoutURI, s.cfg.AllowPrivateNetworkFetch) {
return meta("backchannel_logout_uri must be an https URL")
}
for _, u := range []string{c.ClientURI, c.LogoURI} {
if u != "" && !validWebURL(u, true) {
return meta("client_uri and logo_uri must be http(s) URLs")
}
}
if len(c.ClientName) > 200 {
return meta("client_name too long")
}
return nil
}
// validWebURL accepts https URLs, and http URLs when loose is set or the host is loopback.
func validWebURL(raw string, loose bool) bool {
u, err := url.Parse(raw)
if err != nil || u.Host == "" || u.Fragment != "" {
return false
}
return u.Scheme == "https" || (u.Scheme == "http" && (loose || isLoopbackHost(u.Hostname())))
}
// clientView is the RFC 7591 representation of a client: no secret hash, no registration token hash.
func clientView(c *OAuthServerClient) map[string]any {
v := *c
v.ClientSecretHash, v.RegistrationAccessTokenHash = "", ""
raw, _ := json.Marshal(v)
m := map[string]any{}
_ = json.Unmarshal(raw, &m)
m["client_id"] = c.ClientID
m["token_endpoint_auth_method"] = c.TokenEndpointAuthMethod
m["grant_types"] = c.GrantTypes
m["redirect_uris"] = c.RedirectURIs
if len(c.AllowedScopes) > 0 {
m["scope"] = strings.Join(c.AllowedScopes, " ")
}
return m
}
// --------------------------------------------------------------------------
// RFC 7591 — Dynamic client registration
// --------------------------------------------------------------------------
func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
if s.cfg.InitialAccessToken != "" {
tok, _ := bearerFromHeader(r)
if subtle.ConstantTimeCompare([]byte(tok), []byte(s.cfg.InitialAccessToken)) != 1 {
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`)
writeOAuthError(w, "invalid_token", "an initial access token is required to register clients", http.StatusUnauthorized)
return
}
}
c, e := s.decodeClientMetadata(r)
if e != nil {
e.write(w)
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.
if oauthSliceContains(c.GrantTypes, "client_credentials") && c.TokenEndpointAuthMethod == "none" {
c.TokenEndpointAuthMethod = "client_secret_basic"
}
if e := s.validateClientMetadata(c); e != nil {
e.write(w)
return
}
id, err := randomOAuthToken()
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
c.ClientID = id
c.ClientIDIssuedAt = time.Now().Unix()
var secret string
if strings.HasPrefix(c.TokenEndpointAuthMethod, "client_secret_") {
if secret, err = randomOAuthToken(); err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
c.ClientSecretHash = hashClientSecret(secret)
}
regToken, err := randomOAuthToken()
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
c.RegistrationAccessTokenHash = hashToken(regToken)
if err := s.saveClient(r.Context(), c, false); err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
// RFC 7591 registration response: the plaintext secret and registration token are returned
// exactly once here and never persisted or served again — only their hashes are stored.
resp := clientView(c)
if secret != "" {
resp["client_secret"] = secret
resp["client_secret_expires_at"] = 0
}
resp["registration_access_token"] = regToken
resp["registration_client_uri"] = s.endpoint("/oauth/register/" + c.ClientID)
writeJSON(w, http.StatusCreated, resp)
}
// decodeClientMetadata reads the registration document and strips every field the caller must not
// control.
func (s *OAuthServer) decodeClientMetadata(r *http.Request) (*OAuthServerClient, *oauthError) {
var req struct {
OAuthServerClient
Scope string `json:"scope"`
}
if err := json.NewDecoder(http.MaxBytesReader(nil, r.Body, 64<<10)).Decode(&req); err != nil {
return nil, oerr("invalid_client_metadata", "malformed JSON", http.StatusBadRequest)
}
c := req.OAuthServerClient
c.ClientID, c.ClientSecretHash, c.RegistrationAccessTokenHash = "", "", ""
c.FirstParty, c.ClientSecretExpiresAt, c.ClientIDIssuedAt = false, 0, 0
if len(c.GrantTypes) == 0 {
c.GrantTypes = []string{"authorization_code", "refresh_token"}
}
if len(c.ResponseTypes) == 0 {
c.ResponseTypes = []string{"code"}
}
if len(c.AllowedScopes) == 0 {
if f := strings.Fields(req.Scope); len(f) > 0 {
c.AllowedScopes = f
} else {
c.AllowedScopes = s.cfg.DefaultScopes
}
}
if c.TokenEndpointAuthMethod == "" {
c.TokenEndpointAuthMethod = "none"
}
return &c, nil
}
func bearerFromHeader(r *http.Request) (string, bool) {
h := r.Header.Get("Authorization")
if len(h) > 7 && strings.EqualFold(h[:7], "Bearer ") {
return strings.TrimSpace(h[7:]), true
}
return "", false
}
// registrationClient authenticates a RFC 7592 management request and returns the client.
func (s *OAuthServer) registrationClient(w http.ResponseWriter, r *http.Request) *OAuthServerClient {
deny := func() *OAuthServerClient {
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`)
writeOAuthError(w, "invalid_token", "invalid registration access token", http.StatusUnauthorized)
return nil
}
c, ok := s.lookupOrFetchClient(r.Context(), r.PathValue("id"))
tok, _ := bearerFromHeader(r)
if !ok || c.RegistrationAccessTokenHash == "" || tok == "" ||
subtle.ConstantTimeCompare([]byte(hashToken(tok)), []byte(c.RegistrationAccessTokenHash)) != 1 {
return deny()
}
return c
}
// registrationManageHandler serves GET, PUT and DELETE /oauth/register/{id} (RFC 7592).
func (s *OAuthServer) registrationManageHandler(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodGet, http.MethodPut, http.MethodDelete:
default:
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
cur := s.registrationClient(w, r)
if cur == nil {
return
}
switch r.Method {
case http.MethodGet:
resp := clientView(cur)
resp["registration_client_uri"] = s.endpoint("/oauth/register/" + cur.ClientID)
writeJSON(w, http.StatusOK, resp)
case http.MethodDelete:
if err := s.removeClient(r.Context(), cur.ClientID); err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
w.WriteHeader(http.StatusNoContent)
case http.MethodPut:
next, e := s.decodeClientMetadata(r)
if e != nil {
e.write(w)
return
}
if oauthSliceContains(next.GrantTypes, "client_credentials") && next.TokenEndpointAuthMethod == "none" {
next.TokenEndpointAuthMethod = "client_secret_basic"
}
if e := s.validateClientMetadata(next); e != nil {
e.write(w)
return
}
wantsSecret := strings.HasPrefix(next.TokenEndpointAuthMethod, "client_secret_")
if wantsSecret && cur.ClientSecretHash == "" {
oerr("invalid_client_metadata", "the client has no secret; register a new client instead", http.StatusBadRequest).write(w)
return
}
next.ClientID = cur.ClientID
next.ClientIDIssuedAt = cur.ClientIDIssuedAt
next.RegistrationAccessTokenHash = cur.RegistrationAccessTokenHash
next.FirstParty = cur.FirstParty // only trusted code may change this
next.ClientSecretExpiresAt = cur.ClientSecretExpiresAt
if wantsSecret {
next.ClientSecretHash = cur.ClientSecretHash
}
if err := s.saveClient(r.Context(), next, true); err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
resp := clientView(next)
resp["registration_client_uri"] = s.endpoint("/oauth/register/" + next.ClientID)
writeJSON(w, http.StatusOK, resp)
}
}
// registrationRotateHandler issues a new client secret: POST /oauth/register/{id}/rotate-secret.
func (s *OAuthServer) registrationRotateHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
cur := s.registrationClient(w, r)
if cur == nil {
return
}
if cur.ClientSecretHash == "" {
writeOAuthError(w, "invalid_request", "this client has no secret", http.StatusBadRequest)
return
}
secret, err := randomOAuthToken()
if err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
next := *cur
next.ClientSecretHash = hashClientSecret(secret)
if err := s.saveClient(r.Context(), &next, true); err != nil {
http.Error(w, "server error", http.StatusInternalServerError)
return
}
resp := clientView(&next)
resp["client_secret"] = secret
resp["client_secret_expires_at"] = 0
writeJSON(w, http.StatusOK, resp)
}
+460 -1049
View File
File diff suppressed because it is too large Load Diff
+23
View File
@@ -2,6 +2,8 @@ package security
import (
"context"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// OAuthRegisterClient persists an OAuth2 client registration.
@@ -33,3 +35,24 @@ func (a *DatabaseAuthenticator) OAuthIntrospectToken(ctx context.Context, token
func (a *DatabaseAuthenticator) OAuthRevokeToken(ctx context.Context, token string) error {
return a.src.get().OAuthClient.Revoke(ctx, token)
}
// OAuthGrants returns the store holding consents, managed refresh tokens, device codes,
// pushed authorization requests and the replay cache.
func (a *DatabaseAuthenticator) OAuthGrants() lookup.OAuthGrantStore {
return a.src.get().OAuthGrant
}
// OAuthUpdateClient replaces the registered metadata of a client (RFC 7592).
func (a *DatabaseAuthenticator) OAuthUpdateClient(ctx context.Context, client *OAuthServerClient) error {
return a.src.get().OAuthClient.UpdateClient(ctx, client)
}
// OAuthDeleteClient deactivates a registered client (RFC 7592).
func (a *DatabaseAuthenticator) OAuthDeleteClient(ctx context.Context, clientID string) error {
return a.src.get().OAuthClient.DeleteClient(ctx, clientID)
}
// OAuthGetUser returns the active user with the given id.
func (a *DatabaseAuthenticator) OAuthGetUser(ctx context.Context, userID int) (*UserContext, error) {
return a.src.get().OAuthUser.GetUser(ctx, userID)
}
+154
View File
@@ -0,0 +1,154 @@
package security
import (
"crypto/ecdsa"
"crypto/hmac"
"crypto/rsa"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
"strings"
"time"
)
// deriveOAuthSecret derives the HMAC key for cookies and form state from the default signing key.
func deriveOAuthSecret(kr *oauthKeyring) []byte {
h := sha256.New()
h.Write([]byte("resolvespec-oauth-state-v1"))
switch k := kr.keys[0].signer.(type) {
case *rsa.PrivateKey:
h.Write(k.D.Bytes())
case *ecdsa.PrivateKey:
h.Write(k.D.Bytes())
}
return h.Sum(nil)
}
type sealed struct {
Exp int64 `json:"e"`
Kind string `json:"k"`
V json.RawMessage `json:"v"`
}
func (s *OAuthServer) mac(data []byte) []byte {
m := hmac.New(sha256.New, s.secret)
m.Write(data)
return m.Sum(nil)
}
// seal returns v as a tamper-proof string valid for ttl. kind separates the uses (a sealed login
// form cannot be replayed as a cookie).
func (s *OAuthServer) seal(kind string, v any, ttl time.Duration) (string, error) {
raw, err := json.Marshal(v)
if err != nil {
return "", err
}
body, err := json.Marshal(sealed{Exp: time.Now().Add(ttl).Unix(), Kind: kind, V: raw})
if err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(body) + "." + base64.RawURLEncoding.EncodeToString(s.mac(body)), nil
}
// open verifies and decodes a value produced by seal.
func (s *OAuthServer) open(kind, token string, v any) error {
i := strings.IndexByte(token, '.')
if i < 0 {
return fmt.Errorf("malformed")
}
body, err := base64.RawURLEncoding.DecodeString(token[:i])
if err != nil {
return fmt.Errorf("malformed")
}
sig, err := base64.RawURLEncoding.DecodeString(token[i+1:])
if err != nil || !hmac.Equal(sig, s.mac(body)) {
return fmt.Errorf("bad signature")
}
var env sealed
if err := json.Unmarshal(body, &env); err != nil || env.Kind != kind {
return fmt.Errorf("malformed")
}
if time.Now().Unix() > env.Exp {
return fmt.Errorf("expired")
}
return json.Unmarshal(env.V, v)
}
// ssoSession is the content of the SSO cookie. The login session token never reaches a client.
type ssoSession struct {
Token string `json:"t"` // login session token (user_sessions row)
UserID int `json:"u"`
AuthTime int64 `json:"a"`
SID string `json:"s"` // OIDC session id
Provider string `json:"p,omitempty"` // external provider that authenticated the user
Clients []string `json:"c,omitempty"` // clients that received tokens (back-channel logout)
}
func (s *OAuthServer) cookieSecure() bool {
return !s.cfg.SSOCookie.Insecure && s.issuerURL.Scheme == "https"
}
// ssoFromRequest returns the live SSO session of the request, or nil.
func (s *OAuthServer) ssoFromRequest(r *http.Request) *ssoSession {
if s.cfg.SSOCookie.Disable {
return nil
}
c, err := r.Cookie(s.cfg.SSOCookie.Name)
if err != nil {
return nil
}
var sess ssoSession
if err := s.open("sso", c.Value, &sess); err != nil {
return nil
}
a := s.anyAuth()
if a == nil {
return nil
}
if info, err := a.OAuthIntrospectToken(r.Context(), sess.Token); err != nil || !info.Active {
return nil
}
return &sess
}
func (s *OAuthServer) setSSO(w http.ResponseWriter, sess *ssoSession) {
if s.cfg.SSOCookie.Disable {
return
}
v, err := s.seal("sso", sess, s.cfg.SSOCookie.TTL)
if err != nil {
return
}
http.SetCookie(w, &http.Cookie{ //nolint:gosec // Secure follows the issuer scheme (cookieSecure)
Name: s.cfg.SSOCookie.Name, Value: v, Path: s.cfg.SSOCookie.Path,
MaxAge: int(s.cfg.SSOCookie.TTL.Seconds()), HttpOnly: true, Secure: s.cookieSecure(),
SameSite: s.cfg.SSOCookie.SameSite,
})
}
func (s *OAuthServer) clearSSO(w http.ResponseWriter) {
http.SetCookie(w, &http.Cookie{ //nolint:gosec // Secure follows the issuer scheme (cookieSecure)
Name: s.cfg.SSOCookie.Name, Value: "", Path: s.cfg.SSOCookie.Path, MaxAge: -1,
HttpOnly: true, Secure: s.cookieSecure(), SameSite: s.cfg.SSOCookie.SameSite,
})
}
// newSSO builds the session for a freshly authenticated user.
func (s *OAuthServer) newSSO(token string, userID int, provider string) (*ssoSession, error) {
sid, err := randomOAuthToken()
if err != nil {
return nil, err
}
return &ssoSession{Token: token, UserID: userID, AuthTime: time.Now().Unix(), SID: sid[:22], Provider: provider}, nil
}
// noteClient records that clientID received tokens in this SSO session.
func (s *OAuthServer) noteClient(w http.ResponseWriter, sess *ssoSession, clientID string) {
if sess == nil || oauthSliceContains(sess.Clients, clientID) || len(sess.Clients) >= 12 {
return
}
sess.Clients = append(sess.Clients, clientID)
s.setSSO(w, sess)
}
+144
View File
@@ -0,0 +1,144 @@
package security
import (
"bytes"
"html/template"
"net/http"
)
// OAuthLoginPage is the data of the login page template.
type OAuthLoginPage struct {
Title string
Error string
Action string // form action
State string // opaque, must be posted back as the "req" field
ClientName string
LoginHint string
// Extra hidden fields to post back (device flow).
Hidden map[string]string
}
// OAuthScopeInfo is one line of the consent screen.
type OAuthScopeInfo struct {
Name string
Description string
}
// OAuthConsentPage is the data of the consent page template.
type OAuthConsentPage struct {
Title string
Action string
State string // opaque, must be posted back as the "req" field
ClientName string
ClientURI string
LogoURI string
Scopes []OAuthScopeInfo
User string
Hidden map[string]string
}
type oauthMessagePage struct {
Title string
Message string
Error bool
}
type oauthDevicePage struct {
Title string
Action string
Error string
UserCode string
}
type oauthLogoutPage struct {
Title string
Action string
Hidden map[string]string
}
type oauthTemplates struct {
login, consent *template.Template
base *template.Template
}
const oauthPageCSS = `body{font-family:system-ui,sans-serif;display:flex;justify-content:center;align-items:center;min-height:100vh;margin:0;background:#f5f5f5}
.card{background:#fff;padding:2rem;border-radius:8px;box-shadow:0 2px 8px rgba(0,0,0,.15);width:340px;max-width:92vw}
h2{margin:0 0 1.25rem;font-size:1.25rem}p{color:#444;font-size:.9rem}
label{display:block;margin-bottom:.25rem;font-size:.875rem;color:#555}
input[type=text],input[type=password]{width:100%;box-sizing:border-box;padding:.5rem;border:1px solid #ccc;border-radius:4px;margin-bottom:1rem;font-size:1rem}
button{padding:.6rem 1rem;background:#0070f3;color:#fff;border:none;border-radius:4px;font-size:1rem;cursor:pointer}
button.secondary{background:#e5e5e5;color:#222}button:hover{opacity:.9}.full{width:100%}
.err{color:#d32f2f;margin-bottom:1rem;font-size:.875rem}ul{padding-left:1.2rem}li{margin:.35rem 0;font-size:.9rem}
.row{display:flex;gap:.5rem}.row button{flex:1}.logo{max-height:48px;margin-bottom:1rem}.muted{color:#777;font-size:.8rem}`
const oauthPageTemplates = `
{{define "head"}}<!DOCTYPE html><html><head><meta charset="utf-8"><meta name="viewport" content="width=device-width,initial-scale=1"><title>{{.Title}}</title><style>` + oauthPageCSS + `</style></head><body><div class="card">{{end}}
{{define "foot"}}</div></body></html>{{end}}
{{define "hidden"}}{{range $k, $v := .Hidden}}<input type="hidden" name="{{$k}}" value="{{$v}}">{{end}}{{end}}
{{define "login"}}{{template "head" .}}<h2>{{.Title}}</h2>{{if .ClientName}}<p>to continue to <b>{{.ClientName}}</b></p>{{end}}
{{if .Error}}<div class="err">{{.Error}}</div>{{end}}
<form method="POST" action="{{.Action}}">
<input type="hidden" name="req" value="{{.State}}">{{template "hidden" .}}
<label>Username</label><input type="text" name="username" value="{{.LoginHint}}" autofocus autocomplete="username">
<label>Password</label><input type="password" name="password" autocomplete="current-password">
<button class="full" type="submit">Sign in</button>
</form>{{template "foot"}}{{end}}
{{define "consent"}}{{template "head" .}}{{if .LogoURI}}<img class="logo" src="{{.LogoURI}}" alt="">{{end}}
<h2>{{if .ClientURI}}<a href="{{.ClientURI}}" rel="noopener noreferrer">{{.ClientName}}</a>{{else}}{{.ClientName}}{{end}} wants access</h2>
<p>Signed in as <b>{{.User}}</b>. This application will be able to:</p>
<ul>{{range .Scopes}}<li><b>{{.Name}}</b>{{if .Description}} – {{.Description}}{{end}}</li>{{end}}</ul>
<form method="POST" action="{{.Action}}">
<input type="hidden" name="req" value="{{.State}}">{{template "hidden" .}}
<p class="muted"><label><input type="checkbox" name="remember" value="1" checked> Remember this decision</label></p>
<div class="row"><button class="secondary" type="submit" name="decision" value="deny">Deny</button>
<button type="submit" name="decision" value="allow">Allow</button></div>
</form>{{template "foot"}}{{end}}
{{define "device"}}{{template "head" .}}<h2>{{.Title}}</h2><p>Enter the code shown on your device.</p>
{{if .Error}}<div class="err">{{.Error}}</div>{{end}}
<form method="POST" action="{{.Action}}"><input type="hidden" name="step" value="code">
<input type="text" name="user_code" value="{{.UserCode}}" autofocus autocomplete="off" placeholder="XXXX-XXXX">
<button class="full" type="submit">Continue</button></form>{{template "foot"}}{{end}}
{{define "message"}}{{template "head" .}}<h2>{{.Title}}</h2>{{if .Error}}<div class="err">{{.Message}}</div>{{else}}<p>{{.Message}}</p>{{end}}{{template "foot"}}{{end}}
{{define "logout"}}{{template "head" .}}<h2>{{.Title}}</h2><p>Do you want to sign out?</p>
<form method="POST" action="{{.Action}}">{{template "hidden" .}}
<div class="row"><button class="secondary" type="submit" name="confirm" value="no">Stay signed in</button>
<button type="submit" name="confirm" value="yes">Sign out</button></div></form>{{template "foot"}}{{end}}
`
func newOAuthTemplates(cfg *OAuthServerConfig) oauthTemplates {
base := template.Must(template.New("oauth").Parse(oauthPageTemplates))
return oauthTemplates{base: base, login: cfg.LoginTemplate, consent: cfg.ConsentTemplate}
}
func (s *OAuthServer) renderHTML(w http.ResponseWriter, status int, override *template.Template, name string, data any) {
var buf bytes.Buffer
var err error
if override != nil {
err = override.Execute(&buf, data)
} else {
err = s.tmpl.base.ExecuteTemplate(&buf, name, data)
}
if err != nil {
http.Error(w, "template error", http.StatusInternalServerError)
return
}
h := w.Header()
h.Set("Content-Type", "text/html; charset=utf-8")
h.Set("Cache-Control", "no-store")
h.Set("Pragma", "no-cache")
h.Set("X-Frame-Options", "DENY")
h.Set("X-Content-Type-Options", "nosniff")
h.Set("Referrer-Policy", "no-referrer")
h.Set("Content-Security-Policy", "default-src 'none'; style-src 'unsafe-inline'; img-src https: data:; frame-ancestors 'none'; base-uri 'none'")
w.WriteHeader(status)
w.Write(buf.Bytes()) //nolint:errcheck,gosec // G104: best-effort write, G705: html/template output is escaped
}
func (s *OAuthServer) renderMessage(w http.ResponseWriter, status int, title, msg string, isErr bool) {
s.renderHTML(w, status, nil, "message", oauthMessagePage{Title: title, Message: msg, Error: isErr})
}
+375
View File
@@ -0,0 +1,375 @@
package security
import (
"context"
"net/http"
"strconv"
"strings"
"time"
"github.com/golang-jwt/jwt/v5"
"golang.org/x/oauth2"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
const (
grantDeviceCode = "urn:ietf:params:oauth:grant-type:device_code"
grantTokenExchange = "urn:ietf:params:oauth:grant-type:token-exchange" //nolint:gosec // RFC 8693 URN, not a credential
// accessGrantFamily marks oauth_refresh_tokens rows that record an issued access token
// (scope, client, DPoP binding). They share the table so logout revokes both.
accessGrantPrefix = "at_"
)
// isAccessRecord reports whether a stored token is the record of an access token.
func isAccessRecord(t *lookup.RefreshToken) bool {
return strings.HasPrefix(t.FamilyID, accessGrantPrefix)
}
// accessKey is the token_hash of the record of an access token (keyed by the session token or the
// JWT's jti). The prefix is not hex, so it can never equal the hash of a refresh token.
func accessKey(id string) string { return accessGrantPrefix + hashToken(id)[:61] }
// tokenGrant describes the tokens to issue.
type tokenGrant struct {
Client *OAuthServerClient
UserID int
Scopes []string
Resource []string
Nonce string
AuthTime int64
ACR string
AMR []string
SID string
Claims map[string]any
DPoPJKT string
// LegacyAccess is an existing session token that is used as the access token (codes saved
// by earlier versions).
LegacyAccess string
LegacyRefresh string // refresh token passed through from the DatabaseAuthenticator
IssueRefresh bool // issue a managed refresh token (new family)
FamilyID string // family of the new refresh token
NextRefresh string // already rotated refresh token to return
FamilyExp time.Time // absolute end of the family when rotating
IDToken bool
}
// mintTokens issues the tokens of g and returns the token endpoint response.
func (s *OAuthServer) mintTokens(ctx context.Context, g *tokenGrant) (map[string]any, *oauthError) {
a := s.anyAuth()
if a == nil {
return nil, oerr("server_error", "no authenticator configured", http.StatusInternalServerError)
}
if g.LegacyAccess != "" && g.UserID == 0 {
info, err := a.OAuthIntrospectToken(ctx, g.LegacyAccess)
if err != nil || !info.Active {
return nil, oerr("invalid_grant", "the session of this code has ended", http.StatusBadRequest)
}
g.UserID, _ = strconv.Atoi(info.Sub)
}
now := time.Now()
exp := now.Add(s.cfg.AccessTokenTTL)
sub := strconv.Itoa(g.UserID)
tokenType := "Bearer"
if g.DPoPJKT != "" {
tokenType = "DPoP"
}
var access, recordID string
switch {
case s.cfg.JWTAccessTokens:
jti, err := randomOAuthToken()
if err != nil {
return nil, serverErr()
}
aud := jwt.ClaimStrings{s.cfg.AccessTokenAudience}
if len(g.Resource) > 0 {
aud = g.Resource
}
claims := jwt.MapClaims{
"iss": s.cfg.Issuer, "sub": sub, "aud": aud, "exp": exp.Unix(), "iat": now.Unix(),
"nbf": now.Unix(), "jti": jti, "client_id": g.Client.ClientID,
}
if len(g.Scopes) > 0 {
claims["scope"] = strings.Join(g.Scopes, " ")
}
if g.AuthTime != 0 {
claims["auth_time"] = g.AuthTime
}
if g.SID != "" {
claims["sid"] = g.SID
}
if g.DPoPJKT != "" {
claims["cnf"] = map[string]any{"jkt": g.DPoPJKT}
}
var err2 error
if access, err2 = s.keys.forAlg("").sign(claims, "at+jwt"); err2 != nil {
return nil, serverErr()
}
recordID = jti
case g.LegacyAccess != "":
access, recordID = g.LegacyAccess, g.LegacyAccess
default:
tok, err := randomOAuthToken()
if err != nil {
return nil, serverErr()
}
if err := a.oauth2CreateSession(ctx, tok, g.UserID, &oauth2.Token{AccessToken: tok, TokenType: "Bearer"}, exp, "oauth2_server"); err != nil {
return nil, serverErr()
}
access, recordID = tok, tok
}
if gs := s.grants(); gs != nil {
extra := map[string]any{}
if g.DPoPJKT != "" {
extra["jkt"] = g.DPoPJKT
}
if len(g.Resource) > 0 {
extra["aud"] = g.Resource
}
if len(g.Claims) > 0 {
extra["claims"] = g.Claims
}
if err := gs.SaveRefresh(ctx, lookup.RefreshToken{
TokenHash: accessKey(recordID), FamilyID: accessKey(recordID), ClientID: g.Client.ClientID, UserID: g.UserID,
SessionToken: g.SID, Scopes: g.Scopes, Extra: extra, ExpiresAt: exp,
}); err != nil {
return nil, serverErr()
}
}
resp := map[string]any{
"access_token": access,
"token_type": tokenType,
"expires_in": int64(s.cfg.AccessTokenTTL.Seconds()),
}
if len(g.Scopes) > 0 {
resp["scope"] = strings.Join(g.Scopes, " ")
}
switch {
case g.NextRefresh != "":
resp["refresh_token"] = g.NextRefresh
case g.IssueRefresh && s.cfg.ManagedRefreshTokens:
rt, e := s.newRefreshToken(ctx, g, now)
if e != nil {
return nil, e
}
resp["refresh_token"] = rt
case g.LegacyRefresh != "":
resp["refresh_token"] = g.LegacyRefresh
}
if g.IDToken && oauthSliceContains(g.Scopes, "openid") {
idt, err := s.buildIDToken(ctx, idTokenParams{
Client: g.Client, UserID: g.UserID, Scopes: g.Scopes, Nonce: g.Nonce, AuthTime: g.AuthTime,
ACR: g.ACR, AMR: g.AMR, SID: g.SID, AccessToken: access, Claims: g.Claims,
})
if err != nil {
return nil, oerr("server_error", "could not sign the id_token", http.StatusInternalServerError)
}
resp["id_token"] = idt
}
return resp, nil
}
func (s *OAuthServer) writeTokenResponse(w http.ResponseWriter, resp map[string]any) {
writeJSON(w, http.StatusOK, resp)
}
// --------------------------------------------------------------------------
// Token endpoint — POST /oauth/token
// --------------------------------------------------------------------------
func (s *OAuthServer) tokenHandler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
if err := r.ParseForm(); err != nil {
writeOAuthError(w, "invalid_request", "cannot parse form", http.StatusBadRequest)
return
}
var resp map[string]any
var e *oauthError
switch r.PostFormValue("grant_type") {
case "authorization_code":
resp, e = s.handleAuthCodeGrant(r)
case "refresh_token":
resp, e = s.handleRefreshGrant(r)
case "client_credentials":
resp, e = s.handleClientCredentialsGrant(r)
case grantDeviceCode:
if !s.cfg.EnableDeviceFlow {
e = oerr("unsupported_grant_type", "", http.StatusBadRequest)
} else {
resp, e = s.handleDeviceGrant(r)
}
case grantTokenExchange:
if !s.cfg.EnableTokenExchange {
e = oerr("unsupported_grant_type", "", http.StatusBadRequest)
} else {
resp, e = s.handleTokenExchange(r)
}
default:
e = oerr("unsupported_grant_type", "", http.StatusBadRequest)
}
if e != nil {
e.write(w)
return
}
s.writeTokenResponse(w, resp)
}
// grantAllowed reports whether the client may use grantType. Clients that declare no grant
// types keep every authorization-code style grant.
func grantAllowed(c *OAuthServerClient, grantType string) bool {
if len(c.GrantTypes) == 0 {
return grantType == "authorization_code" || grantType == "refresh_token"
}
return oauthSliceContains(c.GrantTypes, grantType)
}
// dpopForClient validates the DPoP header of a token request and returns the key to bind to.
func (s *OAuthServer) dpopForClient(r *http.Request, c *OAuthServerClient) (string, *oauthError) {
proof, e := s.verifyDPoP(r, "")
if e != nil {
return "", e
}
if proof == nil {
if c.DPoPBoundAccessTokens {
return "", oerr("invalid_dpop_proof", "this client must present a DPoP proof", http.StatusBadRequest)
}
return "", nil
}
return proof.JKT, nil
}
func (s *OAuthServer) handleAuthCodeGrant(r *http.Request) (map[string]any, *oauthError) {
ac, e := s.requireClient(r)
if e != nil {
return nil, e
}
client := ac.Client
if !grantAllowed(client, "authorization_code") {
return nil, oerr("unauthorized_client", "client may not use authorization_code", http.StatusBadRequest)
}
code := r.PostFormValue("code")
verifier := r.PostFormValue("code_verifier")
if code == "" || verifier == "" {
return nil, oerr("invalid_request", "code and code_verifier required", http.StatusBadRequest)
}
jkt, e := s.dpopForClient(r, client)
if e != nil {
return nil, e
}
c, ok := s.takeCode(r.Context(), code)
if !ok {
// A code that is presented twice means it leaked: end the tokens issued from it.
if s.cfg.ManagedRefreshTokens {
if g := s.grants(); g != nil {
_ = g.RevokeRefreshFamily(r.Context(), codeFamily(code))
}
}
return nil, oerr("invalid_grant", "code expired or invalid", http.StatusBadRequest)
}
switch {
case c.ClientID != client.ClientID:
return nil, oerr("invalid_grant", "code was issued to another client", http.StatusBadRequest)
case c.RedirectURI != r.PostFormValue("redirect_uri"):
return nil, oerr("invalid_grant", "redirect_uri mismatch", http.StatusBadRequest)
case !validatePKCESHA256(c.CodeChallenge, verifier):
return nil, oerr("invalid_grant", "code_verifier invalid", http.StatusBadRequest)
case c.DPoPJKT != "" && c.DPoPJKT != jkt:
return nil, oerr("invalid_dpop_proof", "the DPoP key does not match dpop_jkt of the authorization request", http.StatusBadRequest)
}
g := &tokenGrant{
Client: client, UserID: c.UserID, Scopes: c.Scopes, Resource: c.Resource, Nonce: c.Nonce,
AuthTime: c.AuthTime, ACR: c.ACR, AMR: c.AMR, SID: c.SessionID, Claims: c.Claims, DPoPJKT: jkt,
IDToken: true, FamilyID: codeFamily(code),
IssueRefresh: refreshAllowed(client, c.Scopes),
}
if c.UserID == 0 { // saved by an earlier version: the code carries the session itself
g.LegacyAccess, g.LegacyRefresh = c.SessionToken, c.RefreshToken
g.IssueRefresh = false
}
return s.mintTokens(r.Context(), g)
}
// codeFamily derives the refresh family of the tokens issued from a code.
func codeFamily(code string) string { return "c" + hashToken(code)[:30] }
// refreshAllowed reports whether a managed refresh token is issued for the grant.
func refreshAllowed(c *OAuthServerClient, scopes []string) bool {
return oauthSliceContains(scopes, "offline_access") || grantAllowed(c, "refresh_token")
}
// --------------------------------------------------------------------------
// RFC 6749 §4.4 — Client credentials grant
// --------------------------------------------------------------------------
func (s *OAuthServer) handleClientCredentialsGrant(r *http.Request) (map[string]any, *oauthError) {
if s.auth == nil {
return nil, oerr("unsupported_grant_type", "client_credentials requires a local user store", http.StatusBadRequest)
}
ac, e := s.requireClient(r)
if e != nil {
e.WWWAuth = `Basic realm="oauth"`
return nil, e
}
client := ac.Client
if ac.Method == "none" {
return nil, invalidClient("client_credentials requires a confidential client", true)
}
if !oauthSliceContains(client.GrantTypes, "client_credentials") {
return nil, oerr("unauthorized_client", "client is not authorized for client_credentials", http.StatusBadRequest)
}
requested := strings.Fields(r.PostFormValue("scope"))
effective := client.AllowedScopes
if len(requested) > 0 {
effective = nil
for _, sc := range requested {
if oauthSliceContains(client.AllowedScopes, sc) {
effective = append(effective, sc)
}
}
if len(effective) == 0 {
return nil, oerr("invalid_scope", "no requested scope is allowed for this client", http.StatusBadRequest)
}
}
jkt, e := s.dpopForClient(r, client)
if e != nil {
return nil, e
}
// 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.
userID, err := s.auth.oauth2GetOrCreateUser(r.Context(), &UserContext{
UserName: "client:" + client.ClientID,
Email: "oauth-client-" + client.ClientID + "@service.internal",
RemoteID: client.ClientID,
Roles: effective,
}, "oauth2_client")
if err != nil {
return nil, serverErr()
}
// 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.
var resource []string
if res := r.PostForm["resource"]; len(res) > 0 {
resource = res
}
return s.mintTokens(r.Context(), &tokenGrant{
Client: client, UserID: userID, Scopes: effective, Resource: resource, DPoPJKT: jkt,
})
}
+272
View File
@@ -0,0 +1,272 @@
package security
import (
"context"
"crypto/subtle"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"sync"
"time"
"github.com/golang-jwt/jwt/v5"
)
// OIDCConfig configures an OpenID Connect provider found by discovery.
type OIDCConfig struct {
// Issuer is the provider's issuer URL; /.well-known/openid-configuration is fetched from it.
Issuer string
ClientID string
ClientSecret string
RedirectURL string
// Scopes defaults to openid, profile, email.
Scopes []string
ProviderName string // default "oidc"
// Optional, see OAuth2Config.
UserInfoParser func(userInfo map[string]any) (*UserContext, error)
AllowedAlgs []string
AuthStyle string
HTTPClient *http.Client
ClockSkew time.Duration
}
// oidcProvider is the OpenID Connect part of an OAuth2Provider: discovered endpoints and the
// id_token validator.
type oidcProvider struct {
issuer string
clientID string
jwksURL string
endSession string
algs []string
skew time.Duration
client *http.Client
keys *jwksCache
needsLookup bool // endpoints still have to be discovered
mu sync.Mutex
}
func newOIDCProvider(cfg *OAuth2Config) *oidcProvider {
p := &oidcProvider{
issuer: cfg.Issuer,
clientID: cfg.ClientID,
jwksURL: cfg.JWKSURL,
endSession: cfg.EndSessionURL,
algs: cfg.AllowedAlgs,
skew: cfg.ClockSkew,
client: cfg.HTTPClient,
}
if len(p.algs) == 0 {
p.algs = []string{"RS256", "PS256", "ES256", "ES384"}
}
if p.skew == 0 {
p.skew = time.Minute
}
if p.client == nil {
p.client = &http.Client{Timeout: 10 * time.Second}
}
p.keys = newJWKSCache(p.client)
p.needsLookup = cfg.AuthURL == "" || cfg.TokenURL == "" || cfg.JWKSURL == ""
return p
}
// ensureEndpoints runs discovery once when the endpoints were not configured.
func (p *oidcProvider) ensureEndpoints(ctx context.Context, op *OAuth2Provider) error {
p.mu.Lock()
defer p.mu.Unlock()
if !p.needsLookup {
return nil
}
doc, err := fetchOIDCDiscovery(ctx, p.client, p.issuer)
if err != nil {
return err
}
if doc.Issuer != p.issuer {
return fmt.Errorf("discovery issuer mismatch: got %q, want %q", doc.Issuer, p.issuer)
}
if op.config.Endpoint.AuthURL == "" {
op.config.Endpoint.AuthURL = doc.AuthorizationEndpoint
}
if op.config.Endpoint.TokenURL == "" {
op.config.Endpoint.TokenURL = doc.TokenEndpoint
}
if op.userInfoURL == "" {
op.userInfoURL = doc.UserinfoEndpoint
}
if p.jwksURL == "" {
p.jwksURL = doc.JWKSURI
}
if p.endSession == "" {
p.endSession = doc.EndSessionEndpoint
}
if op.config.Endpoint.AuthURL == "" || op.config.Endpoint.TokenURL == "" || p.jwksURL == "" {
return errors.New("discovery document lacks authorization, token or jwks endpoint")
}
p.needsLookup = false
return nil
}
type oidcDiscovery struct {
Issuer string `json:"issuer"`
AuthorizationEndpoint string `json:"authorization_endpoint"`
TokenEndpoint string `json:"token_endpoint"`
UserinfoEndpoint string `json:"userinfo_endpoint"`
JWKSURI string `json:"jwks_uri"`
EndSessionEndpoint string `json:"end_session_endpoint"`
}
func fetchOIDCDiscovery(ctx context.Context, client *http.Client, issuer string) (*oidcDiscovery, error) {
if client == nil {
client = &http.Client{Timeout: 10 * time.Second}
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, strings.TrimRight(issuer, "/")+"/.well-known/openid-configuration", nil)
if err != nil {
return nil, err
}
resp, err := client.Do(req)
if err != nil {
return nil, fmt.Errorf("OIDC discovery: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("OIDC discovery: status %d", resp.StatusCode)
}
var doc oidcDiscovery
if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&doc); err != nil {
return nil, fmt.Errorf("OIDC discovery: %w", err)
}
return &doc, nil
}
// validateIDToken verifies the signature and the claims of an id_token (OIDC Core 3.1.3.7).
// nonce is checked when non-empty (not on refresh); accessToken is checked against at_hash.
func (p *oidcProvider) validateIDToken(ctx context.Context, raw, nonce, accessToken string) (map[string]any, error) {
claims := jwt.MapClaims{}
verify := func(refresh bool) (*jwt.Token, error) {
set, err := p.keys.get(ctx, p.jwksURL, refresh)
if err != nil {
return nil, err
}
return verifyJWTWithSet(raw, set, p.algs, claims,
jwt.WithIssuer(p.issuer), jwt.WithAudience(p.clientID),
jwt.WithExpirationRequired(), jwt.WithLeeway(p.skew))
}
tok, err := verify(false)
if err != nil {
// A rotated key: refetch the key set once.
claims = jwt.MapClaims{}
if tok, err = verify(true); err != nil {
return nil, err
}
}
if sub, _ := claims["sub"].(string); sub == "" {
return nil, errors.New("missing sub")
}
// With several audiences azp must name this client.
if aud, _ := claims.GetAudience(); len(aud) > 1 {
if azp, _ := claims["azp"].(string); azp != p.clientID {
return nil, errors.New("azp does not match the client")
}
}
if azp, ok := claims["azp"].(string); ok && azp != p.clientID {
return nil, errors.New("azp does not match the client")
}
if nonce != "" {
got, _ := claims["nonce"].(string)
if subtle.ConstantTimeCompare([]byte(got), []byte(nonce)) != 1 {
return nil, errors.New("nonce mismatch")
}
}
if accessToken != "" {
if want, ok := claims["at_hash"].(string); ok {
alg, _ := tok.Header["alg"].(string)
if halfHash(alg, accessToken) != want {
return nil, errors.New("at_hash mismatch")
}
}
}
return claims, nil
}
// WithOIDC registers an OpenID Connect provider. The endpoints come from the issuer's discovery
// document. Login uses PKCE and a nonce, and the id_token is validated on callback and refresh.
func (a *DatabaseAuthenticator) WithOIDC(ctx context.Context, cfg OIDCConfig) (*DatabaseAuthenticator, error) {
if cfg.Issuer == "" {
return a, errors.New("OIDC issuer is required")
}
if cfg.ProviderName == "" {
cfg.ProviderName = "oidc"
}
if len(cfg.Scopes) == 0 {
cfg.Scopes = []string{"openid", "profile", "email"}
}
doc, err := fetchOIDCDiscovery(ctx, cfg.HTTPClient, cfg.Issuer)
if err != nil {
return a, err
}
if doc.Issuer != strings.TrimRight(cfg.Issuer, "/") && doc.Issuer != cfg.Issuer {
return a, fmt.Errorf("discovery issuer mismatch: got %q, want %q", doc.Issuer, cfg.Issuer)
}
return a.WithOAuth2(OAuth2Config{
ClientID: cfg.ClientID,
ClientSecret: cfg.ClientSecret,
RedirectURL: cfg.RedirectURL,
Scopes: cfg.Scopes,
AuthURL: doc.AuthorizationEndpoint,
TokenURL: doc.TokenEndpoint,
UserInfoURL: doc.UserinfoEndpoint,
ProviderName: cfg.ProviderName,
UserInfoParser: cfg.UserInfoParser,
Issuer: doc.Issuer,
JWKSURL: doc.JWKSURI,
EndSessionURL: doc.EndSessionEndpoint,
AllowedAlgs: cfg.AllowedAlgs,
AuthStyle: cfg.AuthStyle,
HTTPClient: cfg.HTTPClient,
ClockSkew: cfg.ClockSkew,
}), nil
}
// OAuth2LogoutURL returns the provider's RP-initiated logout URL (OIDC RP-Initiated Logout 1.0).
// idTokenHint is LoginResponse.Meta["id_token"]. It fails when the provider has no end_session_endpoint.
func (a *DatabaseAuthenticator) OAuth2LogoutURL(ctx context.Context, providerName, idTokenHint, postLogoutRedirect, state string) (string, error) {
provider, err := a.getOAuth2Provider(providerName)
if err != nil {
return "", err
}
if provider.oidc == nil {
return "", fmt.Errorf("provider %q is not an OpenID Connect provider", providerName)
}
if err := provider.oidc.ensureEndpoints(provider.withHTTPClient(ctx), provider); err != nil {
return "", err
}
provider.oidc.mu.Lock()
end := provider.oidc.endSession
provider.oidc.mu.Unlock()
if end == "" {
return "", fmt.Errorf("provider %q has no end_session_endpoint", providerName)
}
u, err := url.Parse(end)
if err != nil {
return "", err
}
q := u.Query()
if idTokenHint != "" {
q.Set("id_token_hint", idTokenHint)
}
if postLogoutRedirect != "" {
q.Set("post_logout_redirect_uri", postLogoutRedirect)
if state != "" {
q.Set("state", state)
}
}
q.Set("client_id", provider.config.ClientID)
u.RawQuery = q.Encode()
return u.String(), nil
}
+160
View File
@@ -0,0 +1,160 @@
package security
import (
"context"
"html"
"io"
"net/http"
"net/http/cookiejar"
"net/http/httptest"
"net/url"
"regexp"
"strings"
"testing"
)
var (
reFormAction = regexp.MustCompile(`<form method="POST" action="([^"]*)"`)
reHidden = regexp.MustCompile(`<input type="hidden" name="([^"]*)" value="([^"]*)"`)
)
// browserSubmit posts the first form of an HTML page with extra fields.
func browserSubmit(t *testing.T, c *http.Client, base, page string, extra url.Values) *http.Response {
t.Helper()
m := reFormAction.FindStringSubmatch(page)
if m == nil {
t.Fatalf("no form in page: %s", page)
}
form := url.Values{}
for _, h := range reHidden.FindAllStringSubmatch(page, -1) {
form.Set(h[1], html.UnescapeString(h[2]))
}
for k, v := range extra {
form[k] = v
}
action, _ := url.Parse(html.UnescapeString(m[1]))
b, _ := url.Parse(base)
resp, err := c.PostForm(b.ResolveReference(action).String(), form)
if err != nil {
t.Fatal(err)
}
return resp
}
func bodyOf(t *testing.T, r *http.Response) string {
t.Helper()
defer r.Body.Close()
b, _ := io.ReadAll(r.Body)
return string(b)
}
// TestOIDCClient_FullLoop runs the relying party against our own authorization server:
// discovery, PKCE, nonce, login, consent, code exchange, id_token validation and logout URL.
func TestOIDCClient_FullLoop(t *testing.T) {
var handler http.Handler
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { handler.ServeHTTP(w, r) }))
defer ts.Close()
idpAuth := NewDatabaseAuthenticatorWithOptions(newDirectTestDB(t), DatabaseAuthenticatorOptions{Lookup: directConfig})
srv := NewOAuthServer(OAuthServerConfig{Issuer: ts.URL, PersistCodes: true, RequireConsent: true}, idpAuth)
defer srv.Close()
handler = srv.HTTPHandler()
if _, err := idpAuth.Register(context.Background(), RegisterRequest{Username: "olivia", Password: "pw", Email: "olivia@example.com"}); err != nil {
t.Fatal(err)
}
_, reg := doJSON(t, handler, http.MethodPost, "/oauth/register", map[string]interface{}{
"redirect_uris": []string{"https://rp.example.com/callback"},
})
clientID := reg["client_id"].(string)
rp := NewDatabaseAuthenticatorWithOptions(newDirectTestDB(t), DatabaseAuthenticatorOptions{Lookup: directConfig})
if _, err := rp.WithOIDC(context.Background(), OIDCConfig{
Issuer: ts.URL, ClientID: clientID, RedirectURL: "https://rp.example.com/callback", ProviderName: "idp",
}); err != nil {
t.Fatal(err)
}
authURL, err := rp.OAuth2GetAuthURL("idp", "state-1")
if err != nil {
t.Fatal(err)
}
for _, want := range []string{"code_challenge=", "code_challenge_method=S256", "nonce=", "state=state-1"} {
if !strings.Contains(authURL, want) {
t.Fatalf("auth URL lacks %q: %s", want, authURL)
}
}
jar, _ := cookiejar.New(nil)
browser := &http.Client{Jar: jar, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }}
resp, err := browser.Get(authURL)
if err != nil {
t.Fatal(err)
}
page := bodyOf(t, resp)
resp = browserSubmit(t, browser, authURL, page, url.Values{"username": {"olivia"}, "password": {"pw"}})
page = bodyOf(t, resp)
if resp.StatusCode == http.StatusFound {
t.Fatalf("expected the consent page, got redirect to %s", resp.Header.Get("Location"))
}
resp = browserSubmit(t, browser, authURL, page, url.Values{"decision": {"allow"}, "remember": {"1"}})
loc := resp.Header.Get("Location")
bodyOf(t, resp)
if !strings.HasPrefix(loc, "https://rp.example.com/callback?") {
t.Fatalf("expected redirect to the RP, got %q (status %d)", loc, resp.StatusCode)
}
cb, _ := url.Parse(loc)
if cb.Query().Get("iss") != ts.URL {
t.Errorf("iss = %q", cb.Query().Get("iss"))
}
// A wrong iss (mix-up) is refused, and a replayed state is gone afterwards.
bad := httptest.NewRequest(http.MethodGet, "/cb?code=x&state=state-1&iss=https://evil.example", nil)
if _, err := rp.OAuth2HandleCallbackRequest(context.Background(), "idp", bad); err == nil {
t.Fatal("mismatching iss must fail")
}
authURL, _ = rp.OAuth2GetAuthURL("idp", "state-1")
_ = authURL
// Redo the login for a fresh state/nonce (the failed attempt consumed the first state).
authURL, _ = rp.OAuth2GetAuthURL("idp", "state-2")
resp, _ = browser.Get(authURL)
loc = resp.Header.Get("Location")
bodyOf(t, resp)
if !strings.HasPrefix(loc, "https://rp.example.com/callback?") {
t.Fatalf("SSO + remembered consent should redirect straight to the RP, got %q", loc)
}
cbReq := httptest.NewRequest(http.MethodGet, loc, nil)
login, err := rp.OAuth2HandleCallbackRequest(context.Background(), "idp", cbReq)
if err != nil {
t.Fatalf("callback: %v", err)
}
if login.User == nil || login.User.Email != "olivia@example.com" {
t.Fatalf("unexpected user: %+v", login.User)
}
idToken, _ := login.Meta["id_token"].(string)
if idToken == "" {
t.Fatal("id_token missing in Meta")
}
// Replaying the callback fails: the state is single use.
if _, err := rp.OAuth2HandleCallbackRequest(context.Background(), "idp", cbReq); err == nil {
t.Fatal("replayed callback must fail")
}
out, err := rp.OAuth2LogoutURL(context.Background(), "idp", idToken, "https://rp.example.com/bye", "s")
if err != nil {
t.Fatal(err)
}
if !strings.HasPrefix(out, ts.URL+"/oauth/logout?") || !strings.Contains(out, "id_token_hint=") {
t.Errorf("logout URL = %s", out)
}
}
func TestOIDCClient_RejectsBadIDToken(t *testing.T) {
p := newOIDCProvider(&OAuth2Config{Issuer: "https://idp.example", ClientID: "c", JWKSURL: "http://127.0.0.1:1/jwks", AuthURL: "a", TokenURL: "t"})
if _, err := p.validateIDToken(context.Background(), "not.a.jwt", "n", ""); err == nil {
t.Fatal("garbage id_token must fail")
}
}
+105 -1
View File
@@ -1,6 +1,9 @@
package sectypes
import "time"
import (
"encoding/json"
"time"
)
// OAuthServerClient is a persisted RFC 7591 registered OAuth2 client.
type OAuthServerClient struct {
@@ -11,8 +14,80 @@ type OAuthServerClient struct {
AllowedScopes []string `json:"allowed_scopes,omitempty"`
ClientSecretHash string `json:"client_secret_hash,omitempty"`
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
// The fields below are stored together in the oauth_clients.metadata JSON column, so a
// new field never needs a schema change. See SplitJSON / MergeJSON.
ResponseTypes []string `json:"response_types,omitempty"`
ClientURI string `json:"client_uri,omitempty"`
LogoURI string `json:"logo_uri,omitempty"`
Contacts []string `json:"contacts,omitempty"`
PostLogoutRedirectURIs []string `json:"post_logout_redirect_uris,omitempty"`
BackchannelLogoutURI string `json:"backchannel_logout_uri,omitempty"`
JWKS json.RawMessage `json:"jwks,omitempty"`
JWKSURI string `json:"jwks_uri,omitempty"`
IDTokenSignedResponseAlg string `json:"id_token_signed_response_alg,omitempty"`
UserinfoSignedResponseAlg string `json:"userinfo_signed_response_alg,omitempty"`
TokenEndpointAuthSigningAlg string `json:"token_endpoint_auth_signing_alg,omitempty"`
RequireConsent bool `json:"require_consent,omitempty"`
FirstParty bool `json:"first_party,omitempty"`
RequirePAR bool `json:"require_pushed_authorization_requests,omitempty"`
DPoPBoundAccessTokens bool `json:"dpop_bound_access_tokens,omitempty"`
RegistrationAccessTokenHash string `json:"registration_access_token_hash,omitempty"`
ClientSecretExpiresAt int64 `json:"client_secret_expires_at,omitempty"`
ClientIDIssuedAt int64 `json:"client_id_issued_at,omitempty"`
}
// oauthClientColumns are the keys stored in their own oauth_clients columns; every other
// key of the JSON form is stored in the metadata column.
var oauthClientColumns = []string{
"client_id", "redirect_uris", "client_name", "grant_types", "allowed_scopes",
"client_secret_hash", "token_endpoint_auth_method",
}
// oauthCodeColumns are the keys stored in their own oauth_codes columns.
var oauthCodeColumns = []string{
"code", "client_id", "redirect_uri", "client_state", "code_challenge", "code_challenge_method",
"session_token", "refresh_token", "scopes", "expires_at",
}
// SplitJSON returns the JSON form of v without the keys in columns. It is the value stored in
// a metadata/extra column; an empty object is returned as "".
func SplitJSON(v any, columns []string) (string, error) {
raw, err := json.Marshal(v) //nolint:gosec // G117: client secret hash and tokens are intentionally stored
if err != nil {
return "", err
}
m := map[string]json.RawMessage{}
if err := json.Unmarshal(raw, &m); err != nil {
return "", err
}
for _, c := range columns {
delete(m, c)
}
if len(m) == 0 {
return "", nil
}
out, err := json.Marshal(m)
return string(out), err
}
// MergeJSON applies a metadata/extra JSON document onto dst. Empty input is a no-op.
func MergeJSON(dst any, data string) error {
if data == "" || data == "null" {
return nil
}
return json.Unmarshal([]byte(data), dst)
}
// ClientMetadataJSON returns the value of the oauth_clients.metadata column.
func (c *OAuthServerClient) ClientMetadataJSON() (string, error) {
return SplitJSON(c, oauthClientColumns)
}
// ApplyClientMetadata merges the oauth_clients.metadata column into c.
func (c *OAuthServerClient) ApplyClientMetadata(data string) error { return MergeJSON(c, data) }
// OAuthCode is a short-lived authorization code.
type OAuthCode struct {
Code string `json:"code"`
@@ -25,8 +100,28 @@ type OAuthCode struct {
RefreshToken string `json:"refresh_token,omitempty"`
Scopes []string `json:"scopes,omitempty"`
ExpiresAt time.Time `json:"expires_at"`
// Stored together in the oauth_codes.extra JSON column.
UserID int `json:"user_id,omitempty"`
Nonce string `json:"nonce,omitempty"`
AuthTime int64 `json:"auth_time,omitempty"`
ACR string `json:"acr,omitempty"`
AMR []string `json:"amr,omitempty"`
SessionID string `json:"sid,omitempty"`
Claims map[string]any `json:"claims,omitempty"`
Resource []string `json:"resource,omitempty"`
DPoPJKT string `json:"dpop_jkt,omitempty"`
ResponseType string `json:"response_type,omitempty"`
ConsentedAt int64 `json:"consented_at,omitempty"`
}
// CodeExtraJSON returns the value of the oauth_codes.extra column.
func (c *OAuthCode) CodeExtraJSON() (string, error) { return SplitJSON(c, oauthCodeColumns) }
// ApplyCodeExtra merges the oauth_codes.extra column into c.
func (c *OAuthCode) ApplyCodeExtra(data string) error { return MergeJSON(c, data) }
// OAuthTokenInfo is the RFC 7662 token introspection response.
type OAuthTokenInfo struct {
Active bool `json:"active"`
@@ -37,4 +132,13 @@ type OAuthTokenInfo struct {
Roles []string `json:"roles,omitempty"`
Exp int64 `json:"exp,omitempty"`
Iat int64 `json:"iat,omitempty"`
// Filled in by the OAuth server, not by the stores.
Scope string `json:"scope,omitempty"`
ClientID string `json:"client_id,omitempty"`
TokenType string `json:"token_type,omitempty"`
Iss string `json:"iss,omitempty"`
Aud []string `json:"aud,omitempty"`
Jti string `json:"jti,omitempty"`
Cnf map[string]any `json:"cnf,omitempty"`
}