mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 19:41:57 +00:00
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:
@@ -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).
|
For documentation, see [pkg/security/README.md](pkg/security/README.md) (see "Database access (lookup)" for the SQLite/portable-SQL path).
|
||||||
|
|
||||||
|
It includes a standards-based OAuth 2.1 / OpenID Connect authorization server (consent, rotating refresh tokens, JWT access tokens, DPoP, PAR, device grant, token exchange, logout) and an OIDC relying-party client; see [pkg/security/OAUTH2_SERVER.md](pkg/security/OAUTH2_SERVER.md).
|
||||||
|
|
||||||
#### Middleware
|
#### Middleware
|
||||||
|
|
||||||
HTTP middleware collection for common tasks (CORS, logging, metrics, rate limiting, etc.).
|
HTTP middleware collection for common tasks (CORS, logging, metrics, rate limiting, etc.).
|
||||||
|
|||||||
@@ -179,6 +179,8 @@ It can operate as:
|
|||||||
- **An OAuth2 federation layer** — delegates to external providers (Google, GitHub, Microsoft, etc.)
|
- **An OAuth2 federation layer** — delegates to external providers (Google, GitHub, Microsoft, etc.)
|
||||||
- **Both simultaneously**
|
- **Both simultaneously**
|
||||||
|
|
||||||
|
> The underlying `security.OAuthServer` also supports consent, OpenID Connect, rotating refresh tokens, JWT access tokens, DPoP, PAR, the device grant and token exchange; they are opt-in `OAuthServerConfig` options described in [pkg/security/OAUTH2_SERVER.md](../security/OAUTH2_SERVER.md). The options of `resolvemcp.OAuth2Config` are unchanged.
|
||||||
|
|
||||||
### Standard endpoints served
|
### Standard endpoints served
|
||||||
|
|
||||||
| Path | Spec | Purpose |
|
| Path | Spec | Purpose |
|
||||||
|
|||||||
@@ -4,6 +4,8 @@
|
|||||||
|
|
||||||
The security package provides OAuth2 authentication support for any OAuth2-compliant provider including Google, GitHub, Microsoft, Facebook, and custom providers.
|
The security package provides OAuth2 authentication support for any OAuth2-compliant provider including Google, GitHub, Microsoft, Facebook, and custom providers.
|
||||||
|
|
||||||
|
> **Full OAuth 2.1 / OpenID Connect**: this guide covers the plain OAuth2 client login. For the OIDC relying party (discovery, PKCE, nonce, id_token validation, logout) and the complete authorization server (consent, refresh rotation, DPoP, PAR, device grant, token exchange) see [OAUTH2_SERVER.md](OAUTH2_SERVER.md).
|
||||||
|
|
||||||
## Features
|
## Features
|
||||||
|
|
||||||
- **Universal OAuth2 Support**: Works with any OAuth2 provider
|
- **Universal OAuth2 Support**: Works with any OAuth2 provider
|
||||||
@@ -14,6 +16,7 @@ The security package provides OAuth2 authentication support for any OAuth2-compl
|
|||||||
- **Token Refresh**: Automatic token refresh support
|
- **Token Refresh**: Automatic token refresh support
|
||||||
- **State Validation**: Built-in CSRF protection
|
- **State Validation**: Built-in CSRF protection
|
||||||
- **User Auto-Creation**: Automatically creates users on first login
|
- **User Auto-Creation**: Automatically creates users on first login
|
||||||
|
- **OpenID Connect** (opt-in): `WithOIDC` discovery, PKCE, nonce and id_token validation, RP-initiated logout
|
||||||
- **Unified Authentication**: OAuth2 and traditional auth share same session storage
|
- **Unified Authentication**: OAuth2 and traditional auth share same session storage
|
||||||
|
|
||||||
## Quick Start
|
## Quick Start
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
# OAuth2 Refresh Token - Quick Reference
|
# OAuth2 Refresh Token - Quick Reference
|
||||||
|
|
||||||
|
> This covers refreshing tokens of an upstream provider with `OAuth2RefreshToken`. For refresh tokens issued by `OAuthServer` (rotation, reuse detection, downscoping) see [OAUTH2_SERVER.md](OAUTH2_SERVER.md#refresh-token-rotation).
|
||||||
|
|
||||||
## Quick Setup (3 Steps)
|
## Quick Setup (3 Steps)
|
||||||
|
|
||||||
### 1. Initialize Authenticator
|
### 1. Initialize Authenticator
|
||||||
|
|||||||
@@ -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.
|
||||||
@@ -13,7 +13,7 @@ Type-safe, composable security system for ResolveSpec with support for authentic
|
|||||||
- ✅ **Extensible** - Implement custom providers for your needs
|
- ✅ **Extensible** - Implement custom providers for your needs
|
||||||
- ✅ **Stored Procedures** - Database operations use PostgreSQL stored procedures where available, for security and maintainability
|
- ✅ **Stored Procedures** - Database operations use PostgreSQL stored procedures where available, for security and maintainability
|
||||||
- ✅ **Direct Mode** - Portable Go/SQL fallback for SQLite, MySQL, or Postgres without the stored procedures installed — no code changes required
|
- ✅ **Direct Mode** - Portable Go/SQL fallback for SQLite, MySQL, or Postgres without the stored procedures installed — no code changes required
|
||||||
- ✅ **OAuth2 Authorization Server** - Built-in OAuth 2.1 + PKCE server (RFC 8414, 7591, 7009, 7662) with login form and external provider federation
|
- ✅ **OAuth2 / OpenID Connect** - Built-in OAuth 2.1 + PKCE authorization server and OIDC provider (RFC 8414, 7591/7592, 7009, 7662, 9068, 9126, 9207, 9449, 8628, 8693): consent, rotating refresh tokens, JWT access tokens, logout, federation; plus an OIDC relying-party client. See [OAUTH2_SERVER.md](OAUTH2_SERVER.md)
|
||||||
- ✅ **Password Reset** - Self-service password reset with secure token generation and session invalidation
|
- ✅ **Password Reset** - Self-service password reset with secure token generation and session invalidation
|
||||||
|
|
||||||
## Stored Procedure Architecture
|
## Stored Procedure Architecture
|
||||||
@@ -1068,6 +1068,8 @@ lookup.ProcNames{
|
|||||||
|
|
||||||
## OAuth2 Authorization Server
|
## OAuth2 Authorization Server
|
||||||
|
|
||||||
|
> The complete guide (consent, OIDC, refresh rotation, DPoP, PAR, device grant, token exchange, logout, relying-party client) is in [OAUTH2_SERVER.md](OAUTH2_SERVER.md). The table below lists the original endpoints.
|
||||||
|
|
||||||
`OAuthServer` is a generic OAuth 2.1 + PKCE authorization server. It is not tied to any spec — `pkg/resolvemcp` uses it, but it can be used standalone with any `http.ServeMux`.
|
`OAuthServer` is a generic OAuth 2.1 + PKCE authorization server. It is not tied to any spec — `pkg/resolvemcp` uses it, but it can be used standalone with any `http.ServeMux`.
|
||||||
|
|
||||||
### Endpoints
|
### Endpoints
|
||||||
@@ -1229,12 +1231,13 @@ The main changes:
|
|||||||
|------|-------------|
|
|------|-------------|
|
||||||
| **QUICK_REFERENCE.md** | Quick reference guide with examples |
|
| **QUICK_REFERENCE.md** | Quick reference guide with examples |
|
||||||
| **KEYSTORE.md** | Per-user auth keys and key stores |
|
| **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 |
|
| **OAUTH2_REFRESH_QUICK_REFERENCE.md** / **OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md** | OAuth2 refresh tokens |
|
||||||
| **PASSKEY_QUICK_REFERENCE.md** | WebAuthn passkeys |
|
| **PASSKEY_QUICK_REFERENCE.md** | WebAuthn passkeys |
|
||||||
| **SECURITY_FEATURES.md** | Security feature overview |
|
| **SECURITY_FEATURES.md** | Security feature overview |
|
||||||
| **breaking_changes.md** | Migration notes for the `lookup` refactor |
|
| **breaking_changes.md** | Migration notes (`lookup` refactor, full OAuth2/OIDC schema changes) |
|
||||||
| **examples.go**, **examples_funcspec.go**, **oauth2_examples.go**, **passkey_examples.go** | Working provider implementations |
|
| **examples.go**, **examples_funcspec.go**, **oauth2_examples.go**, **oauth2_full_example.go**, **passkey_examples.go** | Working provider implementations |
|
||||||
|
|
||||||
## API Reference
|
## API Reference
|
||||||
|
|
||||||
|
|||||||
@@ -123,3 +123,37 @@ Behaviour changes:
|
|||||||
- Removed the unexported `password.go` from `pkg/security` (bcrypt helpers live in `lookup/direct`).
|
- 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`,
|
- `README.md`, `KEYSTORE.md` and the root README describe `lookup.Config` instead of `QueryMode`,
|
||||||
`SQLNames` and `TableNames`.
|
`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`.
|
||||||
|
|||||||
@@ -71,6 +71,9 @@ func New(db *sql.DB, cfg lookup.Config, opts Options) (*lookup.Provider, error)
|
|||||||
OAuthUser: &oauthUserRouter{c: c,
|
OAuthUser: &oauthUserRouter{c: c,
|
||||||
proc: procedure.NewOAuthUsers(run, res.Procs),
|
proc: procedure.NewOAuthUsers(run, res.Procs),
|
||||||
direct: direct.NewOAuthUsers(base)},
|
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)},
|
Passkey: &passkeyRouter{c: c, proc: p, direct: direct.NewPasskey(base)},
|
||||||
TOTP: &totpRouter{c: c,
|
TOTP: &totpRouter{c: c,
|
||||||
proc: procedure.NewTOTP(run, res.Procs),
|
proc: procedure.NewTOTP(run, res.Procs),
|
||||||
@@ -90,6 +93,7 @@ func Failed(err error) *lookup.Provider {
|
|||||||
Keys: &keysRouter{c: c},
|
Keys: &keysRouter{c: c},
|
||||||
OAuthClient: &oauthClientRouter{c: c},
|
OAuthClient: &oauthClientRouter{c: c},
|
||||||
OAuthUser: &oauthUserRouter{c: c},
|
OAuthUser: &oauthUserRouter{c: c},
|
||||||
|
OAuthGrant: &oauthGrantRouter{c: c},
|
||||||
Passkey: &passkeyRouter{c: c},
|
Passkey: &passkeyRouter{c: c},
|
||||||
TOTP: &totpRouter{c: c},
|
TOTP: &totpRouter{c: c},
|
||||||
Policy: &policyRouter{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 + "%"
|
like := prefix + "%"
|
||||||
for _, q := range []struct{ table, col string }{
|
for _, q := range []struct{ table, col string }{
|
||||||
{"oauth_codes", "code"},
|
{"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"},
|
{"oauth_clients", "client_id"},
|
||||||
{"token_blacklist", "token"},
|
{"token_blacklist", "token"},
|
||||||
{"sec_column_rules", "schema_name"},
|
{"sec_column_rules", "schema_name"},
|
||||||
|
|||||||
@@ -154,3 +154,27 @@ func TestConformanceMSSQLContainer(t *testing.T) {
|
|||||||
}
|
}
|
||||||
runOnServer(t, "sqlserver", dsn("cf_direct"), "mssql", lookup.Config{}, true)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -199,6 +199,22 @@ func (r *oauthClientRouter) Revoke(ctx context.Context, token string) error {
|
|||||||
return st.Revoke(ctx, token)
|
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 {
|
type oauthUserRouter struct {
|
||||||
c *chooser
|
c *chooser
|
||||||
proc, direct lookup.OAuthUserStore
|
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)
|
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("OAuthClientAndCodes", s.oauthClientAndCodes)
|
||||||
t.Run("OAuthIntrospectRevoke", s.oauthIntrospectRevoke)
|
t.Run("OAuthIntrospectRevoke", s.oauthIntrospectRevoke)
|
||||||
t.Run("OAuthUsers", s.oauthUsers)
|
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("Passkey", s.passkey)
|
||||||
t.Run("TOTP", s.totp)
|
t.Run("TOTP", s.totp)
|
||||||
t.Run("Policy", s.policy)
|
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 }
|
func ptr[T any](v T) *T { return &v }
|
||||||
|
|
||||||
|
// --- OAuth server grant state -------------------------------------------------------------
|
||||||
|
|
||||||
|
func (s *suite) oauthClientMetadata(t *testing.T) {
|
||||||
|
st := s.Provider.OAuthClient
|
||||||
|
id := s.name("meta-client")
|
||||||
|
reg, err := st.RegisterClient(ctx, §ypes.OAuthServerClient{
|
||||||
|
ClientID: id, RedirectURIs: []string{"https://app.example/cb"}, ClientName: "Meta",
|
||||||
|
PostLogoutRedirectURIs: []string{"https://app.example/bye"}, RequireConsent: true, FirstParty: false,
|
||||||
|
IDTokenSignedResponseAlg: "RS256", Contacts: []string{"ops@example.test"}, DPoPBoundAccessTokens: true,
|
||||||
|
})
|
||||||
|
if err != nil || reg.ClientID != id {
|
||||||
|
t.Fatalf("register: %+v %v", reg, err)
|
||||||
|
}
|
||||||
|
got, err := st.GetClient(ctx, id)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("get: %v", err)
|
||||||
|
}
|
||||||
|
if !got.RequireConsent || !got.DPoPBoundAccessTokens || got.IDTokenSignedResponseAlg != "RS256" ||
|
||||||
|
len(got.PostLogoutRedirectURIs) != 1 || got.PostLogoutRedirectURIs[0] != "https://app.example/bye" ||
|
||||||
|
len(got.Contacts) != 1 {
|
||||||
|
t.Fatalf("metadata lost: %+v", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
got.ClientName = "Renamed"
|
||||||
|
got.RequireConsent = false
|
||||||
|
got.RedirectURIs = []string{"https://app.example/cb", "https://app.example/cb2"}
|
||||||
|
if err := st.UpdateClient(ctx, got); err != nil {
|
||||||
|
t.Fatalf("update: %v", err)
|
||||||
|
}
|
||||||
|
again, err := st.GetClient(ctx, id)
|
||||||
|
if err != nil || again.ClientName != "Renamed" || again.RequireConsent || len(again.RedirectURIs) != 2 || !again.DPoPBoundAccessTokens {
|
||||||
|
t.Fatalf("after update: %+v %v", again, err)
|
||||||
|
}
|
||||||
|
if err := st.DeleteClient(ctx, id); err != nil {
|
||||||
|
t.Fatalf("delete: %v", err)
|
||||||
|
}
|
||||||
|
_, err = st.GetClient(ctx, id)
|
||||||
|
rejected(t, "deleted client", err)
|
||||||
|
|
||||||
|
// Code extras round-trip.
|
||||||
|
code := s.name("meta-code")
|
||||||
|
err = st.SaveCode(ctx, §ypes.OAuthCode{
|
||||||
|
Code: code, ClientID: id, RedirectURI: "https://app.example/cb", CodeChallenge: "chal", CodeChallengeMethod: "S256",
|
||||||
|
SessionToken: "sess", Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Minute),
|
||||||
|
Nonce: "n-0S6", AuthTime: 1700000000, ACR: "urn:acr:1", AMR: []string{"pwd"}, UserID: 7,
|
||||||
|
Claims: map[string]any{"id_token": map[string]any{"email": nil}}, Resource: []string{"https://api.example"}, DPoPJKT: "jkt",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("save code: %v", err)
|
||||||
|
}
|
||||||
|
c, err := st.ExchangeCode(ctx, code)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("exchange: %v", err)
|
||||||
|
}
|
||||||
|
if c.Nonce != "n-0S6" || c.AuthTime != 1700000000 || c.ACR != "urn:acr:1" || c.UserID != 7 || c.DPoPJKT != "jkt" ||
|
||||||
|
len(c.AMR) != 1 || len(c.Resource) != 1 || c.Claims["id_token"] == nil {
|
||||||
|
t.Fatalf("code extra lost: %+v", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *suite) oauthGrantConsent(t *testing.T) {
|
||||||
|
g := s.Provider.OAuthGrant
|
||||||
|
client := s.name("consent-client")
|
||||||
|
if _, err := g.GetConsent(ctx, 1, client); !errors.Is(err, lookup.ErrNotFound) {
|
||||||
|
t.Fatalf("missing consent: %v", err)
|
||||||
|
}
|
||||||
|
if err := g.SaveConsent(ctx, lookup.Consent{UserID: 1, ClientID: client, Scopes: []string{"openid", "email"}, ExpiresAt: time.Now().Add(time.Hour)}); err != nil {
|
||||||
|
t.Fatalf("save: %v", err)
|
||||||
|
}
|
||||||
|
c, err := g.GetConsent(ctx, 1, client)
|
||||||
|
if err != nil || len(c.Scopes) != 2 {
|
||||||
|
t.Fatalf("get: %+v %v", c, err)
|
||||||
|
}
|
||||||
|
// Saving again replaces.
|
||||||
|
if err := g.SaveConsent(ctx, lookup.Consent{UserID: 1, ClientID: client, Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Hour)}); err != nil {
|
||||||
|
t.Fatalf("resave: %v", err)
|
||||||
|
}
|
||||||
|
if c, err = g.GetConsent(ctx, 1, client); err != nil || len(c.Scopes) != 1 {
|
||||||
|
t.Fatalf("replaced: %+v %v", c, err)
|
||||||
|
}
|
||||||
|
// Another user is separate.
|
||||||
|
if _, err := g.GetConsent(ctx, 2, client); !errors.Is(err, lookup.ErrNotFound) {
|
||||||
|
t.Fatalf("other user: %v", err)
|
||||||
|
}
|
||||||
|
// Expired consents are not returned.
|
||||||
|
if err := g.SaveConsent(ctx, lookup.Consent{UserID: 3, ClientID: client, Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(-time.Minute)}); err != nil {
|
||||||
|
t.Fatalf("save expired: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.GetConsent(ctx, 3, client); !errors.Is(err, lookup.ErrNotFound) {
|
||||||
|
t.Fatalf("expired consent: %v", err)
|
||||||
|
}
|
||||||
|
if err := g.RevokeConsent(ctx, 1, client); err != nil {
|
||||||
|
t.Fatalf("revoke: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.GetConsent(ctx, 1, client); !errors.Is(err, lookup.ErrNotFound) {
|
||||||
|
t.Fatalf("revoked consent: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *suite) oauthGrantRefresh(t *testing.T) {
|
||||||
|
g := s.Provider.OAuthGrant
|
||||||
|
client := s.name("refresh-client")
|
||||||
|
mk := func(n string, session string) lookup.RefreshToken {
|
||||||
|
return lookup.RefreshToken{TokenHash: s.name(n), FamilyID: s.name("fam-" + n), ClientID: client, UserID: 5,
|
||||||
|
SessionToken: session, Scopes: []string{"openid", "offline_access"},
|
||||||
|
Extra: map[string]any{"nonce": "abc"}, ExpiresAt: time.Now().Add(time.Hour)}
|
||||||
|
}
|
||||||
|
first := mk("r1", s.name("sess1"))
|
||||||
|
if err := g.SaveRefresh(ctx, first); err != nil {
|
||||||
|
t.Fatalf("save: %v", err)
|
||||||
|
}
|
||||||
|
peek, err := g.PeekRefresh(ctx, first.TokenHash)
|
||||||
|
if err != nil || peek.UserID != 5 || peek.ClientID != client || len(peek.Scopes) != 2 || peek.Extra["nonce"] != "abc" {
|
||||||
|
t.Fatalf("peek: %+v %v", peek, err)
|
||||||
|
}
|
||||||
|
if _, err := g.PeekRefresh(ctx, s.name("nope")); !errors.Is(err, lookup.ErrRefreshInvalid) {
|
||||||
|
t.Fatalf("peek unknown: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
next := lookup.RefreshToken{TokenHash: s.name("r2"), ExpiresAt: time.Now().Add(time.Hour), Scopes: []string{"openid"}}
|
||||||
|
old, err := g.RotateRefresh(ctx, first.TokenHash, next)
|
||||||
|
if err != nil || old.FamilyID != first.FamilyID || old.SessionToken != first.SessionToken {
|
||||||
|
t.Fatalf("rotate: %+v %v", old, err)
|
||||||
|
}
|
||||||
|
// The new token belongs to the same family, client and user.
|
||||||
|
n, err := g.PeekRefresh(ctx, next.TokenHash)
|
||||||
|
if err != nil || n.FamilyID != first.FamilyID || n.ClientID != client || n.UserID != 5 || n.SessionToken != first.SessionToken {
|
||||||
|
t.Fatalf("next: %+v %v", n, err)
|
||||||
|
}
|
||||||
|
// The consumed token is still visible to Peek, so that presenting it reaches RotateRefresh.
|
||||||
|
if _, err := g.PeekRefresh(ctx, first.TokenHash); err != nil {
|
||||||
|
t.Fatalf("peek consumed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Presenting the consumed token again is reuse: the family (including the new token) dies.
|
||||||
|
third := lookup.RefreshToken{TokenHash: s.name("r3"), ExpiresAt: time.Now().Add(time.Hour)}
|
||||||
|
reused, err := g.RotateRefresh(ctx, first.TokenHash, third)
|
||||||
|
if !errors.Is(err, lookup.ErrRefreshReused) || reused == nil || reused.FamilyID != first.FamilyID {
|
||||||
|
t.Fatalf("reuse: %+v %v", reused, err)
|
||||||
|
}
|
||||||
|
if _, err := g.PeekRefresh(ctx, next.TokenHash); !errors.Is(err, lookup.ErrRefreshInvalid) {
|
||||||
|
t.Fatalf("family survived reuse: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.RotateRefresh(ctx, next.TokenHash, lookup.RefreshToken{TokenHash: s.name("r4"), ExpiresAt: time.Now().Add(time.Hour)}); !errors.Is(err, lookup.ErrRefreshInvalid) {
|
||||||
|
t.Fatalf("rotate revoked: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.PeekRefresh(ctx, third.TokenHash); err == nil {
|
||||||
|
t.Fatal("a rejected rotation must not store the next token")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Expired tokens cannot rotate.
|
||||||
|
exp := mk("rexp", s.name("sess2"))
|
||||||
|
exp.ExpiresAt = time.Now().Add(-time.Minute)
|
||||||
|
if err := g.SaveRefresh(ctx, exp); err != nil {
|
||||||
|
t.Fatalf("save expired: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.RotateRefresh(ctx, exp.TokenHash, lookup.RefreshToken{TokenHash: s.name("rexp2"), ExpiresAt: time.Now().Add(time.Hour)}); !errors.Is(err, lookup.ErrRefreshInvalid) {
|
||||||
|
t.Fatalf("rotate expired: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Revoking by family and by session.
|
||||||
|
fam := mk("rfam", s.name("sess3"))
|
||||||
|
if err := g.SaveRefresh(ctx, fam); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := g.RevokeRefreshFamily(ctx, fam.FamilyID); err != nil {
|
||||||
|
t.Fatalf("revoke family: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.PeekRefresh(ctx, fam.TokenHash); !errors.Is(err, lookup.ErrRefreshInvalid) {
|
||||||
|
t.Fatalf("family revoked: %v", err)
|
||||||
|
}
|
||||||
|
bs := mk("rsess", s.name("sess4"))
|
||||||
|
if err := g.SaveRefresh(ctx, bs); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := g.RevokeRefreshBySession(ctx, bs.SessionToken); err != nil {
|
||||||
|
t.Fatalf("revoke session: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.PeekRefresh(ctx, bs.TokenHash); !errors.Is(err, lookup.ErrRefreshInvalid) {
|
||||||
|
t.Fatalf("session revoked: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *suite) oauthGrantDevice(t *testing.T) {
|
||||||
|
g := s.Provider.OAuthGrant
|
||||||
|
client := s.name("device-client")
|
||||||
|
mk := func(n string) lookup.DeviceCode {
|
||||||
|
return lookup.DeviceCode{DeviceHash: s.name("dh-" + n), UserCode: strings.ToUpper(s.name("uc-" + n)), ClientID: client,
|
||||||
|
Scopes: []string{"openid"}, Interval: 1, ExpiresAt: time.Now().Add(time.Minute)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// pending -> approved
|
||||||
|
d := mk("a")
|
||||||
|
if err := g.CreateDevice(ctx, d); err != nil {
|
||||||
|
t.Fatalf("create: %v", err)
|
||||||
|
}
|
||||||
|
got, err := g.DeviceByUserCode(ctx, strings.ToLower(d.UserCode)) // user codes are case-insensitive
|
||||||
|
if err != nil || got.ClientID != client || got.DeviceHash != d.DeviceHash || len(got.Scopes) != 1 {
|
||||||
|
t.Fatalf("by user code: %+v %v", got, err)
|
||||||
|
}
|
||||||
|
if _, err := g.DevicePoll(ctx, d.DeviceHash); !errors.Is(err, lookup.ErrDevicePending) {
|
||||||
|
t.Fatalf("first poll: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.DevicePoll(ctx, d.DeviceHash); !errors.Is(err, lookup.ErrDeviceSlowDown) {
|
||||||
|
t.Fatalf("immediate re-poll: %v", err)
|
||||||
|
}
|
||||||
|
if err := g.DeviceDecide(ctx, d.UserCode, true, 9, s.name("dsess")); err != nil {
|
||||||
|
t.Fatalf("approve: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.DeviceByUserCode(ctx, d.UserCode); !errors.Is(err, lookup.ErrNotFound) {
|
||||||
|
t.Fatalf("decided code still pending: %v", err)
|
||||||
|
}
|
||||||
|
if err := g.DeviceDecide(ctx, d.UserCode, true, 9, "x"); !errors.Is(err, lookup.ErrNotFound) {
|
||||||
|
t.Fatalf("second decision: %v", err)
|
||||||
|
}
|
||||||
|
time.Sleep(1100 * time.Millisecond)
|
||||||
|
done, err := g.DevicePoll(ctx, d.DeviceHash)
|
||||||
|
if err != nil || done.UserID != 9 || done.SessionToken != s.name("dsess") || done.ClientID != client {
|
||||||
|
t.Fatalf("approved poll: %+v %v", done, err)
|
||||||
|
}
|
||||||
|
if _, err := g.DevicePoll(ctx, d.DeviceHash); !errors.Is(err, lookup.ErrDeviceExpired) {
|
||||||
|
t.Fatalf("consumed code: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// denied
|
||||||
|
dd := mk("d")
|
||||||
|
if err := g.CreateDevice(ctx, dd); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := g.DeviceDecide(ctx, dd.UserCode, false, 0, ""); err != nil {
|
||||||
|
t.Fatalf("deny: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.DevicePoll(ctx, dd.DeviceHash); !errors.Is(err, lookup.ErrDeviceDenied) {
|
||||||
|
t.Fatalf("denied poll: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// expired
|
||||||
|
de := mk("e")
|
||||||
|
de.ExpiresAt = time.Now().Add(-time.Second)
|
||||||
|
if err := g.CreateDevice(ctx, de); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := g.DevicePoll(ctx, de.DeviceHash); !errors.Is(err, lookup.ErrDeviceExpired) {
|
||||||
|
t.Fatalf("expired poll: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.DeviceByUserCode(ctx, de.UserCode); !errors.Is(err, lookup.ErrNotFound) {
|
||||||
|
t.Fatalf("expired by user code: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := g.DevicePoll(ctx, s.name("unknown")); !errors.Is(err, lookup.ErrDeviceExpired) {
|
||||||
|
t.Fatalf("unknown poll: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *suite) oauthGrantPAR(t *testing.T) {
|
||||||
|
g := s.Provider.OAuthGrant
|
||||||
|
uri := "urn:ietf:params:oauth:request_uri:" + s.name("par")
|
||||||
|
if len(uri) > 255 {
|
||||||
|
t.Fatal("test request_uri too long")
|
||||||
|
}
|
||||||
|
req := lookup.PushedRequest{RequestURI: uri, ClientID: s.name("par-client"),
|
||||||
|
Params: map[string]string{"redirect_uri": "https://app.example/cb", "scope": "openid"}, ExpiresAt: time.Now().Add(time.Minute)}
|
||||||
|
if err := g.SavePushedRequest(ctx, req); err != nil {
|
||||||
|
t.Fatalf("save: %v", err)
|
||||||
|
}
|
||||||
|
got, err := g.ConsumePushedRequest(ctx, uri)
|
||||||
|
if err != nil || got.ClientID != req.ClientID || got.Params["scope"] != "openid" {
|
||||||
|
t.Fatalf("consume: %+v %v", got, err)
|
||||||
|
}
|
||||||
|
if _, err := g.ConsumePushedRequest(ctx, uri); !errors.Is(err, lookup.ErrNotFound) {
|
||||||
|
t.Fatalf("second consume: %v", err)
|
||||||
|
}
|
||||||
|
exp := lookup.PushedRequest{RequestURI: uri + "-exp", ClientID: s.name("par-client"), Params: map[string]string{"a": "b"}, ExpiresAt: time.Now().Add(-time.Minute)}
|
||||||
|
if err := g.SavePushedRequest(ctx, exp); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := g.ConsumePushedRequest(ctx, exp.RequestURI); !errors.Is(err, lookup.ErrNotFound) {
|
||||||
|
t.Fatalf("expired consume: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *suite) oauthGrantJTI(t *testing.T) {
|
||||||
|
g := s.Provider.OAuthGrant
|
||||||
|
key := s.name("jti")
|
||||||
|
seen, err := g.SeenJTI(ctx, key, time.Now().Add(time.Minute))
|
||||||
|
if err != nil || seen {
|
||||||
|
t.Fatalf("first: %v %v", seen, err)
|
||||||
|
}
|
||||||
|
seen, err = g.SeenJTI(ctx, key, time.Now().Add(time.Minute))
|
||||||
|
if err != nil || !seen {
|
||||||
|
t.Fatalf("replay: %v %v", seen, err)
|
||||||
|
}
|
||||||
|
// An expired entry is forgotten.
|
||||||
|
old := s.name("jti-old")
|
||||||
|
if _, err := g.SeenJTI(ctx, old, time.Now().Add(-time.Minute)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
seen, err = g.SeenJTI(ctx, old, time.Now().Add(time.Minute))
|
||||||
|
if err != nil || seen {
|
||||||
|
t.Fatalf("after expiry: %v %v", seen, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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
|
client_secret_hash TEXT, -- sha256 hex of the confidential-client secret; NULL for public clients
|
||||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||||
is_active BOOLEAN DEFAULT true,
|
is_active BOOLEAN DEFAULT true,
|
||||||
|
metadata jsonb, -- every other RFC 7591 field (see sectypes.OAuthServerClient)
|
||||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||||
);
|
);
|
||||||
|
ALTER TABLE oauth_clients ADD COLUMN IF NOT EXISTS metadata jsonb;
|
||||||
|
|
||||||
-- oauth_codes: short-lived authorization codes (for multi-instance deployments)
|
-- oauth_codes: short-lived authorization codes (for multi-instance deployments)
|
||||||
-- Note: client_id is stored without a foreign key so codes can be persisted even
|
-- Note: client_id is stored without a foreign key so codes can be persisted even
|
||||||
@@ -1783,8 +1785,10 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
|
|||||||
refresh_token TEXT,
|
refresh_token TEXT,
|
||||||
scopes TEXT[],
|
scopes TEXT[],
|
||||||
expires_at TIMESTAMP NOT NULL,
|
expires_at TIMESTAMP NOT NULL,
|
||||||
|
extra jsonb, -- nonce, auth_time, acr, claims, user_id, dpop_jkt ... (see sectypes.OAuthCode)
|
||||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||||
);
|
);
|
||||||
|
ALTER TABLE oauth_codes ADD COLUMN IF NOT EXISTS extra jsonb;
|
||||||
CREATE INDEX IF NOT EXISTS idx_oauth_codes_code ON oauth_codes(code);
|
CREATE INDEX IF NOT EXISTS idx_oauth_codes_code ON oauth_codes(code);
|
||||||
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
||||||
|
|
||||||
@@ -1801,7 +1805,7 @@ DECLARE
|
|||||||
BEGIN
|
BEGIN
|
||||||
v_client_id := p_request->>'client_id';
|
v_client_id := p_request->>'client_id';
|
||||||
|
|
||||||
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method)
|
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method, metadata)
|
||||||
VALUES (
|
VALUES (
|
||||||
v_client_id,
|
v_client_id,
|
||||||
ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)),
|
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->'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,
|
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', ''),
|
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;
|
RETURN QUERY SELECT true, null::text, v_row;
|
||||||
EXCEPTION WHEN OTHERS THEN
|
EXCEPTION WHEN OTHERS THEN
|
||||||
@@ -1825,7 +1830,7 @@ LANGUAGE plpgsql AS $$
|
|||||||
DECLARE
|
DECLARE
|
||||||
v_row jsonb;
|
v_row jsonb;
|
||||||
BEGIN
|
BEGIN
|
||||||
SELECT to_jsonb(oauth_clients.*)
|
SELECT (to_jsonb(oauth_clients.*) - 'metadata') || COALESCE(metadata, '{}'::jsonb)
|
||||||
INTO v_row
|
INTO v_row
|
||||||
FROM oauth_clients
|
FROM oauth_clients
|
||||||
WHERE client_id = p_client_id AND is_active = true;
|
WHERE client_id = p_client_id AND is_active = true;
|
||||||
@@ -1842,7 +1847,7 @@ CREATE OR REPLACE FUNCTION resolvespec_oauth_save_code(p_request jsonb)
|
|||||||
RETURNS TABLE(p_success bool, p_error text)
|
RETURNS TABLE(p_success bool, p_error text)
|
||||||
LANGUAGE plpgsql AS $$
|
LANGUAGE plpgsql AS $$
|
||||||
BEGIN
|
BEGIN
|
||||||
INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at)
|
INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at, extra)
|
||||||
VALUES (
|
VALUES (
|
||||||
p_request->>'code',
|
p_request->>'code',
|
||||||
p_request->>'client_id',
|
p_request->>'client_id',
|
||||||
@@ -1853,7 +1858,8 @@ BEGIN
|
|||||||
p_request->>'session_token',
|
p_request->>'session_token',
|
||||||
p_request->>'refresh_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)),
|
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;
|
RETURN QUERY SELECT true, null::text;
|
||||||
@@ -1879,7 +1885,7 @@ BEGIN
|
|||||||
'session_token', session_token,
|
'session_token', session_token,
|
||||||
'refresh_token', refresh_token,
|
'refresh_token', refresh_token,
|
||||||
'scopes', to_jsonb(scopes)
|
'scopes', to_jsonb(scopes)
|
||||||
) INTO v_row;
|
) || COALESCE(extra, '{}'::jsonb) INTO v_row;
|
||||||
|
|
||||||
IF v_row IS NULL THEN
|
IF v_row IS NULL THEN
|
||||||
RETURN QUERY SELECT false, 'invalid or expired code'::text, null::jsonb;
|
RETURN QUERY SELECT false, 'invalid or expired code'::text, null::jsonb;
|
||||||
@@ -1930,3 +1936,374 @@ BEGIN
|
|||||||
RETURN QUERY SELECT true, null::text;
|
RETURN QUERY SELECT true, null::text;
|
||||||
END;
|
END;
|
||||||
$$;
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_update_client(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_rows int;
|
||||||
|
BEGIN
|
||||||
|
UPDATE oauth_clients SET
|
||||||
|
redirect_uris = ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)),
|
||||||
|
client_name = p_request->>'client_name',
|
||||||
|
grant_types = CASE WHEN jsonb_typeof(p_request->'grant_types') = 'array' AND jsonb_array_length(p_request->'grant_types') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'grant_types')) ELSE grant_types END,
|
||||||
|
allowed_scopes = CASE WHEN jsonb_typeof(p_request->'allowed_scopes') = 'array' AND jsonb_array_length(p_request->'allowed_scopes') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'allowed_scopes')) ELSE allowed_scopes END,
|
||||||
|
client_secret_hash = NULLIF(p_request->>'client_secret_hash', ''),
|
||||||
|
token_endpoint_auth_method = COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), token_endpoint_auth_method),
|
||||||
|
metadata = NULLIF(p_request - ARRAY['client_id','redirect_uris','client_name','grant_types','allowed_scopes','client_secret_hash','token_endpoint_auth_method'], '{}'::jsonb)
|
||||||
|
WHERE client_id = p_request->>'client_id' AND is_active = true;
|
||||||
|
GET DIAGNOSTICS v_rows = ROW_COUNT;
|
||||||
|
IF v_rows = 0 THEN
|
||||||
|
RETURN QUERY SELECT false, 'client not found'::text;
|
||||||
|
ELSE
|
||||||
|
RETURN QUERY SELECT true, null::text;
|
||||||
|
END IF;
|
||||||
|
EXCEPTION WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_delete_client(p_client_id text)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
BEGIN
|
||||||
|
UPDATE oauth_clients SET is_active = false WHERE client_id = p_client_id;
|
||||||
|
RETURN QUERY SELECT true, null::text;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
-- ============================================
|
||||||
|
-- OAuth2 Server grant state (consents, refresh tokens, device codes, PAR, replay cache)
|
||||||
|
-- ============================================
|
||||||
|
-- Procedure-backend tables use jsonb for scopes/extra/params. Every procedure takes one jsonb
|
||||||
|
-- request and returns (p_success, p_error, p_data). p_error carries a stable code for the
|
||||||
|
-- failures the Go side maps to errors: not_found, refresh_invalid, refresh_reused,
|
||||||
|
-- device_pending, device_slowdown, device_denied, device_expired.
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS oauth_consents (
|
||||||
|
id SERIAL PRIMARY KEY,
|
||||||
|
user_id INTEGER NOT NULL,
|
||||||
|
client_id VARCHAR(255) NOT NULL,
|
||||||
|
scopes jsonb,
|
||||||
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
expires_at TIMESTAMP NOT NULL
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
|
||||||
|
id SERIAL PRIMARY KEY,
|
||||||
|
token_hash VARCHAR(64) NOT NULL UNIQUE, -- sha256 hex of the raw refresh token
|
||||||
|
family_id VARCHAR(64) NOT NULL,
|
||||||
|
client_id VARCHAR(255) NOT NULL,
|
||||||
|
user_id INTEGER NOT NULL,
|
||||||
|
session_token VARCHAR(255),
|
||||||
|
scopes jsonb,
|
||||||
|
extra jsonb,
|
||||||
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
expires_at TIMESTAMP NOT NULL,
|
||||||
|
used_at TIMESTAMP,
|
||||||
|
revoked_at TIMESTAMP
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS oauth_device_codes (
|
||||||
|
id SERIAL PRIMARY KEY,
|
||||||
|
device_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||||
|
user_code VARCHAR(32) NOT NULL UNIQUE,
|
||||||
|
client_id VARCHAR(255) NOT NULL,
|
||||||
|
scopes jsonb,
|
||||||
|
status VARCHAR(16) NOT NULL DEFAULT 'pending',
|
||||||
|
user_id INTEGER,
|
||||||
|
session_token VARCHAR(255),
|
||||||
|
poll_interval INTEGER NOT NULL DEFAULT 5,
|
||||||
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
expires_at TIMESTAMP NOT NULL,
|
||||||
|
last_polled_at TIMESTAMP
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_oauth_device_expires ON oauth_device_codes(expires_at);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS oauth_par_requests (
|
||||||
|
id SERIAL PRIMARY KEY,
|
||||||
|
request_uri VARCHAR(255) NOT NULL UNIQUE,
|
||||||
|
client_id VARCHAR(255) NOT NULL,
|
||||||
|
params jsonb,
|
||||||
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
expires_at TIMESTAMP NOT NULL
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_oauth_par_expires ON oauth_par_requests(expires_at);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS oauth_jti (
|
||||||
|
id SERIAL PRIMARY KEY,
|
||||||
|
jti_key VARCHAR(255) NOT NULL UNIQUE,
|
||||||
|
expires_at TIMESTAMP NOT NULL
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_oauth_jti_expires ON oauth_jti(expires_at);
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_consent(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
BEGIN
|
||||||
|
DELETE FROM oauth_consents
|
||||||
|
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id';
|
||||||
|
INSERT INTO oauth_consents (user_id, client_id, scopes, expires_at)
|
||||||
|
VALUES ((p_request->>'user_id')::int, p_request->>'client_id', COALESCE(p_request->'scopes', '[]'::jsonb),
|
||||||
|
(p_request->>'expires_at')::timestamptz::timestamp);
|
||||||
|
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||||
|
EXCEPTION WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_get_consent(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_row jsonb;
|
||||||
|
BEGIN
|
||||||
|
SELECT jsonb_build_object('user_id', user_id, 'client_id', client_id, 'scopes', COALESCE(scopes, '[]'::jsonb), 'expires_at', expires_at)
|
||||||
|
INTO v_row
|
||||||
|
FROM oauth_consents
|
||||||
|
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id' AND expires_at > now();
|
||||||
|
IF v_row IS NULL THEN
|
||||||
|
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
|
||||||
|
ELSE
|
||||||
|
RETURN QUERY SELECT true, null::text, v_row;
|
||||||
|
END IF;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_consent(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
BEGIN
|
||||||
|
DELETE FROM oauth_consents
|
||||||
|
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id';
|
||||||
|
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_refresh(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
BEGIN
|
||||||
|
INSERT INTO oauth_refresh_tokens (token_hash, family_id, client_id, user_id, session_token, scopes, extra, expires_at)
|
||||||
|
VALUES (p_request->>'token_hash', p_request->>'family_id', p_request->>'client_id', (p_request->>'user_id')::int,
|
||||||
|
p_request->>'session_token', p_request->'scopes', p_request->'extra',
|
||||||
|
(p_request->>'expires_at')::timestamptz::timestamp);
|
||||||
|
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||||
|
EXCEPTION WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_rotate_refresh(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
r oauth_refresh_tokens%ROWTYPE;
|
||||||
|
v_next jsonb := p_request->'next';
|
||||||
|
v_old jsonb;
|
||||||
|
BEGIN
|
||||||
|
SELECT * INTO r FROM oauth_refresh_tokens WHERE token_hash = p_request->>'old_hash' FOR UPDATE;
|
||||||
|
IF NOT FOUND OR r.revoked_at IS NOT NULL OR r.expires_at <= now() THEN
|
||||||
|
RETURN QUERY SELECT false, 'refresh_invalid'::text, null::jsonb;
|
||||||
|
RETURN;
|
||||||
|
END IF;
|
||||||
|
v_old := jsonb_build_object('token_hash', r.token_hash, 'family_id', r.family_id, 'client_id', r.client_id,
|
||||||
|
'user_id', r.user_id, 'session_token', r.session_token,
|
||||||
|
'scopes', COALESCE(r.scopes, '[]'::jsonb), 'extra', COALESCE(r.extra, '{}'::jsonb),
|
||||||
|
'expires_at', r.expires_at);
|
||||||
|
IF r.used_at IS NOT NULL THEN
|
||||||
|
-- A rotated token came back: revoke the whole family. Returning (not raising) keeps the revoke.
|
||||||
|
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE family_id = r.family_id AND revoked_at IS NULL;
|
||||||
|
RETURN QUERY SELECT false, 'refresh_reused'::text, v_old;
|
||||||
|
RETURN;
|
||||||
|
END IF;
|
||||||
|
UPDATE oauth_refresh_tokens SET used_at = now() WHERE id = r.id;
|
||||||
|
INSERT INTO oauth_refresh_tokens (token_hash, family_id, client_id, user_id, session_token, scopes, extra, expires_at)
|
||||||
|
VALUES (v_next->>'token_hash', r.family_id, r.client_id, r.user_id,
|
||||||
|
COALESCE(NULLIF(v_next->>'session_token', ''), r.session_token),
|
||||||
|
v_next->'scopes', v_next->'extra', (v_next->>'expires_at')::timestamptz::timestamp);
|
||||||
|
RETURN QUERY SELECT true, null::text, v_old;
|
||||||
|
EXCEPTION WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_peek_refresh(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_row jsonb;
|
||||||
|
BEGIN
|
||||||
|
SELECT jsonb_build_object('token_hash', token_hash, 'family_id', family_id, 'client_id', client_id,
|
||||||
|
'user_id', user_id, 'session_token', session_token,
|
||||||
|
'scopes', COALESCE(scopes, '[]'::jsonb), 'extra', COALESCE(extra, '{}'::jsonb),
|
||||||
|
'expires_at', expires_at)
|
||||||
|
INTO v_row
|
||||||
|
FROM oauth_refresh_tokens
|
||||||
|
WHERE token_hash = p_request->>'token_hash' AND revoked_at IS NULL AND expires_at > now();
|
||||||
|
IF v_row IS NULL THEN
|
||||||
|
RETURN QUERY SELECT false, 'refresh_invalid'::text, null::jsonb;
|
||||||
|
ELSE
|
||||||
|
RETURN QUERY SELECT true, null::text, v_row;
|
||||||
|
END IF;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_refresh_family(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
BEGIN
|
||||||
|
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE family_id = p_request->>'family_id' AND revoked_at IS NULL;
|
||||||
|
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_refresh_session(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
BEGIN
|
||||||
|
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE session_token = p_request->>'session_token' AND revoked_at IS NULL;
|
||||||
|
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_create_device(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
BEGIN
|
||||||
|
INSERT INTO oauth_device_codes (device_hash, user_code, client_id, scopes, status, poll_interval, expires_at)
|
||||||
|
VALUES (p_request->>'device_hash', upper(p_request->>'user_code'), p_request->>'client_id', p_request->'scopes',
|
||||||
|
COALESCE(NULLIF(p_request->>'status', ''), 'pending'), COALESCE((p_request->>'interval')::int, 5),
|
||||||
|
(p_request->>'expires_at')::timestamptz::timestamp);
|
||||||
|
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||||
|
EXCEPTION WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_by_user_code(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_row jsonb;
|
||||||
|
BEGIN
|
||||||
|
SELECT jsonb_build_object('device_hash', device_hash, 'user_code', user_code, 'client_id', client_id,
|
||||||
|
'scopes', COALESCE(scopes, '[]'::jsonb), 'status', status, 'interval', poll_interval,
|
||||||
|
'expires_at', expires_at)
|
||||||
|
INTO v_row
|
||||||
|
FROM oauth_device_codes
|
||||||
|
WHERE user_code = upper(p_request->>'user_code') AND status = 'pending' AND expires_at > now();
|
||||||
|
IF v_row IS NULL THEN
|
||||||
|
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
|
||||||
|
ELSE
|
||||||
|
RETURN QUERY SELECT true, null::text, v_row;
|
||||||
|
END IF;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_decide(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_rows int;
|
||||||
|
v_approve boolean := COALESCE((p_request->>'approve')::boolean, false);
|
||||||
|
BEGIN
|
||||||
|
UPDATE oauth_device_codes
|
||||||
|
SET status = CASE WHEN v_approve THEN 'approved' ELSE 'denied' END,
|
||||||
|
user_id = CASE WHEN v_approve THEN (p_request->>'user_id')::int ELSE user_id END,
|
||||||
|
session_token = CASE WHEN v_approve THEN p_request->>'session_token' ELSE session_token END
|
||||||
|
WHERE user_code = upper(p_request->>'user_code') AND status = 'pending' AND expires_at > now();
|
||||||
|
GET DIAGNOSTICS v_rows = ROW_COUNT;
|
||||||
|
IF v_rows = 0 THEN
|
||||||
|
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
|
||||||
|
ELSE
|
||||||
|
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||||
|
END IF;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_poll(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
d oauth_device_codes%ROWTYPE;
|
||||||
|
v_slow boolean;
|
||||||
|
BEGIN
|
||||||
|
SELECT * INTO d FROM oauth_device_codes WHERE device_hash = p_request->>'device_hash' FOR UPDATE;
|
||||||
|
IF NOT FOUND THEN
|
||||||
|
RETURN QUERY SELECT false, 'device_expired'::text, null::jsonb;
|
||||||
|
RETURN;
|
||||||
|
END IF;
|
||||||
|
IF d.expires_at <= now() THEN
|
||||||
|
DELETE FROM oauth_device_codes WHERE id = d.id;
|
||||||
|
RETURN QUERY SELECT false, 'device_expired'::text, null::jsonb;
|
||||||
|
RETURN;
|
||||||
|
END IF;
|
||||||
|
v_slow := d.last_polled_at IS NOT NULL AND (now() - d.last_polled_at) < make_interval(secs => d.poll_interval);
|
||||||
|
UPDATE oauth_device_codes SET last_polled_at = now() WHERE id = d.id;
|
||||||
|
IF v_slow THEN
|
||||||
|
RETURN QUERY SELECT false, 'device_slowdown'::text, null::jsonb;
|
||||||
|
ELSIF d.status = 'denied' THEN
|
||||||
|
DELETE FROM oauth_device_codes WHERE id = d.id;
|
||||||
|
RETURN QUERY SELECT false, 'device_denied'::text, null::jsonb;
|
||||||
|
ELSIF d.status = 'approved' THEN
|
||||||
|
DELETE FROM oauth_device_codes WHERE id = d.id;
|
||||||
|
RETURN QUERY SELECT true, null::text, jsonb_build_object('device_hash', d.device_hash, 'user_code', d.user_code,
|
||||||
|
'client_id', d.client_id, 'scopes', COALESCE(d.scopes, '[]'::jsonb), 'status', d.status,
|
||||||
|
'user_id', d.user_id, 'session_token', d.session_token, 'interval', d.poll_interval, 'expires_at', d.expires_at);
|
||||||
|
ELSE
|
||||||
|
RETURN QUERY SELECT false, 'device_pending'::text, null::jsonb;
|
||||||
|
END IF;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_par(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
BEGIN
|
||||||
|
INSERT INTO oauth_par_requests (request_uri, client_id, params, expires_at)
|
||||||
|
VALUES (p_request->>'request_uri', p_request->>'client_id', p_request->'params',
|
||||||
|
(p_request->>'expires_at')::timestamptz::timestamp);
|
||||||
|
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||||
|
EXCEPTION WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_consume_par(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_row jsonb;
|
||||||
|
BEGIN
|
||||||
|
DELETE FROM oauth_par_requests
|
||||||
|
WHERE request_uri = p_request->>'request_uri' AND expires_at > now()
|
||||||
|
RETURNING jsonb_build_object('request_uri', request_uri, 'client_id', client_id,
|
||||||
|
'params', COALESCE(params, '{}'::jsonb), 'expires_at', expires_at)
|
||||||
|
INTO v_row;
|
||||||
|
IF v_row IS NULL THEN
|
||||||
|
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
|
||||||
|
ELSE
|
||||||
|
RETURN QUERY SELECT true, null::text, v_row;
|
||||||
|
END IF;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_oauth_seen_jti(p_request jsonb)
|
||||||
|
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_rows int;
|
||||||
|
BEGIN
|
||||||
|
DELETE FROM oauth_jti WHERE expires_at < now();
|
||||||
|
INSERT INTO oauth_jti (jti_key, expires_at)
|
||||||
|
VALUES (p_request->>'key', (p_request->>'expires_at')::timestamptz::timestamp)
|
||||||
|
ON CONFLICT (jti_key) DO NOTHING;
|
||||||
|
GET DIAGNOSTICS v_rows = ROW_COUNT;
|
||||||
|
RETURN QUERY SELECT true, null::text, jsonb_build_object('seen', v_rows = 0);
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|||||||
@@ -147,6 +147,7 @@ CREATE TABLE oauth_clients (
|
|||||||
client_secret_hash NVARCHAR(MAX),
|
client_secret_hash NVARCHAR(MAX),
|
||||||
token_endpoint_auth_method NVARCHAR(30) DEFAULT 'none',
|
token_endpoint_auth_method NVARCHAR(30) DEFAULT 'none',
|
||||||
is_active BIT DEFAULT 1,
|
is_active BIT DEFAULT 1,
|
||||||
|
metadata NVARCHAR(MAX),
|
||||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
|
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
|
||||||
);
|
);
|
||||||
|
|
||||||
@@ -166,6 +167,7 @@ CREATE TABLE oauth_codes (
|
|||||||
refresh_token NVARCHAR(MAX),
|
refresh_token NVARCHAR(MAX),
|
||||||
scopes NVARCHAR(MAX),
|
scopes NVARCHAR(MAX),
|
||||||
expires_at DATETIME2 NOT NULL,
|
expires_at DATETIME2 NOT NULL,
|
||||||
|
extra NVARCHAR(MAX),
|
||||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
|
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);
|
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
|
-- key_hash: SHA-256 hex
|
||||||
-- scopes: JSON-encoded array
|
-- scopes: JSON-encoded array
|
||||||
-- meta: JSON-encoded object
|
-- meta: JSON-encoded object
|
||||||
|
|||||||
@@ -126,6 +126,7 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
|
|||||||
client_secret_hash TEXT,
|
client_secret_hash TEXT,
|
||||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||||
is_active TINYINT(1) DEFAULT 1,
|
is_active TINYINT(1) DEFAULT 1,
|
||||||
|
metadata TEXT,
|
||||||
created_at DATETIME NULL
|
created_at DATETIME NULL
|
||||||
);
|
);
|
||||||
|
|
||||||
@@ -144,11 +145,76 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
|
|||||||
refresh_token TEXT,
|
refresh_token TEXT,
|
||||||
scopes TEXT,
|
scopes TEXT,
|
||||||
expires_at DATETIME NOT NULL,
|
expires_at DATETIME NOT NULL,
|
||||||
|
extra TEXT,
|
||||||
created_at DATETIME NULL,
|
created_at DATETIME NULL,
|
||||||
INDEX idx_oauth_codes_expires (expires_at)
|
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
|
-- key_hash: SHA-256 hex
|
||||||
-- scopes: JSON-encoded array
|
-- scopes: JSON-encoded array
|
||||||
-- meta: JSON-encoded object
|
-- meta: JSON-encoded object
|
||||||
|
|||||||
@@ -137,6 +137,7 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
|
|||||||
client_secret_hash TEXT,
|
client_secret_hash TEXT,
|
||||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||||
is_active BOOLEAN DEFAULT true,
|
is_active BOOLEAN DEFAULT true,
|
||||||
|
metadata TEXT,
|
||||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||||
);
|
);
|
||||||
|
|
||||||
@@ -155,12 +156,84 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
|
|||||||
refresh_token TEXT,
|
refresh_token TEXT,
|
||||||
scopes TEXT,
|
scopes TEXT,
|
||||||
expires_at TIMESTAMP NOT NULL,
|
expires_at TIMESTAMP NOT NULL,
|
||||||
|
extra TEXT,
|
||||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||||
);
|
);
|
||||||
|
|
||||||
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
||||||
|
|
||||||
|
|
||||||
|
-- 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
|
-- key_hash: SHA-256 hex
|
||||||
-- scopes: JSON-encoded array
|
-- scopes: JSON-encoded array
|
||||||
-- meta: JSON-encoded object
|
-- meta: JSON-encoded object
|
||||||
|
|||||||
@@ -130,6 +130,7 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
|
|||||||
client_secret_hash TEXT,
|
client_secret_hash TEXT,
|
||||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||||
is_active BOOLEAN DEFAULT 1,
|
is_active BOOLEAN DEFAULT 1,
|
||||||
|
metadata TEXT,
|
||||||
created_at TIMESTAMP
|
created_at TIMESTAMP
|
||||||
);
|
);
|
||||||
|
|
||||||
@@ -148,12 +149,84 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
|
|||||||
refresh_token TEXT,
|
refresh_token TEXT,
|
||||||
scopes TEXT,
|
scopes TEXT,
|
||||||
expires_at TIMESTAMP NOT NULL,
|
expires_at TIMESTAMP NOT NULL,
|
||||||
|
extra TEXT,
|
||||||
created_at TIMESTAMP
|
created_at TIMESTAMP
|
||||||
);
|
);
|
||||||
|
|
||||||
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
||||||
|
|
||||||
|
|
||||||
|
-- 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
|
-- key_hash: SHA-256 hex
|
||||||
-- scopes: JSON-encoded array
|
-- scopes: JSON-encoded array
|
||||||
-- meta: JSON-encoded object
|
-- meta: JSON-encoded object
|
||||||
|
|||||||
@@ -197,6 +197,11 @@ func Gt(c lookup.Column, v any) Cond {
|
|||||||
return func(bl *builder) string { return bl.col(c) + " > " + bl.ph(v) }
|
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`.
|
// IsNull is `col IS NULL`.
|
||||||
func IsNull(c lookup.Column) Cond {
|
func IsNull(c lookup.Column) Cond {
|
||||||
return func(bl *builder) string { return bl.col(c) + " IS NULL" }
|
return func(bl *builder) string { return bl.col(c) + " IS NULL" }
|
||||||
|
|||||||
@@ -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)
|
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 {
|
err = o.do(func(q Querier) error {
|
||||||
return o.Insert(lookup.EntityOAuthClients).Set(
|
return o.Insert(lookup.EntityOAuthClients).Set(
|
||||||
Set(lookup.OAuthClientsClientID, client.ClientID),
|
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.OAuthClientsClientSecretHash, nullIfEmpty(client.ClientSecretHash)),
|
||||||
Set(lookup.OAuthClientsTokenEndpointAuthMethod, authMethod),
|
Set(lookup.OAuthClientsTokenEndpointAuthMethod, authMethod),
|
||||||
Set(lookup.OAuthClientsIsActive, true),
|
Set(lookup.OAuthClientsIsActive, true),
|
||||||
|
Set(lookup.OAuthClientsMetadata, nullIfEmpty(meta)),
|
||||||
Set(lookup.OAuthClientsCreatedAt, o.Now()),
|
Set(lookup.OAuthClientsCreatedAt, o.Now()),
|
||||||
).Exec(ctx, q)
|
).Exec(ctx, q)
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to register client: %w", err)
|
return nil, fmt.Errorf("failed to register client: %w", err)
|
||||||
}
|
}
|
||||||
return §ypes.OAuthServerClient{
|
res := *client
|
||||||
ClientID: client.ClientID,
|
res.GrantTypes = grantTypes
|
||||||
RedirectURIs: client.RedirectURIs,
|
res.AllowedScopes = allowedScopes
|
||||||
ClientName: client.ClientName,
|
res.TokenEndpointAuthMethod = authMethod
|
||||||
GrantTypes: grantTypes,
|
return &res, nil
|
||||||
AllowedScopes: allowedScopes,
|
}
|
||||||
ClientSecretHash: client.ClientSecretHash,
|
|
||||||
TokenEndpointAuthMethod: authMethod,
|
// UpdateClient implements lookup.OAuthClientStore: it rewrites the mutable registration
|
||||||
}, nil
|
// 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.
|
// GetClient implements lookup.OAuthClientStore.
|
||||||
func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
|
func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
|
||||||
var redirects, grants, scopes any
|
var redirects, grants, scopes any
|
||||||
var name, secret, method sql.NullString
|
var name, secret, method sql.NullString
|
||||||
|
var meta any
|
||||||
err := o.do(func(q Querier) error {
|
err := o.do(func(q Querier) error {
|
||||||
return o.From(lookup.EntityOAuthClients).
|
return o.From(lookup.EntityOAuthClients).
|
||||||
Cols(lookup.OAuthClientsRedirectURIs, lookup.OAuthClientsClientName, lookup.OAuthClientsGrantTypes,
|
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)).
|
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 err != nil {
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
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)
|
return nil, fmt.Errorf("failed to get client: %w", err)
|
||||||
}
|
}
|
||||||
res := §ypes.OAuthServerClient{
|
res := §ypes.OAuthServerClient{}
|
||||||
ClientID: clientID,
|
switch v := meta.(type) {
|
||||||
ClientName: name.String,
|
case []byte:
|
||||||
ClientSecretHash: secret.String,
|
_ = res.ApplyClientMetadata(string(v))
|
||||||
TokenEndpointAuthMethod: method.String,
|
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(redirects, &res.RedirectURIs)
|
||||||
_ = o.d.DecodeJSON(grants, &res.GrantTypes)
|
_ = o.d.DecodeJSON(grants, &res.GrantTypes)
|
||||||
_ = o.d.DecodeJSON(scopes, &res.AllowedScopes)
|
_ = o.d.DecodeJSON(scopes, &res.AllowedScopes)
|
||||||
@@ -126,6 +188,10 @@ func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) e
|
|||||||
if method == "" {
|
if method == "" {
|
||||||
method = "S256"
|
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.do(func(q Querier) error {
|
||||||
return o.Insert(lookup.EntityOAuthCodes).Set(
|
return o.Insert(lookup.EntityOAuthCodes).Set(
|
||||||
Set(lookup.OAuthCodesCode, code.Code),
|
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.OAuthCodesRefreshToken, code.RefreshToken),
|
||||||
Set(lookup.OAuthCodesScopes, scopes),
|
Set(lookup.OAuthCodesScopes, scopes),
|
||||||
Set(lookup.OAuthCodesExpiresAt, code.ExpiresAt),
|
Set(lookup.OAuthCodesExpiresAt, code.ExpiresAt),
|
||||||
|
Set(lookup.OAuthCodesExtra, nullIfEmpty(extra)),
|
||||||
Set(lookup.OAuthCodesCreatedAt, o.Now()),
|
Set(lookup.OAuthCodesCreatedAt, o.Now()),
|
||||||
).Exec(ctx, q)
|
).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) {
|
func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) {
|
||||||
var res sectypes.OAuthCode
|
var res sectypes.OAuthCode
|
||||||
var state, refresh sql.NullString
|
var state, refresh sql.NullString
|
||||||
var scopes any
|
var scopes, extra any
|
||||||
err := o.tx(ctx, func(q Querier) error {
|
err := o.tx(ctx, func(q Querier) error {
|
||||||
err := o.From(lookup.EntityOAuthCodes).
|
err := o.From(lookup.EntityOAuthCodes).
|
||||||
Cols(lookup.OAuthCodesClientID, lookup.OAuthCodesRedirectURI, lookup.OAuthCodesClientState,
|
Cols(lookup.OAuthCodesClientID, lookup.OAuthCodesRedirectURI, lookup.OAuthCodesClientState,
|
||||||
lookup.OAuthCodesCodeChallenge, lookup.OAuthCodesCodeChallengeMethod, lookup.OAuthCodesSessionToken,
|
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())).
|
Where(Eq(lookup.OAuthCodesCode, code), Gt(lookup.OAuthCodesExpiresAt, o.Now())).
|
||||||
QueryRow(ctx, q, &res.ClientID, &res.RedirectURI, &state, &res.CodeChallenge, &res.CodeChallengeMethod,
|
QueryRow(ctx, q, &res.ClientID, &res.RedirectURI, &state, &res.CodeChallenge, &res.CodeChallengeMethod,
|
||||||
&res.SessionToken, &refresh, &scopes)
|
&res.SessionToken, &refresh, &scopes, &extra)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -179,6 +246,12 @@ func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes
|
|||||||
res.ClientState = state.String
|
res.ClientState = state.String
|
||||||
res.RefreshToken = refresh.String
|
res.RefreshToken = refresh.String
|
||||||
_ = o.d.DecodeJSON(scopes, &res.Scopes)
|
_ = o.d.DecodeJSON(scopes, &res.Scopes)
|
||||||
|
switch v := extra.(type) {
|
||||||
|
case []byte:
|
||||||
|
_ = res.ApplyCodeExtra(string(v))
|
||||||
|
case string:
|
||||||
|
_ = res.ApplyCodeExtra(v)
|
||||||
|
}
|
||||||
return &res, nil
|
return &res, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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, ¶ms, 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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -62,6 +62,10 @@ type OAuthClientStore interface {
|
|||||||
ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error)
|
ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error)
|
||||||
Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error)
|
Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error)
|
||||||
Revoke(ctx context.Context, token string) 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.
|
// OAuthSession is the session row written after an OAuth2 client login.
|
||||||
@@ -147,6 +151,7 @@ type Provider struct {
|
|||||||
Keys KeyStore
|
Keys KeyStore
|
||||||
OAuthClient OAuthClientStore
|
OAuthClient OAuthClientStore
|
||||||
OAuthUser OAuthUserStore
|
OAuthUser OAuthUserStore
|
||||||
|
OAuthGrant OAuthGrantStore
|
||||||
Passkey PasskeyStore
|
Passkey PasskeyStore
|
||||||
TOTP TOTPStore
|
TOTP TOTPStore
|
||||||
Policy PolicyStore
|
Policy PolicyStore
|
||||||
|
|||||||
@@ -61,6 +61,8 @@ const (
|
|||||||
OpOAuthExchangeCode Op = "oauth_exchange_code"
|
OpOAuthExchangeCode Op = "oauth_exchange_code"
|
||||||
OpOAuthIntrospect Op = "oauth_introspect"
|
OpOAuthIntrospect Op = "oauth_introspect"
|
||||||
OpOAuthRevoke Op = "oauth_revoke"
|
OpOAuthRevoke Op = "oauth_revoke"
|
||||||
|
OpOAuthUpdateClient Op = "oauth_update_client"
|
||||||
|
OpOAuthDeleteClient Op = "oauth_delete_client"
|
||||||
|
|
||||||
OpOAuthGetOrCreateUser Op = "oauth_get_or_create_user"
|
OpOAuthGetOrCreateUser Op = "oauth_get_or_create_user"
|
||||||
OpOAuthCreateSession Op = "oauth_create_session"
|
OpOAuthCreateSession Op = "oauth_create_session"
|
||||||
@@ -68,6 +70,22 @@ const (
|
|||||||
OpOAuthUpdateRefreshToken Op = "oauth_update_refresh_token" //nolint:gosec // operation name, not a credential
|
OpOAuthUpdateRefreshToken Op = "oauth_update_refresh_token" //nolint:gosec // operation name, not a credential
|
||||||
OpOAuthGetUser Op = "oauth_get_user"
|
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"
|
OpPasskeyStore Op = "passkey_store"
|
||||||
OpPasskeyGet Op = "passkey_get"
|
OpPasskeyGet Op = "passkey_get"
|
||||||
OpPasskeyUpdateCounter Op = "passkey_update_counter"
|
OpPasskeyUpdateCounter Op = "passkey_update_counter"
|
||||||
@@ -144,11 +162,28 @@ func AllOps() []Op {
|
|||||||
OpOAuthExchangeCode,
|
OpOAuthExchangeCode,
|
||||||
OpOAuthIntrospect,
|
OpOAuthIntrospect,
|
||||||
OpOAuthRevoke,
|
OpOAuthRevoke,
|
||||||
|
OpOAuthUpdateClient,
|
||||||
|
OpOAuthDeleteClient,
|
||||||
OpOAuthGetOrCreateUser,
|
OpOAuthGetOrCreateUser,
|
||||||
OpOAuthCreateSession,
|
OpOAuthCreateSession,
|
||||||
OpOAuthGetRefreshToken,
|
OpOAuthGetRefreshToken,
|
||||||
OpOAuthUpdateRefreshToken,
|
OpOAuthUpdateRefreshToken,
|
||||||
OpOAuthGetUser,
|
OpOAuthGetUser,
|
||||||
|
OpOAuthSaveConsent,
|
||||||
|
OpOAuthGetConsent,
|
||||||
|
OpOAuthRevokeConsent,
|
||||||
|
OpOAuthSaveRefresh,
|
||||||
|
OpOAuthRotateRefresh,
|
||||||
|
OpOAuthPeekRefresh,
|
||||||
|
OpOAuthRevokeRefreshFamily,
|
||||||
|
OpOAuthRevokeRefreshByUser,
|
||||||
|
OpOAuthCreateDevice,
|
||||||
|
OpOAuthDeviceByUserCode,
|
||||||
|
OpOAuthDeviceDecide,
|
||||||
|
OpOAuthDevicePoll,
|
||||||
|
OpOAuthSavePAR,
|
||||||
|
OpOAuthConsumePAR,
|
||||||
|
OpOAuthSeenJTI,
|
||||||
OpPasskeyStore,
|
OpPasskeyStore,
|
||||||
OpPasskeyGet,
|
OpPasskeyGet,
|
||||||
OpPasskeyUpdateCounter,
|
OpPasskeyUpdateCounter,
|
||||||
|
|||||||
@@ -313,3 +313,38 @@ func (o *OAuthClients) Revoke(ctx context.Context, token string) error {
|
|||||||
}
|
}
|
||||||
return nil
|
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
|
||||||
|
}
|
||||||
@@ -63,6 +63,26 @@ type ProcNames struct {
|
|||||||
OAuthExchangeCode string // default: "resolvespec_oauth_exchange_code"
|
OAuthExchangeCode string // default: "resolvespec_oauth_exchange_code"
|
||||||
OAuthIntrospect string // default: "resolvespec_oauth_introspect"
|
OAuthIntrospect string // default: "resolvespec_oauth_introspect"
|
||||||
OAuthRevoke string // default: "resolvespec_oauth_revoke"
|
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)
|
// Keystore procedures (KeyStore)
|
||||||
KeystoreGetUserKeys string // default: "resolvespec_keystore_get_user_keys"
|
KeystoreGetUserKeys string // default: "resolvespec_keystore_get_user_keys"
|
||||||
@@ -112,6 +132,23 @@ func DefaultProcNames() ProcNames {
|
|||||||
OAuthExchangeCode: "resolvespec_oauth_exchange_code",
|
OAuthExchangeCode: "resolvespec_oauth_exchange_code",
|
||||||
OAuthIntrospect: "resolvespec_oauth_introspect",
|
OAuthIntrospect: "resolvespec_oauth_introspect",
|
||||||
OAuthRevoke: "resolvespec_oauth_revoke",
|
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",
|
KeystoreGetUserKeys: "resolvespec_keystore_get_user_keys",
|
||||||
KeystoreCreateKey: "resolvespec_keystore_create_key",
|
KeystoreCreateKey: "resolvespec_keystore_create_key",
|
||||||
KeystoreDeleteKey: "resolvespec_keystore_delete_key",
|
KeystoreDeleteKey: "resolvespec_keystore_delete_key",
|
||||||
|
|||||||
@@ -26,6 +26,11 @@ const (
|
|||||||
EntityUserPasswordResets Entity = "user_password_resets"
|
EntityUserPasswordResets Entity = "user_password_resets"
|
||||||
EntityOAuthClients Entity = "oauth_clients"
|
EntityOAuthClients Entity = "oauth_clients"
|
||||||
EntityOAuthCodes Entity = "oauth_codes"
|
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"
|
EntityUserKeys Entity = "user_keys"
|
||||||
EntitySecGroupMembers Entity = "sec_group_members"
|
EntitySecGroupMembers Entity = "sec_group_members"
|
||||||
EntitySecColumnRules Entity = "sec_column_rules"
|
EntitySecColumnRules Entity = "sec_column_rules"
|
||||||
@@ -121,6 +126,7 @@ var (
|
|||||||
OAuthClientsClientSecretHash = col(EntityOAuthClients, "client_secret_hash")
|
OAuthClientsClientSecretHash = col(EntityOAuthClients, "client_secret_hash")
|
||||||
OAuthClientsTokenEndpointAuthMethod = col(EntityOAuthClients, "token_endpoint_auth_method")
|
OAuthClientsTokenEndpointAuthMethod = col(EntityOAuthClients, "token_endpoint_auth_method")
|
||||||
OAuthClientsIsActive = col(EntityOAuthClients, "is_active")
|
OAuthClientsIsActive = col(EntityOAuthClients, "is_active")
|
||||||
|
OAuthClientsMetadata = col(EntityOAuthClients, "metadata")
|
||||||
OAuthClientsCreatedAt = col(EntityOAuthClients, "created_at")
|
OAuthClientsCreatedAt = col(EntityOAuthClients, "created_at")
|
||||||
|
|
||||||
OAuthCodesID = col(EntityOAuthCodes, "id")
|
OAuthCodesID = col(EntityOAuthCodes, "id")
|
||||||
@@ -135,6 +141,51 @@ var (
|
|||||||
OAuthCodesScopes = col(EntityOAuthCodes, "scopes")
|
OAuthCodesScopes = col(EntityOAuthCodes, "scopes")
|
||||||
OAuthCodesExpiresAt = col(EntityOAuthCodes, "expires_at")
|
OAuthCodesExpiresAt = col(EntityOAuthCodes, "expires_at")
|
||||||
OAuthCodesCreatedAt = col(EntityOAuthCodes, "created_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")
|
KeysID = col(EntityUserKeys, "id")
|
||||||
KeysUserID = col(EntityUserKeys, "user_id")
|
KeysUserID = col(EntityUserKeys, "user_id")
|
||||||
@@ -190,10 +241,20 @@ var allColumns = []Column{
|
|||||||
ResetsID, ResetsUserID, ResetsTokenHash, ResetsExpiresAt, ResetsCreatedAt, ResetsUsed, ResetsUsedAt,
|
ResetsID, ResetsUserID, ResetsTokenHash, ResetsExpiresAt, ResetsCreatedAt, ResetsUsed, ResetsUsedAt,
|
||||||
OAuthClientsID, OAuthClientsClientID, OAuthClientsRedirectURIs, OAuthClientsClientName, OAuthClientsGrantTypes,
|
OAuthClientsID, OAuthClientsClientID, OAuthClientsRedirectURIs, OAuthClientsClientName, OAuthClientsGrantTypes,
|
||||||
OAuthClientsAllowedScopes, OAuthClientsClientSecretHash, OAuthClientsTokenEndpointAuthMethod,
|
OAuthClientsAllowedScopes, OAuthClientsClientSecretHash, OAuthClientsTokenEndpointAuthMethod,
|
||||||
OAuthClientsIsActive, OAuthClientsCreatedAt,
|
OAuthClientsIsActive, OAuthClientsCreatedAt, OAuthClientsMetadata,
|
||||||
OAuthCodesID, OAuthCodesCode, OAuthCodesClientID, OAuthCodesRedirectURI, OAuthCodesClientState,
|
OAuthCodesID, OAuthCodesCode, OAuthCodesClientID, OAuthCodesRedirectURI, OAuthCodesClientState,
|
||||||
OAuthCodesCodeChallenge, OAuthCodesCodeChallengeMethod, OAuthCodesSessionToken, OAuthCodesRefreshToken,
|
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,
|
KeysID, KeysUserID, KeysKeyType, KeysKeyHash, KeysName, KeysScopes, KeysMeta, KeysExpiresAt,
|
||||||
KeysCreatedAt, KeysLastUsedAt, KeysIsActive,
|
KeysCreatedAt, KeysLastUsedAt, KeysIsActive,
|
||||||
GroupMembersGroupID, GroupMembersUserID,
|
GroupMembersGroupID, GroupMembersUserID,
|
||||||
|
|||||||
@@ -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
@@ -8,6 +8,8 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -32,6 +34,45 @@ type OAuth2Config struct {
|
|||||||
// Optional: Custom user info parser
|
// Optional: Custom user info parser
|
||||||
// If not provided, will use standard claims (sub, email, name)
|
// If not provided, will use standard claims (sub, email, name)
|
||||||
UserInfoParser func(userInfo map[string]any) (*UserContext, error)
|
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
|
// OAuth2Provider holds configuration and state for a single OAuth2 provider
|
||||||
@@ -40,7 +81,10 @@ type OAuth2Provider struct {
|
|||||||
userInfoURL string
|
userInfoURL string
|
||||||
userInfoParser func(userInfo map[string]any) (*UserContext, error)
|
userInfoParser func(userInfo map[string]any) (*UserContext, error)
|
||||||
providerName string
|
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
|
statesMutex sync.RWMutex
|
||||||
stopCh chan struct{} // closed to stop cleanupStates
|
stopCh chan struct{} // closed to stop cleanupStates
|
||||||
stopOnce sync.Once
|
stopOnce sync.Once
|
||||||
@@ -58,6 +102,13 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
|
|||||||
cfg.UserInfoParser = defaultOAuth2UserInfoParser
|
cfg.UserInfoParser = defaultOAuth2UserInfoParser
|
||||||
}
|
}
|
||||||
|
|
||||||
|
authStyle := oauth2.AuthStyleAutoDetect
|
||||||
|
switch cfg.AuthStyle {
|
||||||
|
case "basic":
|
||||||
|
authStyle = oauth2.AuthStyleInHeader
|
||||||
|
case "post":
|
||||||
|
authStyle = oauth2.AuthStyleInParams
|
||||||
|
}
|
||||||
provider := &OAuth2Provider{
|
provider := &OAuth2Provider{
|
||||||
config: &oauth2.Config{
|
config: &oauth2.Config{
|
||||||
ClientID: cfg.ClientID,
|
ClientID: cfg.ClientID,
|
||||||
@@ -65,15 +116,22 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
|
|||||||
RedirectURL: cfg.RedirectURL,
|
RedirectURL: cfg.RedirectURL,
|
||||||
Scopes: cfg.Scopes,
|
Scopes: cfg.Scopes,
|
||||||
Endpoint: oauth2.Endpoint{
|
Endpoint: oauth2.Endpoint{
|
||||||
AuthURL: cfg.AuthURL,
|
AuthURL: cfg.AuthURL,
|
||||||
TokenURL: cfg.TokenURL,
|
TokenURL: cfg.TokenURL,
|
||||||
|
AuthStyle: authStyle,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
userInfoURL: cfg.UserInfoURL,
|
userInfoURL: cfg.UserInfoURL,
|
||||||
userInfoParser: cfg.UserInfoParser,
|
userInfoParser: cfg.UserInfoParser,
|
||||||
providerName: cfg.ProviderName,
|
providerName: cfg.ProviderName,
|
||||||
states: make(map[string]time.Time),
|
states: make(map[string]*oauth2State),
|
||||||
stopCh: make(chan struct{}),
|
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
|
// 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
|
// OAuth2GetAuthURL returns the OAuth2 authorization URL for redirecting users
|
||||||
func (a *DatabaseAuthenticator) OAuth2GetAuthURL(providerName, state string) (string, error) {
|
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)
|
provider, err := a.getOAuth2Provider(providerName)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
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.statesMutex.Lock()
|
||||||
provider.states[state] = time.Now().Add(10 * time.Minute)
|
provider.states[state] = st
|
||||||
provider.statesMutex.Unlock()
|
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
|
// 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
|
// OAuth2HandleCallback handles the OAuth2 callback and exchanges code for token
|
||||||
func (a *DatabaseAuthenticator) OAuth2HandleCallback(ctx context.Context, providerName, code, state string) (*LoginResponse, error) {
|
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)
|
provider, err := a.getOAuth2Provider(providerName)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate state
|
// Validate state
|
||||||
if !provider.validateState(state) {
|
st, ok := provider.validateState(state)
|
||||||
|
if !ok {
|
||||||
return nil, fmt.Errorf("invalid state parameter")
|
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
|
// 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 {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to exchange code: %w", err)
|
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
|
// Fetch user info
|
||||||
client := provider.config.Client(ctx, token)
|
userInfo := map[string]any{}
|
||||||
resp, err := client.Get(provider.userInfoURL)
|
if provider.userInfoURL != "" {
|
||||||
if err != nil {
|
fetched, err := provider.fetchUserInfo(ctx, token)
|
||||||
return nil, fmt.Errorf("failed to fetch user info: %w", err)
|
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()
|
claims := map[string]any{}
|
||||||
|
for k, v := range idClaims {
|
||||||
body, err := io.ReadAll(resp.Body)
|
claims[k] = v
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to read user info: %w", err)
|
|
||||||
}
|
}
|
||||||
|
for k, v := range userInfo {
|
||||||
var userInfo map[string]any
|
claims[k] = v
|
||||||
if err := json.Unmarshal(body, &userInfo); err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to parse user info: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse user info
|
// Parse user info
|
||||||
userCtx, err := provider.userInfoParser(userInfo)
|
userCtx, err := provider.userInfoParser(claims)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
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
|
userCtx.SessionID = sessionToken
|
||||||
|
|
||||||
return &LoginResponse{
|
resp := &LoginResponse{
|
||||||
Token: sessionToken,
|
Token: sessionToken,
|
||||||
RefreshToken: token.RefreshToken,
|
RefreshToken: token.RefreshToken,
|
||||||
User: userCtx,
|
User: userCtx,
|
||||||
ExpiresIn: int64(time.Until(expiresAt).Seconds()),
|
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
|
// 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
|
// validateState validates state using in-memory storage and returns what was remembered with it.
|
||||||
func (p *OAuth2Provider) validateState(state string) bool {
|
func (p *OAuth2Provider) validateState(state string) (*oauth2State, bool) {
|
||||||
p.statesMutex.Lock()
|
p.statesMutex.Lock()
|
||||||
defer p.statesMutex.Unlock()
|
defer p.statesMutex.Unlock()
|
||||||
|
|
||||||
expiry, ok := p.states[state]
|
st, ok := p.states[state]
|
||||||
if !ok {
|
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
|
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
|
// cleanupStates removes expired states periodically
|
||||||
@@ -284,8 +468,8 @@ func (p *OAuth2Provider) cleanupStates() {
|
|||||||
}
|
}
|
||||||
p.statesMutex.Lock()
|
p.statesMutex.Lock()
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
for state, expiry := range p.states {
|
for state, st := range p.states {
|
||||||
if now.After(expiry) {
|
if now.After(st.expiry) {
|
||||||
delete(p.states, state)
|
delete(p.states, state)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -363,7 +547,7 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Use OAuth2 provider to refresh the token
|
// 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()
|
newToken, err := tokenSource.Token()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to refresh token with provider: %w", err)
|
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
|
userCtx.SessionID = newSessionToken
|
||||||
|
|
||||||
return &LoginResponse{
|
resp := &LoginResponse{
|
||||||
Token: newSessionToken,
|
Token: newSessionToken,
|
||||||
RefreshToken: newToken.RefreshToken,
|
RefreshToken: newToken.RefreshToken,
|
||||||
User: userCtx,
|
User: userCtx,
|
||||||
ExpiresIn: int64(time.Until(newToken.Expiry).Seconds()),
|
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
|
// Pre-configured OAuth2 factory methods
|
||||||
@@ -406,10 +599,14 @@ func NewGoogleAuthenticator(clientID, clientSecret, redirectURL string, db *sql.
|
|||||||
ClientSecret: clientSecret,
|
ClientSecret: clientSecret,
|
||||||
RedirectURL: redirectURL,
|
RedirectURL: redirectURL,
|
||||||
Scopes: []string{"openid", "profile", "email"},
|
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",
|
TokenURL: "https://oauth2.googleapis.com/token",
|
||||||
UserInfoURL: "https://www.googleapis.com/oauth2/v2/userinfo",
|
UserInfoURL: "https://openidconnect.googleapis.com/v1/userinfo",
|
||||||
ProviderName: "google",
|
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: "",
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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),
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
@@ -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, " ") }
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -2,6 +2,8 @@ package security
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||||
)
|
)
|
||||||
|
|
||||||
// OAuthRegisterClient persists an OAuth2 client registration.
|
// 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 {
|
func (a *DatabaseAuthenticator) OAuthRevokeToken(ctx context.Context, token string) error {
|
||||||
return a.src.get().OAuthClient.Revoke(ctx, token)
|
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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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})
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,6 +1,9 @@
|
|||||||
package sectypes
|
package sectypes
|
||||||
|
|
||||||
import "time"
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
// OAuthServerClient is a persisted RFC 7591 registered OAuth2 client.
|
// OAuthServerClient is a persisted RFC 7591 registered OAuth2 client.
|
||||||
type OAuthServerClient struct {
|
type OAuthServerClient struct {
|
||||||
@@ -11,8 +14,80 @@ type OAuthServerClient struct {
|
|||||||
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
||||||
ClientSecretHash string `json:"client_secret_hash,omitempty"`
|
ClientSecretHash string `json:"client_secret_hash,omitempty"`
|
||||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,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.
|
// OAuthCode is a short-lived authorization code.
|
||||||
type OAuthCode struct {
|
type OAuthCode struct {
|
||||||
Code string `json:"code"`
|
Code string `json:"code"`
|
||||||
@@ -25,8 +100,28 @@ type OAuthCode struct {
|
|||||||
RefreshToken string `json:"refresh_token,omitempty"`
|
RefreshToken string `json:"refresh_token,omitempty"`
|
||||||
Scopes []string `json:"scopes,omitempty"`
|
Scopes []string `json:"scopes,omitempty"`
|
||||||
ExpiresAt time.Time `json:"expires_at"`
|
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.
|
// OAuthTokenInfo is the RFC 7662 token introspection response.
|
||||||
type OAuthTokenInfo struct {
|
type OAuthTokenInfo struct {
|
||||||
Active bool `json:"active"`
|
Active bool `json:"active"`
|
||||||
@@ -37,4 +132,13 @@ type OAuthTokenInfo struct {
|
|||||||
Roles []string `json:"roles,omitempty"`
|
Roles []string `json:"roles,omitempty"`
|
||||||
Exp int64 `json:"exp,omitempty"`
|
Exp int64 `json:"exp,omitempty"`
|
||||||
Iat int64 `json:"iat,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"`
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user