257 lines
7.5 KiB
Go
257 lines
7.5 KiB
Go
package auth
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"sync"
|
|
|
|
"golang.org/x/oauth2"
|
|
"golang.org/x/oauth2/google"
|
|
|
|
"github.com/gochat/gochat/internal/config"
|
|
"github.com/gochat/gochat/internal/model"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
// Reference: P2E §1.3 — OAuth2 provider integration
|
|
// Replaces Chatwoot's OmniAuth-based Google/GitHub/Facebook login flow.
|
|
|
|
// OAuthProviderType defines supported OAuth2 provider types.
|
|
type OAuthProviderType string
|
|
|
|
const (
|
|
OAuthProviderGoogle OAuthProviderType = "google"
|
|
OAuthProviderGitHub OAuthProviderType = "github"
|
|
OAuthProviderCustom OAuthProviderType = "custom"
|
|
)
|
|
|
|
// OAuthUserInfo represents user info extracted from OAuth2 provider.
|
|
type OAuthUserInfo struct {
|
|
Provider OAuthProviderType
|
|
UID string // provider-specific user ID
|
|
Email string
|
|
Name string
|
|
AvatarURL string
|
|
}
|
|
|
|
// OAuthService manages OAuth2 authentication flows.
|
|
type OAuthService struct {
|
|
db *gorm.DB
|
|
configs map[OAuthProviderType]*oauth2.Config
|
|
mu sync.RWMutex
|
|
}
|
|
|
|
// NewOAuthService creates an OAuth service with configured providers.
|
|
func NewOAuthService(db *gorm.DB, cfg *config.Config) *OAuthService {
|
|
svc := &OAuthService{
|
|
db: db,
|
|
configs: make(map[OAuthProviderType]*oauth2.Config),
|
|
}
|
|
svc.configureProviders(cfg)
|
|
return svc
|
|
}
|
|
|
|
// configureProviders sets up OAuth2 configs from application config.
|
|
func (s *OAuthService) configureProviders(cfg *config.Config) {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
|
|
// Google OAuth2 — ref: Chatwoot Google OmniAuth strategy
|
|
if cfg.OAuth.Google.ClientID != "" {
|
|
s.configs[OAuthProviderGoogle] = &oauth2.Config{
|
|
ClientID: cfg.OAuth.Google.ClientID,
|
|
ClientSecret: cfg.OAuth.Google.ClientSecret,
|
|
RedirectURL: cfg.OAuth.Google.RedirectURL,
|
|
Scopes: []string{"openid", "email", "profile"},
|
|
Endpoint: google.Endpoint,
|
|
}
|
|
}
|
|
|
|
// Generic/GitHub OAuth2 — ref: Chatwoot GitHub OmniAuth strategy
|
|
if cfg.OAuth.GitHub.ClientID != "" {
|
|
s.configs[OAuthProviderGitHub] = &oauth2.Config{
|
|
ClientID: cfg.OAuth.GitHub.ClientID,
|
|
ClientSecret: cfg.OAuth.GitHub.ClientSecret,
|
|
RedirectURL: cfg.OAuth.GitHub.RedirectURL,
|
|
Scopes: []string{"user:email", "read:user"},
|
|
Endpoint: oauth2.Endpoint{
|
|
AuthURL: "https://github.com/login/oauth/authorize",
|
|
TokenURL: "https://github.com/login/oauth/access_token",
|
|
},
|
|
}
|
|
}
|
|
}
|
|
|
|
// GetAuthURL generates the OAuth2 authorization URL for a provider.
|
|
// This is the first step of the OAuth2 flow: redirect user to provider login.
|
|
func (s *OAuthService) GetAuthURL(provider OAuthProviderType, state string) (string, error) {
|
|
s.mu.RLock()
|
|
config, ok := s.configs[provider]
|
|
s.mu.RUnlock()
|
|
|
|
if !ok {
|
|
return "", fmt.Errorf("oauth provider %s not configured", provider)
|
|
}
|
|
|
|
return config.AuthCodeURL(state, oauth2.AccessTypeOffline), nil
|
|
}
|
|
|
|
// ExchangeCode exchanges an OAuth2 authorization code for user info.
|
|
// This is the callback step: provider redirects back with a code.
|
|
func (s *OAuthService) ExchangeCode(ctx context.Context, provider OAuthProviderType, code string) (*OAuthUserInfo, error) {
|
|
s.mu.RLock()
|
|
config, ok := s.configs[provider]
|
|
s.mu.RUnlock()
|
|
|
|
if !ok {
|
|
return nil, fmt.Errorf("oauth provider %s not configured", provider)
|
|
}
|
|
|
|
token, err := config.Exchange(ctx, code)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("oauth token exchange failed: %w", err)
|
|
}
|
|
|
|
// Extract user info based on provider
|
|
switch provider {
|
|
case OAuthProviderGoogle:
|
|
return s.extractGoogleUserInfo(ctx, config, token)
|
|
case OAuthProviderGitHub:
|
|
return s.extractGitHubUserInfo(ctx, config, token)
|
|
default:
|
|
return nil, fmt.Errorf("unsupported oauth provider: %s", provider)
|
|
}
|
|
}
|
|
|
|
// FindOrCreateUser finds an existing user by OAuth provider+UID, or creates a new one.
|
|
// Ref: Chatwoot's after_sign_in_path_for OmniAuth callback logic.
|
|
func (s *OAuthService) FindOrCreateUser(info *OAuthUserInfo) (*model.User, error) {
|
|
var user model.User
|
|
|
|
// Try to find existing user by provider + UID
|
|
err := s.db.Where("provider = ? AND uid = ?", string(info.Provider), info.UID).First(&user).Error
|
|
if err == nil {
|
|
// Update last sign-in info
|
|
user.SignInCount++
|
|
s.db.Save(&user)
|
|
return &user, nil
|
|
}
|
|
|
|
// If not found by provider+UID, try by email (link existing account)
|
|
if info.Email != "" {
|
|
err = s.db.Where("email = ?", info.Email).First(&user).Error
|
|
if err == nil {
|
|
// Link OAuth to existing account
|
|
user.Provider = string(info.Provider)
|
|
user.UID = info.UID
|
|
if info.AvatarURL != "" && user.AvatarURL == "" {
|
|
user.AvatarURL = info.AvatarURL
|
|
}
|
|
user.SignInCount++
|
|
s.db.Save(&user)
|
|
return &user, nil
|
|
}
|
|
}
|
|
|
|
// Create new user from OAuth info
|
|
user = model.User{
|
|
Name: info.Name,
|
|
Email: info.Email,
|
|
Provider: string(info.Provider),
|
|
UID: info.UID,
|
|
AvatarURL: info.AvatarURL,
|
|
Password: "oauth_no_password", // OAuth users don't need a password
|
|
}
|
|
|
|
if err := s.db.Create(&user).Error; err != nil {
|
|
return nil, fmt.Errorf("failed to create oauth user: %w", err)
|
|
}
|
|
|
|
return &user, nil
|
|
}
|
|
|
|
// extractGoogleUserInfo extracts user info from Google OAuth2 token.
|
|
func (s *OAuthService) extractGoogleUserInfo(ctx context.Context, config *oauth2.Config, token *oauth2.Token) (*OAuthUserInfo, error) {
|
|
client := config.Client(ctx, token)
|
|
|
|
resp, err := client.Get("https://www.googleapis.com/oauth2/v2/userinfo")
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to fetch google user info: %w", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
|
|
if resp.StatusCode != 200 {
|
|
return nil, fmt.Errorf("google userinfo endpoint returned status %d", resp.StatusCode)
|
|
}
|
|
|
|
var gUser struct {
|
|
ID string `json:"id"`
|
|
Email string `json:"email"`
|
|
Name string `json:"name"`
|
|
Picture string `json:"picture"`
|
|
}
|
|
|
|
if err := json.NewDecoder(resp.Body).Decode(&gUser); err != nil {
|
|
return nil, fmt.Errorf("failed to decode google user info: %w", err)
|
|
}
|
|
|
|
return &OAuthUserInfo{
|
|
Provider: OAuthProviderGoogle,
|
|
UID: gUser.ID,
|
|
Email: gUser.Email,
|
|
Name: gUser.Name,
|
|
AvatarURL: gUser.Picture,
|
|
}, nil
|
|
}
|
|
|
|
// extractGitHubUserInfo extracts user info from GitHub OAuth2 token.
|
|
func (s *OAuthService) extractGitHubUserInfo(ctx context.Context, config *oauth2.Config, token *oauth2.Token) (*OAuthUserInfo, error) {
|
|
client := config.Client(ctx, token)
|
|
|
|
resp, err := client.Get("https://api.github.com/user")
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to fetch github user info: %w", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
|
|
if resp.StatusCode != 200 {
|
|
return nil, fmt.Errorf("github user endpoint returned status %d", resp.StatusCode)
|
|
}
|
|
|
|
var ghUser struct {
|
|
ID int `json:"id"`
|
|
Login string `json:"login"`
|
|
Email string `json:"email"`
|
|
Name string `json:"name"`
|
|
AvatarURL string `json:"avatar_url"`
|
|
}
|
|
|
|
if err := json.NewDecoder(resp.Body).Decode(&ghUser); err != nil {
|
|
return nil, fmt.Errorf("failed to decode github user info: %w", err)
|
|
}
|
|
|
|
return &OAuthUserInfo{
|
|
Provider: OAuthProviderGitHub,
|
|
UID: fmt.Sprintf("%d", ghUser.ID),
|
|
Email: ghUser.Email,
|
|
Name: ghUser.Name,
|
|
AvatarURL: ghUser.AvatarURL,
|
|
}, nil
|
|
}
|
|
|
|
// RegisterGenericProvider allows registering a custom OAuth2 provider at runtime.
|
|
// Ref: Chatwoot's configurable OmniAuth strategies in config/initializers/omniauth.rb
|
|
func (s *OAuthService) RegisterGenericProvider(name OAuthProviderType, cfg *oauth2.Config) {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
s.configs[name] = cfg
|
|
}
|
|
|
|
// IsProviderConfigured checks if a provider has been configured.
|
|
func (s *OAuthService) IsProviderConfigured(provider OAuthProviderType) bool {
|
|
s.mu.RLock()
|
|
defer s.mu.RUnlock()
|
|
_, ok := s.configs[provider]
|
|
return ok
|
|
} |