mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 19:20:31 +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).
|
||||
|
||||
It includes a standards-based OAuth 2.1 / OpenID Connect authorization server (consent, rotating refresh tokens, JWT access tokens, DPoP, PAR, device grant, token exchange, logout) and an OIDC relying-party client; see [pkg/security/OAUTH2_SERVER.md](pkg/security/OAUTH2_SERVER.md).
|
||||
|
||||
#### Middleware
|
||||
|
||||
HTTP middleware collection for common tasks (CORS, logging, metrics, rate limiting, etc.).
|
||||
|
||||
@@ -179,6 +179,8 @@ It can operate as:
|
||||
- **An OAuth2 federation layer** — delegates to external providers (Google, GitHub, Microsoft, etc.)
|
||||
- **Both simultaneously**
|
||||
|
||||
> The underlying `security.OAuthServer` also supports consent, OpenID Connect, rotating refresh tokens, JWT access tokens, DPoP, PAR, the device grant and token exchange; they are opt-in `OAuthServerConfig` options described in [pkg/security/OAUTH2_SERVER.md](../security/OAUTH2_SERVER.md). The options of `resolvemcp.OAuth2Config` are unchanged.
|
||||
|
||||
### Standard endpoints served
|
||||
|
||||
| Path | Spec | Purpose |
|
||||
|
||||
@@ -4,6 +4,8 @@
|
||||
|
||||
The security package provides OAuth2 authentication support for any OAuth2-compliant provider including Google, GitHub, Microsoft, Facebook, and custom providers.
|
||||
|
||||
> **Full OAuth 2.1 / OpenID Connect**: this guide covers the plain OAuth2 client login. For the OIDC relying party (discovery, PKCE, nonce, id_token validation, logout) and the complete authorization server (consent, refresh rotation, DPoP, PAR, device grant, token exchange) see [OAUTH2_SERVER.md](OAUTH2_SERVER.md).
|
||||
|
||||
## Features
|
||||
|
||||
- **Universal OAuth2 Support**: Works with any OAuth2 provider
|
||||
@@ -14,6 +16,7 @@ The security package provides OAuth2 authentication support for any OAuth2-compl
|
||||
- **Token Refresh**: Automatic token refresh support
|
||||
- **State Validation**: Built-in CSRF protection
|
||||
- **User Auto-Creation**: Automatically creates users on first login
|
||||
- **OpenID Connect** (opt-in): `WithOIDC` discovery, PKCE, nonce and id_token validation, RP-initiated logout
|
||||
- **Unified Authentication**: OAuth2 and traditional auth share same session storage
|
||||
|
||||
## Quick Start
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# OAuth2 Refresh Token - Quick Reference
|
||||
|
||||
> This covers refreshing tokens of an upstream provider with `OAuth2RefreshToken`. For refresh tokens issued by `OAuthServer` (rotation, reuse detection, downscoping) see [OAUTH2_SERVER.md](OAUTH2_SERVER.md#refresh-token-rotation).
|
||||
|
||||
## Quick Setup (3 Steps)
|
||||
|
||||
### 1. Initialize Authenticator
|
||||
|
||||
@@ -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
|
||||
- ✅ **Stored Procedures** - Database operations use PostgreSQL stored procedures where available, for security and maintainability
|
||||
- ✅ **Direct Mode** - Portable Go/SQL fallback for SQLite, MySQL, or Postgres without the stored procedures installed — no code changes required
|
||||
- ✅ **OAuth2 Authorization Server** - Built-in OAuth 2.1 + PKCE server (RFC 8414, 7591, 7009, 7662) with login form and external provider federation
|
||||
- ✅ **OAuth2 / OpenID Connect** - Built-in OAuth 2.1 + PKCE authorization server and OIDC provider (RFC 8414, 7591/7592, 7009, 7662, 9068, 9126, 9207, 9449, 8628, 8693): consent, rotating refresh tokens, JWT access tokens, logout, federation; plus an OIDC relying-party client. See [OAUTH2_SERVER.md](OAUTH2_SERVER.md)
|
||||
- ✅ **Password Reset** - Self-service password reset with secure token generation and session invalidation
|
||||
|
||||
## Stored Procedure Architecture
|
||||
@@ -1068,6 +1068,8 @@ lookup.ProcNames{
|
||||
|
||||
## OAuth2 Authorization Server
|
||||
|
||||
> The complete guide (consent, OIDC, refresh rotation, DPoP, PAR, device grant, token exchange, logout, relying-party client) is in [OAUTH2_SERVER.md](OAUTH2_SERVER.md). The table below lists the original endpoints.
|
||||
|
||||
`OAuthServer` is a generic OAuth 2.1 + PKCE authorization server. It is not tied to any spec — `pkg/resolvemcp` uses it, but it can be used standalone with any `http.ServeMux`.
|
||||
|
||||
### Endpoints
|
||||
@@ -1229,12 +1231,13 @@ The main changes:
|
||||
|------|-------------|
|
||||
| **QUICK_REFERENCE.md** | Quick reference guide with examples |
|
||||
| **KEYSTORE.md** | Per-user auth keys and key stores |
|
||||
| **OAUTH2.md** | OAuth2 client login and the authorization server |
|
||||
| **OAUTH2.md** | OAuth2 client login |
|
||||
| **OAUTH2_SERVER.md** | OAuth 2.1 / OpenID Connect server and relying-party client (full guide) |
|
||||
| **OAUTH2_REFRESH_QUICK_REFERENCE.md** / **OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md** | OAuth2 refresh tokens |
|
||||
| **PASSKEY_QUICK_REFERENCE.md** | WebAuthn passkeys |
|
||||
| **SECURITY_FEATURES.md** | Security feature overview |
|
||||
| **breaking_changes.md** | Migration notes for the `lookup` refactor |
|
||||
| **examples.go**, **examples_funcspec.go**, **oauth2_examples.go**, **passkey_examples.go** | Working provider implementations |
|
||||
| **breaking_changes.md** | Migration notes (`lookup` refactor, full OAuth2/OIDC schema changes) |
|
||||
| **examples.go**, **examples_funcspec.go**, **oauth2_examples.go**, **oauth2_full_example.go**, **passkey_examples.go** | Working provider implementations |
|
||||
|
||||
## API Reference
|
||||
|
||||
|
||||
@@ -123,3 +123,37 @@ Behaviour changes:
|
||||
- Removed the unexported `password.go` from `pkg/security` (bcrypt helpers live in `lookup/direct`).
|
||||
- `README.md`, `KEYSTORE.md` and the root README describe `lookup.Config` instead of `QueryMode`,
|
||||
`SQLNames` and `TableNames`.
|
||||
|
||||
## Step 8: full OAuth2 / OpenID Connect
|
||||
|
||||
Full guide: [OAUTH2_SERVER.md](OAUTH2_SERVER.md). New features are opt-in; the items below are what existing installs must do or notice.
|
||||
|
||||
### Schema (existing installs)
|
||||
|
||||
Fresh installs use `lookup/database_schema.sql` or `lookup/ddl/<dialect>.sql`. Existing databases need:
|
||||
|
||||
- `ALTER TABLE oauth_clients ADD COLUMN metadata <json>` (client metadata: logout URIs, jwks, require_consent, first_party, dpop_bound, signing algs, ...)
|
||||
- `ALTER TABLE oauth_codes ADD COLUMN extra <json>` (nonce, auth_time, acr, amr, claims, user_id, dpop_jkt, resource)
|
||||
- New tables `oauth_consents`, `oauth_refresh_tokens`, `oauth_device_codes`, `oauth_par_requests`, `oauth_jti` (copy them from the schema files). Access-grant records are stored in `oauth_refresh_tokens`.
|
||||
- Postgres procedure mode: reapply `lookup/database_schema.sql` (new `resolvespec_oauth_*` functions, listed in `lookup/procs.go`).
|
||||
|
||||
`<json>` is `jsonb` on Postgres, `TEXT` on SQLite, `JSON` on MySQL and `NVARCHAR(MAX)` on SQL Server. New lookup operations and `lookup.OAuthGrantStore` (`Provider.OAuthGrant`) are added to the procedure and direct backends and to the conformance suite; custom `lookup.Config.Procs` overrides gain the new names.
|
||||
|
||||
### New API (no action needed)
|
||||
|
||||
`OAuthServerConfig` options (see the guide), `OAuthSigningKey`, `OAuthServer.RegisterTrustedClient`, `VerifyAccessToken`, `OAuthClaimsProvider`; `OIDCConfig`, `DatabaseAuthenticator.WithOIDC`, `OAuth2GetAuthURLWithOptions`, `OAuth2HandleCallbackRequest`, `OAuth2LogoutURL`; `OAuth2Config` gains `Issuer`, `JWKSURL`, `EndSessionURL`, `UsePKCE`, `AllowedAlgs`, `AuthStyle`, `HTTPClient`, `ClockSkew`. `DatabaseAuthenticator` gains `OAuthUpdateClient`, `OAuthDeleteClient`, `OAuthGetUser`, `OAuthGrants`.
|
||||
|
||||
### Behaviour changes
|
||||
|
||||
- `/oauth/introspect` and `/oauth/revoke` require client authentication. Set `AllowAnonymousIntrospection` for the old behaviour.
|
||||
- Once the `redirect_uri` is validated, authorization errors are redirected to the client (`error`, `state`, `iss`) instead of being returned as JSON. Authorization responses carry `iss` (RFC 9207).
|
||||
- Only PKCE `S256` is accepted.
|
||||
- The login form is an `html/template` page with a signed state field; direct form POSTs of earlier versions are still accepted.
|
||||
- Default grant types of a dynamically registered client include `refresh_token`.
|
||||
- Authorization-code grants mint a fresh session for the grant. Tokens saved directly with `OAuthSaveCode(SessionToken: ...)` keep working.
|
||||
- `OAuth2Provider` keeps its PKCE verifier and nonce with the `state`; `Google` preset now validates id_tokens and uses the OpenID Connect endpoints.
|
||||
- Unauthenticated `userinfo` and discovery routes are unchanged; `userinfo` also answers POST and releases only the claims the granted scopes allow.
|
||||
|
||||
### Not supported
|
||||
|
||||
`client_secret_jwt`, signed request objects, the DPoP server nonce, `c_hash`, encrypted id_tokens and `actor_token`.
|
||||
|
||||
@@ -71,6 +71,9 @@ func New(db *sql.DB, cfg lookup.Config, opts Options) (*lookup.Provider, error)
|
||||
OAuthUser: &oauthUserRouter{c: c,
|
||||
proc: procedure.NewOAuthUsers(run, res.Procs),
|
||||
direct: direct.NewOAuthUsers(base)},
|
||||
OAuthGrant: &oauthGrantRouter{c: c,
|
||||
proc: procedure.NewOAuthGrants(run, res.Procs),
|
||||
direct: direct.NewOAuthGrants(base)},
|
||||
Passkey: &passkeyRouter{c: c, proc: p, direct: direct.NewPasskey(base)},
|
||||
TOTP: &totpRouter{c: c,
|
||||
proc: procedure.NewTOTP(run, res.Procs),
|
||||
@@ -90,6 +93,7 @@ func Failed(err error) *lookup.Provider {
|
||||
Keys: &keysRouter{c: c},
|
||||
OAuthClient: &oauthClientRouter{c: c},
|
||||
OAuthUser: &oauthUserRouter{c: c},
|
||||
OAuthGrant: &oauthGrantRouter{c: c},
|
||||
Passkey: &passkeyRouter{c: c},
|
||||
TOTP: &totpRouter{c: c},
|
||||
Policy: &policyRouter{c: c},
|
||||
|
||||
@@ -113,6 +113,11 @@ func cleanup(t *testing.T, db *sql.DB, d dialect.Dialect, prefix string) {
|
||||
like := prefix + "%"
|
||||
for _, q := range []struct{ table, col string }{
|
||||
{"oauth_codes", "code"},
|
||||
{"oauth_consents", "client_id"},
|
||||
{"oauth_refresh_tokens", "client_id"},
|
||||
{"oauth_device_codes", "client_id"},
|
||||
{"oauth_par_requests", "client_id"},
|
||||
{"oauth_jti", "jti_key"},
|
||||
{"oauth_clients", "client_id"},
|
||||
{"token_blacklist", "token"},
|
||||
{"sec_column_rules", "schema_name"},
|
||||
|
||||
@@ -154,3 +154,27 @@ func TestConformanceMSSQLContainer(t *testing.T) {
|
||||
}
|
||||
runOnServer(t, "sqlserver", dsn("cf_direct"), "mssql", lookup.Config{}, true)
|
||||
}
|
||||
|
||||
// TestContainerLifecycle checks the start/stop plumbing the container tests rely on: the
|
||||
// container comes up and accepts connections, and after stop it is gone (it runs with --rm).
|
||||
func TestContainerLifecycle(t *testing.T) {
|
||||
rt := containerRuntime(t)
|
||||
port := startContainer(t, rt, "docker.io/library/postgres:16-alpine", "5432", map[string]string{"POSTGRES_PASSWORD": containerPassword})
|
||||
waitReady(t, "pgx", fmt.Sprintf("postgres://postgres:%s@127.0.0.1:%s/postgres?sslmode=disable", containerPassword, port), 90*time.Second)
|
||||
|
||||
listed := func() string {
|
||||
return run(t, 30*time.Second, rt, "ps", "-q", "--filter", "ancestor=docker.io/library/postgres:16-alpine")
|
||||
}
|
||||
id := listed()
|
||||
if id == "" {
|
||||
t.Fatal("container is not running after start")
|
||||
}
|
||||
run(t, time.Minute, rt, "stop", "-t", "2", id)
|
||||
deadline := time.Now().Add(30 * time.Second)
|
||||
for listed() != "" {
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatal("container still present after stop")
|
||||
}
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -199,6 +199,22 @@ func (r *oauthClientRouter) Revoke(ctx context.Context, token string) error {
|
||||
return st.Revoke(ctx, token)
|
||||
}
|
||||
|
||||
func (r *oauthClientRouter) UpdateClient(ctx context.Context, client *sectypes.OAuthServerClient) error {
|
||||
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthUpdateClient, r.c.procs.OAuthUpdateClient, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.UpdateClient(ctx, client)
|
||||
}
|
||||
|
||||
func (r *oauthClientRouter) DeleteClient(ctx context.Context, clientID string) error {
|
||||
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthDeleteClient, r.c.procs.OAuthDeleteClient, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.DeleteClient(ctx, clientID)
|
||||
}
|
||||
|
||||
type oauthUserRouter struct {
|
||||
c *chooser
|
||||
proc, direct lookup.OAuthUserStore
|
||||
@@ -394,3 +410,134 @@ func (r *policyRouter) RowSecurity(ctx context.Context, userRef any, schema, tab
|
||||
}
|
||||
return st.RowSecurity(ctx, userRef, schema, table)
|
||||
}
|
||||
|
||||
type oauthGrantRouter struct {
|
||||
c *chooser
|
||||
proc, direct lookup.OAuthGrantStore
|
||||
}
|
||||
|
||||
var _ lookup.OAuthGrantStore = (*oauthGrantRouter)(nil)
|
||||
|
||||
func (r *oauthGrantRouter) pick(ctx context.Context, op lookup.Op, proc string) (lookup.OAuthGrantStore, error) {
|
||||
return pick[lookup.OAuthGrantStore](r.c, ctx, op, proc, r.proc, r.direct)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) SaveConsent(ctx context.Context, c lookup.Consent) error {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthSaveConsent, r.c.procs.OAuthSaveConsent)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.SaveConsent(ctx, c)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) GetConsent(ctx context.Context, userID int, clientID string) (*lookup.Consent, error) {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthGetConsent, r.c.procs.OAuthGetConsent)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.GetConsent(ctx, userID, clientID)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) RevokeConsent(ctx context.Context, userID int, clientID string) error {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthRevokeConsent, r.c.procs.OAuthRevokeConsent)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.RevokeConsent(ctx, userID, clientID)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) SaveRefresh(ctx context.Context, t lookup.RefreshToken) error {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthSaveRefresh, r.c.procs.OAuthSaveRefresh)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.SaveRefresh(ctx, t)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) RotateRefresh(ctx context.Context, oldHash string, next lookup.RefreshToken) (*lookup.RefreshToken, error) {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthRotateRefresh, r.c.procs.OAuthRotateRefresh)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.RotateRefresh(ctx, oldHash, next)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) PeekRefresh(ctx context.Context, hash string) (*lookup.RefreshToken, error) {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthPeekRefresh, r.c.procs.OAuthPeekRefresh)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.PeekRefresh(ctx, hash)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) RevokeRefreshFamily(ctx context.Context, familyID string) error {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthRevokeRefreshFamily, r.c.procs.OAuthRevokeRefreshFamily)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.RevokeRefreshFamily(ctx, familyID)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) RevokeRefreshBySession(ctx context.Context, sessionToken string) error {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthRevokeRefreshByUser, r.c.procs.OAuthRevokeRefreshByUser)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.RevokeRefreshBySession(ctx, sessionToken)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) CreateDevice(ctx context.Context, d lookup.DeviceCode) error {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthCreateDevice, r.c.procs.OAuthCreateDevice)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.CreateDevice(ctx, d)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) DeviceByUserCode(ctx context.Context, userCode string) (*lookup.DeviceCode, error) {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthDeviceByUserCode, r.c.procs.OAuthDeviceByUserCode)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.DeviceByUserCode(ctx, userCode)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) DeviceDecide(ctx context.Context, userCode string, approve bool, userID int, sessionToken string) error {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthDeviceDecide, r.c.procs.OAuthDeviceDecide)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.DeviceDecide(ctx, userCode, approve, userID, sessionToken)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) DevicePoll(ctx context.Context, deviceHash string) (*lookup.DeviceCode, error) {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthDevicePoll, r.c.procs.OAuthDevicePoll)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.DevicePoll(ctx, deviceHash)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) SavePushedRequest(ctx context.Context, req lookup.PushedRequest) error {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthSavePAR, r.c.procs.OAuthSavePAR)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.SavePushedRequest(ctx, req)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) ConsumePushedRequest(ctx context.Context, requestURI string) (*lookup.PushedRequest, error) {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthConsumePAR, r.c.procs.OAuthConsumePAR)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.ConsumePushedRequest(ctx, requestURI)
|
||||
}
|
||||
|
||||
func (r *oauthGrantRouter) SeenJTI(ctx context.Context, key string, expires time.Time) (bool, error) {
|
||||
st, err := r.pick(ctx, lookup.OpOAuthSeenJTI, r.c.procs.OAuthSeenJTI)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return st.SeenJTI(ctx, key, expires)
|
||||
}
|
||||
|
||||
@@ -52,6 +52,12 @@ func Run(t *testing.T, env Env) {
|
||||
t.Run("OAuthClientAndCodes", s.oauthClientAndCodes)
|
||||
t.Run("OAuthIntrospectRevoke", s.oauthIntrospectRevoke)
|
||||
t.Run("OAuthUsers", s.oauthUsers)
|
||||
t.Run("OAuthClientMetadata", s.oauthClientMetadata)
|
||||
t.Run("OAuthGrantConsent", s.oauthGrantConsent)
|
||||
t.Run("OAuthGrantRefresh", s.oauthGrantRefresh)
|
||||
t.Run("OAuthGrantDevice", s.oauthGrantDevice)
|
||||
t.Run("OAuthGrantPAR", s.oauthGrantPAR)
|
||||
t.Run("OAuthGrantJTI", s.oauthGrantJTI)
|
||||
t.Run("Passkey", s.passkey)
|
||||
t.Run("TOTP", s.totp)
|
||||
t.Run("Policy", s.policy)
|
||||
@@ -562,3 +568,305 @@ func (s *suite) policy(t *testing.T) {
|
||||
}
|
||||
|
||||
func ptr[T any](v T) *T { return &v }
|
||||
|
||||
// --- OAuth server grant state -------------------------------------------------------------
|
||||
|
||||
func (s *suite) oauthClientMetadata(t *testing.T) {
|
||||
st := s.Provider.OAuthClient
|
||||
id := s.name("meta-client")
|
||||
reg, err := st.RegisterClient(ctx, §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
|
||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||
is_active BOOLEAN DEFAULT true,
|
||||
metadata jsonb, -- every other RFC 7591 field (see sectypes.OAuthServerClient)
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
ALTER TABLE oauth_clients ADD COLUMN IF NOT EXISTS metadata jsonb;
|
||||
|
||||
-- oauth_codes: short-lived authorization codes (for multi-instance deployments)
|
||||
-- Note: client_id is stored without a foreign key so codes can be persisted even
|
||||
@@ -1783,8 +1785,10 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
|
||||
refresh_token TEXT,
|
||||
scopes TEXT[],
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
extra jsonb, -- nonce, auth_time, acr, claims, user_id, dpop_jkt ... (see sectypes.OAuthCode)
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
ALTER TABLE oauth_codes ADD COLUMN IF NOT EXISTS extra jsonb;
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_codes_code ON oauth_codes(code);
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
||||
|
||||
@@ -1801,7 +1805,7 @@ DECLARE
|
||||
BEGIN
|
||||
v_client_id := p_request->>'client_id';
|
||||
|
||||
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method)
|
||||
INSERT INTO oauth_clients (client_id, redirect_uris, client_name, grant_types, allowed_scopes, client_secret_hash, token_endpoint_auth_method, metadata)
|
||||
VALUES (
|
||||
v_client_id,
|
||||
ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)),
|
||||
@@ -1809,9 +1813,10 @@ BEGIN
|
||||
CASE WHEN jsonb_typeof(p_request->'grant_types') = 'array' AND jsonb_array_length(p_request->'grant_types') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'grant_types')) ELSE ARRAY['authorization_code'] END,
|
||||
CASE WHEN jsonb_typeof(p_request->'allowed_scopes') = 'array' AND jsonb_array_length(p_request->'allowed_scopes') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'allowed_scopes')) ELSE ARRAY['openid','profile','email'] END,
|
||||
NULLIF(p_request->>'client_secret_hash', ''),
|
||||
COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), 'none')
|
||||
COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), 'none'),
|
||||
NULLIF(p_request - ARRAY['client_id','redirect_uris','client_name','grant_types','allowed_scopes','client_secret_hash','token_endpoint_auth_method'], '{}'::jsonb)
|
||||
)
|
||||
RETURNING to_jsonb(oauth_clients.*) INTO v_row;
|
||||
RETURNING (to_jsonb(oauth_clients.*) - 'metadata') || COALESCE(metadata, '{}'::jsonb) INTO v_row;
|
||||
|
||||
RETURN QUERY SELECT true, null::text, v_row;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
@@ -1825,7 +1830,7 @@ LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_row jsonb;
|
||||
BEGIN
|
||||
SELECT to_jsonb(oauth_clients.*)
|
||||
SELECT (to_jsonb(oauth_clients.*) - 'metadata') || COALESCE(metadata, '{}'::jsonb)
|
||||
INTO v_row
|
||||
FROM oauth_clients
|
||||
WHERE client_id = p_client_id AND is_active = true;
|
||||
@@ -1842,7 +1847,7 @@ CREATE OR REPLACE FUNCTION resolvespec_oauth_save_code(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text)
|
||||
LANGUAGE plpgsql AS $$
|
||||
BEGIN
|
||||
INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at)
|
||||
INSERT INTO oauth_codes (code, client_id, redirect_uri, client_state, code_challenge, code_challenge_method, session_token, refresh_token, scopes, expires_at, extra)
|
||||
VALUES (
|
||||
p_request->>'code',
|
||||
p_request->>'client_id',
|
||||
@@ -1853,7 +1858,8 @@ BEGIN
|
||||
p_request->>'session_token',
|
||||
p_request->>'refresh_token',
|
||||
ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'scopes') = 'array' THEN p_request->'scopes' ELSE '[]'::jsonb END)),
|
||||
(p_request->>'expires_at')::timestamptz::timestamp
|
||||
(p_request->>'expires_at')::timestamptz::timestamp,
|
||||
NULLIF(p_request - ARRAY['code','client_id','redirect_uri','client_state','code_challenge','code_challenge_method','session_token','refresh_token','scopes','expires_at'], '{}'::jsonb)
|
||||
);
|
||||
|
||||
RETURN QUERY SELECT true, null::text;
|
||||
@@ -1879,7 +1885,7 @@ BEGIN
|
||||
'session_token', session_token,
|
||||
'refresh_token', refresh_token,
|
||||
'scopes', to_jsonb(scopes)
|
||||
) INTO v_row;
|
||||
) || COALESCE(extra, '{}'::jsonb) INTO v_row;
|
||||
|
||||
IF v_row IS NULL THEN
|
||||
RETURN QUERY SELECT false, 'invalid or expired code'::text, null::jsonb;
|
||||
@@ -1930,3 +1936,374 @@ BEGIN
|
||||
RETURN QUERY SELECT true, null::text;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_update_client(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_rows int;
|
||||
BEGIN
|
||||
UPDATE oauth_clients SET
|
||||
redirect_uris = ARRAY(SELECT jsonb_array_elements_text(CASE WHEN jsonb_typeof(p_request->'redirect_uris') = 'array' THEN p_request->'redirect_uris' ELSE '[]'::jsonb END)),
|
||||
client_name = p_request->>'client_name',
|
||||
grant_types = CASE WHEN jsonb_typeof(p_request->'grant_types') = 'array' AND jsonb_array_length(p_request->'grant_types') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'grant_types')) ELSE grant_types END,
|
||||
allowed_scopes = CASE WHEN jsonb_typeof(p_request->'allowed_scopes') = 'array' AND jsonb_array_length(p_request->'allowed_scopes') > 0 THEN ARRAY(SELECT jsonb_array_elements_text(p_request->'allowed_scopes')) ELSE allowed_scopes END,
|
||||
client_secret_hash = NULLIF(p_request->>'client_secret_hash', ''),
|
||||
token_endpoint_auth_method = COALESCE(NULLIF(p_request->>'token_endpoint_auth_method', ''), token_endpoint_auth_method),
|
||||
metadata = NULLIF(p_request - ARRAY['client_id','redirect_uris','client_name','grant_types','allowed_scopes','client_secret_hash','token_endpoint_auth_method'], '{}'::jsonb)
|
||||
WHERE client_id = p_request->>'client_id' AND is_active = true;
|
||||
GET DIAGNOSTICS v_rows = ROW_COUNT;
|
||||
IF v_rows = 0 THEN
|
||||
RETURN QUERY SELECT false, 'client not found'::text;
|
||||
ELSE
|
||||
RETURN QUERY SELECT true, null::text;
|
||||
END IF;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_delete_client(p_client_id text)
|
||||
RETURNS TABLE(p_success bool, p_error text)
|
||||
LANGUAGE plpgsql AS $$
|
||||
BEGIN
|
||||
UPDATE oauth_clients SET is_active = false WHERE client_id = p_client_id;
|
||||
RETURN QUERY SELECT true, null::text;
|
||||
END;
|
||||
$$;
|
||||
|
||||
-- ============================================
|
||||
-- OAuth2 Server grant state (consents, refresh tokens, device codes, PAR, replay cache)
|
||||
-- ============================================
|
||||
-- Procedure-backend tables use jsonb for scopes/extra/params. Every procedure takes one jsonb
|
||||
-- request and returns (p_success, p_error, p_data). p_error carries a stable code for the
|
||||
-- failures the Go side maps to errors: not_found, refresh_invalid, refresh_reused,
|
||||
-- device_pending, device_slowdown, device_denied, device_expired.
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_consents (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
scopes jsonb,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
|
||||
id SERIAL PRIMARY KEY,
|
||||
token_hash VARCHAR(64) NOT NULL UNIQUE, -- sha256 hex of the raw refresh token
|
||||
family_id VARCHAR(64) NOT NULL,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
user_id INTEGER NOT NULL,
|
||||
session_token VARCHAR(255),
|
||||
scopes jsonb,
|
||||
extra jsonb,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
used_at TIMESTAMP,
|
||||
revoked_at TIMESTAMP
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_device_codes (
|
||||
id SERIAL PRIMARY KEY,
|
||||
device_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
user_code VARCHAR(32) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
scopes jsonb,
|
||||
status VARCHAR(16) NOT NULL DEFAULT 'pending',
|
||||
user_id INTEGER,
|
||||
session_token VARCHAR(255),
|
||||
poll_interval INTEGER NOT NULL DEFAULT 5,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
last_polled_at TIMESTAMP
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_device_expires ON oauth_device_codes(expires_at);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_par_requests (
|
||||
id SERIAL PRIMARY KEY,
|
||||
request_uri VARCHAR(255) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
params jsonb,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_par_expires ON oauth_par_requests(expires_at);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_jti (
|
||||
id SERIAL PRIMARY KEY,
|
||||
jti_key VARCHAR(255) NOT NULL UNIQUE,
|
||||
expires_at TIMESTAMP NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_jti_expires ON oauth_jti(expires_at);
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_consent(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
BEGIN
|
||||
DELETE FROM oauth_consents
|
||||
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id';
|
||||
INSERT INTO oauth_consents (user_id, client_id, scopes, expires_at)
|
||||
VALUES ((p_request->>'user_id')::int, p_request->>'client_id', COALESCE(p_request->'scopes', '[]'::jsonb),
|
||||
(p_request->>'expires_at')::timestamptz::timestamp);
|
||||
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_get_consent(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_row jsonb;
|
||||
BEGIN
|
||||
SELECT jsonb_build_object('user_id', user_id, 'client_id', client_id, 'scopes', COALESCE(scopes, '[]'::jsonb), 'expires_at', expires_at)
|
||||
INTO v_row
|
||||
FROM oauth_consents
|
||||
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id' AND expires_at > now();
|
||||
IF v_row IS NULL THEN
|
||||
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
|
||||
ELSE
|
||||
RETURN QUERY SELECT true, null::text, v_row;
|
||||
END IF;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_consent(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
BEGIN
|
||||
DELETE FROM oauth_consents
|
||||
WHERE user_id = (p_request->>'user_id')::int AND client_id = p_request->>'client_id';
|
||||
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_refresh(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
BEGIN
|
||||
INSERT INTO oauth_refresh_tokens (token_hash, family_id, client_id, user_id, session_token, scopes, extra, expires_at)
|
||||
VALUES (p_request->>'token_hash', p_request->>'family_id', p_request->>'client_id', (p_request->>'user_id')::int,
|
||||
p_request->>'session_token', p_request->'scopes', p_request->'extra',
|
||||
(p_request->>'expires_at')::timestamptz::timestamp);
|
||||
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_rotate_refresh(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
r oauth_refresh_tokens%ROWTYPE;
|
||||
v_next jsonb := p_request->'next';
|
||||
v_old jsonb;
|
||||
BEGIN
|
||||
SELECT * INTO r FROM oauth_refresh_tokens WHERE token_hash = p_request->>'old_hash' FOR UPDATE;
|
||||
IF NOT FOUND OR r.revoked_at IS NOT NULL OR r.expires_at <= now() THEN
|
||||
RETURN QUERY SELECT false, 'refresh_invalid'::text, null::jsonb;
|
||||
RETURN;
|
||||
END IF;
|
||||
v_old := jsonb_build_object('token_hash', r.token_hash, 'family_id', r.family_id, 'client_id', r.client_id,
|
||||
'user_id', r.user_id, 'session_token', r.session_token,
|
||||
'scopes', COALESCE(r.scopes, '[]'::jsonb), 'extra', COALESCE(r.extra, '{}'::jsonb),
|
||||
'expires_at', r.expires_at);
|
||||
IF r.used_at IS NOT NULL THEN
|
||||
-- A rotated token came back: revoke the whole family. Returning (not raising) keeps the revoke.
|
||||
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE family_id = r.family_id AND revoked_at IS NULL;
|
||||
RETURN QUERY SELECT false, 'refresh_reused'::text, v_old;
|
||||
RETURN;
|
||||
END IF;
|
||||
UPDATE oauth_refresh_tokens SET used_at = now() WHERE id = r.id;
|
||||
INSERT INTO oauth_refresh_tokens (token_hash, family_id, client_id, user_id, session_token, scopes, extra, expires_at)
|
||||
VALUES (v_next->>'token_hash', r.family_id, r.client_id, r.user_id,
|
||||
COALESCE(NULLIF(v_next->>'session_token', ''), r.session_token),
|
||||
v_next->'scopes', v_next->'extra', (v_next->>'expires_at')::timestamptz::timestamp);
|
||||
RETURN QUERY SELECT true, null::text, v_old;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_peek_refresh(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_row jsonb;
|
||||
BEGIN
|
||||
SELECT jsonb_build_object('token_hash', token_hash, 'family_id', family_id, 'client_id', client_id,
|
||||
'user_id', user_id, 'session_token', session_token,
|
||||
'scopes', COALESCE(scopes, '[]'::jsonb), 'extra', COALESCE(extra, '{}'::jsonb),
|
||||
'expires_at', expires_at)
|
||||
INTO v_row
|
||||
FROM oauth_refresh_tokens
|
||||
WHERE token_hash = p_request->>'token_hash' AND revoked_at IS NULL AND expires_at > now();
|
||||
IF v_row IS NULL THEN
|
||||
RETURN QUERY SELECT false, 'refresh_invalid'::text, null::jsonb;
|
||||
ELSE
|
||||
RETURN QUERY SELECT true, null::text, v_row;
|
||||
END IF;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_refresh_family(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
BEGIN
|
||||
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE family_id = p_request->>'family_id' AND revoked_at IS NULL;
|
||||
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_revoke_refresh_session(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
BEGIN
|
||||
UPDATE oauth_refresh_tokens SET revoked_at = now() WHERE session_token = p_request->>'session_token' AND revoked_at IS NULL;
|
||||
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_create_device(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
BEGIN
|
||||
INSERT INTO oauth_device_codes (device_hash, user_code, client_id, scopes, status, poll_interval, expires_at)
|
||||
VALUES (p_request->>'device_hash', upper(p_request->>'user_code'), p_request->>'client_id', p_request->'scopes',
|
||||
COALESCE(NULLIF(p_request->>'status', ''), 'pending'), COALESCE((p_request->>'interval')::int, 5),
|
||||
(p_request->>'expires_at')::timestamptz::timestamp);
|
||||
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_by_user_code(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_row jsonb;
|
||||
BEGIN
|
||||
SELECT jsonb_build_object('device_hash', device_hash, 'user_code', user_code, 'client_id', client_id,
|
||||
'scopes', COALESCE(scopes, '[]'::jsonb), 'status', status, 'interval', poll_interval,
|
||||
'expires_at', expires_at)
|
||||
INTO v_row
|
||||
FROM oauth_device_codes
|
||||
WHERE user_code = upper(p_request->>'user_code') AND status = 'pending' AND expires_at > now();
|
||||
IF v_row IS NULL THEN
|
||||
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
|
||||
ELSE
|
||||
RETURN QUERY SELECT true, null::text, v_row;
|
||||
END IF;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_decide(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_rows int;
|
||||
v_approve boolean := COALESCE((p_request->>'approve')::boolean, false);
|
||||
BEGIN
|
||||
UPDATE oauth_device_codes
|
||||
SET status = CASE WHEN v_approve THEN 'approved' ELSE 'denied' END,
|
||||
user_id = CASE WHEN v_approve THEN (p_request->>'user_id')::int ELSE user_id END,
|
||||
session_token = CASE WHEN v_approve THEN p_request->>'session_token' ELSE session_token END
|
||||
WHERE user_code = upper(p_request->>'user_code') AND status = 'pending' AND expires_at > now();
|
||||
GET DIAGNOSTICS v_rows = ROW_COUNT;
|
||||
IF v_rows = 0 THEN
|
||||
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
|
||||
ELSE
|
||||
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||
END IF;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_device_poll(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
d oauth_device_codes%ROWTYPE;
|
||||
v_slow boolean;
|
||||
BEGIN
|
||||
SELECT * INTO d FROM oauth_device_codes WHERE device_hash = p_request->>'device_hash' FOR UPDATE;
|
||||
IF NOT FOUND THEN
|
||||
RETURN QUERY SELECT false, 'device_expired'::text, null::jsonb;
|
||||
RETURN;
|
||||
END IF;
|
||||
IF d.expires_at <= now() THEN
|
||||
DELETE FROM oauth_device_codes WHERE id = d.id;
|
||||
RETURN QUERY SELECT false, 'device_expired'::text, null::jsonb;
|
||||
RETURN;
|
||||
END IF;
|
||||
v_slow := d.last_polled_at IS NOT NULL AND (now() - d.last_polled_at) < make_interval(secs => d.poll_interval);
|
||||
UPDATE oauth_device_codes SET last_polled_at = now() WHERE id = d.id;
|
||||
IF v_slow THEN
|
||||
RETURN QUERY SELECT false, 'device_slowdown'::text, null::jsonb;
|
||||
ELSIF d.status = 'denied' THEN
|
||||
DELETE FROM oauth_device_codes WHERE id = d.id;
|
||||
RETURN QUERY SELECT false, 'device_denied'::text, null::jsonb;
|
||||
ELSIF d.status = 'approved' THEN
|
||||
DELETE FROM oauth_device_codes WHERE id = d.id;
|
||||
RETURN QUERY SELECT true, null::text, jsonb_build_object('device_hash', d.device_hash, 'user_code', d.user_code,
|
||||
'client_id', d.client_id, 'scopes', COALESCE(d.scopes, '[]'::jsonb), 'status', d.status,
|
||||
'user_id', d.user_id, 'session_token', d.session_token, 'interval', d.poll_interval, 'expires_at', d.expires_at);
|
||||
ELSE
|
||||
RETURN QUERY SELECT false, 'device_pending'::text, null::jsonb;
|
||||
END IF;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_save_par(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
BEGIN
|
||||
INSERT INTO oauth_par_requests (request_uri, client_id, params, expires_at)
|
||||
VALUES (p_request->>'request_uri', p_request->>'client_id', p_request->'params',
|
||||
(p_request->>'expires_at')::timestamptz::timestamp);
|
||||
RETURN QUERY SELECT true, null::text, null::jsonb;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM, null::jsonb;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_consume_par(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_row jsonb;
|
||||
BEGIN
|
||||
DELETE FROM oauth_par_requests
|
||||
WHERE request_uri = p_request->>'request_uri' AND expires_at > now()
|
||||
RETURNING jsonb_build_object('request_uri', request_uri, 'client_id', client_id,
|
||||
'params', COALESCE(params, '{}'::jsonb), 'expires_at', expires_at)
|
||||
INTO v_row;
|
||||
IF v_row IS NULL THEN
|
||||
RETURN QUERY SELECT false, 'not_found'::text, null::jsonb;
|
||||
ELSE
|
||||
RETURN QUERY SELECT true, null::text, v_row;
|
||||
END IF;
|
||||
END;
|
||||
$$;
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_oauth_seen_jti(p_request jsonb)
|
||||
RETURNS TABLE(p_success bool, p_error text, p_data jsonb)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_rows int;
|
||||
BEGIN
|
||||
DELETE FROM oauth_jti WHERE expires_at < now();
|
||||
INSERT INTO oauth_jti (jti_key, expires_at)
|
||||
VALUES (p_request->>'key', (p_request->>'expires_at')::timestamptz::timestamp)
|
||||
ON CONFLICT (jti_key) DO NOTHING;
|
||||
GET DIAGNOSTICS v_rows = ROW_COUNT;
|
||||
RETURN QUERY SELECT true, null::text, jsonb_build_object('seen', v_rows = 0);
|
||||
END;
|
||||
$$;
|
||||
|
||||
@@ -147,6 +147,7 @@ CREATE TABLE oauth_clients (
|
||||
client_secret_hash NVARCHAR(MAX),
|
||||
token_endpoint_auth_method NVARCHAR(30) DEFAULT 'none',
|
||||
is_active BIT DEFAULT 1,
|
||||
metadata NVARCHAR(MAX),
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
|
||||
);
|
||||
|
||||
@@ -166,6 +167,7 @@ CREATE TABLE oauth_codes (
|
||||
refresh_token NVARCHAR(MAX),
|
||||
scopes NVARCHAR(MAX),
|
||||
expires_at DATETIME2 NOT NULL,
|
||||
extra NVARCHAR(MAX),
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
|
||||
);
|
||||
|
||||
@@ -173,6 +175,89 @@ IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_codes_expires'
|
||||
CREATE INDEX idx_oauth_codes_expires ON oauth_codes(expires_at);
|
||||
|
||||
|
||||
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
|
||||
|
||||
IF OBJECT_ID(N'oauth_consents', N'U') IS NULL
|
||||
CREATE TABLE oauth_consents (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
user_id INT NOT NULL,
|
||||
client_id NVARCHAR(255) NOT NULL,
|
||||
scopes NVARCHAR(MAX),
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
expires_at DATETIME2 NOT NULL
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_consents_user_client' AND object_id = OBJECT_ID(N'oauth_consents'))
|
||||
CREATE INDEX idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
|
||||
|
||||
IF OBJECT_ID(N'oauth_refresh_tokens', N'U') IS NULL
|
||||
CREATE TABLE oauth_refresh_tokens (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
token_hash NVARCHAR(64) NOT NULL UNIQUE,
|
||||
family_id NVARCHAR(64) NOT NULL,
|
||||
client_id NVARCHAR(255) NOT NULL,
|
||||
user_id INT NOT NULL,
|
||||
session_token NVARCHAR(255),
|
||||
scopes NVARCHAR(MAX),
|
||||
extra NVARCHAR(MAX),
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
expires_at DATETIME2 NOT NULL,
|
||||
used_at DATETIME2,
|
||||
revoked_at DATETIME2
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_refresh_family' AND object_id = OBJECT_ID(N'oauth_refresh_tokens'))
|
||||
CREATE INDEX idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_refresh_session' AND object_id = OBJECT_ID(N'oauth_refresh_tokens'))
|
||||
CREATE INDEX idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_refresh_expires' AND object_id = OBJECT_ID(N'oauth_refresh_tokens'))
|
||||
CREATE INDEX idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
|
||||
|
||||
IF OBJECT_ID(N'oauth_device_codes', N'U') IS NULL
|
||||
CREATE TABLE oauth_device_codes (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
device_hash NVARCHAR(64) NOT NULL UNIQUE,
|
||||
user_code NVARCHAR(32) NOT NULL UNIQUE,
|
||||
client_id NVARCHAR(255) NOT NULL,
|
||||
scopes NVARCHAR(MAX),
|
||||
status NVARCHAR(16) NOT NULL DEFAULT 'pending',
|
||||
user_id INT,
|
||||
session_token NVARCHAR(255),
|
||||
poll_interval INT NOT NULL DEFAULT 5,
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
expires_at DATETIME2 NOT NULL,
|
||||
last_polled_at DATETIME2
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_device_expires' AND object_id = OBJECT_ID(N'oauth_device_codes'))
|
||||
CREATE INDEX idx_oauth_device_expires ON oauth_device_codes(expires_at);
|
||||
|
||||
IF OBJECT_ID(N'oauth_par_requests', N'U') IS NULL
|
||||
CREATE TABLE oauth_par_requests (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
request_uri NVARCHAR(255) NOT NULL UNIQUE,
|
||||
client_id NVARCHAR(255) NOT NULL,
|
||||
params NVARCHAR(MAX),
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
expires_at DATETIME2 NOT NULL
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_par_expires' AND object_id = OBJECT_ID(N'oauth_par_requests'))
|
||||
CREATE INDEX idx_oauth_par_expires ON oauth_par_requests(expires_at);
|
||||
|
||||
IF OBJECT_ID(N'oauth_jti', N'U') IS NULL
|
||||
CREATE TABLE oauth_jti (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
jti_key NVARCHAR(255) NOT NULL UNIQUE,
|
||||
expires_at DATETIME2 NOT NULL
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_jti_expires' AND object_id = OBJECT_ID(N'oauth_jti'))
|
||||
CREATE INDEX idx_oauth_jti_expires ON oauth_jti(expires_at);
|
||||
|
||||
|
||||
-- key_hash: SHA-256 hex
|
||||
-- scopes: JSON-encoded array
|
||||
-- meta: JSON-encoded object
|
||||
|
||||
@@ -126,6 +126,7 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||
client_secret_hash TEXT,
|
||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||
is_active TINYINT(1) DEFAULT 1,
|
||||
metadata TEXT,
|
||||
created_at DATETIME NULL
|
||||
);
|
||||
|
||||
@@ -144,11 +145,76 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
|
||||
refresh_token TEXT,
|
||||
scopes TEXT,
|
||||
expires_at DATETIME NOT NULL,
|
||||
extra TEXT,
|
||||
created_at DATETIME NULL,
|
||||
INDEX idx_oauth_codes_expires (expires_at)
|
||||
);
|
||||
|
||||
|
||||
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_consents (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
user_id INT NOT NULL,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
scopes TEXT,
|
||||
created_at DATETIME NULL,
|
||||
expires_at DATETIME NOT NULL,
|
||||
INDEX idx_oauth_consents_user_client (user_id, client_id)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
token_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
family_id VARCHAR(64) NOT NULL,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
user_id INT NOT NULL,
|
||||
session_token VARCHAR(255),
|
||||
scopes TEXT,
|
||||
extra TEXT,
|
||||
created_at DATETIME NULL,
|
||||
expires_at DATETIME NOT NULL,
|
||||
used_at DATETIME,
|
||||
revoked_at DATETIME,
|
||||
INDEX idx_oauth_refresh_family (family_id),
|
||||
INDEX idx_oauth_refresh_session (session_token),
|
||||
INDEX idx_oauth_refresh_expires (expires_at)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_device_codes (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
device_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
user_code VARCHAR(32) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
scopes TEXT,
|
||||
status VARCHAR(16) NOT NULL DEFAULT 'pending',
|
||||
user_id INT,
|
||||
session_token VARCHAR(255),
|
||||
poll_interval INT NOT NULL DEFAULT 5,
|
||||
created_at DATETIME NULL,
|
||||
expires_at DATETIME NOT NULL,
|
||||
last_polled_at DATETIME,
|
||||
INDEX idx_oauth_device_expires (expires_at)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_par_requests (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
request_uri VARCHAR(255) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
params TEXT,
|
||||
created_at DATETIME NULL,
|
||||
expires_at DATETIME NOT NULL,
|
||||
INDEX idx_oauth_par_expires (expires_at)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_jti (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
jti_key VARCHAR(255) NOT NULL UNIQUE,
|
||||
expires_at DATETIME NOT NULL,
|
||||
INDEX idx_oauth_jti_expires (expires_at)
|
||||
);
|
||||
|
||||
|
||||
-- key_hash: SHA-256 hex
|
||||
-- scopes: JSON-encoded array
|
||||
-- meta: JSON-encoded object
|
||||
|
||||
@@ -137,6 +137,7 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||
client_secret_hash TEXT,
|
||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||
is_active BOOLEAN DEFAULT true,
|
||||
metadata TEXT,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
@@ -155,12 +156,84 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
|
||||
refresh_token TEXT,
|
||||
scopes TEXT,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
extra TEXT,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
||||
|
||||
|
||||
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_consents (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
scopes TEXT,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
|
||||
id SERIAL PRIMARY KEY,
|
||||
token_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
family_id VARCHAR(64) NOT NULL,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
user_id INTEGER NOT NULL,
|
||||
session_token VARCHAR(255),
|
||||
scopes TEXT,
|
||||
extra TEXT,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
used_at TIMESTAMP,
|
||||
revoked_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_device_codes (
|
||||
id SERIAL PRIMARY KEY,
|
||||
device_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
user_code VARCHAR(32) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
scopes TEXT,
|
||||
status VARCHAR(16) NOT NULL DEFAULT 'pending',
|
||||
user_id INTEGER,
|
||||
session_token VARCHAR(255),
|
||||
poll_interval INTEGER NOT NULL DEFAULT 5,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
last_polled_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_device_expires ON oauth_device_codes(expires_at);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_par_requests (
|
||||
id SERIAL PRIMARY KEY,
|
||||
request_uri VARCHAR(255) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
params TEXT,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_par_expires ON oauth_par_requests(expires_at);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_jti (
|
||||
id SERIAL PRIMARY KEY,
|
||||
jti_key VARCHAR(255) NOT NULL UNIQUE,
|
||||
expires_at TIMESTAMP NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_jti_expires ON oauth_jti(expires_at);
|
||||
|
||||
|
||||
-- key_hash: SHA-256 hex
|
||||
-- scopes: JSON-encoded array
|
||||
-- meta: JSON-encoded object
|
||||
|
||||
@@ -130,6 +130,7 @@ CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||
client_secret_hash TEXT,
|
||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||
is_active BOOLEAN DEFAULT 1,
|
||||
metadata TEXT,
|
||||
created_at TIMESTAMP
|
||||
);
|
||||
|
||||
@@ -148,12 +149,84 @@ CREATE TABLE IF NOT EXISTS oauth_codes (
|
||||
refresh_token TEXT,
|
||||
scopes TEXT,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
extra TEXT,
|
||||
created_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
||||
|
||||
|
||||
-- oauth_consents.scopes / oauth_refresh_tokens.scopes+extra / oauth_device_codes.scopes / oauth_par_requests.params: JSON text
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_consents (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
scopes TEXT,
|
||||
created_at TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_consents_user_client ON oauth_consents(user_id, client_id);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
token_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
family_id VARCHAR(64) NOT NULL,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
user_id INTEGER NOT NULL,
|
||||
session_token VARCHAR(255),
|
||||
scopes TEXT,
|
||||
extra TEXT,
|
||||
created_at TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
used_at TIMESTAMP,
|
||||
revoked_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_family ON oauth_refresh_tokens(family_id);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_session ON oauth_refresh_tokens(session_token);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_refresh_expires ON oauth_refresh_tokens(expires_at);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_device_codes (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
device_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
user_code VARCHAR(32) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
scopes TEXT,
|
||||
status VARCHAR(16) NOT NULL DEFAULT 'pending',
|
||||
user_id INTEGER,
|
||||
session_token VARCHAR(255),
|
||||
poll_interval INTEGER NOT NULL DEFAULT 5,
|
||||
created_at TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
last_polled_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_device_expires ON oauth_device_codes(expires_at);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_par_requests (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
request_uri VARCHAR(255) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
params TEXT,
|
||||
created_at TIMESTAMP,
|
||||
expires_at TIMESTAMP NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_par_expires ON oauth_par_requests(expires_at);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_jti (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
jti_key VARCHAR(255) NOT NULL UNIQUE,
|
||||
expires_at TIMESTAMP NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_jti_expires ON oauth_jti(expires_at);
|
||||
|
||||
|
||||
-- key_hash: SHA-256 hex
|
||||
-- scopes: JSON-encoded array
|
||||
-- meta: JSON-encoded object
|
||||
|
||||
@@ -197,6 +197,11 @@ func Gt(c lookup.Column, v any) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " > " + bl.ph(v) }
|
||||
}
|
||||
|
||||
// Lt is `col < value`.
|
||||
func Lt(c lookup.Column, v any) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " < " + bl.ph(v) }
|
||||
}
|
||||
|
||||
// IsNull is `col IS NULL`.
|
||||
func IsNull(c lookup.Column) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " IS NULL" }
|
||||
|
||||
@@ -60,6 +60,10 @@ func (o *OAuthClients) RegisterClient(ctx context.Context, client *sectypes.OAut
|
||||
return nil, fmt.Errorf("failed to marshal allowed_scopes: %w", err)
|
||||
}
|
||||
|
||||
meta, err := client.ClientMetadataJSON()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal client metadata: %w", err)
|
||||
}
|
||||
err = o.do(func(q Querier) error {
|
||||
return o.Insert(lookup.EntityOAuthClients).Set(
|
||||
Set(lookup.OAuthClientsClientID, client.ClientID),
|
||||
@@ -70,33 +74,86 @@ func (o *OAuthClients) RegisterClient(ctx context.Context, client *sectypes.OAut
|
||||
Set(lookup.OAuthClientsClientSecretHash, nullIfEmpty(client.ClientSecretHash)),
|
||||
Set(lookup.OAuthClientsTokenEndpointAuthMethod, authMethod),
|
||||
Set(lookup.OAuthClientsIsActive, true),
|
||||
Set(lookup.OAuthClientsMetadata, nullIfEmpty(meta)),
|
||||
Set(lookup.OAuthClientsCreatedAt, o.Now()),
|
||||
).Exec(ctx, q)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to register client: %w", err)
|
||||
}
|
||||
return §ypes.OAuthServerClient{
|
||||
ClientID: client.ClientID,
|
||||
RedirectURIs: client.RedirectURIs,
|
||||
ClientName: client.ClientName,
|
||||
GrantTypes: grantTypes,
|
||||
AllowedScopes: allowedScopes,
|
||||
ClientSecretHash: client.ClientSecretHash,
|
||||
TokenEndpointAuthMethod: authMethod,
|
||||
}, nil
|
||||
res := *client
|
||||
res.GrantTypes = grantTypes
|
||||
res.AllowedScopes = allowedScopes
|
||||
res.TokenEndpointAuthMethod = authMethod
|
||||
return &res, nil
|
||||
}
|
||||
|
||||
// UpdateClient implements lookup.OAuthClientStore: it rewrites the mutable registration
|
||||
// fields of an existing client (RFC 7592 management).
|
||||
func (o *OAuthClients) UpdateClient(ctx context.Context, client *sectypes.OAuthServerClient) error {
|
||||
redirects, err := o.d.EncodeJSON(client.RedirectURIs)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if redirects == nil {
|
||||
redirects = "[]"
|
||||
}
|
||||
grants, err := o.d.EncodeJSON(client.GrantTypes)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
scopes, err := o.d.EncodeJSON(client.AllowedScopes)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
meta, err := client.ClientMetadataJSON()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var n int64
|
||||
err = o.do(func(q Querier) error {
|
||||
var err error
|
||||
n, err = o.Update(lookup.EntityOAuthClients).Set(
|
||||
Set(lookup.OAuthClientsRedirectURIs, redirects),
|
||||
Set(lookup.OAuthClientsClientName, client.ClientName),
|
||||
Set(lookup.OAuthClientsGrantTypes, grants),
|
||||
Set(lookup.OAuthClientsAllowedScopes, scopes),
|
||||
Set(lookup.OAuthClientsClientSecretHash, nullIfEmpty(client.ClientSecretHash)),
|
||||
Set(lookup.OAuthClientsTokenEndpointAuthMethod, client.TokenEndpointAuthMethod),
|
||||
Set(lookup.OAuthClientsMetadata, nullIfEmpty(meta)),
|
||||
).Where(Eq(lookup.OAuthClientsClientID, client.ClientID), Eq(lookup.OAuthClientsIsActive, true)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update client: %w", err)
|
||||
}
|
||||
if n == 0 {
|
||||
return fmt.Errorf("client not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteClient implements lookup.OAuthClientStore: the client is deactivated.
|
||||
func (o *OAuthClients) DeleteClient(ctx context.Context, clientID string) error {
|
||||
return o.do(func(q Querier) error {
|
||||
_, err := o.Update(lookup.EntityOAuthClients).Set(Set(lookup.OAuthClientsIsActive, false)).
|
||||
Where(Eq(lookup.OAuthClientsClientID, clientID)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
// GetClient implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
|
||||
var redirects, grants, scopes any
|
||||
var name, secret, method sql.NullString
|
||||
var meta any
|
||||
err := o.do(func(q Querier) error {
|
||||
return o.From(lookup.EntityOAuthClients).
|
||||
Cols(lookup.OAuthClientsRedirectURIs, lookup.OAuthClientsClientName, lookup.OAuthClientsGrantTypes,
|
||||
lookup.OAuthClientsAllowedScopes, lookup.OAuthClientsClientSecretHash, lookup.OAuthClientsTokenEndpointAuthMethod).
|
||||
lookup.OAuthClientsAllowedScopes, lookup.OAuthClientsClientSecretHash, lookup.OAuthClientsTokenEndpointAuthMethod,
|
||||
lookup.OAuthClientsMetadata).
|
||||
Where(Eq(lookup.OAuthClientsClientID, clientID), Eq(lookup.OAuthClientsIsActive, true)).
|
||||
QueryRow(ctx, q, &redirects, &name, &grants, &scopes, &secret, &method)
|
||||
QueryRow(ctx, q, &redirects, &name, &grants, &scopes, &secret, &method, &meta)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
@@ -104,12 +161,17 @@ func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectype
|
||||
}
|
||||
return nil, fmt.Errorf("failed to get client: %w", err)
|
||||
}
|
||||
res := §ypes.OAuthServerClient{
|
||||
ClientID: clientID,
|
||||
ClientName: name.String,
|
||||
ClientSecretHash: secret.String,
|
||||
TokenEndpointAuthMethod: method.String,
|
||||
res := §ypes.OAuthServerClient{}
|
||||
switch v := meta.(type) {
|
||||
case []byte:
|
||||
_ = res.ApplyClientMetadata(string(v))
|
||||
case string:
|
||||
_ = res.ApplyClientMetadata(v)
|
||||
}
|
||||
res.ClientID = clientID
|
||||
res.ClientName = name.String
|
||||
res.ClientSecretHash = secret.String
|
||||
res.TokenEndpointAuthMethod = method.String
|
||||
_ = o.d.DecodeJSON(redirects, &res.RedirectURIs)
|
||||
_ = o.d.DecodeJSON(grants, &res.GrantTypes)
|
||||
_ = o.d.DecodeJSON(scopes, &res.AllowedScopes)
|
||||
@@ -126,6 +188,10 @@ func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) e
|
||||
if method == "" {
|
||||
method = "S256"
|
||||
}
|
||||
extra, err := code.CodeExtraJSON()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal code extra: %w", err)
|
||||
}
|
||||
return o.do(func(q Querier) error {
|
||||
return o.Insert(lookup.EntityOAuthCodes).Set(
|
||||
Set(lookup.OAuthCodesCode, code.Code),
|
||||
@@ -138,6 +204,7 @@ func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) e
|
||||
Set(lookup.OAuthCodesRefreshToken, code.RefreshToken),
|
||||
Set(lookup.OAuthCodesScopes, scopes),
|
||||
Set(lookup.OAuthCodesExpiresAt, code.ExpiresAt),
|
||||
Set(lookup.OAuthCodesExtra, nullIfEmpty(extra)),
|
||||
Set(lookup.OAuthCodesCreatedAt, o.Now()),
|
||||
).Exec(ctx, q)
|
||||
})
|
||||
@@ -148,15 +215,15 @@ func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) e
|
||||
func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) {
|
||||
var res sectypes.OAuthCode
|
||||
var state, refresh sql.NullString
|
||||
var scopes any
|
||||
var scopes, extra any
|
||||
err := o.tx(ctx, func(q Querier) error {
|
||||
err := o.From(lookup.EntityOAuthCodes).
|
||||
Cols(lookup.OAuthCodesClientID, lookup.OAuthCodesRedirectURI, lookup.OAuthCodesClientState,
|
||||
lookup.OAuthCodesCodeChallenge, lookup.OAuthCodesCodeChallengeMethod, lookup.OAuthCodesSessionToken,
|
||||
lookup.OAuthCodesRefreshToken, lookup.OAuthCodesScopes).
|
||||
lookup.OAuthCodesRefreshToken, lookup.OAuthCodesScopes, lookup.OAuthCodesExtra).
|
||||
Where(Eq(lookup.OAuthCodesCode, code), Gt(lookup.OAuthCodesExpiresAt, o.Now())).
|
||||
QueryRow(ctx, q, &res.ClientID, &res.RedirectURI, &state, &res.CodeChallenge, &res.CodeChallengeMethod,
|
||||
&res.SessionToken, &refresh, &scopes)
|
||||
&res.SessionToken, &refresh, &scopes, &extra)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -179,6 +246,12 @@ func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes
|
||||
res.ClientState = state.String
|
||||
res.RefreshToken = refresh.String
|
||||
_ = o.d.DecodeJSON(scopes, &res.Scopes)
|
||||
switch v := extra.(type) {
|
||||
case []byte:
|
||||
_ = res.ApplyCodeExtra(string(v))
|
||||
case string:
|
||||
_ = res.ApplyCodeExtra(v)
|
||||
}
|
||||
return &res, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error)
|
||||
Revoke(ctx context.Context, token string) error
|
||||
// UpdateClient rewrites the registration fields of an existing client (RFC 7592).
|
||||
UpdateClient(ctx context.Context, client *sectypes.OAuthServerClient) error
|
||||
// DeleteClient deactivates a client.
|
||||
DeleteClient(ctx context.Context, clientID string) error
|
||||
}
|
||||
|
||||
// OAuthSession is the session row written after an OAuth2 client login.
|
||||
@@ -147,6 +151,7 @@ type Provider struct {
|
||||
Keys KeyStore
|
||||
OAuthClient OAuthClientStore
|
||||
OAuthUser OAuthUserStore
|
||||
OAuthGrant OAuthGrantStore
|
||||
Passkey PasskeyStore
|
||||
TOTP TOTPStore
|
||||
Policy PolicyStore
|
||||
|
||||
@@ -61,6 +61,8 @@ const (
|
||||
OpOAuthExchangeCode Op = "oauth_exchange_code"
|
||||
OpOAuthIntrospect Op = "oauth_introspect"
|
||||
OpOAuthRevoke Op = "oauth_revoke"
|
||||
OpOAuthUpdateClient Op = "oauth_update_client"
|
||||
OpOAuthDeleteClient Op = "oauth_delete_client"
|
||||
|
||||
OpOAuthGetOrCreateUser Op = "oauth_get_or_create_user"
|
||||
OpOAuthCreateSession Op = "oauth_create_session"
|
||||
@@ -68,6 +70,22 @@ const (
|
||||
OpOAuthUpdateRefreshToken Op = "oauth_update_refresh_token" //nolint:gosec // operation name, not a credential
|
||||
OpOAuthGetUser Op = "oauth_get_user"
|
||||
|
||||
OpOAuthSaveConsent Op = "oauth_save_consent"
|
||||
OpOAuthGetConsent Op = "oauth_get_consent"
|
||||
OpOAuthRevokeConsent Op = "oauth_revoke_consent"
|
||||
OpOAuthSaveRefresh Op = "oauth_save_refresh" //nolint:gosec // operation name, not a credential
|
||||
OpOAuthRotateRefresh Op = "oauth_rotate_refresh" //nolint:gosec // operation name, not a credential
|
||||
OpOAuthPeekRefresh Op = "oauth_peek_refresh" //nolint:gosec // operation name, not a credential
|
||||
OpOAuthRevokeRefreshFamily Op = "oauth_revoke_refresh_family" //nolint:gosec // operation name, not a credential
|
||||
OpOAuthRevokeRefreshByUser Op = "oauth_revoke_refresh_session" //nolint:gosec // operation name, not a credential
|
||||
OpOAuthCreateDevice Op = "oauth_create_device"
|
||||
OpOAuthDeviceByUserCode Op = "oauth_device_by_user_code"
|
||||
OpOAuthDeviceDecide Op = "oauth_device_decide"
|
||||
OpOAuthDevicePoll Op = "oauth_device_poll"
|
||||
OpOAuthSavePAR Op = "oauth_save_par"
|
||||
OpOAuthConsumePAR Op = "oauth_consume_par"
|
||||
OpOAuthSeenJTI Op = "oauth_seen_jti"
|
||||
|
||||
OpPasskeyStore Op = "passkey_store"
|
||||
OpPasskeyGet Op = "passkey_get"
|
||||
OpPasskeyUpdateCounter Op = "passkey_update_counter"
|
||||
@@ -144,11 +162,28 @@ func AllOps() []Op {
|
||||
OpOAuthExchangeCode,
|
||||
OpOAuthIntrospect,
|
||||
OpOAuthRevoke,
|
||||
OpOAuthUpdateClient,
|
||||
OpOAuthDeleteClient,
|
||||
OpOAuthGetOrCreateUser,
|
||||
OpOAuthCreateSession,
|
||||
OpOAuthGetRefreshToken,
|
||||
OpOAuthUpdateRefreshToken,
|
||||
OpOAuthGetUser,
|
||||
OpOAuthSaveConsent,
|
||||
OpOAuthGetConsent,
|
||||
OpOAuthRevokeConsent,
|
||||
OpOAuthSaveRefresh,
|
||||
OpOAuthRotateRefresh,
|
||||
OpOAuthPeekRefresh,
|
||||
OpOAuthRevokeRefreshFamily,
|
||||
OpOAuthRevokeRefreshByUser,
|
||||
OpOAuthCreateDevice,
|
||||
OpOAuthDeviceByUserCode,
|
||||
OpOAuthDeviceDecide,
|
||||
OpOAuthDevicePoll,
|
||||
OpOAuthSavePAR,
|
||||
OpOAuthConsumePAR,
|
||||
OpOAuthSeenJTI,
|
||||
OpPasskeyStore,
|
||||
OpPasskeyGet,
|
||||
OpPasskeyUpdateCounter,
|
||||
|
||||
@@ -313,3 +313,38 @@ func (o *OAuthClients) Revoke(ctx context.Context, token string) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateClient implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) UpdateClient(ctx context.Context, client *sectypes.OAuthServerClient) error {
|
||||
input, err := json.Marshal(client)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal client: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errMsg sql.NullString
|
||||
err = o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error
|
||||
FROM %s($1::jsonb)
|
||||
`, o.procs.OAuthUpdateClient), input).Scan(&success, &errMsg)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update client: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return failure(errMsg, "failed to update client")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteClient implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) DeleteClient(ctx context.Context, clientID string) error {
|
||||
ok, errMsg, err := o.callNoData(ctx, o.procs.OAuthDeleteClient, clientID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete client: %w", err)
|
||||
}
|
||||
if !ok {
|
||||
return failure(errMsg, "failed to delete client")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,217 @@
|
||||
package procedure
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
)
|
||||
|
||||
// OAuthGrants implements lookup.OAuthGrantStore with the resolvespec_oauth_* grant procedures.
|
||||
// Every procedure takes one jsonb request and returns (p_success, p_error, p_data). A failure
|
||||
// that maps to a lookup sentinel carries a stable code in p_error (see the grantErrors table).
|
||||
type OAuthGrants struct {
|
||||
run Runner
|
||||
procs lookup.ProcNames
|
||||
}
|
||||
|
||||
var _ lookup.OAuthGrantStore = (*OAuthGrants)(nil)
|
||||
|
||||
// NewOAuthGrants creates the procedure-backed OAuthGrantStore.
|
||||
func NewOAuthGrants(run Runner, procs lookup.ProcNames) *OAuthGrants {
|
||||
return &OAuthGrants{run: run, procs: procs}
|
||||
}
|
||||
|
||||
// grantErrors maps the codes a grant procedure puts in p_error to the lookup sentinels.
|
||||
var grantErrors = map[string]error{
|
||||
"not_found": lookup.ErrNotFound,
|
||||
"refresh_invalid": lookup.ErrRefreshInvalid,
|
||||
"refresh_reused": lookup.ErrRefreshReused,
|
||||
"device_pending": lookup.ErrDevicePending,
|
||||
"device_slowdown": lookup.ErrDeviceSlowDown,
|
||||
"device_denied": lookup.ErrDeviceDenied,
|
||||
"device_expired": lookup.ErrDeviceExpired,
|
||||
}
|
||||
|
||||
// call runs proc with the JSON-encoded request. The returned data is the p_data of the
|
||||
// procedure, also when it reports a failure (rotate returns the reused token that way).
|
||||
func (o *OAuthGrants) call(ctx context.Context, proc string, req any) (data []byte, err error) {
|
||||
input, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errMsg sql.NullString
|
||||
err = o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error, p_data::text
|
||||
FROM %s($1::jsonb)
|
||||
`, proc), input).Scan(&success, &errMsg, &data)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", proc, err)
|
||||
}
|
||||
if success {
|
||||
return data, nil
|
||||
}
|
||||
if e, ok := grantErrors[errMsg.String]; ok {
|
||||
return data, e
|
||||
}
|
||||
return data, failure(errMsg, proc+" failed")
|
||||
}
|
||||
|
||||
// SaveConsent implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) SaveConsent(ctx context.Context, c lookup.Consent) error {
|
||||
_, err := o.call(ctx, o.procs.OAuthSaveConsent, c)
|
||||
return err
|
||||
}
|
||||
|
||||
// GetConsent implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) GetConsent(ctx context.Context, userID int, clientID string) (*lookup.Consent, error) {
|
||||
data, err := o.call(ctx, o.procs.OAuthGetConsent, map[string]any{"user_id": userID, "client_id": clientID})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var c lookup.Consent
|
||||
if err := json.Unmarshal(normalizeTimes(data), &c); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse consent: %w", err)
|
||||
}
|
||||
return &c, nil
|
||||
}
|
||||
|
||||
// RevokeConsent implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) RevokeConsent(ctx context.Context, userID int, clientID string) error {
|
||||
_, err := o.call(ctx, o.procs.OAuthRevokeConsent, map[string]any{"user_id": userID, "client_id": clientID})
|
||||
return err
|
||||
}
|
||||
|
||||
// SaveRefresh implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) SaveRefresh(ctx context.Context, t lookup.RefreshToken) error {
|
||||
_, err := o.call(ctx, o.procs.OAuthSaveRefresh, t)
|
||||
return err
|
||||
}
|
||||
|
||||
func parseRefresh(data []byte) (*lookup.RefreshToken, error) {
|
||||
if len(data) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
var t lookup.RefreshToken
|
||||
if err := json.Unmarshal(normalizeTimes(data), &t); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse refresh token: %w", err)
|
||||
}
|
||||
return &t, nil
|
||||
}
|
||||
|
||||
// RotateRefresh implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) RotateRefresh(ctx context.Context, oldHash string, next lookup.RefreshToken) (*lookup.RefreshToken, error) {
|
||||
data, err := o.call(ctx, o.procs.OAuthRotateRefresh, map[string]any{"old_hash": oldHash, "next": next})
|
||||
if err != nil && err != lookup.ErrRefreshReused { //nolint:errorlint // sentinel returned unwrapped by call
|
||||
return nil, err
|
||||
}
|
||||
t, perr := parseRefresh(data)
|
||||
if perr != nil {
|
||||
return nil, perr
|
||||
}
|
||||
return t, err
|
||||
}
|
||||
|
||||
// PeekRefresh implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) PeekRefresh(ctx context.Context, hash string) (*lookup.RefreshToken, error) {
|
||||
data, err := o.call(ctx, o.procs.OAuthPeekRefresh, map[string]any{"token_hash": hash})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return parseRefresh(data)
|
||||
}
|
||||
|
||||
// RevokeRefreshFamily implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) RevokeRefreshFamily(ctx context.Context, familyID string) error {
|
||||
_, err := o.call(ctx, o.procs.OAuthRevokeRefreshFamily, map[string]any{"family_id": familyID})
|
||||
return err
|
||||
}
|
||||
|
||||
// RevokeRefreshBySession implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) RevokeRefreshBySession(ctx context.Context, sessionToken string) error {
|
||||
_, err := o.call(ctx, o.procs.OAuthRevokeRefreshByUser, map[string]any{"session_token": sessionToken})
|
||||
return err
|
||||
}
|
||||
|
||||
// CreateDevice implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) CreateDevice(ctx context.Context, d lookup.DeviceCode) error {
|
||||
if d.Status == "" {
|
||||
d.Status = lookup.DevicePending
|
||||
}
|
||||
_, err := o.call(ctx, o.procs.OAuthCreateDevice, d)
|
||||
return err
|
||||
}
|
||||
|
||||
func parseDevice(data []byte) (*lookup.DeviceCode, error) {
|
||||
var d lookup.DeviceCode
|
||||
if err := json.Unmarshal(normalizeTimes(data), &d); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse device code: %w", err)
|
||||
}
|
||||
return &d, nil
|
||||
}
|
||||
|
||||
// DeviceByUserCode implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) DeviceByUserCode(ctx context.Context, userCode string) (*lookup.DeviceCode, error) {
|
||||
data, err := o.call(ctx, o.procs.OAuthDeviceByUserCode, map[string]any{"user_code": userCode})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return parseDevice(data)
|
||||
}
|
||||
|
||||
// DeviceDecide implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) DeviceDecide(ctx context.Context, userCode string, approve bool, userID int, sessionToken string) error {
|
||||
_, err := o.call(ctx, o.procs.OAuthDeviceDecide, map[string]any{
|
||||
"user_code": userCode, "approve": approve, "user_id": userID, "session_token": sessionToken,
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
// DevicePoll implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) DevicePoll(ctx context.Context, deviceHash string) (*lookup.DeviceCode, error) {
|
||||
data, err := o.call(ctx, o.procs.OAuthDevicePoll, map[string]any{"device_hash": deviceHash})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return parseDevice(data)
|
||||
}
|
||||
|
||||
// SavePushedRequest implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) SavePushedRequest(ctx context.Context, r lookup.PushedRequest) error {
|
||||
_, err := o.call(ctx, o.procs.OAuthSavePAR, r)
|
||||
return err
|
||||
}
|
||||
|
||||
// ConsumePushedRequest implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) ConsumePushedRequest(ctx context.Context, requestURI string) (*lookup.PushedRequest, error) {
|
||||
data, err := o.call(ctx, o.procs.OAuthConsumePAR, map[string]any{"request_uri": requestURI})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var r lookup.PushedRequest
|
||||
if err := json.Unmarshal(normalizeTimes(data), &r); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse pushed request: %w", err)
|
||||
}
|
||||
return &r, nil
|
||||
}
|
||||
|
||||
// SeenJTI implements lookup.OAuthGrantStore.
|
||||
func (o *OAuthGrants) SeenJTI(ctx context.Context, key string, expires time.Time) (bool, error) {
|
||||
data, err := o.call(ctx, o.procs.OAuthSeenJTI, map[string]any{"key": key, "expires_at": expires})
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
var out struct {
|
||||
Seen bool `json:"seen"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &out); err != nil {
|
||||
return false, fmt.Errorf("failed to parse jti result: %w", err)
|
||||
}
|
||||
return out.Seen, nil
|
||||
}
|
||||
@@ -63,6 +63,26 @@ type ProcNames struct {
|
||||
OAuthExchangeCode string // default: "resolvespec_oauth_exchange_code"
|
||||
OAuthIntrospect string // default: "resolvespec_oauth_introspect"
|
||||
OAuthRevoke string // default: "resolvespec_oauth_revoke"
|
||||
OAuthUpdateClient string // default: "resolvespec_oauth_update_client"
|
||||
OAuthDeleteClient string // default: "resolvespec_oauth_delete_client"
|
||||
|
||||
// OAuth2 server grant procedures (consents, refresh tokens, device codes, PAR, replay cache).
|
||||
// Each takes a jsonb request and returns (p_success, p_error, p_data).
|
||||
OAuthSaveConsent string // default: "resolvespec_oauth_save_consent"
|
||||
OAuthGetConsent string // default: "resolvespec_oauth_get_consent"
|
||||
OAuthRevokeConsent string // default: "resolvespec_oauth_revoke_consent"
|
||||
OAuthSaveRefresh string // default: "resolvespec_oauth_save_refresh"
|
||||
OAuthRotateRefresh string // default: "resolvespec_oauth_rotate_refresh"
|
||||
OAuthPeekRefresh string // default: "resolvespec_oauth_peek_refresh"
|
||||
OAuthRevokeRefreshFamily string // default: "resolvespec_oauth_revoke_refresh_family"
|
||||
OAuthRevokeRefreshByUser string // default: "resolvespec_oauth_revoke_refresh_session"
|
||||
OAuthCreateDevice string // default: "resolvespec_oauth_create_device"
|
||||
OAuthDeviceByUserCode string // default: "resolvespec_oauth_device_by_user_code"
|
||||
OAuthDeviceDecide string // default: "resolvespec_oauth_device_decide"
|
||||
OAuthDevicePoll string // default: "resolvespec_oauth_device_poll"
|
||||
OAuthSavePAR string // default: "resolvespec_oauth_save_par"
|
||||
OAuthConsumePAR string // default: "resolvespec_oauth_consume_par"
|
||||
OAuthSeenJTI string // default: "resolvespec_oauth_seen_jti"
|
||||
|
||||
// Keystore procedures (KeyStore)
|
||||
KeystoreGetUserKeys string // default: "resolvespec_keystore_get_user_keys"
|
||||
@@ -112,6 +132,23 @@ func DefaultProcNames() ProcNames {
|
||||
OAuthExchangeCode: "resolvespec_oauth_exchange_code",
|
||||
OAuthIntrospect: "resolvespec_oauth_introspect",
|
||||
OAuthRevoke: "resolvespec_oauth_revoke",
|
||||
OAuthUpdateClient: "resolvespec_oauth_update_client",
|
||||
OAuthDeleteClient: "resolvespec_oauth_delete_client",
|
||||
OAuthSaveConsent: "resolvespec_oauth_save_consent",
|
||||
OAuthGetConsent: "resolvespec_oauth_get_consent",
|
||||
OAuthRevokeConsent: "resolvespec_oauth_revoke_consent",
|
||||
OAuthSaveRefresh: "resolvespec_oauth_save_refresh",
|
||||
OAuthRotateRefresh: "resolvespec_oauth_rotate_refresh",
|
||||
OAuthPeekRefresh: "resolvespec_oauth_peek_refresh",
|
||||
OAuthRevokeRefreshFamily: "resolvespec_oauth_revoke_refresh_family",
|
||||
OAuthRevokeRefreshByUser: "resolvespec_oauth_revoke_refresh_session",
|
||||
OAuthCreateDevice: "resolvespec_oauth_create_device",
|
||||
OAuthDeviceByUserCode: "resolvespec_oauth_device_by_user_code",
|
||||
OAuthDeviceDecide: "resolvespec_oauth_device_decide",
|
||||
OAuthDevicePoll: "resolvespec_oauth_device_poll",
|
||||
OAuthSavePAR: "resolvespec_oauth_save_par",
|
||||
OAuthConsumePAR: "resolvespec_oauth_consume_par",
|
||||
OAuthSeenJTI: "resolvespec_oauth_seen_jti",
|
||||
KeystoreGetUserKeys: "resolvespec_keystore_get_user_keys",
|
||||
KeystoreCreateKey: "resolvespec_keystore_create_key",
|
||||
KeystoreDeleteKey: "resolvespec_keystore_delete_key",
|
||||
|
||||
@@ -26,6 +26,11 @@ const (
|
||||
EntityUserPasswordResets Entity = "user_password_resets"
|
||||
EntityOAuthClients Entity = "oauth_clients"
|
||||
EntityOAuthCodes Entity = "oauth_codes"
|
||||
EntityOAuthConsents Entity = "oauth_consents"
|
||||
EntityOAuthRefreshTokens Entity = "oauth_refresh_tokens" //nolint:gosec // table name, not a credential
|
||||
EntityOAuthDeviceCodes Entity = "oauth_device_codes"
|
||||
EntityOAuthPARRequests Entity = "oauth_par_requests"
|
||||
EntityOAuthJTI Entity = "oauth_jti"
|
||||
EntityUserKeys Entity = "user_keys"
|
||||
EntitySecGroupMembers Entity = "sec_group_members"
|
||||
EntitySecColumnRules Entity = "sec_column_rules"
|
||||
@@ -121,6 +126,7 @@ var (
|
||||
OAuthClientsClientSecretHash = col(EntityOAuthClients, "client_secret_hash")
|
||||
OAuthClientsTokenEndpointAuthMethod = col(EntityOAuthClients, "token_endpoint_auth_method")
|
||||
OAuthClientsIsActive = col(EntityOAuthClients, "is_active")
|
||||
OAuthClientsMetadata = col(EntityOAuthClients, "metadata")
|
||||
OAuthClientsCreatedAt = col(EntityOAuthClients, "created_at")
|
||||
|
||||
OAuthCodesID = col(EntityOAuthCodes, "id")
|
||||
@@ -135,6 +141,51 @@ var (
|
||||
OAuthCodesScopes = col(EntityOAuthCodes, "scopes")
|
||||
OAuthCodesExpiresAt = col(EntityOAuthCodes, "expires_at")
|
||||
OAuthCodesCreatedAt = col(EntityOAuthCodes, "created_at")
|
||||
OAuthCodesExtra = col(EntityOAuthCodes, "extra")
|
||||
|
||||
OAuthConsentsID = col(EntityOAuthConsents, "id")
|
||||
OAuthConsentsUserID = col(EntityOAuthConsents, "user_id")
|
||||
OAuthConsentsClientID = col(EntityOAuthConsents, "client_id")
|
||||
OAuthConsentsScopes = col(EntityOAuthConsents, "scopes")
|
||||
OAuthConsentsCreatedAt = col(EntityOAuthConsents, "created_at")
|
||||
OAuthConsentsExpiresAt = col(EntityOAuthConsents, "expires_at")
|
||||
|
||||
OAuthRefreshID = col(EntityOAuthRefreshTokens, "id")
|
||||
OAuthRefreshTokenHash = col(EntityOAuthRefreshTokens, "token_hash")
|
||||
OAuthRefreshFamilyID = col(EntityOAuthRefreshTokens, "family_id")
|
||||
OAuthRefreshClientID = col(EntityOAuthRefreshTokens, "client_id")
|
||||
OAuthRefreshUserID = col(EntityOAuthRefreshTokens, "user_id")
|
||||
OAuthRefreshSessionToken = col(EntityOAuthRefreshTokens, "session_token")
|
||||
OAuthRefreshScopes = col(EntityOAuthRefreshTokens, "scopes")
|
||||
OAuthRefreshExtra = col(EntityOAuthRefreshTokens, "extra")
|
||||
OAuthRefreshCreatedAt = col(EntityOAuthRefreshTokens, "created_at")
|
||||
OAuthRefreshExpiresAt = col(EntityOAuthRefreshTokens, "expires_at")
|
||||
OAuthRefreshUsedAt = col(EntityOAuthRefreshTokens, "used_at")
|
||||
OAuthRefreshRevokedAt = col(EntityOAuthRefreshTokens, "revoked_at")
|
||||
|
||||
OAuthDeviceID = col(EntityOAuthDeviceCodes, "id")
|
||||
OAuthDeviceHash = col(EntityOAuthDeviceCodes, "device_hash")
|
||||
OAuthDeviceUserCode = col(EntityOAuthDeviceCodes, "user_code")
|
||||
OAuthDeviceClientID = col(EntityOAuthDeviceCodes, "client_id")
|
||||
OAuthDeviceScopes = col(EntityOAuthDeviceCodes, "scopes")
|
||||
OAuthDeviceStatus = col(EntityOAuthDeviceCodes, "status")
|
||||
OAuthDeviceUserID = col(EntityOAuthDeviceCodes, "user_id")
|
||||
OAuthDeviceSessionToken = col(EntityOAuthDeviceCodes, "session_token")
|
||||
OAuthDeviceInterval = col(EntityOAuthDeviceCodes, "poll_interval")
|
||||
OAuthDeviceCreatedAt = col(EntityOAuthDeviceCodes, "created_at")
|
||||
OAuthDeviceExpiresAt = col(EntityOAuthDeviceCodes, "expires_at")
|
||||
OAuthDeviceLastPolledAt = col(EntityOAuthDeviceCodes, "last_polled_at")
|
||||
|
||||
OAuthPARID = col(EntityOAuthPARRequests, "id")
|
||||
OAuthPARRequestURI = col(EntityOAuthPARRequests, "request_uri")
|
||||
OAuthPARClientID = col(EntityOAuthPARRequests, "client_id")
|
||||
OAuthPARParams = col(EntityOAuthPARRequests, "params")
|
||||
OAuthPARCreatedAt = col(EntityOAuthPARRequests, "created_at")
|
||||
OAuthPARExpiresAt = col(EntityOAuthPARRequests, "expires_at")
|
||||
|
||||
OAuthJTIID = col(EntityOAuthJTI, "id")
|
||||
OAuthJTIKey = col(EntityOAuthJTI, "jti_key")
|
||||
OAuthJTIExpiresAt = col(EntityOAuthJTI, "expires_at")
|
||||
|
||||
KeysID = col(EntityUserKeys, "id")
|
||||
KeysUserID = col(EntityUserKeys, "user_id")
|
||||
@@ -190,10 +241,20 @@ var allColumns = []Column{
|
||||
ResetsID, ResetsUserID, ResetsTokenHash, ResetsExpiresAt, ResetsCreatedAt, ResetsUsed, ResetsUsedAt,
|
||||
OAuthClientsID, OAuthClientsClientID, OAuthClientsRedirectURIs, OAuthClientsClientName, OAuthClientsGrantTypes,
|
||||
OAuthClientsAllowedScopes, OAuthClientsClientSecretHash, OAuthClientsTokenEndpointAuthMethod,
|
||||
OAuthClientsIsActive, OAuthClientsCreatedAt,
|
||||
OAuthClientsIsActive, OAuthClientsCreatedAt, OAuthClientsMetadata,
|
||||
OAuthCodesID, OAuthCodesCode, OAuthCodesClientID, OAuthCodesRedirectURI, OAuthCodesClientState,
|
||||
OAuthCodesCodeChallenge, OAuthCodesCodeChallengeMethod, OAuthCodesSessionToken, OAuthCodesRefreshToken,
|
||||
OAuthCodesScopes, OAuthCodesExpiresAt, OAuthCodesCreatedAt,
|
||||
OAuthCodesScopes, OAuthCodesExpiresAt, OAuthCodesCreatedAt, OAuthCodesExtra,
|
||||
OAuthConsentsID, OAuthConsentsUserID, OAuthConsentsClientID, OAuthConsentsScopes, OAuthConsentsCreatedAt,
|
||||
OAuthConsentsExpiresAt,
|
||||
OAuthRefreshID, OAuthRefreshTokenHash, OAuthRefreshFamilyID, OAuthRefreshClientID, OAuthRefreshUserID,
|
||||
OAuthRefreshSessionToken, OAuthRefreshScopes, OAuthRefreshExtra, OAuthRefreshCreatedAt, OAuthRefreshExpiresAt,
|
||||
OAuthRefreshUsedAt, OAuthRefreshRevokedAt,
|
||||
OAuthDeviceID, OAuthDeviceHash, OAuthDeviceUserCode, OAuthDeviceClientID, OAuthDeviceScopes, OAuthDeviceStatus,
|
||||
OAuthDeviceUserID, OAuthDeviceSessionToken, OAuthDeviceInterval, OAuthDeviceCreatedAt, OAuthDeviceExpiresAt,
|
||||
OAuthDeviceLastPolledAt,
|
||||
OAuthPARID, OAuthPARRequestURI, OAuthPARClientID, OAuthPARParams, OAuthPARCreatedAt, OAuthPARExpiresAt,
|
||||
OAuthJTIID, OAuthJTIKey, OAuthJTIExpiresAt,
|
||||
KeysID, KeysUserID, KeysKeyType, KeysKeyHash, KeysName, KeysScopes, KeysMeta, KeysExpiresAt,
|
||||
KeysCreatedAt, KeysLastUsedAt, KeysIsActive,
|
||||
GroupMembersGroupID, GroupMembersUserID,
|
||||
|
||||
@@ -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"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -32,6 +34,45 @@ type OAuth2Config struct {
|
||||
// Optional: Custom user info parser
|
||||
// If not provided, will use standard claims (sub, email, name)
|
||||
UserInfoParser func(userInfo map[string]any) (*UserContext, error)
|
||||
|
||||
// --- OpenID Connect (see oidc_client.go) ---
|
||||
|
||||
// Issuer turns the provider into an OpenID Connect provider: PKCE and a nonce are used and
|
||||
// the id_token returned by the token endpoint is validated (signature, iss, aud, exp, nonce,
|
||||
// at_hash). WithOIDC fills the endpoints in by discovery; with WithOAuth2 set JWKSURL
|
||||
// as well. UserInfoURL stays optional: the id_token claims are used when it is empty.
|
||||
Issuer string
|
||||
// JWKSURL is the provider's key set. Only needed with WithOAuth2; WithOIDC discovers it.
|
||||
JWKSURL string
|
||||
// EndSessionURL is the provider's RP-initiated logout endpoint (discovered by WithOIDC).
|
||||
EndSessionURL string
|
||||
// UsePKCE sends a PKCE S256 challenge for a provider that is not OIDC. It is always on in OIDC mode.
|
||||
UsePKCE bool
|
||||
// AllowedAlgs lists the id_token signature algorithms to accept. Default: RS256, PS256, ES256, ES384.
|
||||
AllowedAlgs []string
|
||||
// AuthStyle selects how the client authenticates at the token endpoint: "basic", "post" or ""
|
||||
// (try basic, fall back to post).
|
||||
AuthStyle string
|
||||
// HTTPClient is used for discovery, JWKS, token and userinfo requests.
|
||||
HTTPClient *http.Client
|
||||
// ClockSkew tolerates clock differences when validating the id_token. Default 1 minute.
|
||||
ClockSkew time.Duration
|
||||
}
|
||||
|
||||
// OAuth2AuthOptions are optional OpenID Connect authentication request parameters.
|
||||
type OAuth2AuthOptions struct {
|
||||
LoginHint string
|
||||
Prompt string // none, login, consent, select_account
|
||||
MaxAge *int
|
||||
ACRValues string
|
||||
Extra map[string]string
|
||||
}
|
||||
|
||||
// oauth2State is what the login redirect remembers until the callback.
|
||||
type oauth2State struct {
|
||||
expiry time.Time
|
||||
verifier string // PKCE code_verifier
|
||||
nonce string
|
||||
}
|
||||
|
||||
// OAuth2Provider holds configuration and state for a single OAuth2 provider
|
||||
@@ -40,7 +81,10 @@ type OAuth2Provider struct {
|
||||
userInfoURL string
|
||||
userInfoParser func(userInfo map[string]any) (*UserContext, error)
|
||||
providerName string
|
||||
states map[string]time.Time // state -> expiry time
|
||||
states map[string]*oauth2State
|
||||
oidc *oidcProvider // nil for plain OAuth2
|
||||
usePKCE bool
|
||||
httpClient *http.Client
|
||||
statesMutex sync.RWMutex
|
||||
stopCh chan struct{} // closed to stop cleanupStates
|
||||
stopOnce sync.Once
|
||||
@@ -58,6 +102,13 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
|
||||
cfg.UserInfoParser = defaultOAuth2UserInfoParser
|
||||
}
|
||||
|
||||
authStyle := oauth2.AuthStyleAutoDetect
|
||||
switch cfg.AuthStyle {
|
||||
case "basic":
|
||||
authStyle = oauth2.AuthStyleInHeader
|
||||
case "post":
|
||||
authStyle = oauth2.AuthStyleInParams
|
||||
}
|
||||
provider := &OAuth2Provider{
|
||||
config: &oauth2.Config{
|
||||
ClientID: cfg.ClientID,
|
||||
@@ -65,15 +116,22 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
|
||||
RedirectURL: cfg.RedirectURL,
|
||||
Scopes: cfg.Scopes,
|
||||
Endpoint: oauth2.Endpoint{
|
||||
AuthURL: cfg.AuthURL,
|
||||
TokenURL: cfg.TokenURL,
|
||||
AuthURL: cfg.AuthURL,
|
||||
TokenURL: cfg.TokenURL,
|
||||
AuthStyle: authStyle,
|
||||
},
|
||||
},
|
||||
userInfoURL: cfg.UserInfoURL,
|
||||
userInfoParser: cfg.UserInfoParser,
|
||||
providerName: cfg.ProviderName,
|
||||
states: make(map[string]time.Time),
|
||||
states: make(map[string]*oauth2State),
|
||||
stopCh: make(chan struct{}),
|
||||
usePKCE: cfg.UsePKCE,
|
||||
httpClient: cfg.HTTPClient,
|
||||
}
|
||||
if cfg.Issuer != "" {
|
||||
provider.oidc = newOIDCProvider(&cfg)
|
||||
provider.usePKCE = true
|
||||
}
|
||||
|
||||
// Initialize providers map if needed
|
||||
@@ -97,17 +155,56 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
|
||||
|
||||
// OAuth2GetAuthURL returns the OAuth2 authorization URL for redirecting users
|
||||
func (a *DatabaseAuthenticator) OAuth2GetAuthURL(providerName, state string) (string, error) {
|
||||
return a.OAuth2GetAuthURLWithOptions(providerName, state, OAuth2AuthOptions{})
|
||||
}
|
||||
|
||||
// OAuth2GetAuthURLWithOptions is OAuth2GetAuthURL with OpenID Connect request parameters. For an
|
||||
// OIDC provider (and with UsePKCE) it also creates the PKCE verifier and the nonce, which are
|
||||
// kept with the state until the callback.
|
||||
func (a *DatabaseAuthenticator) OAuth2GetAuthURLWithOptions(providerName, state string, opts OAuth2AuthOptions) (string, error) {
|
||||
provider, err := a.getOAuth2Provider(providerName)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if provider.oidc != nil {
|
||||
dctx, cancel := context.WithTimeout(provider.withHTTPClient(context.Background()), 15*time.Second)
|
||||
defer cancel()
|
||||
if err := provider.oidc.ensureEndpoints(dctx, provider); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
st := &oauth2State{expiry: time.Now().Add(10 * time.Minute)}
|
||||
var params []oauth2.AuthCodeOption
|
||||
if provider.usePKCE {
|
||||
st.verifier = oauth2.GenerateVerifier()
|
||||
params = append(params, oauth2.S256ChallengeOption(st.verifier))
|
||||
}
|
||||
if provider.oidc != nil {
|
||||
if st.nonce, err = randomOAuthToken(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
params = append(params, oauth2.SetAuthURLParam("nonce", st.nonce))
|
||||
}
|
||||
set := func(k, v string) {
|
||||
if v != "" {
|
||||
params = append(params, oauth2.SetAuthURLParam(k, v))
|
||||
}
|
||||
}
|
||||
set("login_hint", opts.LoginHint)
|
||||
set("prompt", opts.Prompt)
|
||||
set("acr_values", opts.ACRValues)
|
||||
if opts.MaxAge != nil {
|
||||
set("max_age", strconv.Itoa(*opts.MaxAge))
|
||||
}
|
||||
for k, v := range opts.Extra {
|
||||
set(k, v)
|
||||
}
|
||||
|
||||
// Store state for validation
|
||||
provider.statesMutex.Lock()
|
||||
provider.states[state] = time.Now().Add(10 * time.Minute)
|
||||
provider.states[state] = st
|
||||
provider.statesMutex.Unlock()
|
||||
|
||||
return provider.config.AuthCodeURL(state), nil
|
||||
return provider.config.AuthCodeURL(state, params...), nil
|
||||
}
|
||||
|
||||
// OAuth2GenerateState generates a random state string for CSRF protection
|
||||
@@ -121,42 +218,97 @@ func (a *DatabaseAuthenticator) OAuth2GenerateState() (string, error) {
|
||||
|
||||
// OAuth2HandleCallback handles the OAuth2 callback and exchanges code for token
|
||||
func (a *DatabaseAuthenticator) OAuth2HandleCallback(ctx context.Context, providerName, code, state string) (*LoginResponse, error) {
|
||||
return a.oauth2Callback(ctx, providerName, code, state, "")
|
||||
}
|
||||
|
||||
// OAuth2HandleCallbackRequest is OAuth2HandleCallback for the redirect request itself. Besides
|
||||
// code and state it honours the error parameters and the RFC 9207 "iss" parameter, which
|
||||
// protects against mix-up attacks when several providers are in use.
|
||||
func (a *DatabaseAuthenticator) OAuth2HandleCallbackRequest(ctx context.Context, providerName string, r *http.Request) (*LoginResponse, error) {
|
||||
q := r.URL.Query()
|
||||
if e := q.Get("error"); e != "" {
|
||||
return nil, fmt.Errorf("provider returned an error: %s %s", e, q.Get("error_description"))
|
||||
}
|
||||
return a.oauth2Callback(ctx, providerName, q.Get("code"), q.Get("state"), q.Get("iss"))
|
||||
}
|
||||
|
||||
func (a *DatabaseAuthenticator) oauth2Callback(ctx context.Context, providerName, code, state, iss string) (*LoginResponse, error) {
|
||||
provider, err := a.getOAuth2Provider(providerName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Validate state
|
||||
if !provider.validateState(state) {
|
||||
st, ok := provider.validateState(state)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("invalid state parameter")
|
||||
}
|
||||
if code == "" {
|
||||
return nil, fmt.Errorf("missing authorization code")
|
||||
}
|
||||
if provider.oidc != nil && iss != "" && iss != provider.oidc.issuer {
|
||||
return nil, fmt.Errorf("authorization response issuer mismatch")
|
||||
}
|
||||
if ctx = provider.withHTTPClient(ctx); ctx == nil {
|
||||
return nil, fmt.Errorf("no context")
|
||||
}
|
||||
|
||||
// Exchange code for token
|
||||
token, err := provider.config.Exchange(ctx, code)
|
||||
var exchange []oauth2.AuthCodeOption
|
||||
if st.verifier != "" {
|
||||
exchange = append(exchange, oauth2.VerifierOption(st.verifier))
|
||||
}
|
||||
if provider.oidc != nil {
|
||||
if err := provider.oidc.ensureEndpoints(ctx, provider); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
token, err := provider.config.Exchange(ctx, code, exchange...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to exchange code: %w", err)
|
||||
}
|
||||
|
||||
// OpenID Connect: validate the id_token.
|
||||
var rawIDToken string
|
||||
var idClaims map[string]any
|
||||
if provider.oidc != nil {
|
||||
rawIDToken, _ = token.Extra("id_token").(string)
|
||||
if rawIDToken == "" && oauthSliceContains(provider.config.Scopes, "openid") {
|
||||
return nil, fmt.Errorf("token response contains no id_token")
|
||||
}
|
||||
if rawIDToken != "" {
|
||||
if idClaims, err = provider.oidc.validateIDToken(ctx, rawIDToken, st.nonce, token.AccessToken); err != nil {
|
||||
return nil, fmt.Errorf("invalid id_token: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Fetch user info
|
||||
client := provider.config.Client(ctx, token)
|
||||
resp, err := client.Get(provider.userInfoURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to fetch user info: %w", err)
|
||||
userInfo := map[string]any{}
|
||||
if provider.userInfoURL != "" {
|
||||
fetched, err := provider.fetchUserInfo(ctx, token)
|
||||
switch {
|
||||
case err == nil:
|
||||
if sub, _ := idClaims["sub"].(string); sub != "" {
|
||||
if us, _ := fetched["sub"].(string); us != "" && us != sub {
|
||||
return nil, fmt.Errorf("userinfo subject does not match the id_token")
|
||||
}
|
||||
}
|
||||
userInfo = fetched
|
||||
case provider.oidc == nil || idClaims == nil:
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read user info: %w", err)
|
||||
claims := map[string]any{}
|
||||
for k, v := range idClaims {
|
||||
claims[k] = v
|
||||
}
|
||||
|
||||
var userInfo map[string]any
|
||||
if err := json.Unmarshal(body, &userInfo); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse user info: %w", err)
|
||||
for k, v := range userInfo {
|
||||
claims[k] = v
|
||||
}
|
||||
|
||||
// Parse user info
|
||||
userCtx, err := provider.userInfoParser(userInfo)
|
||||
userCtx, err := provider.userInfoParser(claims)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
||||
}
|
||||
@@ -187,12 +339,47 @@ func (a *DatabaseAuthenticator) OAuth2HandleCallback(ctx context.Context, provid
|
||||
|
||||
userCtx.SessionID = sessionToken
|
||||
|
||||
return &LoginResponse{
|
||||
resp := &LoginResponse{
|
||||
Token: sessionToken,
|
||||
RefreshToken: token.RefreshToken,
|
||||
User: userCtx,
|
||||
ExpiresIn: int64(time.Until(expiresAt).Seconds()),
|
||||
}, nil
|
||||
}
|
||||
if rawIDToken != "" {
|
||||
// Keep the id_token: it is the id_token_hint of OAuth2LogoutURL.
|
||||
resp.Meta = map[string]any{"id_token": rawIDToken}
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// fetchUserInfo calls the provider's userinfo endpoint with the access token.
|
||||
func (p *OAuth2Provider) fetchUserInfo(ctx context.Context, token *oauth2.Token) (map[string]any, error) {
|
||||
resp, err := p.config.Client(ctx, token).Get(p.userInfoURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to fetch user info: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read user info: %w", err)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("user info request failed: status %d", resp.StatusCode)
|
||||
}
|
||||
var userInfo map[string]any
|
||||
if err := json.Unmarshal(body, &userInfo); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse user info: %w", err)
|
||||
}
|
||||
return userInfo, nil
|
||||
}
|
||||
|
||||
// withHTTPClient makes oauth2 use the provider's HTTP client.
|
||||
func (p *OAuth2Provider) withHTTPClient(ctx context.Context) context.Context {
|
||||
if p.httpClient == nil {
|
||||
return ctx
|
||||
}
|
||||
return context.WithValue(ctx, oauth2.HTTPClient, p.httpClient)
|
||||
}
|
||||
|
||||
// OAuth2GetProviders returns list of configured OAuth2 provider names
|
||||
@@ -251,23 +438,20 @@ func (a *DatabaseAuthenticator) oauth2CreateSession(ctx context.Context, session
|
||||
})
|
||||
}
|
||||
|
||||
// validateState validates state using in-memory storage
|
||||
func (p *OAuth2Provider) validateState(state string) bool {
|
||||
// validateState validates state using in-memory storage and returns what was remembered with it.
|
||||
func (p *OAuth2Provider) validateState(state string) (*oauth2State, bool) {
|
||||
p.statesMutex.Lock()
|
||||
defer p.statesMutex.Unlock()
|
||||
|
||||
expiry, ok := p.states[state]
|
||||
st, ok := p.states[state]
|
||||
if !ok {
|
||||
return false
|
||||
return nil, false
|
||||
}
|
||||
|
||||
if time.Now().After(expiry) {
|
||||
delete(p.states, state)
|
||||
return false
|
||||
}
|
||||
|
||||
delete(p.states, state) // One-time use
|
||||
return true
|
||||
if time.Now().After(st.expiry) {
|
||||
return nil, false
|
||||
}
|
||||
return st, true
|
||||
}
|
||||
|
||||
// cleanupStates removes expired states periodically
|
||||
@@ -284,8 +468,8 @@ func (p *OAuth2Provider) cleanupStates() {
|
||||
}
|
||||
p.statesMutex.Lock()
|
||||
now := time.Now()
|
||||
for state, expiry := range p.states {
|
||||
if now.After(expiry) {
|
||||
for state, st := range p.states {
|
||||
if now.After(st.expiry) {
|
||||
delete(p.states, state)
|
||||
}
|
||||
}
|
||||
@@ -363,7 +547,7 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
|
||||
}
|
||||
|
||||
// Use OAuth2 provider to refresh the token
|
||||
tokenSource := provider.config.TokenSource(ctx, oldToken)
|
||||
tokenSource := provider.config.TokenSource(provider.withHTTPClient(ctx), oldToken)
|
||||
newToken, err := tokenSource.Token()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to refresh token with provider: %w", err)
|
||||
@@ -388,12 +572,21 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
|
||||
|
||||
userCtx.SessionID = newSessionToken
|
||||
|
||||
return &LoginResponse{
|
||||
resp := &LoginResponse{
|
||||
Token: newSessionToken,
|
||||
RefreshToken: newToken.RefreshToken,
|
||||
User: userCtx,
|
||||
ExpiresIn: int64(time.Until(newToken.Expiry).Seconds()),
|
||||
}, nil
|
||||
}
|
||||
if provider.oidc != nil {
|
||||
if raw, _ := newToken.Extra("id_token").(string); raw != "" {
|
||||
if _, err := provider.oidc.validateIDToken(provider.withHTTPClient(ctx), raw, "", newToken.AccessToken); err != nil {
|
||||
return nil, fmt.Errorf("invalid id_token in refresh response: %w", err)
|
||||
}
|
||||
resp.Meta = map[string]any{"id_token": raw}
|
||||
}
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// Pre-configured OAuth2 factory methods
|
||||
@@ -406,10 +599,14 @@ func NewGoogleAuthenticator(clientID, clientSecret, redirectURL string, db *sql.
|
||||
ClientSecret: clientSecret,
|
||||
RedirectURL: redirectURL,
|
||||
Scopes: []string{"openid", "profile", "email"},
|
||||
AuthURL: "https://accounts.google.com/o/oauth2/auth",
|
||||
AuthURL: "https://accounts.google.com/o/oauth2/v2/auth",
|
||||
TokenURL: "https://oauth2.googleapis.com/token",
|
||||
UserInfoURL: "https://www.googleapis.com/oauth2/v2/userinfo",
|
||||
UserInfoURL: "https://openidconnect.googleapis.com/v1/userinfo",
|
||||
ProviderName: "google",
|
||||
// OpenID Connect: PKCE, nonce and id_token validation against Google's published keys.
|
||||
Issuer: "https://accounts.google.com",
|
||||
JWKSURL: "https://www.googleapis.com/oauth2/v3/certs",
|
||||
EndSessionURL: "",
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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 (
|
||||
"context"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
)
|
||||
|
||||
// OAuthRegisterClient persists an OAuth2 client registration.
|
||||
@@ -33,3 +35,24 @@ func (a *DatabaseAuthenticator) OAuthIntrospectToken(ctx context.Context, token
|
||||
func (a *DatabaseAuthenticator) OAuthRevokeToken(ctx context.Context, token string) error {
|
||||
return a.src.get().OAuthClient.Revoke(ctx, token)
|
||||
}
|
||||
|
||||
// OAuthGrants returns the store holding consents, managed refresh tokens, device codes,
|
||||
// pushed authorization requests and the replay cache.
|
||||
func (a *DatabaseAuthenticator) OAuthGrants() lookup.OAuthGrantStore {
|
||||
return a.src.get().OAuthGrant
|
||||
}
|
||||
|
||||
// OAuthUpdateClient replaces the registered metadata of a client (RFC 7592).
|
||||
func (a *DatabaseAuthenticator) OAuthUpdateClient(ctx context.Context, client *OAuthServerClient) error {
|
||||
return a.src.get().OAuthClient.UpdateClient(ctx, client)
|
||||
}
|
||||
|
||||
// OAuthDeleteClient deactivates a registered client (RFC 7592).
|
||||
func (a *DatabaseAuthenticator) OAuthDeleteClient(ctx context.Context, clientID string) error {
|
||||
return a.src.get().OAuthClient.DeleteClient(ctx, clientID)
|
||||
}
|
||||
|
||||
// OAuthGetUser returns the active user with the given id.
|
||||
func (a *DatabaseAuthenticator) OAuthGetUser(ctx context.Context, userID int) (*UserContext, error) {
|
||||
return a.src.get().OAuthUser.GetUser(ctx, userID)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
import "time"
|
||||
import (
|
||||
"encoding/json"
|
||||
"time"
|
||||
)
|
||||
|
||||
// OAuthServerClient is a persisted RFC 7591 registered OAuth2 client.
|
||||
type OAuthServerClient struct {
|
||||
@@ -11,8 +14,80 @@ type OAuthServerClient struct {
|
||||
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
||||
ClientSecretHash string `json:"client_secret_hash,omitempty"`
|
||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
|
||||
|
||||
// The fields below are stored together in the oauth_clients.metadata JSON column, so a
|
||||
// new field never needs a schema change. See SplitJSON / MergeJSON.
|
||||
|
||||
ResponseTypes []string `json:"response_types,omitempty"`
|
||||
ClientURI string `json:"client_uri,omitempty"`
|
||||
LogoURI string `json:"logo_uri,omitempty"`
|
||||
Contacts []string `json:"contacts,omitempty"`
|
||||
PostLogoutRedirectURIs []string `json:"post_logout_redirect_uris,omitempty"`
|
||||
BackchannelLogoutURI string `json:"backchannel_logout_uri,omitempty"`
|
||||
JWKS json.RawMessage `json:"jwks,omitempty"`
|
||||
JWKSURI string `json:"jwks_uri,omitempty"`
|
||||
IDTokenSignedResponseAlg string `json:"id_token_signed_response_alg,omitempty"`
|
||||
UserinfoSignedResponseAlg string `json:"userinfo_signed_response_alg,omitempty"`
|
||||
TokenEndpointAuthSigningAlg string `json:"token_endpoint_auth_signing_alg,omitempty"`
|
||||
RequireConsent bool `json:"require_consent,omitempty"`
|
||||
FirstParty bool `json:"first_party,omitempty"`
|
||||
RequirePAR bool `json:"require_pushed_authorization_requests,omitempty"`
|
||||
DPoPBoundAccessTokens bool `json:"dpop_bound_access_tokens,omitempty"`
|
||||
RegistrationAccessTokenHash string `json:"registration_access_token_hash,omitempty"`
|
||||
ClientSecretExpiresAt int64 `json:"client_secret_expires_at,omitempty"`
|
||||
ClientIDIssuedAt int64 `json:"client_id_issued_at,omitempty"`
|
||||
}
|
||||
|
||||
// oauthClientColumns are the keys stored in their own oauth_clients columns; every other
|
||||
// key of the JSON form is stored in the metadata column.
|
||||
var oauthClientColumns = []string{
|
||||
"client_id", "redirect_uris", "client_name", "grant_types", "allowed_scopes",
|
||||
"client_secret_hash", "token_endpoint_auth_method",
|
||||
}
|
||||
|
||||
// oauthCodeColumns are the keys stored in their own oauth_codes columns.
|
||||
var oauthCodeColumns = []string{
|
||||
"code", "client_id", "redirect_uri", "client_state", "code_challenge", "code_challenge_method",
|
||||
"session_token", "refresh_token", "scopes", "expires_at",
|
||||
}
|
||||
|
||||
// SplitJSON returns the JSON form of v without the keys in columns. It is the value stored in
|
||||
// a metadata/extra column; an empty object is returned as "".
|
||||
func SplitJSON(v any, columns []string) (string, error) {
|
||||
raw, err := json.Marshal(v) //nolint:gosec // G117: client secret hash and tokens are intentionally stored
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
m := map[string]json.RawMessage{}
|
||||
if err := json.Unmarshal(raw, &m); err != nil {
|
||||
return "", err
|
||||
}
|
||||
for _, c := range columns {
|
||||
delete(m, c)
|
||||
}
|
||||
if len(m) == 0 {
|
||||
return "", nil
|
||||
}
|
||||
out, err := json.Marshal(m)
|
||||
return string(out), err
|
||||
}
|
||||
|
||||
// MergeJSON applies a metadata/extra JSON document onto dst. Empty input is a no-op.
|
||||
func MergeJSON(dst any, data string) error {
|
||||
if data == "" || data == "null" {
|
||||
return nil
|
||||
}
|
||||
return json.Unmarshal([]byte(data), dst)
|
||||
}
|
||||
|
||||
// ClientMetadataJSON returns the value of the oauth_clients.metadata column.
|
||||
func (c *OAuthServerClient) ClientMetadataJSON() (string, error) {
|
||||
return SplitJSON(c, oauthClientColumns)
|
||||
}
|
||||
|
||||
// ApplyClientMetadata merges the oauth_clients.metadata column into c.
|
||||
func (c *OAuthServerClient) ApplyClientMetadata(data string) error { return MergeJSON(c, data) }
|
||||
|
||||
// OAuthCode is a short-lived authorization code.
|
||||
type OAuthCode struct {
|
||||
Code string `json:"code"`
|
||||
@@ -25,8 +100,28 @@ type OAuthCode struct {
|
||||
RefreshToken string `json:"refresh_token,omitempty"`
|
||||
Scopes []string `json:"scopes,omitempty"`
|
||||
ExpiresAt time.Time `json:"expires_at"`
|
||||
|
||||
// Stored together in the oauth_codes.extra JSON column.
|
||||
|
||||
UserID int `json:"user_id,omitempty"`
|
||||
Nonce string `json:"nonce,omitempty"`
|
||||
AuthTime int64 `json:"auth_time,omitempty"`
|
||||
ACR string `json:"acr,omitempty"`
|
||||
AMR []string `json:"amr,omitempty"`
|
||||
SessionID string `json:"sid,omitempty"`
|
||||
Claims map[string]any `json:"claims,omitempty"`
|
||||
Resource []string `json:"resource,omitempty"`
|
||||
DPoPJKT string `json:"dpop_jkt,omitempty"`
|
||||
ResponseType string `json:"response_type,omitempty"`
|
||||
ConsentedAt int64 `json:"consented_at,omitempty"`
|
||||
}
|
||||
|
||||
// CodeExtraJSON returns the value of the oauth_codes.extra column.
|
||||
func (c *OAuthCode) CodeExtraJSON() (string, error) { return SplitJSON(c, oauthCodeColumns) }
|
||||
|
||||
// ApplyCodeExtra merges the oauth_codes.extra column into c.
|
||||
func (c *OAuthCode) ApplyCodeExtra(data string) error { return MergeJSON(c, data) }
|
||||
|
||||
// OAuthTokenInfo is the RFC 7662 token introspection response.
|
||||
type OAuthTokenInfo struct {
|
||||
Active bool `json:"active"`
|
||||
@@ -37,4 +132,13 @@ type OAuthTokenInfo struct {
|
||||
Roles []string `json:"roles,omitempty"`
|
||||
Exp int64 `json:"exp,omitempty"`
|
||||
Iat int64 `json:"iat,omitempty"`
|
||||
|
||||
// Filled in by the OAuth server, not by the stores.
|
||||
Scope string `json:"scope,omitempty"`
|
||||
ClientID string `json:"client_id,omitempty"`
|
||||
TokenType string `json:"token_type,omitempty"`
|
||||
Iss string `json:"iss,omitempty"`
|
||||
Aud []string `json:"aud,omitempty"`
|
||||
Jti string `json:"jti,omitempty"`
|
||||
Cnf map[string]any `json:"cnf,omitempty"`
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user