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