Reorganize repo: backend/, deploy/, docs/ layout + AGENTS.md
Restructure the monorepo into clear top-level directories: - backend/: Go module root (cmd, internal, pkg, configs, migrations, docs/swagger, scripts, tests, go.mod, Makefile, .air.toml) - deploy/: Docker (Dockerfile, docker-compose*), quickstart, fluentd - docs/: project documentation + reports/ (moved from repo root) - AGENTS.md: new AI coding-agent guide at repo root Update all references to the new layout: - Dockerfile: COPY backend/go.mod, COPY backend/ (context = repo root) - docker-compose files: context ../.., dockerfile deploy/docker/Dockerfile, env_file ../../.env, volume mounts ../../backend:/app - deploy/quickstart/compose.yaml: dockerfile deploy/docker/Dockerfile - CI: working-directory: backend for go commands, file deploy/docker/Dockerfile, coverage path backend/coverage.out, health_check backend/scripts/ - backend/Makefile: docker target uses -f ../deploy/docker/Dockerfile ../ - README: architecture tree, quickstart, config paths updated Move root stray scripts (rename_models.*, run_m11_tests.sh, verify_build.sh, gorm_bool_main.go) to backend/scripts/legacy/. All moves via git mv to preserve history. Build, vet, SQLite tests, and docker compose config verified.
This commit is contained in:
@@ -0,0 +1,825 @@
|
||||
package auth
|
||||
|
||||
// Reference: M13 §4.3 — OIDC (OpenID Connect) enterprise authentication service
|
||||
// Provides OIDC/OAuth2 authorization code flow, ID token validation, userinfo extraction,
|
||||
// and role/claim mapping for enterprise identity providers.
|
||||
// Supports per-account OIDC configuration for multi-tenant identity isolation.
|
||||
//
|
||||
// Enterprise feature: GoChat extends beyond Chatwoot's SAML-only SSO by adding
|
||||
// OIDC support for modern enterprise IdPs (Google Workspace, Auth0, Keycloak,
|
||||
// Azure AD, Okta, and any OIDC-compliant provider).
|
||||
//
|
||||
// Authentication flow (Authorization Code + PKCE):
|
||||
// 1. Client redirects to GoChat OIDC authorize endpoint
|
||||
// 2. GoChat generates PKCE code_verifier + code_challenge
|
||||
// 3. GoChat redirects user to IdP authorization URL with code_challenge
|
||||
// 4. User authenticates at IdP and consents
|
||||
// 5. IdP redirects back to GoChat callback with authorization code
|
||||
// 6. GoChat exchanges code + code_verifier for tokens at IdP token endpoint
|
||||
// 7. GoChat validates ID token (signature, issuer, audience, expiry)
|
||||
// 8. GoChat optionally calls userinfo endpoint for additional claims
|
||||
// 9. GoChat maps OIDC claims to GoChat user fields using attribute mapping
|
||||
// 10. Auto-provision GoChat user if configured
|
||||
// 11. Issue JWT token + create SSO session
|
||||
//
|
||||
// OIDC Discovery:
|
||||
// If AuthorizationURL/TokenURL/JWKSURL are not set per-account, GoChat performs
|
||||
// OIDC discovery by fetching the issuer's .well-known/openid-configuration endpoint
|
||||
// to auto-discover all required endpoints.
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"golang.org/x/oauth2"
|
||||
|
||||
"github.com/gochat/gochat/internal/config"
|
||||
"github.com/gochat/gochat/internal/model"
|
||||
applogger "github.com/gochat/gochat/pkg/logger"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// OIDC errors
|
||||
var (
|
||||
ErrOIDCDisabled = fmt.Errorf("oidc authentication is not enabled")
|
||||
ErrOIDCInvalidConfig = fmt.Errorf("oidc configuration is invalid")
|
||||
ErrOIDCDiscovery = fmt.Errorf("oidc provider discovery failed")
|
||||
ErrOIDCTokenExchange = fmt.Errorf("oidc token exchange failed")
|
||||
ErrOIDCTokenValidation = fmt.Errorf("oidc id token validation failed")
|
||||
ErrOIDCUserInfo = fmt.Errorf("oidc userinfo retrieval failed")
|
||||
)
|
||||
|
||||
// OIDCUserInfo represents user info extracted from OIDC ID token and userinfo endpoint.
|
||||
type OIDCUserInfo struct {
|
||||
Subject string // sub claim — unique user identifier from IdP
|
||||
Email string // email claim
|
||||
EmailVerified bool // email_verified claim
|
||||
Name string // name claim
|
||||
FirstName string // given_name claim
|
||||
LastName string // family_name claim
|
||||
AvatarURL string // picture claim
|
||||
Groups []string // groups claim (custom — varies by IdP)
|
||||
Claims map[string]interface{} // all claims from ID token + userinfo
|
||||
}
|
||||
|
||||
// OIDCDiscoveryDocument represents an OIDC provider's discovery document
|
||||
// (fetched from .well-known/openid-configuration).
|
||||
type OIDCDiscoveryDocument 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"`
|
||||
ScopesSupported []string `json:"scopes_supported"`
|
||||
ResponseTypesSupported []string `json:"response_types_supported"`
|
||||
SubjectTypesSupported []string `json:"subject_types_supported"`
|
||||
IDTokenSigningAlgValuesSupported []string `json:"id_token_signing_alg_values_supported"`
|
||||
EndSessionEndpoint string `json:"end_session_endpoint,omitempty"`
|
||||
}
|
||||
|
||||
// OIDCState stores the state parameter for an OIDC authorization request.
|
||||
// Stored in Redis with TTL to prevent CSRF and replay attacks.
|
||||
type OIDCState struct {
|
||||
AccountID uint `json:"account_id"`
|
||||
CodeVerifier string `json:"code_verifier"` // PKCE code verifier
|
||||
RedirectPath string `json:"redirect_path"` // original client redirect after auth
|
||||
ProviderHint string `json:"provider_hint"` // hint about which IdP (e.g. "keycloak")
|
||||
CreatedAt int64 `json:"created_at"` // timestamp for TTL validation
|
||||
}
|
||||
|
||||
// OIDCService provides OIDC/OAuth2 enterprise authentication.
|
||||
type OIDCService struct {
|
||||
cfg *config.OIDCConfig
|
||||
rdb redis.Cmdable
|
||||
db *gorm.DB
|
||||
httpClient *http.Client
|
||||
mu sync.RWMutex
|
||||
discovery map[uint]*OIDCDiscoveryDocument // cached per-account discovery documents
|
||||
oauthConfigs map[uint]*oauth2.Config // per-account OAuth2 configs
|
||||
}
|
||||
|
||||
// NewOIDCService creates an OIDC service with configuration.
|
||||
func NewOIDCService(cfg *config.OIDCConfig, rdb redis.Cmdable, db *gorm.DB) (*OIDCService, error) {
|
||||
if !cfg.Enabled {
|
||||
applogger.L().Info("OIDC service initialized (disabled)")
|
||||
return &OIDCService{
|
||||
cfg: cfg,
|
||||
rdb: rdb,
|
||||
db: db,
|
||||
httpClient: &http.Client{Timeout: 15 * time.Second},
|
||||
discovery: make(map[uint]*OIDCDiscoveryDocument),
|
||||
oauthConfigs: make(map[uint]*oauth2.Config),
|
||||
}, nil
|
||||
}
|
||||
|
||||
svc := &OIDCService{
|
||||
cfg: cfg,
|
||||
rdb: rdb,
|
||||
db: db,
|
||||
httpClient: &http.Client{Timeout: 15 * time.Second},
|
||||
discovery: make(map[uint]*OIDCDiscoveryDocument),
|
||||
oauthConfigs: make(map[uint]*oauth2.Config),
|
||||
}
|
||||
|
||||
// If global default issuer is set, perform discovery immediately
|
||||
if cfg.DefaultIssuerURL != "" {
|
||||
_, err := svc.discoverProvider(context.Background(), cfg.DefaultIssuerURL)
|
||||
if err != nil {
|
||||
applogger.L().Warnf("OIDC discovery for default issuer %s failed: %v (will retry on first auth)", cfg.DefaultIssuerURL, err)
|
||||
} else {
|
||||
applogger.L().Infof("OIDC discovery successful for default issuer: %s", cfg.DefaultIssuerURL)
|
||||
}
|
||||
}
|
||||
|
||||
applogger.L().Infof("OIDC service initialized (enabled, default_issuer=%s)", cfg.DefaultIssuerURL)
|
||||
return svc, nil
|
||||
}
|
||||
|
||||
// GetAuthorizationURL generates the OIDC authorization URL for an account.
|
||||
// Uses PKCE (code_challenge_method=S256) for enhanced security.
|
||||
// Returns: authorization URL, state parameter (stored in Redis), error.
|
||||
func (s *OIDCService) GetAuthorizationURL(ctx context.Context, accountID uint, redirectPath string) (string, string, error) {
|
||||
if !s.cfg.Enabled {
|
||||
return "", "", ErrOIDCDisabled
|
||||
}
|
||||
|
||||
// Get per-account OIDC settings
|
||||
settings, err := s.getAccountSettings(accountID)
|
||||
if err != nil || settings == nil {
|
||||
return "", "", ErrOIDCDisabled
|
||||
}
|
||||
|
||||
// Get or create OAuth2 config for this account
|
||||
oauthConfig, err := s.getOAuthConfig(ctx, settings)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("failed to build OAuth2 config: %w", err)
|
||||
}
|
||||
|
||||
// Generate PKCE code verifier (43-128 chars, base64url-encoded)
|
||||
codeVerifier, err := generatePKCECodeVerifier()
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("failed to generate PKCE code verifier: %w", err)
|
||||
}
|
||||
|
||||
// Compute code challenge from verifier (S256 = SHA256 + base64url)
|
||||
codeChallenge := computePKCECodeChallenge(codeVerifier)
|
||||
|
||||
// Generate state parameter (random nonce for CSRF protection)
|
||||
state, err := generateOIDCState()
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("failed to generate OIDC state: %w", err)
|
||||
}
|
||||
|
||||
// Store state + code_verifier in Redis with 10-minute TTL
|
||||
stateData := OIDCState{
|
||||
AccountID: accountID,
|
||||
CodeVerifier: codeVerifier,
|
||||
RedirectPath: redirectPath,
|
||||
CreatedAt: time.Now().Unix(),
|
||||
}
|
||||
stateJSON, _ := json.Marshal(stateData)
|
||||
if err := s.rdb.Set(ctx, fmt.Sprintf("oidc:state:%s", state), stateJSON, 10*time.Minute).Err(); err != nil {
|
||||
return "", "", fmt.Errorf("failed to store OIDC state in Redis: %w", err)
|
||||
}
|
||||
|
||||
// Build authorization URL with PKCE parameters
|
||||
authURL := oauthConfig.AuthCodeURL(state,
|
||||
oauth2.SetAuthURLParam("code_challenge", codeChallenge),
|
||||
oauth2.SetAuthURLParam("code_challenge_method", "S256"),
|
||||
)
|
||||
|
||||
applogger.L().Infof("OIDC authorization URL generated (account=%d, state=%s)", accountID, state)
|
||||
return authURL, state, nil
|
||||
}
|
||||
|
||||
// HandleCallback processes the OIDC authorization callback.
|
||||
// Exchanges the authorization code for tokens, validates the ID token,
|
||||
// extracts user info, and optionally provisions the GoChat user.
|
||||
func (s *OIDCService) HandleCallback(ctx context.Context, state, code string) (*OIDCUserInfo, *OIDCState, error) {
|
||||
if !s.cfg.Enabled {
|
||||
return nil, nil, ErrOIDCDisabled
|
||||
}
|
||||
|
||||
// Retrieve state from Redis
|
||||
stateData, err := s.retrieveState(ctx, state)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("invalid or expired OIDC state: %w", err)
|
||||
}
|
||||
|
||||
// Get per-account OIDC settings
|
||||
settings, err := s.getAccountSettings(stateData.AccountID)
|
||||
if err != nil || settings == nil {
|
||||
return nil, nil, ErrOIDCDisabled
|
||||
}
|
||||
|
||||
// Get OAuth2 config for this account
|
||||
oauthConfig, err := s.getOAuthConfig(ctx, settings)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to build OAuth2 config: %w", err)
|
||||
}
|
||||
|
||||
// Exchange code for tokens (with PKCE code_verifier)
|
||||
token, err := oauthConfig.Exchange(ctx, code,
|
||||
oauth2.SetAuthURLParam("code_verifier", stateData.CodeVerifier),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, stateData, fmt.Errorf("%w: %v", ErrOIDCTokenExchange, err)
|
||||
}
|
||||
|
||||
// Extract ID token claims
|
||||
idTokenClaims, err := s.validateAndExtractIDToken(token, settings)
|
||||
if err != nil {
|
||||
return nil, stateData, fmt.Errorf("%w: %v", ErrOIDCTokenValidation, err)
|
||||
}
|
||||
|
||||
// Optionally call userinfo endpoint for additional claims
|
||||
userInfoClaims := idTokenClaims
|
||||
if settings.UserInfoURL != "" && token.AccessToken != "" {
|
||||
extraClaims, err := s.fetchUserInfo(ctx, token.AccessToken, settings.UserInfoURL)
|
||||
if err != nil {
|
||||
applogger.L().Warnf("OIDC userinfo retrieval failed (account=%d): %v", stateData.AccountID, err)
|
||||
} else {
|
||||
// Merge userinfo claims into ID token claims (userinfo takes precedence)
|
||||
for k, v := range extraClaims {
|
||||
userInfoClaims[k] = v
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Map claims to OIDCUserInfo
|
||||
userInfo := s.mapClaimsToUserInfo(userInfoClaims, settings)
|
||||
|
||||
applogger.L().Infof("OIDC authentication successful (account=%d, subject=%s)", stateData.AccountID, userInfo.Subject)
|
||||
return userInfo, stateData, nil
|
||||
}
|
||||
|
||||
// GetLogoutURL generates the OIDC logout URL for an account (if end_session_endpoint is available).
|
||||
func (s *OIDCService) GetLogoutURL(ctx context.Context, accountID uint, idTokenHint, postLogoutRedirectURI string) (string, error) {
|
||||
if !s.cfg.Enabled {
|
||||
return "", ErrOIDCDisabled
|
||||
}
|
||||
|
||||
settings, err := s.getAccountSettings(accountID)
|
||||
if err != nil || settings == nil {
|
||||
return "", ErrOIDCDisabled
|
||||
}
|
||||
|
||||
// Get discovery document for this account
|
||||
doc, err := s.getDiscovery(ctx, settings)
|
||||
if err != nil || doc == nil {
|
||||
return "", fmt.Errorf("OIDC discovery not available for account %d", accountID)
|
||||
}
|
||||
|
||||
if doc.EndSessionEndpoint == "" {
|
||||
return "", fmt.Errorf("OIDC provider does not support end_session_endpoint")
|
||||
}
|
||||
|
||||
// Build logout URL with id_token_hint and post_logout_redirect_uri
|
||||
logoutURL := doc.EndSessionEndpoint
|
||||
params := []string{}
|
||||
if idTokenHint != "" {
|
||||
params = append(params, fmt.Sprintf("id_token_hint=%s", idTokenHint))
|
||||
}
|
||||
if postLogoutRedirectURI != "" {
|
||||
params = append(params, fmt.Sprintf("post_logout_redirect_uri=%s", postLogoutRedirectURI))
|
||||
}
|
||||
if len(params) > 0 {
|
||||
logoutURL += "?" + strings.Join(params, "&")
|
||||
}
|
||||
|
||||
return logoutURL, nil
|
||||
}
|
||||
|
||||
// --- Exported wrappers for internal methods (used by OIDC handler) ---
|
||||
|
||||
// GetAccountSettings loads per-account OIDC configuration from DB.
|
||||
// Exported wrapper for the internal getAccountSettings method.
|
||||
func (s *OIDCService) GetAccountSettings(accountID uint) (*model.AccountOIDCSettings, error) {
|
||||
return s.getAccountSettings(accountID)
|
||||
}
|
||||
|
||||
// GetDiscoveryDocument fetches and caches the OIDC discovery document for an account.
|
||||
// Exported wrapper for the internal getDiscovery method.
|
||||
// Requires the account's OIDC settings to determine the issuer URL.
|
||||
func (s *OIDCService) GetDiscoveryDocument(ctx context.Context, accountID uint) (*OIDCDiscoveryDocument, error) {
|
||||
settings, err := s.getAccountSettings(accountID)
|
||||
if err != nil || settings == nil {
|
||||
return nil, ErrOIDCDisabled
|
||||
}
|
||||
return s.getDiscovery(ctx, settings)
|
||||
}
|
||||
|
||||
// --- Internal methods ---
|
||||
|
||||
// getAccountSettings loads per-account OIDC configuration from DB.
|
||||
// Falls back to global defaults for fields not overridden per-account.
|
||||
func (s *OIDCService) getAccountSettings(accountID uint) (*model.AccountOIDCSettings, error) {
|
||||
if s.db == nil {
|
||||
return nil, fmt.Errorf("database not available")
|
||||
}
|
||||
|
||||
var settings model.AccountOIDCSettings
|
||||
err := s.db.Where("account_id = ? AND active = ?", accountID, true).First(&settings).Error
|
||||
if err != nil {
|
||||
if err == gorm.ErrRecordNotFound {
|
||||
// No per-account settings — check if global defaults exist
|
||||
if s.cfg.DefaultIssuerURL == "" && s.cfg.DefaultClientID == "" {
|
||||
return nil, nil // OIDC not configured for this account
|
||||
}
|
||||
// Use global defaults as fallback
|
||||
defaultScopes, _ := json.Marshal(s.cfg.DefaultScopes)
|
||||
if len(defaultScopes) == 0 {
|
||||
defaultScopes = []byte(`["openid","profile","email"]`)
|
||||
}
|
||||
settings = model.AccountOIDCSettings{
|
||||
AccountID: accountID,
|
||||
ClientID: s.cfg.DefaultClientID,
|
||||
ClientSecret: s.cfg.DefaultClientSecret,
|
||||
RedirectURL: s.cfg.DefaultRedirectURL,
|
||||
IssuerURL: s.cfg.DefaultIssuerURL,
|
||||
AuthorizationURL: s.cfg.DefaultAuthorizationURL,
|
||||
TokenURL: s.cfg.DefaultTokenURL,
|
||||
UserInfoURL: s.cfg.DefaultUserInfoURL,
|
||||
JWKSURL: s.cfg.DefaultJWKSURL,
|
||||
Scopes: defaultScopes,
|
||||
AutoProvision: true,
|
||||
Active: true,
|
||||
}
|
||||
return &settings, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Override empty per-account fields with global defaults
|
||||
if settings.IssuerURL == "" && s.cfg.DefaultIssuerURL != "" {
|
||||
settings.IssuerURL = s.cfg.DefaultIssuerURL
|
||||
}
|
||||
if settings.ClientID == "" && s.cfg.DefaultClientID != "" {
|
||||
settings.ClientID = s.cfg.DefaultClientID
|
||||
}
|
||||
if settings.ClientSecret == "" && s.cfg.DefaultClientSecret != "" {
|
||||
settings.ClientSecret = s.cfg.DefaultClientSecret
|
||||
}
|
||||
if settings.RedirectURL == "" && s.cfg.DefaultRedirectURL != "" {
|
||||
settings.RedirectURL = s.cfg.DefaultRedirectURL
|
||||
}
|
||||
if settings.AuthorizationURL == "" && s.cfg.DefaultAuthorizationURL != "" {
|
||||
settings.AuthorizationURL = s.cfg.DefaultAuthorizationURL
|
||||
}
|
||||
if settings.TokenURL == "" && s.cfg.DefaultTokenURL != "" {
|
||||
settings.TokenURL = s.cfg.DefaultTokenURL
|
||||
}
|
||||
if settings.UserInfoURL == "" && s.cfg.DefaultUserInfoURL != "" {
|
||||
settings.UserInfoURL = s.cfg.DefaultUserInfoURL
|
||||
}
|
||||
if settings.JWKSURL == "" && s.cfg.DefaultJWKSURL != "" {
|
||||
settings.JWKSURL = s.cfg.DefaultJWKSURL
|
||||
}
|
||||
|
||||
return &settings, nil
|
||||
}
|
||||
|
||||
// getOAuthConfig builds an OAuth2.Config for the given account settings.
|
||||
// If endpoints are not set, discovers them from the issuer URL.
|
||||
func (s *OIDCService) getOAuthConfig(ctx context.Context, settings *model.AccountOIDCSettings) (*oauth2.Config, error) {
|
||||
accountID := settings.AccountID
|
||||
|
||||
// Check cache
|
||||
s.mu.RLock()
|
||||
if cfg, ok := s.oauthConfigs[accountID]; ok {
|
||||
s.mu.RUnlock()
|
||||
return cfg, nil
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
|
||||
// Resolve endpoint URLs — either from settings or discovery
|
||||
authURL := settings.AuthorizationURL
|
||||
tokenURL := settings.TokenURL
|
||||
|
||||
if authURL == "" || tokenURL == "" {
|
||||
// Perform OIDC discovery
|
||||
doc, err := s.getDiscovery(ctx, settings)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrOIDCDiscovery, err)
|
||||
}
|
||||
if authURL == "" {
|
||||
authURL = doc.AuthorizationEndpoint
|
||||
}
|
||||
if tokenURL == "" {
|
||||
tokenURL = doc.TokenEndpoint
|
||||
}
|
||||
}
|
||||
|
||||
// Parse scopes from JSON
|
||||
scopes := []string{"openid"} // openid is always required
|
||||
if settings.Scopes != nil {
|
||||
var parsedScopes []string
|
||||
if err := json.Unmarshal(settings.Scopes, &parsedScopes); err == nil {
|
||||
// Deduplicate: ensure openid is present but not duplicated
|
||||
for _, sc := range parsedScopes {
|
||||
if sc != "openid" && sc != "" {
|
||||
scopes = append(scopes, sc)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Build OAuth2 config
|
||||
oauthConfig := &oauth2.Config{
|
||||
ClientID: settings.ClientID,
|
||||
ClientSecret: settings.ClientSecret,
|
||||
RedirectURL: settings.RedirectURL,
|
||||
Endpoint: oauth2.Endpoint{
|
||||
AuthURL: authURL,
|
||||
TokenURL: tokenURL,
|
||||
},
|
||||
Scopes: scopes,
|
||||
}
|
||||
|
||||
// Cache the config
|
||||
s.mu.Lock()
|
||||
s.oauthConfigs[accountID] = oauthConfig
|
||||
s.mu.Unlock()
|
||||
|
||||
return oauthConfig, nil
|
||||
}
|
||||
|
||||
// getDiscovery fetches and caches the OIDC discovery document for the given account.
|
||||
func (s *OIDCService) getDiscovery(ctx context.Context, settings *model.AccountOIDCSettings) (*OIDCDiscoveryDocument, error) {
|
||||
accountID := settings.AccountID
|
||||
|
||||
// Check cache
|
||||
s.mu.RLock()
|
||||
if doc, ok := s.discovery[accountID]; ok {
|
||||
s.mu.RUnlock()
|
||||
return doc, nil
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
|
||||
if settings.IssuerURL == "" {
|
||||
return nil, fmt.Errorf("issuer URL not configured for account %d", accountID)
|
||||
}
|
||||
|
||||
doc, err := s.discoverProvider(ctx, settings.IssuerURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Cache the discovery document
|
||||
s.mu.Lock()
|
||||
s.discovery[accountID] = doc
|
||||
s.mu.Unlock()
|
||||
|
||||
return doc, nil
|
||||
}
|
||||
|
||||
// discoverProvider fetches the OIDC discovery document from an issuer URL.
|
||||
func (s *OIDCService) discoverProvider(ctx context.Context, issuerURL string) (*OIDCDiscoveryDocument, error) {
|
||||
// Normalize issuer URL — remove trailing slash
|
||||
issuerURL = strings.TrimSuffix(issuerURL, "/")
|
||||
|
||||
discoveryURL := issuerURL + "/.well-known/openid-configuration"
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, discoveryURL, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create discovery request: %w", err)
|
||||
}
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("discovery request to %s failed: %w", discoveryURL, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("discovery endpoint returned status %d: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var doc OIDCDiscoveryDocument
|
||||
if err := json.NewDecoder(resp.Body).Decode(&doc); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse discovery document: %w", err)
|
||||
}
|
||||
|
||||
// Validate: issuer in document must match requested issuer
|
||||
if doc.Issuer != issuerURL {
|
||||
return nil, fmt.Errorf("discovery issuer mismatch: expected %s, got %s", issuerURL, doc.Issuer)
|
||||
}
|
||||
|
||||
return &doc, nil
|
||||
}
|
||||
|
||||
// retrieveState retrieves and validates the OIDC state from Redis.
|
||||
// Deletes the state after retrieval (one-time use to prevent replay attacks).
|
||||
func (s *OIDCService) retrieveState(ctx context.Context, state string) (*OIDCState, error) {
|
||||
key := fmt.Sprintf("oidc:state:%s", state)
|
||||
val, err := s.rdb.Get(ctx, key).Bytes()
|
||||
if err != nil {
|
||||
if err == redis.Nil {
|
||||
return nil, fmt.Errorf("state not found or expired")
|
||||
}
|
||||
return nil, fmt.Errorf("failed to retrieve state: %w", err)
|
||||
}
|
||||
|
||||
// Delete state after retrieval (one-time use)
|
||||
s.rdb.Del(ctx, key)
|
||||
|
||||
var stateData OIDCState
|
||||
if err := json.Unmarshal(val, &stateData); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse state data: %w", err)
|
||||
}
|
||||
|
||||
// Validate state age (max 10 minutes)
|
||||
if time.Since(time.Unix(stateData.CreatedAt, 0)) > 10*time.Minute {
|
||||
return nil, fmt.Errorf("state expired")
|
||||
}
|
||||
|
||||
return &stateData, nil
|
||||
}
|
||||
|
||||
// validateAndExtractIDToken validates the OIDC ID token and extracts claims.
|
||||
// In a full implementation, this would use coreos/go-oidc/v3 for proper
|
||||
// JWT signature verification via JWKS. This implementation extracts claims
|
||||
// from the raw JWT payload and performs basic validation.
|
||||
func (s *OIDCService) validateAndExtractIDToken(token *oauth2.Token, settings *model.AccountOIDCSettings) (map[string]interface{}, error) {
|
||||
rawIDToken := token.Extra("id_token")
|
||||
if rawIDToken == nil {
|
||||
// Some providers don't return ID token — use userinfo instead
|
||||
return make(map[string]interface{}), nil
|
||||
}
|
||||
|
||||
idTokenStr, ok := rawIDToken.(string)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("id_token is not a string")
|
||||
}
|
||||
|
||||
// Parse JWT payload (middle segment between dots)
|
||||
// JWT format: header.payload.signature
|
||||
parts := strings.Split(idTokenStr, ".")
|
||||
if len(parts) != 3 {
|
||||
return nil, fmt.Errorf("invalid id_token format")
|
||||
}
|
||||
|
||||
// Decode payload (base64url)
|
||||
payload, err := base64urlDecode(parts[1])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to decode id_token payload: %w", err)
|
||||
}
|
||||
|
||||
var claims map[string]interface{}
|
||||
if err := json.Unmarshal(payload, &claims); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse id_token claims: %w", err)
|
||||
}
|
||||
|
||||
// Basic validation: check issuer and audience
|
||||
// Note: Full signature verification via JWKS requires coreos/go-oidc/v3.
|
||||
// This basic validation checks issuer match and audience but does NOT
|
||||
// verify the JWT signature. Add coreos/go-oidc/v3 dependency for production.
|
||||
|
||||
if issuer, ok := claims["iss"].(string); ok {
|
||||
if issuer != settings.IssuerURL {
|
||||
return nil, fmt.Errorf("id_token issuer mismatch: expected %s, got %s", settings.IssuerURL, issuer)
|
||||
}
|
||||
}
|
||||
|
||||
// Check audience (aud) — must contain our client ID
|
||||
if aud, ok := claims["aud"]; ok {
|
||||
switch v := aud.(type) {
|
||||
case string:
|
||||
if v != settings.ClientID {
|
||||
return nil, fmt.Errorf("id_token audience mismatch: expected %s, got %s", settings.ClientID, v)
|
||||
}
|
||||
case []interface{}:
|
||||
found := false
|
||||
for _, a := range v {
|
||||
if aStr, ok := a.(string); ok && aStr == settings.ClientID {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return nil, fmt.Errorf("id_token audience does not contain client ID %s", settings.ClientID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Check expiry (exp)
|
||||
if exp, ok := claims["exp"].(float64); ok {
|
||||
if time.Unix(int64(exp), 0).Before(time.Now()) {
|
||||
return nil, fmt.Errorf("id_token expired")
|
||||
}
|
||||
}
|
||||
|
||||
// Check issued-at (iat) — should not be too far in the future
|
||||
if iat, ok := claims["iat"].(float64); ok {
|
||||
if time.Unix(int64(iat), 0).After(time.Now().Add(5 * time.Minute)) {
|
||||
return nil, fmt.Errorf("id_token issued-at is in the future")
|
||||
}
|
||||
}
|
||||
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
// fetchUserInfo calls the OIDC userinfo endpoint with the access token.
|
||||
func (s *OIDCService) fetchUserInfo(ctx context.Context, accessToken, userInfoURL string) (map[string]interface{}, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, userInfoURL, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create userinfo request: %w", err)
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("userinfo request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("%w: status %d, body: %s", ErrOIDCUserInfo, resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var claims map[string]interface{}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&claims); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse userinfo response: %w", err)
|
||||
}
|
||||
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
// mapClaimsToUserInfo maps OIDC claims to OIDCUserInfo using the account's attribute mapping.
|
||||
func (s *OIDCService) mapClaimsToUserInfo(claims map[string]interface{}, settings *model.AccountOIDCSettings) *OIDCUserInfo {
|
||||
// Default attribute mapping (OIDC standard claims)
|
||||
defaultMapping := map[string]string{
|
||||
"subject": "sub",
|
||||
"email": "email",
|
||||
"email_verified": "email_verified",
|
||||
"name": "name",
|
||||
"firstName": "given_name",
|
||||
"lastName": "family_name",
|
||||
"avatar": "picture",
|
||||
"groups": "groups",
|
||||
}
|
||||
|
||||
// Merge with custom attribute mapping if provided
|
||||
if settings.AttributeMapping != "" {
|
||||
var customMapping map[string]string
|
||||
if err := json.Unmarshal([]byte(settings.AttributeMapping), &customMapping); err == nil {
|
||||
for k, v := range customMapping {
|
||||
defaultMapping[k] = v
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
info := &OIDCUserInfo{
|
||||
Claims: claims,
|
||||
}
|
||||
|
||||
// Extract fields using mapping
|
||||
info.Subject = getClaimString(claims, defaultMapping["subject"])
|
||||
info.Email = getClaimString(claims, defaultMapping["email"])
|
||||
info.EmailVerified = getClaimBool(claims, defaultMapping["email_verified"])
|
||||
info.Name = getClaimString(claims, defaultMapping["name"])
|
||||
info.FirstName = getClaimString(claims, defaultMapping["firstName"])
|
||||
info.LastName = getClaimString(claims, defaultMapping["lastName"])
|
||||
info.AvatarURL = getClaimString(claims, defaultMapping["avatar"])
|
||||
info.Groups = getClaimStringSlice(claims, defaultMapping["groups"])
|
||||
|
||||
return info
|
||||
}
|
||||
|
||||
// MapOIDCGroupsToRoles maps OIDC groups/roles to GoChat roles.
|
||||
func (s *OIDCService) MapOIDCGroupsToRoles(settings *model.AccountOIDCSettings, groups []string) string {
|
||||
if settings.RoleMappings == nil || len(settings.RoleMappings) == 0 {
|
||||
return "agent" // default role
|
||||
}
|
||||
|
||||
var mappings map[string]string
|
||||
if err := json.Unmarshal(settings.RoleMappings, &mappings); err != nil {
|
||||
applogger.L().Warnf("Invalid OIDC role mappings JSON: %v", err)
|
||||
return "agent"
|
||||
}
|
||||
|
||||
rolePriority := map[string]int{
|
||||
"administrator": 4,
|
||||
"admin": 3,
|
||||
"supervisor": 2,
|
||||
"agent": 1,
|
||||
}
|
||||
|
||||
bestRole := "agent"
|
||||
bestPriority := 1
|
||||
|
||||
for _, group := range groups {
|
||||
if mappedRole, ok := mappings[group]; ok {
|
||||
if p, ok := rolePriority[mappedRole]; ok && p > bestPriority {
|
||||
bestRole = mappedRole
|
||||
bestPriority = p
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return bestRole
|
||||
}
|
||||
|
||||
// --- PKCE helpers ---
|
||||
|
||||
// generatePKCECodeVerifier generates a random PKCE code verifier.
|
||||
// Per RFC 7636: 43-128 characters, unreserved chars (A-Z, a-z, 0-9, -, ., _, ~).
|
||||
func generatePKCECodeVerifier() (string, error) {
|
||||
b := make([]byte, 32) // 32 random bytes → 43 base64url chars
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64urlEncode(b), nil
|
||||
}
|
||||
|
||||
// computePKCECodeChallenge computes S256 code challenge from verifier.
|
||||
// challenge = BASE64URL(SHA256(verifier))
|
||||
func computePKCECodeChallenge(verifier string) string {
|
||||
h := sha256.Sum256([]byte(verifier))
|
||||
return base64urlEncode(h[:])
|
||||
}
|
||||
|
||||
// generateOIDCState generates a random state parameter for CSRF protection.
|
||||
func generateOIDCState() (string, error) {
|
||||
b := make([]byte, 16)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64urlEncode(b), nil
|
||||
}
|
||||
|
||||
// --- Encoding helpers ---
|
||||
|
||||
// base64urlEncode encodes bytes using base64url (no padding) per RFC 7636.
|
||||
func base64urlEncode(b []byte) string {
|
||||
return base64.RawURLEncoding.EncodeToString(b)
|
||||
}
|
||||
|
||||
// base64urlDecode decodes a base64url string (no padding).
|
||||
func base64urlDecode(s string) ([]byte, error) {
|
||||
// Add padding if needed
|
||||
switch len(s) % 4 {
|
||||
case 2:
|
||||
s += "=="
|
||||
case 3:
|
||||
s += "="
|
||||
}
|
||||
return base64.URLEncoding.DecodeString(s)
|
||||
}
|
||||
|
||||
// --- Claim extraction helpers ---
|
||||
|
||||
func getClaimString(claims map[string]interface{}, key string) string {
|
||||
if v, ok := claims[key]; ok {
|
||||
switch val := v.(type) {
|
||||
case string:
|
||||
return val
|
||||
case float64:
|
||||
return fmt.Sprintf("%v", val)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func getClaimBool(claims map[string]interface{}, key string) bool {
|
||||
if v, ok := claims[key]; ok {
|
||||
switch val := v.(type) {
|
||||
case bool:
|
||||
return val
|
||||
case string:
|
||||
return val == "true"
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func getClaimStringSlice(claims map[string]interface{}, key string) []string {
|
||||
if v, ok := claims[key]; ok {
|
||||
switch val := v.(type) {
|
||||
case []interface{}:
|
||||
result := []string{}
|
||||
for _, item := range val {
|
||||
if s, ok := item.(string); ok {
|
||||
result = append(result, s)
|
||||
}
|
||||
}
|
||||
return result
|
||||
case string:
|
||||
// Some IdPs return groups as comma-separated string
|
||||
return strings.Split(val, ",")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
Reference in New Issue
Block a user