diff --git a/README.md b/README.md index a4bfe6e..8e12992 100644 --- a/README.md +++ b/README.md @@ -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.). diff --git a/pkg/resolvemcp/README.md b/pkg/resolvemcp/README.md index 545c62b..0f7c5a8 100644 --- a/pkg/resolvemcp/README.md +++ b/pkg/resolvemcp/README.md @@ -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 | diff --git a/pkg/security/OAUTH2.md b/pkg/security/OAUTH2.md index 7072c14..942e379 100644 --- a/pkg/security/OAUTH2.md +++ b/pkg/security/OAUTH2.md @@ -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 diff --git a/pkg/security/OAUTH2_REFRESH_QUICK_REFERENCE.md b/pkg/security/OAUTH2_REFRESH_QUICK_REFERENCE.md index 7d34ce6..fdcd2d9 100644 --- a/pkg/security/OAUTH2_REFRESH_QUICK_REFERENCE.md +++ b/pkg/security/OAUTH2_REFRESH_QUICK_REFERENCE.md @@ -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 diff --git a/pkg/security/OAUTH2_SERVER.md b/pkg/security/OAUTH2_SERVER.md new file mode 100644 index 0000000..42f53c7 --- /dev/null +++ b/pkg/security/OAUTH2_SERVER.md @@ -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=` (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 ` 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. diff --git a/pkg/security/README.md b/pkg/security/README.md index ff0f3ac..a6c36a0 100644 --- a/pkg/security/README.md +++ b/pkg/security/README.md @@ -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 diff --git a/pkg/security/breaking_changes.md b/pkg/security/breaking_changes.md index 6348f17..d81d766 100644 --- a/pkg/security/breaking_changes.md +++ b/pkg/security/breaking_changes.md @@ -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/.sql`. Existing databases need: + +- `ALTER TABLE oauth_clients ADD COLUMN metadata ` (client metadata: logout URIs, jwks, require_consent, first_party, dpop_bound, signing algs, ...) +- `ALTER TABLE oauth_codes ADD COLUMN extra ` (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`). + +`` 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`. diff --git a/pkg/security/lookup/backends/backends.go b/pkg/security/lookup/backends/backends.go index 73b766a..bb1e225 100644 --- a/pkg/security/lookup/backends/backends.go +++ b/pkg/security/lookup/backends/backends.go @@ -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}, diff --git a/pkg/security/lookup/backends/conformance_test.go b/pkg/security/lookup/backends/conformance_test.go index 70c27e5..d407e85 100644 --- a/pkg/security/lookup/backends/conformance_test.go +++ b/pkg/security/lookup/backends/conformance_test.go @@ -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"}, diff --git a/pkg/security/lookup/backends/container_test.go b/pkg/security/lookup/backends/container_test.go index 06b2cb6..4272528 100644 --- a/pkg/security/lookup/backends/container_test.go +++ b/pkg/security/lookup/backends/container_test.go @@ -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) + } +} diff --git a/pkg/security/lookup/backends/routers.go b/pkg/security/lookup/backends/routers.go index 566e686..623e949 100644 --- a/pkg/security/lookup/backends/routers.go +++ b/pkg/security/lookup/backends/routers.go @@ -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) +} diff --git a/pkg/security/lookup/conformance/conformance.go b/pkg/security/lookup/conformance/conformance.go index 006b34b..104e616 100644 --- a/pkg/security/lookup/conformance/conformance.go +++ b/pkg/security/lookup/conformance/conformance.go @@ -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) + } +} diff --git a/pkg/security/lookup/database_schema.sql b/pkg/security/lookup/database_schema.sql index 214804f..b918884 100644 --- a/pkg/security/lookup/database_schema.sql +++ b/pkg/security/lookup/database_schema.sql @@ -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; +$$; diff --git a/pkg/security/lookup/ddl/mssql.sql b/pkg/security/lookup/ddl/mssql.sql index 7c3ec25..b739f3d 100644 --- a/pkg/security/lookup/ddl/mssql.sql +++ b/pkg/security/lookup/ddl/mssql.sql @@ -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 diff --git a/pkg/security/lookup/ddl/mysql.sql b/pkg/security/lookup/ddl/mysql.sql index fd7959a..c8d4202 100644 --- a/pkg/security/lookup/ddl/mysql.sql +++ b/pkg/security/lookup/ddl/mysql.sql @@ -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 diff --git a/pkg/security/lookup/ddl/postgres.sql b/pkg/security/lookup/ddl/postgres.sql index 813872b..8a76dae 100644 --- a/pkg/security/lookup/ddl/postgres.sql +++ b/pkg/security/lookup/ddl/postgres.sql @@ -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 diff --git a/pkg/security/lookup/ddl/sqlite.sql b/pkg/security/lookup/ddl/sqlite.sql index c3802bf..100e7a1 100644 --- a/pkg/security/lookup/ddl/sqlite.sql +++ b/pkg/security/lookup/ddl/sqlite.sql @@ -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 diff --git a/pkg/security/lookup/direct/base.go b/pkg/security/lookup/direct/base.go index cccf60e..c2fa692 100644 --- a/pkg/security/lookup/direct/base.go +++ b/pkg/security/lookup/direct/base.go @@ -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" } diff --git a/pkg/security/lookup/direct/oauth.go b/pkg/security/lookup/direct/oauth.go index 132f239..e005c0e 100644 --- a/pkg/security/lookup/direct/oauth.go +++ b/pkg/security/lookup/direct/oauth.go @@ -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 } diff --git a/pkg/security/lookup/direct/oauth_grant.go b/pkg/security/lookup/direct/oauth_grant.go new file mode 100644 index 0000000..6dd2086 --- /dev/null +++ b/pkg/security/lookup/direct/oauth_grant.go @@ -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 +} diff --git a/pkg/security/lookup/grant.go b/pkg/security/lookup/grant.go new file mode 100644 index 0000000..9f5c3b4 --- /dev/null +++ b/pkg/security/lookup/grant.go @@ -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) +} diff --git a/pkg/security/lookup/lookup.go b/pkg/security/lookup/lookup.go index 7daf022..086f64c 100644 --- a/pkg/security/lookup/lookup.go +++ b/pkg/security/lookup/lookup.go @@ -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 diff --git a/pkg/security/lookup/mode.go b/pkg/security/lookup/mode.go index e73f528..4f577ee 100644 --- a/pkg/security/lookup/mode.go +++ b/pkg/security/lookup/mode.go @@ -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, diff --git a/pkg/security/lookup/procedure/oauth.go b/pkg/security/lookup/procedure/oauth.go index 7ea810f..17c42c6 100644 --- a/pkg/security/lookup/procedure/oauth.go +++ b/pkg/security/lookup/procedure/oauth.go @@ -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 +} diff --git a/pkg/security/lookup/procedure/oauth_grant.go b/pkg/security/lookup/procedure/oauth_grant.go new file mode 100644 index 0000000..8f3d9fe --- /dev/null +++ b/pkg/security/lookup/procedure/oauth_grant.go @@ -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 +} diff --git a/pkg/security/lookup/procs.go b/pkg/security/lookup/procs.go index 4db0e64..f7751fa 100644 --- a/pkg/security/lookup/procs.go +++ b/pkg/security/lookup/procs.go @@ -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", diff --git a/pkg/security/lookup/schema.go b/pkg/security/lookup/schema.go index b2d0e17..bb5ae6f 100644 --- a/pkg/security/lookup/schema.go +++ b/pkg/security/lookup/schema.go @@ -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, diff --git a/pkg/security/oauth2_full_example.go b/pkg/security/oauth2_full_example.go new file mode 100644 index 0000000..12d0bb5 --- /dev/null +++ b/pkg/security/oauth2_full_example.go @@ -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/.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()) +} diff --git a/pkg/security/oauth2_methods.go b/pkg/security/oauth2_methods.go index 40da412..aa658f7 100644 --- a/pkg/security/oauth2_methods.go +++ b/pkg/security/oauth2_methods.go @@ -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: "", }) } diff --git a/pkg/security/oauth_authorize.go b/pkg/security/oauth_authorize.go new file mode 100644 index 0000000..47edb4d --- /dev/null +++ b/pkg/security/oauth_authorize.go @@ -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(`
`) + for k, v := range params { + b.WriteString(``) + } + b.WriteString(`
`) + 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 +} diff --git a/pkg/security/oauth_clientauth.go b/pkg/security/oauth_clientauth.go new file mode 100644 index 0000000..6f88f1c --- /dev/null +++ b/pkg/security/oauth_clientauth.go @@ -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 +} diff --git a/pkg/security/oauth_consent.go b/pkg/security/oauth_consent.go new file mode 100644 index 0000000..c141d18 --- /dev/null +++ b/pkg/security/oauth_consent.go @@ -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) +} diff --git a/pkg/security/oauth_device.go b/pkg/security/oauth_device.go new file mode 100644 index 0000000..50de142 --- /dev/null +++ b/pkg/security/oauth_device.go @@ -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), + }) +} diff --git a/pkg/security/oauth_dpop.go b/pkg/security/oauth_dpop.go new file mode 100644 index 0000000..7b819f7 --- /dev/null +++ b/pkg/security/oauth_dpop.go @@ -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) +} diff --git a/pkg/security/oauth_exchange.go b/pkg/security/oauth_exchange.go new file mode 100644 index 0000000..9c3b79a --- /dev/null +++ b/pkg/security/oauth_exchange.go @@ -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 +} diff --git a/pkg/security/oauth_flow_test.go b/pkg/security/oauth_flow_test.go new file mode 100644 index 0000000..d12c863 --- /dev/null +++ b/pkg/security/oauth_flow_test.go @@ -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") + } +} diff --git a/pkg/security/oauth_introspect.go b/pkg/security/oauth_introspect.go new file mode 100644 index 0000000..eaeb9c2 --- /dev/null +++ b/pkg/security/oauth_introspect.go @@ -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) +} diff --git a/pkg/security/oauth_jwk.go b/pkg/security/oauth_jwk.go new file mode 100644 index 0000000..038a7d4 --- /dev/null +++ b/pkg/security/oauth_jwk.go @@ -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 +} diff --git a/pkg/security/oauth_logout.go b/pkg/security/oauth_logout.go new file mode 100644 index 0000000..ff4539c --- /dev/null +++ b/pkg/security/oauth_logout.go @@ -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() + } + }() +} diff --git a/pkg/security/oauth_oidc.go b/pkg/security/oauth_oidc.go new file mode 100644 index 0000000..6131f1c --- /dev/null +++ b/pkg/security/oauth_oidc.go @@ -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/) 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, " ") } diff --git a/pkg/security/oauth_par.go b/pkg/security/oauth_par.go new file mode 100644 index 0000000..ab35a11 --- /dev/null +++ b/pkg/security/oauth_par.go @@ -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 +} diff --git a/pkg/security/oauth_refresh.go b/pkg/security/oauth_refresh.go new file mode 100644 index 0000000..aae0309 --- /dev/null +++ b/pkg/security/oauth_refresh.go @@ -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, + }) +} diff --git a/pkg/security/oauth_register.go b/pkg/security/oauth_register.go new file mode 100644 index 0000000..e5d5ac2 --- /dev/null +++ b/pkg/security/oauth_register.go @@ -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) +} diff --git a/pkg/security/oauth_server.go b/pkg/security/oauth_server.go index f2c6fff..41d46d7 100644 --- a/pkg/security/oauth_server.go +++ b/pkg/security/oauth_server.go @@ -9,21 +9,25 @@ import ( "encoding/base64" "encoding/hex" "encoding/json" - "fmt" + "html/template" "net/http" "net/url" "strings" "sync" "time" - "github.com/golang-jwt/jwt/v5" - "golang.org/x/oauth2" + "github.com/bitechdev/ResolveSpec/pkg/security/lookup" ) -// OAuthServerConfig configures the MCP-standard OAuth2 authorization server. +// OAuthServerConfig configures the OAuth2 / OpenID Connect authorization server. +// +// Every field except Issuer is optional. The zero value of each new option keeps the +// behaviour of earlier versions; see OAUTH2_SERVER.md for the full guide. type OAuthServerConfig struct { - // Issuer is the public base URL of this server (e.g. "https://api.example.com"). - // Used in /.well-known/oauth-authorization-server and to build endpoint URLs. + // Issuer is the public base URL of this server (e.g. "https://api.example.com"). It is + // the "iss" of every token and the base of every endpoint URL. A path is allowed + // ("https://example.com/auth"); the server then also answers the RFC 8414 path-insertion + // well-known URLs. Issuer string // ProviderCallbackPath is the path on this server that external OAuth2 providers @@ -42,7 +46,8 @@ type OAuthServerConfig struct { // Useful for multi-instance deployments. Defaults to in-memory. PersistCodes bool - // DefaultScopes lists scopes advertised in server metadata. Defaults to ["openid","profile","email"]. + // DefaultScopes lists scopes advertised in server metadata and granted to clients that + // register without allowed_scopes. Defaults to ["openid","profile","email"]. DefaultScopes []string // AccessTokenTTL is the issued token lifetime. Defaults to 24h. @@ -55,41 +60,152 @@ type OAuthServerConfig struct { // RFC 9728 metadata. Defaults to Issuer. ResourceIdentifier string - // SigningKey signs id_tokens (RS256) and is exposed via the JWKS endpoint. If nil, an - // RSA-2048 key is generated in memory when the server starts. Supply a persistent key - // for multi-instance deployments so id_tokens remain verifiable across restarts/instances. + // SigningKey signs id_tokens (RS256) and is exposed via the JWKS endpoint. If nil and + // SigningKeys is empty, an RSA-2048 key is generated in memory when the server starts. + // Supply a persistent key for multi-instance deployments so tokens remain verifiable + // across restarts and instances. SigningKey *rsa.PrivateKey + + // SigningKeys supersedes SigningKey. The first key is the default; all are published in + // the JWKS, which is how a key is rotated (add the new key first, publish, then remove the + // old one once its tokens have expired). RSA and ECDSA (P-256/P-384) keys are supported. + SigningKeys []OAuthSigningKey + + // CookieSecret keys the HMAC that protects the SSO cookie and the state carried through + // the login and consent forms. Defaults to a value derived from the first signing key, so + // instances sharing a signing key share sessions. + CookieSecret []byte + + // SSOCookie configures the browser session cookie that makes prompt=none, max_age, + // single sign-on and logout work. + SSOCookie OAuthSSOCookieConfig + + // --- Consent and scopes --- + + // RequireConsent shows a consent screen for every client that is not first-party and has + // no stored consent covering the requested scopes. A client can also opt in with its + // require_consent metadata. + RequireConsent bool + + // ConsentTTL is how long a stored consent is honoured. Defaults to 90 days. + ConsentTTL time.Duration + + // --- Tokens --- + + // ManagedRefreshTokens makes the server issue and rotate its own refresh tokens (stored + // hashed, one family per grant, reuse of a rotated token revokes the family). When false, + // the refresh token of the underlying DatabaseAuthenticator is passed through as before. + ManagedRefreshTokens bool + + // RefreshTokenTTL is the absolute lifetime of a refresh token family. Defaults to 30 days. + RefreshTokenTTL time.Duration + + // JWTAccessTokens issues RFC 9068 JWT access tokens instead of opaque session tokens. + // Resource servers can verify them locally (VerifyAccessToken). + JWTAccessTokens bool + + // AccessTokenAudience is the "aud" of JWT access tokens that were not requested for a + // specific resource. Defaults to ResourceIdentifier. + AccessTokenAudience string + + // EnableDPoP accepts RFC 9449 DPoP proofs at the token and userinfo endpoints and binds + // the issued tokens to the proof key. + EnableDPoP bool + + // EnablePAR serves the RFC 9126 pushed authorization request endpoint; RequirePAR makes it + // mandatory for every client. + EnablePAR bool + RequirePAR bool + PARTTL time.Duration // default 90 seconds + + // EnableDeviceFlow serves the RFC 8628 device authorization grant. + EnableDeviceFlow bool + DeviceCodeTTL time.Duration // default 10 minutes + DevicePollSeconds int // minimum poll interval, default 5 + + // EnableTokenExchange serves the RFC 8693 token exchange grant. + EnableTokenExchange bool + + // --- OpenID Connect --- + + // ClaimsProvider supplies the user claims for id_tokens and UserInfo (profile, email, + // address, phone, custom claims). The default returns sub, preferred_username and email. + ClaimsProvider OAuthClaimsProvider + + // SupportedACR lists the authentication context class references advertised and accepted. + SupportedACR []string + + // DisableLogout does not serve the RP-initiated logout endpoint. + DisableLogout bool + + // --- Hardening --- + + // InitialAccessToken, when set, must be presented as a Bearer token to register a client. + InitialAccessToken string + + // AllowAnonymousIntrospection lets callers without client credentials use the revocation + // and introspection endpoints. By default they must authenticate as a client. + AllowAnonymousIntrospection bool + + // RateLimiter, when set, is called for every request to a token-issuing endpoint + // ("authorize", "token", "par", "device", "register", "introspect", "revoke", "userinfo", + // "logout"). Returning false answers 429. + RateLimiter func(r *http.Request, endpoint string) bool + + // AllowPrivateNetworkFetch lets the server fetch client jwks_uri documents from loopback and + // private addresses. Leave false in production (SSRF protection). + AllowPrivateNetworkFetch bool + + // ScopeDescriptions are shown next to each scope on the consent screen. Built-in + // descriptions exist for openid, profile, email and offline_access. + ScopeDescriptions map[string]string + + // LoginTemplate and ConsentTemplate replace the built-in pages. They receive + // OAuthLoginPage and OAuthConsentPage. + LoginTemplate *template.Template + ConsentTemplate *template.Template } -// oauthClient is a dynamically registered OAuth2 client (RFC 7591). -type oauthClient struct { - ClientID string `json:"client_id"` - RedirectURIs []string `json:"redirect_uris"` - ClientName string `json:"client_name,omitempty"` - GrantTypes []string `json:"grant_types"` - AllowedScopes []string `json:"allowed_scopes,omitempty"` - ClientSecretHash string `json:"client_secret_hash,omitempty"` - TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"` +// OAuthSSOCookieConfig configures the SSO cookie. +type OAuthSSOCookieConfig struct { + Name string // default "resolvespec_sso" + Path string // default "/" + TTL time.Duration // default 8h + SameSite http.SameSite // default Lax + // Insecure sends the cookie over plain HTTP. By default Secure is on whenever Issuer is https. + Insecure bool + Disable bool // never set a cookie: every authorization request authenticates again } -// isConfidential reports whether the client has a registered secret and must -// authenticate itself at the token endpoint. -func (c *oauthClient) isConfidential() bool { - return c.ClientSecretHash != "" +// OAuthClaimsRequest is the input of an OAuthClaimsProvider. +type OAuthClaimsRequest struct { + UserID int + Sub string + Scopes []string + // Requested holds the names the client asked for individually through the OIDC "claims" + // request parameter for this destination. + Requested []string + // Destination is "id_token" or "userinfo". + Destination string + // Base are the claims the server already knows (sub, preferred_username, email). + Base map[string]any } -// pendingAuth tracks an in-progress authorization code exchange. +// OAuthClaimsProvider returns the claims for a user. The server only includes the standard claims +// the granted scopes entitle the client to (profile, email, address, phone) plus every claim the +// client requested explicitly; claims outside those sets are dropped. +type OAuthClaimsProvider func(ctx context.Context, req OAuthClaimsRequest) (map[string]any, error) + +// pendingAuth tracks an authorization request that is waiting for an external provider. type pendingAuth struct { - ClientID string - RedirectURI string - ClientState string - CodeChallenge string - CodeChallengeMethod string - ProviderName string // empty = password login - ExpiresAt time.Time - SessionToken string // set after authentication completes - RefreshToken string // set after authentication completes when refresh tokens are issued - Scopes []string // requested scopes + Req *authzRequest + Provider string + ExpiresAt time.Time +} + +type cachedClient struct { + c *OAuthServerClient + at time.Time // zero for clients that only exist in memory } // externalProvider pairs a DatabaseAuthenticator with its provider name. @@ -98,52 +214,54 @@ type externalProvider struct { providerName string } -// OAuthServer implements the MCP-standard OAuth2 authorization server (OAuth 2.1 + PKCE). +// OAuthServer is an OAuth 2.1 authorization server and OpenID Connect provider. // // It can act as both: // - A direct identity provider using DatabaseAuthenticator username/password login // - A federation layer that delegates authentication to external OAuth2 providers // (Google, GitHub, Microsoft, etc.) registered via RegisterExternalProvider // -// The server exposes these RFC-compliant endpoints: +// Endpoints (see OAUTH2_SERVER.md for parameters and examples): // -// GET /.well-known/oauth-authorization-server RFC 8414 — server metadata discovery -// GET /.well-known/openid-configuration OIDC discovery (superset of the above) -// GET /.well-known/oauth-protected-resource RFC 9728 — protected resource metadata -// POST /oauth/register RFC 7591 — dynamic client registration -// GET /oauth/authorize OAuth 2.1 + PKCE — start authorization -// POST /oauth/authorize Direct login form submission -// POST /oauth/token Token exchange: authorization_code, -// refresh_token, client_credentials (RFC 6749 §4.4) -// POST /oauth/revoke RFC 7009 — token revocation -// POST /oauth/introspect RFC 7662 — token introspection -// GET /oauth/userinfo OIDC UserInfo endpoint -// GET /oauth/jwks.json JWKS — id_token verification keys -// GET {ProviderCallbackPath} Internal — external provider callback -// -// Confidential clients (registered with token_endpoint_auth_method other than "none", or -// any grant_types including client_credentials) authenticate at /oauth/token via -// client_secret_basic or client_secret_post. Public clients keep relying on PKCE alone. -// -// When the granted scope includes "openid", authorization_code and refresh_token responses -// include an RS256-signed id_token (see OAuthServerConfig.SigningKey). +// GET /.well-known/oauth-authorization-server RFC 8414 server metadata +// GET /.well-known/openid-configuration OIDC Discovery +// GET /.well-known/oauth-protected-resource RFC 9728 protected resource metadata +// POST /oauth/register RFC 7591 dynamic client registration +// GET|PUT|DELETE /oauth/register/{client_id} RFC 7592 client management +// GET|POST /oauth/authorize authorization endpoint (PKCE S256 required) +// POST /oauth/token authorization_code, refresh_token, client_credentials, +// device_code and token-exchange grants +// POST /oauth/par RFC 9126 pushed authorization requests +// POST /oauth/device_authorization RFC 8628 device authorization +// GET|POST /oauth/device device verification page +// POST /oauth/revoke RFC 7009 token revocation +// POST /oauth/introspect RFC 7662 token introspection +// GET|POST /oauth/userinfo OIDC UserInfo +// GET /oauth/jwks.json JWKS +// GET|POST /oauth/logout OIDC RP-initiated logout +// GET {ProviderCallbackPath} external provider callback type OAuthServer struct { cfg OAuthServerConfig auth *DatabaseAuthenticator // nil = only external providers providers []externalProvider mu sync.RWMutex - clients map[string]*oauthClient - pending map[string]*pendingAuth // provider_state → pending (external flow) - codes map[string]*pendingAuth // auth_code → pending (post-auth) + clients map[string]*cachedClient + pending map[string]*pendingAuth // provider_state → request waiting for the provider + codes map[string]*OAuthCode // auth code → code (when PersistCodes is false) - signingKey *rsa.PrivateKey - signingKeyID string + keys *oauthKeyring + secret []byte + jwks *jwksCache + tmpl oauthTemplates + issuerURL *url.URL + + bcWG sync.WaitGroup // in-flight back-channel logout notifications done chan struct{} // closed by Close() to stop background goroutines } -// NewOAuthServer creates a new MCP OAuth2 authorization server. +// NewOAuthServer creates a new OAuth2 / OIDC authorization server. // // Pass a DatabaseAuthenticator to enable direct username/password login (the server // acts as its own identity provider). Pass nil to use only external providers. @@ -159,6 +277,9 @@ func NewOAuthServer(cfg OAuthServerConfig, auth *DatabaseAuthenticator) *OAuthSe } if len(cfg.DefaultScopes) == 0 { cfg.DefaultScopes = []string{"openid", "profile", "email"} + if cfg.ManagedRefreshTokens { + cfg.DefaultScopes = append(cfg.DefaultScopes, "offline_access") + } } if cfg.AccessTokenTTL == 0 { cfg.AccessTokenTTL = 24 * time.Hour @@ -166,35 +287,70 @@ func NewOAuthServer(cfg OAuthServerConfig, auth *DatabaseAuthenticator) *OAuthSe if cfg.AuthCodeTTL == 0 { cfg.AuthCodeTTL = 2 * time.Minute } + if cfg.ConsentTTL == 0 { + cfg.ConsentTTL = 90 * 24 * time.Hour + } + if cfg.RefreshTokenTTL == 0 { + cfg.RefreshTokenTTL = 30 * 24 * time.Hour + } + if cfg.PARTTL == 0 { + cfg.PARTTL = 90 * time.Second + } + if cfg.DeviceCodeTTL == 0 { + cfg.DeviceCodeTTL = 10 * time.Minute + } + if cfg.DevicePollSeconds == 0 { + cfg.DevicePollSeconds = 5 + } + if cfg.SSOCookie.Name == "" { + cfg.SSOCookie.Name = "resolvespec_sso" + } + if cfg.SSOCookie.Path == "" { + cfg.SSOCookie.Path = "/" + } + if cfg.SSOCookie.TTL == 0 { + cfg.SSOCookie.TTL = 8 * time.Hour + } + if cfg.SSOCookie.SameSite == 0 { + cfg.SSOCookie.SameSite = http.SameSiteLaxMode + } // Normalize issuer: remove trailing slash to ensure consistent endpoint URL construction. cfg.Issuer = strings.TrimSuffix(cfg.Issuer, "/") if cfg.ResourceIdentifier == "" { cfg.ResourceIdentifier = cfg.Issuer } + if cfg.AccessTokenAudience == "" { + cfg.AccessTokenAudience = cfg.ResourceIdentifier + } - signingKey := cfg.SigningKey - if signingKey == nil { - var err error - signingKey, err = rsa.GenerateKey(rand.Reader, 2048) - if err != nil { - // Signing keys are only required for id_token issuance (OIDC "openid" scope); - // leaving signingKey nil degrades gracefully by omitting id_token/JWKS support. - signingKey = nil - } + keys, err := newOAuthKeyring(&cfg) + if err != nil { + // Only a bad explicit key reaches here (generating one cannot realistically fail). + // Fall back to a fresh key so the server still starts; id_tokens then do not survive a restart. + cfg.SigningKeys, cfg.SigningKey = nil, nil + keys, _ = newOAuthKeyring(&cfg) + } + issuerURL, _ := url.Parse(cfg.Issuer) + if issuerURL == nil { + issuerURL = &url.URL{} } s := &OAuthServer{ - cfg: cfg, - auth: auth, - clients: make(map[string]*oauthClient), - pending: make(map[string]*pendingAuth), - codes: make(map[string]*pendingAuth), - signingKey: signingKey, - done: make(chan struct{}), + cfg: cfg, + auth: auth, + clients: make(map[string]*cachedClient), + pending: make(map[string]*pendingAuth), + codes: make(map[string]*OAuthCode), + keys: keys, + jwks: newJWKSCache(publicHTTPClient(cfg.AllowPrivateNetworkFetch)), + issuerURL: issuerURL, + done: make(chan struct{}), } - if signingKey != nil { - s.signingKeyID = rsaKeyID(&signingKey.PublicKey) + s.secret = cfg.CookieSecret + if len(s.secret) == 0 { + s.secret = deriveOAuthSecret(keys) } + s.tmpl = newOAuthTemplates(&cfg) go s.cleanupExpired() return s } @@ -208,6 +364,7 @@ func (s *OAuthServer) Close() { default: close(s.done) } + s.bcWG.Wait() } // RegisterExternalProvider adds an external OAuth2 provider (Google, GitHub, Microsoft, etc.) @@ -234,20 +391,55 @@ func (s *OAuthServer) ProviderCallbackPath() string { // mux.Handle("/mcp/", mcpTransport) func (s *OAuthServer) HTTPHandler() http.Handler { mux := http.NewServeMux() - mux.HandleFunc("/.well-known/oauth-authorization-server", s.metadataHandler) - mux.HandleFunc("/.well-known/openid-configuration", s.openIDConfigurationHandler) - mux.HandleFunc("/.well-known/oauth-protected-resource", s.protectedResourceHandler) - mux.HandleFunc("/oauth/register", s.registerHandler) - mux.HandleFunc("/oauth/authorize", s.authorizeHandler) - mux.HandleFunc("/oauth/token", s.tokenHandler) - mux.HandleFunc("/oauth/revoke", s.revokeHandler) - mux.HandleFunc("/oauth/introspect", s.introspectHandler) - mux.HandleFunc("/oauth/userinfo", s.userinfoHandler) + handle := func(pattern, endpoint string, h http.HandlerFunc) { + mux.HandleFunc(pattern, s.limited(endpoint, h)) + } + for _, name := range []string{"oauth-authorization-server", "openid-configuration", "oauth-protected-resource"} { + h := s.metadataHandler(name) + mux.HandleFunc("/.well-known/"+name, h) + mux.HandleFunc("/.well-known/"+name+"/{path...}", h) + if p := strings.TrimSuffix(s.issuerURL.Path, "/"); p != "" { + mux.HandleFunc(p+"/.well-known/"+name, h) + } + } + handle("/oauth/register", "register", s.registerHandler) + handle("/oauth/register/{id}", "register", s.registrationManageHandler) + handle("/oauth/register/{id}/rotate-secret", "register", s.registrationRotateHandler) + handle("/oauth/authorize", "authorize", s.authorizeHandler) + handle("/oauth/token", "token", s.tokenHandler) + handle("/oauth/revoke", "revoke", s.revokeHandler) + handle("/oauth/introspect", "introspect", s.introspectHandler) + handle("/oauth/userinfo", "userinfo", s.userinfoHandler) mux.HandleFunc("/oauth/jwks.json", s.jwksHandler) + if s.cfg.EnablePAR || s.cfg.RequirePAR { + handle("/oauth/par", "par", s.parHandler) + } + if s.cfg.EnableDeviceFlow { + handle("/oauth/device_authorization", "device", s.deviceAuthorizationHandler) + handle("/oauth/device", "device", s.deviceVerificationHandler) + } + if !s.cfg.DisableLogout { + handle("/oauth/logout", "logout", s.logoutHandler) + } mux.HandleFunc(s.cfg.ProviderCallbackPath, s.providerCallbackHandler) return mux } +// limited applies OAuthServerConfig.RateLimiter. +func (s *OAuthServer) limited(endpoint string, h http.HandlerFunc) http.HandlerFunc { + if s.cfg.RateLimiter == nil { + return h + } + return func(w http.ResponseWriter, r *http.Request) { + if !s.cfg.RateLimiter(r, endpoint) { + w.Header().Set("Retry-After", "5") + writeOAuthError(w, "temporarily_unavailable", "rate limit exceeded", http.StatusTooManyRequests) + return + } + h(w, r) + } +} + // cleanupExpired removes stale pending auths and codes every 5 minutes. func (s *OAuthServer) cleanupExpired() { ticker := time.NewTicker(5 * time.Minute) @@ -264,8 +456,8 @@ func (s *OAuthServer) cleanupExpired() { delete(s.pending, k) } } - for k, p := range s.codes { - if now.After(p.ExpiresAt) { + for k, c := range s.codes { + if now.After(c.ExpiresAt) { delete(s.codes, k) } } @@ -275,868 +467,33 @@ func (s *OAuthServer) cleanupExpired() { } // -------------------------------------------------------------------------- -// RFC 8414 — Server metadata +// Collaborators // -------------------------------------------------------------------------- -// serverMetadata builds the fields shared by RFC 8414 authorization-server -// metadata and OIDC discovery metadata. -func (s *OAuthServer) serverMetadata() map[string]interface{} { - issuer := s.cfg.Issuer - grantTypes := []string{"authorization_code", "refresh_token"} +// anyAuth returns the authenticator used for sessions: the primary one, or the first provider's. +func (s *OAuthServer) anyAuth() *DatabaseAuthenticator { if s.auth != nil { - grantTypes = append(grantTypes, "client_credentials") + return s.auth } - return map[string]interface{}{ - "issuer": issuer, - "authorization_endpoint": issuer + "/oauth/authorize", - "token_endpoint": issuer + "/oauth/token", - "registration_endpoint": issuer + "/oauth/register", - "revocation_endpoint": issuer + "/oauth/revoke", - "introspection_endpoint": issuer + "/oauth/introspect", - "userinfo_endpoint": issuer + "/oauth/userinfo", - "jwks_uri": issuer + "/oauth/jwks.json", - "scopes_supported": s.cfg.DefaultScopes, - "response_types_supported": []string{"code"}, - "grant_types_supported": grantTypes, - "code_challenge_methods_supported": []string{"S256"}, - "token_endpoint_auth_methods_supported": []string{"none", "client_secret_basic", "client_secret_post"}, - } -} - -func (s *OAuthServer) metadataHandler(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(s.serverMetadata()) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored -} - -// -------------------------------------------------------------------------- -// OIDC discovery — GET /.well-known/openid-configuration -// -------------------------------------------------------------------------- - -func (s *OAuthServer) openIDConfigurationHandler(w http.ResponseWriter, r *http.Request) { - meta := s.serverMetadata() - meta["subject_types_supported"] = []string{"public"} - meta["id_token_signing_alg_values_supported"] = []string{"RS256"} - meta["claims_supported"] = []string{"sub", "iss", "aud", "exp", "iat", "email", "preferred_username"} - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(meta) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored -} - -// -------------------------------------------------------------------------- -// RFC 9728 — Protected Resource Metadata -// -------------------------------------------------------------------------- - -func (s *OAuthServer) protectedResourceHandler(w http.ResponseWriter, r *http.Request) { - meta := map[string]interface{}{ - "resource": s.cfg.ResourceIdentifier, - "authorization_servers": []string{s.cfg.Issuer}, - "scopes_supported": s.cfg.DefaultScopes, - "bearer_methods_supported": []string{"header"}, - } - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(meta) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored -} - -// -------------------------------------------------------------------------- -// JWKS — GET /oauth/jwks.json -// -------------------------------------------------------------------------- - -func (s *OAuthServer) jwksHandler(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - if s.signingKey == nil { - json.NewEncoder(w).Encode(map[string]interface{}{"keys": []interface{}{}}) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored - return - } - pub := s.signingKey.PublicKey - jwk := map[string]interface{}{ - "kty": "RSA", - "use": "sig", - "alg": "RS256", - "kid": s.signingKeyID, - "n": base64.RawURLEncoding.EncodeToString(pub.N.Bytes()), - "e": base64.RawURLEncoding.EncodeToString(bigEndianBytes(pub.E)), - } - json.NewEncoder(w).Encode(map[string]interface{}{"keys": []interface{}{jwk}}) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored -} - -// -------------------------------------------------------------------------- -// Userinfo — GET/POST /oauth/userinfo -// -------------------------------------------------------------------------- - -func (s *OAuthServer) userinfoHandler(w http.ResponseWriter, r *http.Request) { - auth := r.Header.Get("Authorization") - token := strings.TrimPrefix(auth, "Bearer ") - if token == "" || token == auth { - w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`) - writeOAuthError(w, "invalid_token", "missing bearer token", http.StatusUnauthorized) - return - } - - authToUse := s.auth - if authToUse == nil { - s.mu.RLock() - if len(s.providers) > 0 { - authToUse = s.providers[0].auth - } - s.mu.RUnlock() - } - if authToUse == nil { - writeOAuthError(w, "invalid_token", "no authenticator configured", http.StatusUnauthorized) - return - } - - info, err := authToUse.OAuthIntrospectToken(r.Context(), token) - if err != nil || !info.Active { - w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`) - writeOAuthError(w, "invalid_token", "token is inactive or invalid", http.StatusUnauthorized) - return - } - - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(map[string]interface{}{ //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored - "sub": info.Sub, - "preferred_username": info.Username, - "email": info.Email, - }) -} - -// -------------------------------------------------------------------------- -// 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 - } - var req struct { - RedirectURIs []string `json:"redirect_uris"` - ClientName string `json:"client_name"` - GrantTypes []string `json:"grant_types"` - AllowedScopes []string `json:"allowed_scopes"` - TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"` - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - writeOAuthError(w, "invalid_request", "malformed JSON", http.StatusBadRequest) - return - } - if len(req.RedirectURIs) == 0 { - writeOAuthError(w, "invalid_request", "redirect_uris required", http.StatusBadRequest) - return - } - grantTypes := req.GrantTypes - if len(grantTypes) == 0 { - grantTypes = []string{"authorization_code"} - } - allowedScopes := req.AllowedScopes - if len(allowedScopes) == 0 { - allowedScopes = s.cfg.DefaultScopes - } - clientID, err := randomOAuthToken() - if err != nil { - http.Error(w, "server error", http.StatusInternalServerError) - return - } - - // client_credentials is a machine-to-machine grant and requires a confidential - // client (RFC 6749 §4.4), so it always forces secret issuance regardless of the - // requested auth method. - authMethod := req.TokenEndpointAuthMethod - if authMethod == "" { - authMethod = "none" - } - if oauthSliceContains(grantTypes, "client_credentials") && authMethod == "none" { - authMethod = "client_secret_basic" - } - - var plaintextSecret string - var secretHash string - if authMethod != "none" { - plaintextSecret, err = randomOAuthToken() - if err != nil { - http.Error(w, "server error", http.StatusInternalServerError) - return - } - secretHash = hashClientSecret(plaintextSecret) - } - - client := &oauthClient{ - ClientID: clientID, - RedirectURIs: req.RedirectURIs, - ClientName: req.ClientName, - GrantTypes: grantTypes, - AllowedScopes: allowedScopes, - ClientSecretHash: secretHash, - TokenEndpointAuthMethod: authMethod, - } - - if s.cfg.PersistClients && s.auth != nil { - dbClient := &OAuthServerClient{ - ClientID: client.ClientID, - RedirectURIs: client.RedirectURIs, - ClientName: client.ClientName, - GrantTypes: client.GrantTypes, - AllowedScopes: client.AllowedScopes, - ClientSecretHash: client.ClientSecretHash, - TokenEndpointAuthMethod: client.TokenEndpointAuthMethod, - } - if _, err := s.auth.OAuthRegisterClient(r.Context(), dbClient); err != nil { - http.Error(w, "server error", http.StatusInternalServerError) - return - } - } - - s.mu.Lock() - s.clients[clientID] = client - s.mu.Unlock() - - // RFC 7591 registration response: the plaintext secret is returned exactly once here - // and never persisted or served again — only its hash (client.ClientSecretHash) is - // stored, and that hash is deliberately excluded from this response. - resp := map[string]interface{}{ - "client_id": client.ClientID, - "redirect_uris": client.RedirectURIs, - "client_name": client.ClientName, - "grant_types": client.GrantTypes, - "allowed_scopes": client.AllowedScopes, - "token_endpoint_auth_method": client.TokenEndpointAuthMethod, - } - if plaintextSecret != "" { - resp["client_secret"] = plaintextSecret - resp["client_secret_expires_at"] = 0 - } - - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusCreated) - json.NewEncoder(w).Encode(resp) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored -} - -// -------------------------------------------------------------------------- -// 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) - } -} - -// authorizeGet validates the request and either: -// - Redirects to an external provider (if providers are registered) -// - Renders a login form (if the server is its own identity provider) -func (s *OAuthServer) authorizeGet(w http.ResponseWriter, r *http.Request) { - q := r.URL.Query() - clientID := q.Get("client_id") - redirectURI := q.Get("redirect_uri") - clientState := q.Get("state") - codeChallenge := q.Get("code_challenge") - codeChallengeMethod := q.Get("code_challenge_method") - providerName := q.Get("provider") - scopeStr := q.Get("scope") - var scopes []string - if scopeStr != "" { - scopes = strings.Fields(scopeStr) - } - - if q.Get("response_type") != "code" { - writeOAuthError(w, "unsupported_response_type", "only 'code' is supported", http.StatusBadRequest) - return - } - if codeChallenge == "" { - writeOAuthError(w, "invalid_request", "code_challenge required (PKCE S256)", http.StatusBadRequest) - return - } - if codeChallengeMethod != "" && codeChallengeMethod != "S256" { - writeOAuthError(w, "invalid_request", "only S256 code_challenge_method is supported", http.StatusBadRequest) - return - } - client, ok := s.lookupOrFetchClient(r.Context(), clientID) - if !ok { - writeOAuthError(w, "invalid_client", "unknown client_id", http.StatusBadRequest) - return - } - if !oauthSliceContains(client.RedirectURIs, redirectURI) { - writeOAuthError(w, "invalid_request", "redirect_uri not registered", http.StatusBadRequest) - return - } - - // External provider path - if len(s.providers) > 0 { - s.redirectToExternalProvider(w, r, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, providerName, scopes) - return - } - - // Direct login form path (server is its own identity provider) - if s.auth == nil { - http.Error(w, "no authentication provider configured", http.StatusInternalServerError) - return - } - s.renderLoginForm(w, r, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, scopeStr, "") -} - -// authorizePost handles login form submission for the direct login flow. -func (s *OAuthServer) authorizePost(w http.ResponseWriter, r *http.Request) { - if err := r.ParseForm(); err != nil { - http.Error(w, "invalid form", http.StatusBadRequest) - return - } - - clientID := r.FormValue("client_id") - redirectURI := r.FormValue("redirect_uri") - clientState := r.FormValue("client_state") - codeChallenge := r.FormValue("code_challenge") - codeChallengeMethod := r.FormValue("code_challenge_method") - username := r.FormValue("username") - password := r.FormValue("password") - scopeStr := r.FormValue("scope") - var scopes []string - if scopeStr != "" { - scopes = strings.Fields(scopeStr) - } - - client, ok := s.lookupOrFetchClient(r.Context(), clientID) - if !ok || !oauthSliceContains(client.RedirectURIs, redirectURI) { - http.Error(w, "invalid client or redirect_uri", http.StatusBadRequest) - return - } - if s.auth == nil { - http.Error(w, "no authentication provider configured", http.StatusInternalServerError) - return - } - - loginResp, err := s.auth.Login(r.Context(), LoginRequest{ - Username: username, - Password: password, - }) - if err != nil { - s.renderLoginForm(w, r, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, scopeStr, "Invalid username or password") - return - } - - s.issueCodeAndRedirect(w, r, loginResp.Token, loginResp.RefreshToken, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, "", scopes) -} - -// redirectToExternalProvider stores the pending auth and redirects to the configured provider. -func (s *OAuthServer) redirectToExternalProvider(w http.ResponseWriter, r *http.Request, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, providerName string, scopes []string) { - var provider *externalProvider - if providerName != "" { - for i := range s.providers { - if s.providers[i].providerName == providerName { - provider = &s.providers[i] - break - } - } - if provider == nil { - http.Error(w, fmt.Sprintf("provider %q not found", providerName), http.StatusBadRequest) - return - } - } else { - provider = &s.providers[0] - } - - providerState, err := randomOAuthToken() - if err != nil { - http.Error(w, "server error", http.StatusInternalServerError) - return - } - - pending := &pendingAuth{ - ClientID: clientID, - RedirectURI: redirectURI, - ClientState: clientState, - CodeChallenge: codeChallenge, - CodeChallengeMethod: codeChallengeMethod, - ProviderName: provider.providerName, - ExpiresAt: time.Now().Add(10 * time.Minute), - Scopes: scopes, - } - s.mu.Lock() - s.pending[providerState] = pending - 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") - - if code == "" { - http.Error(w, "missing code", http.StatusBadRequest) - return - } - - 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 - } - - provider := s.providerByName(pending.ProviderName) - if provider == nil { - http.Error(w, fmt.Sprintf("provider %q not found", pending.ProviderName), http.StatusInternalServerError) - return - } - - loginResp, err := provider.auth.OAuth2HandleCallback(r.Context(), pending.ProviderName, code, providerState) - if err != nil { - http.Error(w, err.Error(), http.StatusUnauthorized) - return - } - - s.issueCodeAndRedirect(w, r, loginResp.Token, loginResp.RefreshToken, - pending.ClientID, pending.RedirectURI, pending.ClientState, - pending.CodeChallenge, pending.CodeChallengeMethod, pending.ProviderName, pending.Scopes) -} - -// issueCodeAndRedirect generates a short-lived auth code and redirects to the MCP client. -func (s *OAuthServer) issueCodeAndRedirect(w http.ResponseWriter, r *http.Request, sessionToken, refreshToken, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, providerName string, scopes []string) { - authCode, err := randomOAuthToken() - if err != nil { - http.Error(w, "server error", http.StatusInternalServerError) - return - } - - pending := &pendingAuth{ - ClientID: clientID, - RedirectURI: redirectURI, - ClientState: clientState, - CodeChallenge: codeChallenge, - CodeChallengeMethod: codeChallengeMethod, - ProviderName: providerName, - SessionToken: sessionToken, - RefreshToken: refreshToken, - ExpiresAt: time.Now().Add(s.cfg.AuthCodeTTL), - Scopes: scopes, - } - - if s.cfg.PersistCodes && s.auth != nil { - oauthCode := &OAuthCode{ - Code: authCode, - ClientID: clientID, - RedirectURI: redirectURI, - ClientState: clientState, - CodeChallenge: codeChallenge, - CodeChallengeMethod: codeChallengeMethod, - SessionToken: sessionToken, - RefreshToken: refreshToken, - Scopes: scopes, - ExpiresAt: pending.ExpiresAt, - } - if err := s.auth.OAuthSaveCode(r.Context(), oauthCode); err != nil { - http.Error(w, "server error", http.StatusInternalServerError) - return - } - } else { - s.mu.Lock() - s.codes[authCode] = pending - s.mu.Unlock() - } - - redirectURL, err := url.Parse(redirectURI) - if err != nil { - http.Error(w, "invalid redirect_uri", http.StatusInternalServerError) - return - } - qp := redirectURL.Query() - qp.Set("code", authCode) - if clientState != "" { - qp.Set("state", clientState) - } - redirectURL.RawQuery = qp.Encode() - http.Redirect(w, r, redirectURL.String(), http.StatusFound) -} - -// -------------------------------------------------------------------------- -// 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 - } - switch r.FormValue("grant_type") { - case "authorization_code": - s.handleAuthCodeGrant(w, r) - case "refresh_token": - s.handleRefreshGrant(w, r) - case "client_credentials": - s.handleClientCredentialsGrant(w, r) - default: - writeOAuthError(w, "unsupported_grant_type", "", http.StatusBadRequest) - } -} - -func (s *OAuthServer) handleAuthCodeGrant(w http.ResponseWriter, r *http.Request) { - code := r.FormValue("code") - redirectURI := r.FormValue("redirect_uri") - clientID := r.FormValue("client_id") - codeVerifier := r.FormValue("code_verifier") - - if code == "" || codeVerifier == "" { - writeOAuthError(w, "invalid_request", "code and code_verifier required", http.StatusBadRequest) - return - } - - // Confidential clients (those registered with a client_secret) must authenticate; - // public clients keep relying on PKCE alone, unchanged from prior behavior. - if client, ok := s.lookupOrFetchClient(r.Context(), clientID); ok && client.isConfidential() { - if _, err := s.authenticateClient(r); err != nil { - writeOAuthError(w, "invalid_client", err.Error(), http.StatusUnauthorized) - return - } - } - - var sessionToken string - var refreshToken string - var scopes []string - - if s.cfg.PersistCodes && s.auth != nil { - oauthCode, err := s.auth.OAuthExchangeCode(r.Context(), code) - if err != nil { - writeOAuthError(w, "invalid_grant", "code expired or invalid", http.StatusBadRequest) - return - } - if oauthCode.ClientID != clientID { - writeOAuthError(w, "invalid_client", "", http.StatusBadRequest) - return - } - if oauthCode.RedirectURI != redirectURI { - writeOAuthError(w, "invalid_grant", "redirect_uri mismatch", http.StatusBadRequest) - return - } - if !validatePKCESHA256(oauthCode.CodeChallenge, codeVerifier) { - writeOAuthError(w, "invalid_grant", "code_verifier invalid", http.StatusBadRequest) - return - } - sessionToken = oauthCode.SessionToken - refreshToken = oauthCode.RefreshToken - scopes = oauthCode.Scopes - } else { - s.mu.Lock() - pending, ok := s.codes[code] - if ok { - delete(s.codes, code) - } - s.mu.Unlock() - - if !ok || time.Now().After(pending.ExpiresAt) { - writeOAuthError(w, "invalid_grant", "code expired or invalid", http.StatusBadRequest) - return - } - if pending.ClientID != clientID { - writeOAuthError(w, "invalid_client", "", http.StatusBadRequest) - return - } - if pending.RedirectURI != redirectURI { - writeOAuthError(w, "invalid_grant", "redirect_uri mismatch", http.StatusBadRequest) - return - } - if !validatePKCESHA256(pending.CodeChallenge, codeVerifier) { - writeOAuthError(w, "invalid_grant", "code_verifier invalid", http.StatusBadRequest) - return - } - sessionToken = pending.SessionToken - refreshToken = pending.RefreshToken - scopes = pending.Scopes - } - - s.writeOAuthToken(w, r, sessionToken, refreshToken, clientID, scopes, true) -} - -func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request) { - refreshToken := r.FormValue("refresh_token") - providerName := r.FormValue("provider") - clientID := r.FormValue("client_id") - if refreshToken == "" { - writeOAuthError(w, "invalid_request", "refresh_token required", http.StatusBadRequest) - return - } - - // Try external providers first, then fall back to DatabaseAuthenticator - provider := s.providerByName(providerName) - if provider != nil { - loginResp, err := provider.auth.OAuth2RefreshToken(r.Context(), refreshToken, providerName) - if err != nil { - writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest) - return - } - s.writeOAuthToken(w, r, loginResp.Token, loginResp.RefreshToken, clientID, nil, false) - return - } - - if s.auth != nil { - loginResp, err := s.auth.RefreshToken(r.Context(), refreshToken) - if err != nil { - writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest) - return - } - s.writeOAuthToken(w, r, loginResp.Token, loginResp.RefreshToken, clientID, nil, false) - return - } - - writeOAuthError(w, "invalid_grant", "no provider available for refresh", http.StatusBadRequest) -} - -// -------------------------------------------------------------------------- -// RFC 6749 §4.4 — Client credentials grant -// -------------------------------------------------------------------------- - -func (s *OAuthServer) handleClientCredentialsGrant(w http.ResponseWriter, r *http.Request) { - if s.auth == nil { - writeOAuthError(w, "unsupported_grant_type", "client_credentials requires a local user store", http.StatusBadRequest) - return - } - - client, err := s.authenticateClient(r) - if err != nil { - w.Header().Set("WWW-Authenticate", `Basic realm="oauth"`) - writeOAuthError(w, "invalid_client", err.Error(), http.StatusUnauthorized) - return - } - if !oauthSliceContains(client.GrantTypes, "client_credentials") { - writeOAuthError(w, "unauthorized_client", "client is not authorized for client_credentials", http.StatusBadRequest) - return - } - - requested := strings.Fields(r.FormValue("scope")) - effectiveScopes := client.AllowedScopes - if len(requested) > 0 { - effectiveScopes = nil - for _, sc := range requested { - if oauthSliceContains(client.AllowedScopes, sc) { - effectiveScopes = append(effectiveScopes, sc) - } - } - if len(effectiveScopes) == 0 { - writeOAuthError(w, "invalid_scope", "no requested scope is allowed for this client", http.StatusBadRequest) - return - } - } - - // client_credentials tokens have no end user, but the rest of the stack (RLS-scoping - // hooks, introspection) expects every access token to resolve to a user_sessions row - // with a user_id. Represent the client as a deterministic synthetic "service account" - // user so the existing get-or-create/create-session/introspection pipeline handles it - // unchanged — no new tables or code paths required. - userCtx := &UserContext{ - UserName: "client:" + client.ClientID, - Email: "oauth-client-" + client.ClientID + "@service.internal", - RemoteID: client.ClientID, - Roles: effectiveScopes, - } - userID, err := s.auth.oauth2GetOrCreateUser(r.Context(), userCtx, "oauth2_client") - if err != nil { - http.Error(w, "server error", http.StatusInternalServerError) - return - } - - sessionToken, err := randomOAuthToken() - if err != nil { - http.Error(w, "server error", http.StatusInternalServerError) - return - } - expiresAt := time.Now().Add(s.cfg.AccessTokenTTL) - err = s.auth.oauth2CreateSession(r.Context(), sessionToken, userID, &oauth2.Token{ - AccessToken: sessionToken, - TokenType: "Bearer", - }, expiresAt, "oauth2_client") - if err != nil { - http.Error(w, "server error", http.StatusInternalServerError) - return - } - - // No refresh token per RFC 6749 §4.4.3, and no id_token — client_credentials has no - // end-user subject to represent in OIDC terms. - s.writeOAuthToken(w, r, sessionToken, "", client.ClientID, effectiveScopes, false) -} - -// -------------------------------------------------------------------------- -// RFC 7009 — Token revocation -// -------------------------------------------------------------------------- - -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 { - w.WriteHeader(http.StatusOK) - return - } - token := r.FormValue("token") - if token == "" { - w.WriteHeader(http.StatusOK) - return - } - - if s.auth != nil { - s.auth.OAuthRevokeToken(r.Context(), token) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored - } else { - // In external-provider-only mode, attempt revocation via the first provider's auth. - s.mu.RLock() - var providerAuth *DatabaseAuthenticator - if len(s.providers) > 0 { - providerAuth = s.providers[0].auth - } - s.mu.RUnlock() - if providerAuth != nil { - providerAuth.OAuthRevokeToken(r.Context(), token) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored - } - } - w.WriteHeader(http.StatusOK) -} - -// -------------------------------------------------------------------------- -// 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 { - w.Header().Set("Content-Type", "application/json") - w.Write([]byte(`{"active":false}`)) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored - return - } - token := r.FormValue("token") - w.Header().Set("Content-Type", "application/json") - - if token == "" { - w.Write([]byte(`{"active":false}`)) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored - return - } - - // Resolve the authenticator to use: prefer the primary auth, then the first provider's auth. - authToUse := s.auth - if authToUse == nil { - s.mu.RLock() - if len(s.providers) > 0 { - authToUse = s.providers[0].auth - } - s.mu.RUnlock() - } - if authToUse == nil { - w.Write([]byte(`{"active":false}`)) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored - return - } - - info, err := authToUse.OAuthIntrospectToken(r.Context(), token) - if err != nil { - w.Write([]byte(`{"active":false}`)) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored - return - } - json.NewEncoder(w).Encode(info) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored -} - -// -------------------------------------------------------------------------- -// Login form (direct identity provider mode) -// -------------------------------------------------------------------------- - -func (s *OAuthServer) renderLoginForm(w http.ResponseWriter, r *http.Request, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, scope, errMsg string) { - w.Header().Set("Content-Type", "text/html; charset=utf-8") - errHTML := "" - if errMsg != "" { - errHTML = `

` + htmlEscape(errMsg) + `

` - } - fmt.Fprintf(w, loginFormHTML, //nolint:gosec // G705: output is HTML-escaped - htmlEscape(s.cfg.LoginTitle), - htmlEscape(s.cfg.LoginTitle), - errHTML, - htmlEscape(clientID), - htmlEscape(redirectURI), - htmlEscape(clientState), - htmlEscape(codeChallenge), - htmlEscape(codeChallengeMethod), - htmlEscape(scope), - ) -} - -const loginFormHTML = ` -%s - -
-

%s

%s -
- - - - - - - - - -
` - -// -------------------------------------------------------------------------- -// Helpers -// -------------------------------------------------------------------------- - -// lookupOrFetchClient checks in-memory first, then DB if PersistClients is enabled. -func (s *OAuthServer) lookupOrFetchClient(ctx context.Context, clientID string) (*oauthClient, bool) { s.mu.RLock() - c, ok := s.clients[clientID] - s.mu.RUnlock() - if ok { - return c, true + defer s.mu.RUnlock() + if len(s.providers) > 0 { + return s.providers[0].auth } + return nil +} - if !s.cfg.PersistClients || s.auth == nil { - return nil, false +// grants returns the grant store, or nil when no authenticator is configured. +func (s *OAuthServer) grants() lookup.OAuthGrantStore { + if a := s.anyAuth(); a != nil { + return a.OAuthGrants() } - - dbClient, err := s.auth.OAuthGetClient(ctx, clientID) - if err != nil { - return nil, false - } - - c = &oauthClient{ - ClientID: dbClient.ClientID, - RedirectURIs: dbClient.RedirectURIs, - ClientName: dbClient.ClientName, - GrantTypes: dbClient.GrantTypes, - AllowedScopes: dbClient.AllowedScopes, - } - s.mu.Lock() - s.clients[clientID] = c - s.mu.Unlock() - return c, true + return nil } func (s *OAuthServer) providerByName(name string) *externalProvider { + s.mu.RLock() + defer s.mu.RUnlock() for i := range s.providers { if s.providers[i].providerName == name { return &s.providers[i] @@ -1149,9 +506,139 @@ func (s *OAuthServer) providerByName(name string) *externalProvider { return nil } +func (s *OAuthServer) hasProviders() bool { + s.mu.RLock() + defer s.mu.RUnlock() + return len(s.providers) > 0 +} + +// endpoint returns the absolute URL of a path below the issuer. +func (s *OAuthServer) endpoint(path string) string { return s.cfg.Issuer + path } + +// -------------------------------------------------------------------------- +// Clients +// -------------------------------------------------------------------------- + +// needsClientAuth reports whether the client must authenticate at the token endpoint. +func needsClientAuth(c *OAuthServerClient) bool { + return c.ClientSecretHash != "" || c.TokenEndpointAuthMethod == "private_key_jwt" +} + +// clientCacheTTL bounds how long a persisted client is served from memory, so changes made by +// other instances (or directly in the database) become visible. +const clientCacheTTL = 30 * time.Second + +// lookupOrFetchClient checks in-memory first, then DB if PersistClients is enabled. +func (s *OAuthServer) lookupOrFetchClient(ctx context.Context, clientID string) (*OAuthServerClient, bool) { + if clientID == "" { + return nil, false + } + s.mu.RLock() + c, ok := s.clients[clientID] + s.mu.RUnlock() + if ok && (c.at.IsZero() || time.Since(c.at) < clientCacheTTL) { + return c.c, true + } + + if !s.cfg.PersistClients || s.auth == nil { + if ok { + return c.c, true + } + return nil, false + } + + dbClient, err := s.auth.OAuthGetClient(ctx, clientID) + if err != nil { + s.mu.Lock() + delete(s.clients, clientID) + s.mu.Unlock() + return nil, false + } + s.mu.Lock() + s.clients[clientID] = &cachedClient{c: dbClient, at: time.Now()} + s.mu.Unlock() + return dbClient, true +} + +// saveClient stores a new client in memory and, with PersistClients, in the database. +func (s *OAuthServer) saveClient(ctx context.Context, c *OAuthServerClient, update bool) error { + cc := &cachedClient{c: c} + if s.cfg.PersistClients && s.auth != nil { + var err error + if update { + err = s.auth.OAuthUpdateClient(ctx, c) + } else { + _, err = s.auth.OAuthRegisterClient(ctx, c) + } + if err != nil { + return err + } + cc.at = time.Now() + } + s.mu.Lock() + s.clients[c.ClientID] = cc + s.mu.Unlock() + return nil +} + +func (s *OAuthServer) removeClient(ctx context.Context, clientID string) error { + if s.cfg.PersistClients && s.auth != nil { + if err := s.auth.OAuthDeleteClient(ctx, clientID); err != nil { + return err + } + } + s.mu.Lock() + delete(s.clients, clientID) + s.mu.Unlock() + return nil +} + +// RegisterTrustedClient registers a client programmatically and returns its plaintext secret +// (empty for a public client). Unlike dynamic registration it may set FirstParty (skips the +// consent screen), RequireConsent and RequirePAR, which a remote caller must not control. +// ClientID is generated when empty. Leave ClientSecretHash empty and set +// TokenEndpointAuthMethod to a client_secret_* method to have a secret generated. +func (s *OAuthServer) RegisterTrustedClient(ctx context.Context, c OAuthServerClient) (registered *OAuthServerClient, plainSecret string, err error) { + if c.ClientID == "" { + id, err := randomOAuthToken() + if err != nil { + return nil, "", err + } + c.ClientID = id + } + if len(c.GrantTypes) == 0 { + c.GrantTypes = []string{"authorization_code", "refresh_token"} + } + if len(c.AllowedScopes) == 0 { + c.AllowedScopes = s.cfg.DefaultScopes + } + if c.TokenEndpointAuthMethod == "" { + c.TokenEndpointAuthMethod = "none" + } + var secret string + if c.ClientSecretHash == "" && strings.HasPrefix(c.TokenEndpointAuthMethod, "client_secret_") { + var err error + if secret, err = randomOAuthToken(); err != nil { + return nil, "", err + } + c.ClientSecretHash = hashClientSecret(secret) + } + if c.ClientIDIssuedAt == 0 { + c.ClientIDIssuedAt = time.Now().Unix() + } + if err := s.saveClient(ctx, &c, false); err != nil { + return nil, "", err + } + return &c, secret, nil +} + +// -------------------------------------------------------------------------- +// Secrets and encoding helpers +// -------------------------------------------------------------------------- + func validatePKCESHA256(challenge, verifier string) bool { h := sha256.Sum256([]byte(verifier)) - return base64.RawURLEncoding.EncodeToString(h[:]) == challenge + return subtle.ConstantTimeCompare([]byte(base64.RawURLEncoding.EncodeToString(h[:])), []byte(challenge)) == 1 } func randomOAuthToken() (string, error) { @@ -1171,120 +658,35 @@ func oauthSliceContains(slice []string, s string) bool { return false } -// writeOAuthToken writes the token response. When issueIDToken is true and the granted -// scopes include "openid", an RS256-signed id_token is included (OIDC); client_credentials -// responses always pass issueIDToken=false since that grant has no end-user subject. -func (s *OAuthServer) writeOAuthToken(w http.ResponseWriter, r *http.Request, accessToken, refreshToken, clientID string, scopes []string, issueIDToken bool) { - expiresIn := int64(s.cfg.AccessTokenTTL.Seconds()) - resp := map[string]interface{}{ - "access_token": accessToken, - "token_type": "Bearer", - "expires_in": expiresIn, - } - if refreshToken != "" { - resp["refresh_token"] = refreshToken - } - if len(scopes) > 0 { - resp["scope"] = strings.Join(scopes, " ") - } - if issueIDToken && oauthSliceContains(scopes, "openid") { - if idToken, err := s.buildIDToken(r.Context(), accessToken, clientID, scopes); err == nil { - resp["id_token"] = idToken - } - } - w.Header().Set("Content-Type", "application/json") - w.Header().Set("Cache-Control", "no-store") - w.Header().Set("Pragma", "no-cache") - json.NewEncoder(w).Encode(resp) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored -} - -// buildIDToken issues an OIDC id_token for the just-issued access token by reusing the -// existing introspection pipeline to resolve the subject's claims. -func (s *OAuthServer) buildIDToken(ctx context.Context, accessToken, clientID string, scopes []string) (string, error) { - if s.signingKey == nil { - return "", fmt.Errorf("no signing key configured") - } - authToUse := s.auth - if authToUse == nil { - s.mu.RLock() - if len(s.providers) > 0 { - authToUse = s.providers[0].auth - } - s.mu.RUnlock() - } - if authToUse == nil { - return "", fmt.Errorf("no authenticator configured") - } - info, err := authToUse.OAuthIntrospectToken(ctx, accessToken) - if err != nil || !info.Active { - return "", fmt.Errorf("token not active") - } - - now := time.Now() - claims := jwt.MapClaims{ - "iss": s.cfg.Issuer, - "sub": info.Sub, - "aud": clientID, - "exp": now.Add(s.cfg.AccessTokenTTL).Unix(), - "iat": now.Unix(), - } - if oauthSliceContains(scopes, "profile") && info.Username != "" { - claims["preferred_username"] = info.Username - } - if oauthSliceContains(scopes, "email") && info.Email != "" { - claims["email"] = info.Email - } - - token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims) - token.Header["kid"] = s.signingKeyID - return token.SignedString(s.signingKey) -} - -// authenticateClient validates client_secret_basic (Authorization: Basic) or -// client_secret_post (client_id/client_secret form fields) credentials against a -// registered confidential client's stored secret hash. -func (s *OAuthServer) authenticateClient(r *http.Request) (*oauthClient, error) { - clientID, clientSecret, ok := r.BasicAuth() - if !ok { - clientID = r.FormValue("client_id") - clientSecret = r.FormValue("client_secret") - } - if clientID == "" || clientSecret == "" { - return nil, fmt.Errorf("client authentication required") - } - client, ok := s.lookupOrFetchClient(r.Context(), clientID) - if !ok || !client.isConfidential() { - return nil, fmt.Errorf("invalid client credentials") - } - if subtle.ConstantTimeCompare([]byte(hashClientSecret(clientSecret)), []byte(client.ClientSecretHash)) != 1 { - return nil, fmt.Errorf("invalid client credentials") - } - return client, nil -} - func hashClientSecret(secret string) string { sum := sha256.Sum256([]byte(secret)) return hex.EncodeToString(sum[:]) } -// rsaKeyID derives a stable JWKS "kid" from an RSA public key's modulus. -func rsaKeyID(pub *rsa.PublicKey) string { - sum := sha256.Sum256(pub.N.Bytes()) - return base64.RawURLEncoding.EncodeToString(sum[:8]) +// hashToken is the stored form of refresh tokens and access-grant keys. +func hashToken(token string) string { return hashClientSecret(token) } + +type oauthError struct { + Code string + Desc string + Status int + // WWWAuth is set as the WWW-Authenticate header. + WWWAuth string } -// bigEndianBytes encodes a small positive int (e.g. an RSA public exponent) as -// minimal big-endian bytes for JWK "e" encoding. -func bigEndianBytes(n int) []byte { - if n == 0 { - return []byte{0} +func (e *oauthError) write(w http.ResponseWriter) { + if e.WWWAuth != "" { + w.Header().Set("WWW-Authenticate", e.WWWAuth) } - var b []byte - for n > 0 { - b = append([]byte{byte(n & 0xff)}, b...) - n >>= 8 - } - return b + writeOAuthError(w, e.Code, e.Desc, e.Status) +} + +func oerr(code, desc string, status int) *oauthError { + return &oauthError{Code: code, Desc: desc, Status: status} +} + +func serverErr() *oauthError { + return oerr("server_error", "internal error", http.StatusInternalServerError) } func writeOAuthError(w http.ResponseWriter, errCode, description string, status int) { @@ -1293,10 +695,19 @@ func writeOAuthError(w http.ResponseWriter, errCode, description string, status resp["error_description"] = description } w.Header().Set("Content-Type", "application/json") + w.Header().Set("Cache-Control", "no-store") w.WriteHeader(status) json.NewEncoder(w).Encode(resp) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored } +func writeJSON(w http.ResponseWriter, status int, v any) { + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("Pragma", "no-cache") + w.WriteHeader(status) + json.NewEncoder(w).Encode(v) //nolint:errcheck,gosec // G104: best-effort write, error intentionally ignored +} + func htmlEscape(s string) string { s = strings.ReplaceAll(s, "&", "&") s = strings.ReplaceAll(s, `"`, """) diff --git a/pkg/security/oauth_server_db.go b/pkg/security/oauth_server_db.go index 330dc4f..5e02c45 100644 --- a/pkg/security/oauth_server_db.go +++ b/pkg/security/oauth_server_db.go @@ -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) +} diff --git a/pkg/security/oauth_session.go b/pkg/security/oauth_session.go new file mode 100644 index 0000000..48d2397 --- /dev/null +++ b/pkg/security/oauth_session.go @@ -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) +} diff --git a/pkg/security/oauth_templates.go b/pkg/security/oauth_templates.go new file mode 100644 index 0000000..5a65418 --- /dev/null +++ b/pkg/security/oauth_templates.go @@ -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"}}{{.Title}}
{{end}} +{{define "foot"}}
{{end}} +{{define "hidden"}}{{range $k, $v := .Hidden}}{{end}}{{end}} + +{{define "login"}}{{template "head" .}}

{{.Title}}

{{if .ClientName}}

to continue to {{.ClientName}}

{{end}} +{{if .Error}}
{{.Error}}
{{end}} +
+{{template "hidden" .}} + + + +
{{template "foot"}}{{end}} + +{{define "consent"}}{{template "head" .}}{{if .LogoURI}}{{end}} +

{{if .ClientURI}}{{.ClientName}}{{else}}{{.ClientName}}{{end}} wants access

+

Signed in as {{.User}}. This application will be able to:

+
    {{range .Scopes}}
  • {{.Name}}{{if .Description}} – {{.Description}}{{end}}
  • {{end}}
+
+{{template "hidden" .}} +

+
+
+
{{template "foot"}}{{end}} + +{{define "device"}}{{template "head" .}}

{{.Title}}

Enter the code shown on your device.

+{{if .Error}}
{{.Error}}
{{end}} +
+ +
{{template "foot"}}{{end}} + +{{define "message"}}{{template "head" .}}

{{.Title}}

{{if .Error}}
{{.Message}}
{{else}}

{{.Message}}

{{end}}{{template "foot"}}{{end}} + +{{define "logout"}}{{template "head" .}}

{{.Title}}

Do you want to sign out?

+
{{template "hidden" .}} +
+
{{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}) +} diff --git a/pkg/security/oauth_token.go b/pkg/security/oauth_token.go new file mode 100644 index 0000000..a40a5e2 --- /dev/null +++ b/pkg/security/oauth_token.go @@ -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, + }) +} diff --git a/pkg/security/oidc_client.go b/pkg/security/oidc_client.go new file mode 100644 index 0000000..2d085ba --- /dev/null +++ b/pkg/security/oidc_client.go @@ -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 +} diff --git a/pkg/security/oidc_client_test.go b/pkg/security/oidc_client_test.go new file mode 100644 index 0000000..076c950 --- /dev/null +++ b/pkg/security/oidc_client_test.go @@ -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(`