mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 03:22:09 +00:00
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.
669 lines
22 KiB
Go
669 lines
22 KiB
Go
package security
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"database/sql"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
|
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
|
|
|
"golang.org/x/oauth2"
|
|
)
|
|
|
|
// OAuth2Config contains configuration for OAuth2 authentication
|
|
type OAuth2Config struct {
|
|
ClientID string
|
|
ClientSecret string
|
|
RedirectURL string
|
|
Scopes []string
|
|
AuthURL string
|
|
TokenURL string
|
|
UserInfoURL string
|
|
ProviderName string
|
|
|
|
// 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
|
|
type OAuth2Provider struct {
|
|
config *oauth2.Config
|
|
userInfoURL string
|
|
userInfoParser func(userInfo map[string]any) (*UserContext, error)
|
|
providerName string
|
|
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
|
|
}
|
|
|
|
// WithOAuth2 configures OAuth2 support for the DatabaseAuthenticator
|
|
// Can be called multiple times to add multiple OAuth2 providers
|
|
// Returns the same DatabaseAuthenticator instance for method chaining
|
|
func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthenticator {
|
|
if cfg.ProviderName == "" {
|
|
cfg.ProviderName = "oauth2"
|
|
}
|
|
|
|
if cfg.UserInfoParser == nil {
|
|
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,
|
|
ClientSecret: cfg.ClientSecret,
|
|
RedirectURL: cfg.RedirectURL,
|
|
Scopes: cfg.Scopes,
|
|
Endpoint: oauth2.Endpoint{
|
|
AuthURL: cfg.AuthURL,
|
|
TokenURL: cfg.TokenURL,
|
|
AuthStyle: authStyle,
|
|
},
|
|
},
|
|
userInfoURL: cfg.UserInfoURL,
|
|
userInfoParser: cfg.UserInfoParser,
|
|
providerName: cfg.ProviderName,
|
|
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
|
|
a.oauth2ProvidersMutex.Lock()
|
|
if a.oauth2Providers == nil {
|
|
a.oauth2Providers = make(map[string]*OAuth2Provider)
|
|
}
|
|
|
|
// Register provider
|
|
if old := a.oauth2Providers[cfg.ProviderName]; old != nil {
|
|
old.stop() // replaced provider: stop its cleanup goroutine
|
|
}
|
|
a.oauth2Providers[cfg.ProviderName] = provider
|
|
a.oauth2ProvidersMutex.Unlock()
|
|
|
|
// Start state cleanup goroutine for this provider
|
|
go provider.cleanupStates()
|
|
|
|
return a
|
|
}
|
|
|
|
// 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)
|
|
}
|
|
|
|
provider.statesMutex.Lock()
|
|
provider.states[state] = st
|
|
provider.statesMutex.Unlock()
|
|
|
|
return provider.config.AuthCodeURL(state, params...), nil
|
|
}
|
|
|
|
// OAuth2GenerateState generates a random state string for CSRF protection
|
|
func (a *DatabaseAuthenticator) OAuth2GenerateState() (string, error) {
|
|
b := make([]byte, 32)
|
|
if _, err := rand.Read(b); err != nil {
|
|
return "", err
|
|
}
|
|
return base64.URLEncoding.EncodeToString(b), nil
|
|
}
|
|
|
|
// 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
|
|
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
|
|
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
|
|
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
|
|
}
|
|
}
|
|
claims := map[string]any{}
|
|
for k, v := range idClaims {
|
|
claims[k] = v
|
|
}
|
|
for k, v := range userInfo {
|
|
claims[k] = v
|
|
}
|
|
|
|
// Parse user info
|
|
userCtx, err := provider.userInfoParser(claims)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
|
}
|
|
|
|
// Get or create user in database
|
|
userID, err := a.oauth2GetOrCreateUser(ctx, userCtx, providerName)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to get or create user: %w", err)
|
|
}
|
|
userCtx.UserID = userID
|
|
|
|
// Create session token
|
|
sessionToken, err := a.OAuth2GenerateState()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to generate session token: %w", err)
|
|
}
|
|
|
|
expiresAt := time.Now().Add(24 * time.Hour)
|
|
if token.Expiry.After(time.Now()) {
|
|
expiresAt = token.Expiry
|
|
}
|
|
|
|
// Store session in database
|
|
err = a.oauth2CreateSession(ctx, sessionToken, userCtx.UserID, token, expiresAt, providerName)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to create session: %w", err)
|
|
}
|
|
|
|
userCtx.SessionID = sessionToken
|
|
|
|
resp := &LoginResponse{
|
|
Token: sessionToken,
|
|
RefreshToken: token.RefreshToken,
|
|
User: userCtx,
|
|
ExpiresIn: int64(time.Until(expiresAt).Seconds()),
|
|
}
|
|
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
|
|
func (a *DatabaseAuthenticator) OAuth2GetProviders() []string {
|
|
a.oauth2ProvidersMutex.RLock()
|
|
defer a.oauth2ProvidersMutex.RUnlock()
|
|
|
|
if a.oauth2Providers == nil {
|
|
return nil
|
|
}
|
|
|
|
providers := make([]string, 0, len(a.oauth2Providers))
|
|
for name := range a.oauth2Providers {
|
|
providers = append(providers, name)
|
|
}
|
|
return providers
|
|
}
|
|
|
|
// getOAuth2Provider retrieves a registered OAuth2 provider by name
|
|
func (a *DatabaseAuthenticator) getOAuth2Provider(providerName string) (*OAuth2Provider, error) {
|
|
a.oauth2ProvidersMutex.RLock()
|
|
defer a.oauth2ProvidersMutex.RUnlock()
|
|
|
|
if a.oauth2Providers == nil {
|
|
return nil, fmt.Errorf("OAuth2 not configured - call WithOAuth2() first")
|
|
}
|
|
|
|
provider, ok := a.oauth2Providers[providerName]
|
|
if !ok {
|
|
// Build provider list without calling OAuth2GetProviders to avoid recursion
|
|
providerNames := make([]string, 0, len(a.oauth2Providers))
|
|
for name := range a.oauth2Providers {
|
|
providerNames = append(providerNames, name)
|
|
}
|
|
return nil, fmt.Errorf("OAuth2 provider '%s' not found - available providers: %v", providerName, providerNames)
|
|
}
|
|
|
|
return provider, nil
|
|
}
|
|
|
|
// oauth2GetOrCreateUser finds or creates a user based on OAuth2 info using stored procedure
|
|
func (a *DatabaseAuthenticator) oauth2GetOrCreateUser(ctx context.Context, userCtx *UserContext, providerName string) (int, error) {
|
|
return a.src.get().OAuthUser.GetOrCreateUser(ctx, userCtx, providerName)
|
|
}
|
|
|
|
// oauth2CreateSession creates a new OAuth2 session using stored procedure
|
|
func (a *DatabaseAuthenticator) oauth2CreateSession(ctx context.Context, sessionToken string, userID int, token *oauth2.Token, expiresAt time.Time, providerName string) error {
|
|
return a.src.get().OAuthUser.CreateSession(ctx, lookup.OAuthSession{
|
|
SessionToken: sessionToken,
|
|
UserID: userID,
|
|
AccessToken: token.AccessToken,
|
|
RefreshToken: token.RefreshToken,
|
|
TokenType: token.TokenType,
|
|
ExpiresAt: expiresAt,
|
|
Provider: providerName,
|
|
})
|
|
}
|
|
|
|
// 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()
|
|
|
|
st, ok := p.states[state]
|
|
if !ok {
|
|
return nil, false
|
|
}
|
|
delete(p.states, state) // One-time use
|
|
if time.Now().After(st.expiry) {
|
|
return nil, false
|
|
}
|
|
return st, true
|
|
}
|
|
|
|
// cleanupStates removes expired states periodically
|
|
func (p *OAuth2Provider) cleanupStates() {
|
|
defer logger.CatchPanic("OAuth2Provider.cleanupStates")()
|
|
ticker := time.NewTicker(5 * time.Minute)
|
|
defer ticker.Stop()
|
|
|
|
for {
|
|
select {
|
|
case <-p.stopCh:
|
|
return
|
|
case <-ticker.C:
|
|
}
|
|
p.statesMutex.Lock()
|
|
now := time.Now()
|
|
for state, st := range p.states {
|
|
if now.After(st.expiry) {
|
|
delete(p.states, state)
|
|
}
|
|
}
|
|
p.statesMutex.Unlock()
|
|
}
|
|
}
|
|
|
|
// stop terminates the cleanup goroutine; safe to call more than once.
|
|
func (p *OAuth2Provider) stop() {
|
|
p.stopOnce.Do(func() { close(p.stopCh) })
|
|
}
|
|
|
|
// Close stops the background OAuth2 state cleanup goroutines and waits for
|
|
// in-flight session activity updates. It is safe to call more than once.
|
|
func (a *DatabaseAuthenticator) Close() error {
|
|
a.oauth2ProvidersMutex.RLock()
|
|
for _, p := range a.oauth2Providers {
|
|
p.stop()
|
|
}
|
|
a.oauth2ProvidersMutex.RUnlock()
|
|
a.activityWG.Wait()
|
|
return nil
|
|
}
|
|
|
|
// defaultOAuth2UserInfoParser parses standard OAuth2 user info claims
|
|
func defaultOAuth2UserInfoParser(userInfo map[string]any) (*UserContext, error) {
|
|
ctx := &UserContext{
|
|
Claims: userInfo,
|
|
Roles: []string{"user"},
|
|
}
|
|
|
|
// Extract standard claims
|
|
if sub, ok := userInfo["sub"].(string); ok {
|
|
ctx.RemoteID = sub
|
|
}
|
|
if email, ok := userInfo["email"].(string); ok {
|
|
ctx.Email = email
|
|
// Use email as username if name not available
|
|
ctx.UserName = strings.Split(email, "@")[0]
|
|
}
|
|
if name, ok := userInfo["name"].(string); ok {
|
|
ctx.UserName = name
|
|
}
|
|
if login, ok := userInfo["login"].(string); ok {
|
|
ctx.UserName = login // GitHub uses "login"
|
|
}
|
|
|
|
if ctx.UserName == "" {
|
|
return nil, fmt.Errorf("could not extract username from user info")
|
|
}
|
|
|
|
return ctx, nil
|
|
}
|
|
|
|
// OAuth2RefreshToken refreshes an expired OAuth2 access token using the refresh token
|
|
// Takes the refresh token and returns a new LoginResponse with updated tokens
|
|
func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshToken, providerName string) (*LoginResponse, error) {
|
|
provider, err := a.getOAuth2Provider(providerName)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Get session by refresh token from database
|
|
session, err := a.src.get().OAuthUser.GetByRefreshToken(ctx, refreshToken)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Create oauth2.Token from stored data
|
|
oldToken := &oauth2.Token{
|
|
AccessToken: session.AccessToken,
|
|
TokenType: session.TokenType,
|
|
RefreshToken: refreshToken,
|
|
Expiry: session.Expiry,
|
|
}
|
|
|
|
// Use OAuth2 provider to refresh the token
|
|
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)
|
|
}
|
|
|
|
// Generate new session token
|
|
newSessionToken, err := a.OAuth2GenerateState()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to generate new session token: %w", err)
|
|
}
|
|
|
|
// Update session in database with new tokens
|
|
if err := a.src.get().OAuthUser.UpdateRefreshToken(ctx, session.UserID, refreshToken, newSessionToken, newToken.AccessToken, newToken.RefreshToken, newToken.Expiry); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Get user data
|
|
userCtx, err := a.src.get().OAuthUser.GetUser(ctx, session.UserID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
userCtx.SessionID = newSessionToken
|
|
|
|
resp := &LoginResponse{
|
|
Token: newSessionToken,
|
|
RefreshToken: newToken.RefreshToken,
|
|
User: userCtx,
|
|
ExpiresIn: int64(time.Until(newToken.Expiry).Seconds()),
|
|
}
|
|
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
|
|
|
|
// NewGoogleAuthenticator creates a DatabaseAuthenticator configured for Google OAuth2
|
|
func NewGoogleAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
|
|
auth := NewDatabaseAuthenticator(db)
|
|
return auth.WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
|
|
ClientID: clientID,
|
|
ClientSecret: clientSecret,
|
|
RedirectURL: redirectURL,
|
|
Scopes: []string{"openid", "profile", "email"},
|
|
AuthURL: "https://accounts.google.com/o/oauth2/v2/auth",
|
|
TokenURL: "https://oauth2.googleapis.com/token",
|
|
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: "",
|
|
})
|
|
}
|
|
|
|
// NewGitHubAuthenticator creates a DatabaseAuthenticator configured for GitHub OAuth2
|
|
func NewGitHubAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
|
|
auth := NewDatabaseAuthenticator(db)
|
|
return auth.WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
|
|
ClientID: clientID,
|
|
ClientSecret: clientSecret,
|
|
RedirectURL: redirectURL,
|
|
Scopes: []string{"user:email"},
|
|
AuthURL: "https://github.com/login/oauth/authorize",
|
|
TokenURL: "https://github.com/login/oauth/access_token",
|
|
UserInfoURL: "https://api.github.com/user",
|
|
ProviderName: "github",
|
|
})
|
|
}
|
|
|
|
// NewMicrosoftAuthenticator creates a DatabaseAuthenticator configured for Microsoft OAuth2
|
|
func NewMicrosoftAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
|
|
auth := NewDatabaseAuthenticator(db)
|
|
return auth.WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
|
|
ClientID: clientID,
|
|
ClientSecret: clientSecret,
|
|
RedirectURL: redirectURL,
|
|
Scopes: []string{"openid", "profile", "email"},
|
|
AuthURL: "https://login.microsoftonline.com/common/oauth2/v2.0/authorize",
|
|
TokenURL: "https://login.microsoftonline.com/common/oauth2/v2.0/token",
|
|
UserInfoURL: "https://graph.microsoft.com/v1.0/me",
|
|
ProviderName: "microsoft",
|
|
})
|
|
}
|
|
|
|
// NewFacebookAuthenticator creates a DatabaseAuthenticator configured for Facebook OAuth2
|
|
func NewFacebookAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
|
|
auth := NewDatabaseAuthenticator(db)
|
|
return auth.WithOAuth2(OAuth2Config{ //nolint:gosec // G101: false positive: identifier/example, not a credential
|
|
ClientID: clientID,
|
|
ClientSecret: clientSecret,
|
|
RedirectURL: redirectURL,
|
|
Scopes: []string{"email"},
|
|
AuthURL: "https://www.facebook.com/v12.0/dialog/oauth",
|
|
TokenURL: "https://graph.facebook.com/v12.0/oauth/access_token",
|
|
UserInfoURL: "https://graph.facebook.com/me?fields=id,name,email",
|
|
ProviderName: "facebook",
|
|
})
|
|
}
|
|
|
|
// NewMultiProviderAuthenticator creates a DatabaseAuthenticator with all major OAuth2 providers configured
|
|
func NewMultiProviderAuthenticator(db *sql.DB, configs map[string]OAuth2Config) *DatabaseAuthenticator {
|
|
auth := NewDatabaseAuthenticator(db)
|
|
|
|
//nolint:gocritic // OAuth2Config is copied but kept for API simplicity
|
|
for _, cfg := range configs {
|
|
auth.WithOAuth2(cfg)
|
|
}
|
|
|
|
return auth
|
|
}
|