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

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

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

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

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

Docs: OAUTH2_SERVER.md, oauth2_full_example.go, breaking_changes.md step 8.
This commit is contained in:
Hein
2026-10-01 14:42:12 +02:00
parent f54b707040
commit 640faeeeaf
51 changed files with 8682 additions and 1125 deletions
+240 -43
View File
@@ -8,6 +8,8 @@ import (
"encoding/json"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"sync"
"time"
@@ -32,6 +34,45 @@ type OAuth2Config struct {
// Optional: Custom user info parser
// If not provided, will use standard claims (sub, email, name)
UserInfoParser func(userInfo map[string]any) (*UserContext, error)
// --- OpenID Connect (see oidc_client.go) ---
// Issuer turns the provider into an OpenID Connect provider: PKCE and a nonce are used and
// the id_token returned by the token endpoint is validated (signature, iss, aud, exp, nonce,
// at_hash). WithOIDC fills the endpoints in by discovery; with WithOAuth2 set JWKSURL
// as well. UserInfoURL stays optional: the id_token claims are used when it is empty.
Issuer string
// JWKSURL is the provider's key set. Only needed with WithOAuth2; WithOIDC discovers it.
JWKSURL string
// EndSessionURL is the provider's RP-initiated logout endpoint (discovered by WithOIDC).
EndSessionURL string
// UsePKCE sends a PKCE S256 challenge for a provider that is not OIDC. It is always on in OIDC mode.
UsePKCE bool
// AllowedAlgs lists the id_token signature algorithms to accept. Default: RS256, PS256, ES256, ES384.
AllowedAlgs []string
// AuthStyle selects how the client authenticates at the token endpoint: "basic", "post" or ""
// (try basic, fall back to post).
AuthStyle string
// HTTPClient is used for discovery, JWKS, token and userinfo requests.
HTTPClient *http.Client
// ClockSkew tolerates clock differences when validating the id_token. Default 1 minute.
ClockSkew time.Duration
}
// OAuth2AuthOptions are optional OpenID Connect authentication request parameters.
type OAuth2AuthOptions struct {
LoginHint string
Prompt string // none, login, consent, select_account
MaxAge *int
ACRValues string
Extra map[string]string
}
// oauth2State is what the login redirect remembers until the callback.
type oauth2State struct {
expiry time.Time
verifier string // PKCE code_verifier
nonce string
}
// OAuth2Provider holds configuration and state for a single OAuth2 provider
@@ -40,7 +81,10 @@ type OAuth2Provider struct {
userInfoURL string
userInfoParser func(userInfo map[string]any) (*UserContext, error)
providerName string
states map[string]time.Time // state -> expiry time
states map[string]*oauth2State
oidc *oidcProvider // nil for plain OAuth2
usePKCE bool
httpClient *http.Client
statesMutex sync.RWMutex
stopCh chan struct{} // closed to stop cleanupStates
stopOnce sync.Once
@@ -58,6 +102,13 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
cfg.UserInfoParser = defaultOAuth2UserInfoParser
}
authStyle := oauth2.AuthStyleAutoDetect
switch cfg.AuthStyle {
case "basic":
authStyle = oauth2.AuthStyleInHeader
case "post":
authStyle = oauth2.AuthStyleInParams
}
provider := &OAuth2Provider{
config: &oauth2.Config{
ClientID: cfg.ClientID,
@@ -65,15 +116,22 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
RedirectURL: cfg.RedirectURL,
Scopes: cfg.Scopes,
Endpoint: oauth2.Endpoint{
AuthURL: cfg.AuthURL,
TokenURL: cfg.TokenURL,
AuthURL: cfg.AuthURL,
TokenURL: cfg.TokenURL,
AuthStyle: authStyle,
},
},
userInfoURL: cfg.UserInfoURL,
userInfoParser: cfg.UserInfoParser,
providerName: cfg.ProviderName,
states: make(map[string]time.Time),
states: make(map[string]*oauth2State),
stopCh: make(chan struct{}),
usePKCE: cfg.UsePKCE,
httpClient: cfg.HTTPClient,
}
if cfg.Issuer != "" {
provider.oidc = newOIDCProvider(&cfg)
provider.usePKCE = true
}
// Initialize providers map if needed
@@ -97,17 +155,56 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
// OAuth2GetAuthURL returns the OAuth2 authorization URL for redirecting users
func (a *DatabaseAuthenticator) OAuth2GetAuthURL(providerName, state string) (string, error) {
return a.OAuth2GetAuthURLWithOptions(providerName, state, OAuth2AuthOptions{})
}
// OAuth2GetAuthURLWithOptions is OAuth2GetAuthURL with OpenID Connect request parameters. For an
// OIDC provider (and with UsePKCE) it also creates the PKCE verifier and the nonce, which are
// kept with the state until the callback.
func (a *DatabaseAuthenticator) OAuth2GetAuthURLWithOptions(providerName, state string, opts OAuth2AuthOptions) (string, error) {
provider, err := a.getOAuth2Provider(providerName)
if err != nil {
return "", err
}
if provider.oidc != nil {
dctx, cancel := context.WithTimeout(provider.withHTTPClient(context.Background()), 15*time.Second)
defer cancel()
if err := provider.oidc.ensureEndpoints(dctx, provider); err != nil {
return "", err
}
}
st := &oauth2State{expiry: time.Now().Add(10 * time.Minute)}
var params []oauth2.AuthCodeOption
if provider.usePKCE {
st.verifier = oauth2.GenerateVerifier()
params = append(params, oauth2.S256ChallengeOption(st.verifier))
}
if provider.oidc != nil {
if st.nonce, err = randomOAuthToken(); err != nil {
return "", err
}
params = append(params, oauth2.SetAuthURLParam("nonce", st.nonce))
}
set := func(k, v string) {
if v != "" {
params = append(params, oauth2.SetAuthURLParam(k, v))
}
}
set("login_hint", opts.LoginHint)
set("prompt", opts.Prompt)
set("acr_values", opts.ACRValues)
if opts.MaxAge != nil {
set("max_age", strconv.Itoa(*opts.MaxAge))
}
for k, v := range opts.Extra {
set(k, v)
}
// Store state for validation
provider.statesMutex.Lock()
provider.states[state] = time.Now().Add(10 * time.Minute)
provider.states[state] = st
provider.statesMutex.Unlock()
return provider.config.AuthCodeURL(state), nil
return provider.config.AuthCodeURL(state, params...), nil
}
// OAuth2GenerateState generates a random state string for CSRF protection
@@ -121,42 +218,97 @@ func (a *DatabaseAuthenticator) OAuth2GenerateState() (string, error) {
// OAuth2HandleCallback handles the OAuth2 callback and exchanges code for token
func (a *DatabaseAuthenticator) OAuth2HandleCallback(ctx context.Context, providerName, code, state string) (*LoginResponse, error) {
return a.oauth2Callback(ctx, providerName, code, state, "")
}
// OAuth2HandleCallbackRequest is OAuth2HandleCallback for the redirect request itself. Besides
// code and state it honours the error parameters and the RFC 9207 "iss" parameter, which
// protects against mix-up attacks when several providers are in use.
func (a *DatabaseAuthenticator) OAuth2HandleCallbackRequest(ctx context.Context, providerName string, r *http.Request) (*LoginResponse, error) {
q := r.URL.Query()
if e := q.Get("error"); e != "" {
return nil, fmt.Errorf("provider returned an error: %s %s", e, q.Get("error_description"))
}
return a.oauth2Callback(ctx, providerName, q.Get("code"), q.Get("state"), q.Get("iss"))
}
func (a *DatabaseAuthenticator) oauth2Callback(ctx context.Context, providerName, code, state, iss string) (*LoginResponse, error) {
provider, err := a.getOAuth2Provider(providerName)
if err != nil {
return nil, err
}
// Validate state
if !provider.validateState(state) {
st, ok := provider.validateState(state)
if !ok {
return nil, fmt.Errorf("invalid state parameter")
}
if code == "" {
return nil, fmt.Errorf("missing authorization code")
}
if provider.oidc != nil && iss != "" && iss != provider.oidc.issuer {
return nil, fmt.Errorf("authorization response issuer mismatch")
}
if ctx = provider.withHTTPClient(ctx); ctx == nil {
return nil, fmt.Errorf("no context")
}
// Exchange code for token
token, err := provider.config.Exchange(ctx, code)
var exchange []oauth2.AuthCodeOption
if st.verifier != "" {
exchange = append(exchange, oauth2.VerifierOption(st.verifier))
}
if provider.oidc != nil {
if err := provider.oidc.ensureEndpoints(ctx, provider); err != nil {
return nil, err
}
}
token, err := provider.config.Exchange(ctx, code, exchange...)
if err != nil {
return nil, fmt.Errorf("failed to exchange code: %w", err)
}
// OpenID Connect: validate the id_token.
var rawIDToken string
var idClaims map[string]any
if provider.oidc != nil {
rawIDToken, _ = token.Extra("id_token").(string)
if rawIDToken == "" && oauthSliceContains(provider.config.Scopes, "openid") {
return nil, fmt.Errorf("token response contains no id_token")
}
if rawIDToken != "" {
if idClaims, err = provider.oidc.validateIDToken(ctx, rawIDToken, st.nonce, token.AccessToken); err != nil {
return nil, fmt.Errorf("invalid id_token: %w", err)
}
}
}
// Fetch user info
client := provider.config.Client(ctx, token)
resp, err := client.Get(provider.userInfoURL)
if err != nil {
return nil, fmt.Errorf("failed to fetch user info: %w", err)
userInfo := map[string]any{}
if provider.userInfoURL != "" {
fetched, err := provider.fetchUserInfo(ctx, token)
switch {
case err == nil:
if sub, _ := idClaims["sub"].(string); sub != "" {
if us, _ := fetched["sub"].(string); us != "" && us != sub {
return nil, fmt.Errorf("userinfo subject does not match the id_token")
}
}
userInfo = fetched
case provider.oidc == nil || idClaims == nil:
return nil, err
}
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("failed to read user info: %w", err)
claims := map[string]any{}
for k, v := range idClaims {
claims[k] = v
}
var userInfo map[string]any
if err := json.Unmarshal(body, &userInfo); err != nil {
return nil, fmt.Errorf("failed to parse user info: %w", err)
for k, v := range userInfo {
claims[k] = v
}
// Parse user info
userCtx, err := provider.userInfoParser(userInfo)
userCtx, err := provider.userInfoParser(claims)
if err != nil {
return nil, fmt.Errorf("failed to parse user context: %w", err)
}
@@ -187,12 +339,47 @@ func (a *DatabaseAuthenticator) OAuth2HandleCallback(ctx context.Context, provid
userCtx.SessionID = sessionToken
return &LoginResponse{
resp := &LoginResponse{
Token: sessionToken,
RefreshToken: token.RefreshToken,
User: userCtx,
ExpiresIn: int64(time.Until(expiresAt).Seconds()),
}, nil
}
if rawIDToken != "" {
// Keep the id_token: it is the id_token_hint of OAuth2LogoutURL.
resp.Meta = map[string]any{"id_token": rawIDToken}
}
return resp, nil
}
// fetchUserInfo calls the provider's userinfo endpoint with the access token.
func (p *OAuth2Provider) fetchUserInfo(ctx context.Context, token *oauth2.Token) (map[string]any, error) {
resp, err := p.config.Client(ctx, token).Get(p.userInfoURL)
if err != nil {
return nil, fmt.Errorf("failed to fetch user info: %w", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
if err != nil {
return nil, fmt.Errorf("failed to read user info: %w", err)
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("user info request failed: status %d", resp.StatusCode)
}
var userInfo map[string]any
if err := json.Unmarshal(body, &userInfo); err != nil {
return nil, fmt.Errorf("failed to parse user info: %w", err)
}
return userInfo, nil
}
// withHTTPClient makes oauth2 use the provider's HTTP client.
func (p *OAuth2Provider) withHTTPClient(ctx context.Context) context.Context {
if p.httpClient == nil {
return ctx
}
return context.WithValue(ctx, oauth2.HTTPClient, p.httpClient)
}
// OAuth2GetProviders returns list of configured OAuth2 provider names
@@ -251,23 +438,20 @@ func (a *DatabaseAuthenticator) oauth2CreateSession(ctx context.Context, session
})
}
// validateState validates state using in-memory storage
func (p *OAuth2Provider) validateState(state string) bool {
// validateState validates state using in-memory storage and returns what was remembered with it.
func (p *OAuth2Provider) validateState(state string) (*oauth2State, bool) {
p.statesMutex.Lock()
defer p.statesMutex.Unlock()
expiry, ok := p.states[state]
st, ok := p.states[state]
if !ok {
return false
return nil, false
}
if time.Now().After(expiry) {
delete(p.states, state)
return false
}
delete(p.states, state) // One-time use
return true
if time.Now().After(st.expiry) {
return nil, false
}
return st, true
}
// cleanupStates removes expired states periodically
@@ -284,8 +468,8 @@ func (p *OAuth2Provider) cleanupStates() {
}
p.statesMutex.Lock()
now := time.Now()
for state, expiry := range p.states {
if now.After(expiry) {
for state, st := range p.states {
if now.After(st.expiry) {
delete(p.states, state)
}
}
@@ -363,7 +547,7 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
}
// Use OAuth2 provider to refresh the token
tokenSource := provider.config.TokenSource(ctx, oldToken)
tokenSource := provider.config.TokenSource(provider.withHTTPClient(ctx), oldToken)
newToken, err := tokenSource.Token()
if err != nil {
return nil, fmt.Errorf("failed to refresh token with provider: %w", err)
@@ -388,12 +572,21 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshT
userCtx.SessionID = newSessionToken
return &LoginResponse{
resp := &LoginResponse{
Token: newSessionToken,
RefreshToken: newToken.RefreshToken,
User: userCtx,
ExpiresIn: int64(time.Until(newToken.Expiry).Seconds()),
}, nil
}
if provider.oidc != nil {
if raw, _ := newToken.Extra("id_token").(string); raw != "" {
if _, err := provider.oidc.validateIDToken(provider.withHTTPClient(ctx), raw, "", newToken.AccessToken); err != nil {
return nil, fmt.Errorf("invalid id_token in refresh response: %w", err)
}
resp.Meta = map[string]any{"id_token": raw}
}
}
return resp, nil
}
// Pre-configured OAuth2 factory methods
@@ -406,10 +599,14 @@ func NewGoogleAuthenticator(clientID, clientSecret, redirectURL string, db *sql.
ClientSecret: clientSecret,
RedirectURL: redirectURL,
Scopes: []string{"openid", "profile", "email"},
AuthURL: "https://accounts.google.com/o/oauth2/auth",
AuthURL: "https://accounts.google.com/o/oauth2/v2/auth",
TokenURL: "https://oauth2.googleapis.com/token",
UserInfoURL: "https://www.googleapis.com/oauth2/v2/userinfo",
UserInfoURL: "https://openidconnect.googleapis.com/v1/userinfo",
ProviderName: "google",
// OpenID Connect: PKCE, nonce and id_token validation against Google's published keys.
Issuer: "https://accounts.google.com",
JWKSURL: "https://www.googleapis.com/oauth2/v3/certs",
EndSessionURL: "",
})
}