mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-01 22:46:38 +08:00
510 lines
16 KiB
Go
510 lines
16 KiB
Go
// Copyright 2026 Arctel.net
|
|
// SPDX-License-Identifier: Apache-2.0
|
|
|
|
package auth
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"net/url"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/Rain-kl/Wavelet/internal/apps/openflare/compat"
|
|
"github.com/Rain-kl/Wavelet/internal/db"
|
|
"github.com/Rain-kl/Wavelet/internal/listener"
|
|
"github.com/Rain-kl/Wavelet/internal/model"
|
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
|
"github.com/gin-contrib/sessions"
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/google/uuid"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
const pendingExternalAccountSessionKey = "pending_external_account"
|
|
|
|
// OAuthCallbackResult is the legacy OAuth callback payload.
|
|
type OAuthCallbackResult struct {
|
|
Status string `json:"status"`
|
|
User *LegacyUser `json:"user,omitempty"`
|
|
}
|
|
|
|
// PendingExternalAccount stores OAuth bind-pending state in session.
|
|
type PendingExternalAccount struct {
|
|
AuthSourceID uint64 `json:"auth_source_id"`
|
|
ExternalID string `json:"external_id"`
|
|
ExternalUsername string `json:"external_username"`
|
|
DisplayName string `json:"display_name"`
|
|
Email string `json:"email"`
|
|
}
|
|
|
|
// OAuthAuthorize builds an authorize URL for a legacy auth source route param.
|
|
func OAuthAuthorize(ctx context.Context, c *gin.Context, sourceKey string) (string, error) {
|
|
source, err := resolveAuthSourceByRoute(ctx, sourceKey)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
if !source.IsActive {
|
|
return "", errors.New(errAuthSourceDisabled)
|
|
}
|
|
if err := source.Validate(); err != nil {
|
|
return "", err
|
|
}
|
|
|
|
state := uuid.NewString()
|
|
session := sessions.Default(c)
|
|
session.Set(oauthStateSessionKey(source.ID), state)
|
|
if err := session.Save(); err != nil {
|
|
return "", errors.New(errSaveSessionFailed)
|
|
}
|
|
|
|
redirectURL := legacyOAuthCallbackURL(c, source)
|
|
return buildLegacyAuthorizeURL(ctx, source, redirectURL, state)
|
|
}
|
|
|
|
// OAuthCallback handles GET /oauth/:source/callback for the legacy frontend.
|
|
func OAuthCallback(ctx context.Context, c *gin.Context, sourceKey string) (OAuthCallbackResult, error) {
|
|
source, err := resolveAuthSourceByRoute(ctx, sourceKey)
|
|
if err != nil {
|
|
return OAuthCallbackResult{}, err
|
|
}
|
|
if !source.IsActive {
|
|
return OAuthCallbackResult{}, errors.New(errAuthSourceDisabled)
|
|
}
|
|
|
|
session := sessions.Default(c)
|
|
expectedState, _ := session.Get(oauthStateSessionKey(source.ID)).(string)
|
|
state := c.Query("state")
|
|
if expectedState == "" || state == "" || state != expectedState {
|
|
return OAuthCallbackResult{}, errors.New("授权状态无效,请重新登录")
|
|
}
|
|
session.Delete(oauthStateSessionKey(source.ID))
|
|
if err := session.Save(); err != nil {
|
|
return OAuthCallbackResult{}, errors.New(errSaveSessionFailed)
|
|
}
|
|
if oauthError := c.Query("error"); oauthError != "" {
|
|
description := c.Query("error_description")
|
|
if description == "" {
|
|
description = oauthError
|
|
}
|
|
return OAuthCallbackResult{}, errors.New(description)
|
|
}
|
|
|
|
redirectURL := legacyOAuthCallbackURL(c, source)
|
|
userInfo, err := exchangeLegacyOAuthProfile(ctx, source, c.Query("code"), state, redirectURL)
|
|
if err != nil {
|
|
return OAuthCallbackResult{}, err
|
|
}
|
|
|
|
var currentUserID *uint64
|
|
if current := currentUserFromLegacyToken(ctx, c); current != nil {
|
|
currentUserID = ¤t.ID
|
|
}
|
|
|
|
result, pending, err := completeLegacyOAuthLogin(ctx, source, userInfo, currentUserID)
|
|
if err != nil {
|
|
return OAuthCallbackResult{}, err
|
|
}
|
|
if pending != nil {
|
|
raw, marshalErr := json.Marshal(pending)
|
|
if marshalErr != nil {
|
|
return OAuthCallbackResult{}, marshalErr
|
|
}
|
|
session.Set(pendingExternalAccountSessionKey, string(raw))
|
|
if err := session.Save(); err != nil {
|
|
return OAuthCallbackResult{}, errors.New(errSaveSessionFailed)
|
|
}
|
|
return result, nil
|
|
}
|
|
if result.User != nil {
|
|
var dbUser model.User
|
|
if err := db.DB(ctx).Where("id = ?", result.User.ID).First(&dbUser).Error; err != nil {
|
|
return OAuthCallbackResult{}, err
|
|
}
|
|
if err := setLoginSession(ctx, c, &dbUser); err != nil {
|
|
return OAuthCallbackResult{}, errors.New(errSaveSessionFailed)
|
|
}
|
|
token, tokenErr := issueLegacyAccessToken(ctx, &dbUser)
|
|
if tokenErr != nil {
|
|
return OAuthCallbackResult{}, tokenErr
|
|
}
|
|
legacy := ToLegacyUser(&dbUser, token)
|
|
result.User = &legacy
|
|
listener.EmitAdminLoggedIn(ctx, &dbUser, c.ClientIP())
|
|
}
|
|
return result, nil
|
|
}
|
|
|
|
// LinkExistingOAuthAccount binds a pending external account to an existing user.
|
|
func LinkExistingOAuthAccount(ctx context.Context, c *gin.Context, input LinkExistingInput) (OAuthCallbackResult, error) {
|
|
session := sessions.Default(c)
|
|
raw, _ := session.Get(pendingExternalAccountSessionKey).(string)
|
|
if raw == "" {
|
|
return OAuthCallbackResult{}, errors.New(errPendingOAuthExpired)
|
|
}
|
|
var pending PendingExternalAccount
|
|
if err := json.Unmarshal([]byte(raw), &pending); err != nil {
|
|
return OAuthCallbackResult{}, errors.New(errPendingOAuthInvalid)
|
|
}
|
|
|
|
user, err := linkPendingExternalAccount(ctx, &pending, input)
|
|
if err != nil {
|
|
return OAuthCallbackResult{}, err
|
|
}
|
|
|
|
session.Delete(pendingExternalAccountSessionKey)
|
|
if err := session.Save(); err != nil {
|
|
return OAuthCallbackResult{}, errors.New(errSaveSessionFailed)
|
|
}
|
|
if err := setLoginSession(ctx, c, user); err != nil {
|
|
return OAuthCallbackResult{}, errors.New(errSaveSessionFailed)
|
|
}
|
|
token, err := issueLegacyAccessToken(ctx, user)
|
|
if err != nil {
|
|
return OAuthCallbackResult{}, err
|
|
}
|
|
legacy := ToLegacyUser(user, token)
|
|
return OAuthCallbackResult{Status: "linked", User: &legacy}, nil
|
|
}
|
|
|
|
func resolveAuthSourceByRoute(ctx context.Context, raw string) (*model.AuthSource, error) {
|
|
raw = strings.TrimSpace(raw)
|
|
if raw == "" {
|
|
return nil, errors.New("认证源不能为空")
|
|
}
|
|
if parsed, err := parseUint64(raw); err == nil && parsed > 0 {
|
|
return model.GetAuthSourceByID(ctx, parsed)
|
|
}
|
|
return model.GetAuthSourceByName(ctx, raw)
|
|
}
|
|
|
|
func oauthStateSessionKey(sourceID uint64) string {
|
|
return fmt.Sprintf("oauth_state_%d", sourceID)
|
|
}
|
|
|
|
func legacyOAuthCallbackURL(c *gin.Context, source *model.AuthSource) string {
|
|
ctx := c.Request.Context()
|
|
base := ""
|
|
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress); err == nil {
|
|
base = strings.TrimRight(sc.Value, "/")
|
|
}
|
|
if base == "" {
|
|
scheme := "http"
|
|
if c.Request.TLS != nil || c.GetHeader("X-Forwarded-Proto") == "https" {
|
|
scheme = "https"
|
|
}
|
|
host := c.Request.Host
|
|
if forwardedHost := c.GetHeader("X-Forwarded-Host"); forwardedHost != "" {
|
|
host = forwardedHost
|
|
}
|
|
base = scheme + "://" + host
|
|
}
|
|
sourceName := source.Name
|
|
if sourceName == "" {
|
|
sourceName = fmt.Sprintf("%d", source.ID)
|
|
}
|
|
callback, _ := url.JoinPath(base, "oauth", sourceName)
|
|
return callback
|
|
}
|
|
|
|
func buildLegacyAuthorizeURL(ctx context.Context, source *model.AuthSource, redirectURL, state string) (string, error) {
|
|
payloadValue, err := encodeLegacyOAuthState(source.Name, state)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
stateKey := fmt.Sprintf("of_oauth_state:%s", state)
|
|
if err := db.Redis.Set(ctx, db.PrefixedKey(stateKey), payloadValue, 10*time.Minute).Err(); err != nil {
|
|
return "", err
|
|
}
|
|
return oauthBuildAuthorizeURL(ctx, source, redirectURL, state)
|
|
}
|
|
|
|
func encodeLegacyOAuthState(sourceName, state string) (string, error) {
|
|
payload := map[string]string{
|
|
"source_name": sourceName,
|
|
"state": state,
|
|
}
|
|
raw, err := json.Marshal(payload)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
return string(raw), nil
|
|
}
|
|
|
|
func legacyFrontendLoginRedirectURL(ctx context.Context, source *model.AuthSource) (string, error) {
|
|
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
|
|
if err != nil || strings.TrimSpace(sc.Value) == "" {
|
|
return "", errors.New("server_address 未配置")
|
|
}
|
|
base := strings.TrimRight(sc.Value, "/")
|
|
name := source.Name
|
|
if name == "" {
|
|
name = fmt.Sprintf("%d", source.ID)
|
|
}
|
|
return base + "/oauth/" + url.PathEscape(name), nil
|
|
}
|
|
|
|
func exchangeLegacyOAuthProfile(ctx context.Context, source *model.AuthSource, code, state, redirectURL string) (*model.OAuthUserInfo, error) {
|
|
if strings.TrimSpace(code) == "" {
|
|
return nil, errors.New("授权 code 不能为空")
|
|
}
|
|
// Validate state from Redis cache written during authorize.
|
|
stateKey := fmt.Sprintf("of_oauth_state:%s", state)
|
|
payloadRaw, err := db.Redis.Get(ctx, db.PrefixedKey(stateKey)).Result()
|
|
if err != nil {
|
|
return nil, errors.New("授权状态无效,请重新登录")
|
|
}
|
|
_ = db.Redis.Del(ctx, db.PrefixedKey(stateKey)).Err()
|
|
|
|
var payload map[string]string
|
|
if err := json.Unmarshal([]byte(payloadRaw), &payload); err != nil {
|
|
return nil, err
|
|
}
|
|
if payload["source_name"] != source.Name {
|
|
return nil, errors.New("授权状态无效,请重新登录")
|
|
}
|
|
|
|
userInfo, err := buildOAuthUserInfo(ctx, source, code, state, redirectURL)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if err := normalizeOAuthUserInfo(userInfo); err != nil {
|
|
return nil, err
|
|
}
|
|
if userInfo.Sub == "" {
|
|
userInfo.Sub = userInfo.Username
|
|
}
|
|
return userInfo, nil
|
|
}
|
|
|
|
func completeLegacyOAuthLogin(ctx context.Context, source *model.AuthSource, profile *model.OAuthUserInfo, currentUserID *uint64) (OAuthCallbackResult, *PendingExternalAccount, error) {
|
|
if source == nil || profile == nil || strings.TrimSpace(profile.Sub) == "" {
|
|
return OAuthCallbackResult{}, nil, errors.New("第三方账号资料不完整")
|
|
}
|
|
|
|
account, err := model.FindExternalAccount(ctx, source.ID, profile.Sub)
|
|
if err == nil {
|
|
var user model.User
|
|
if err := db.DB(ctx).Where("id = ?", account.UserID).First(&user).Error; err != nil {
|
|
return OAuthCallbackResult{}, nil, err
|
|
}
|
|
if !user.IsActive {
|
|
return OAuthCallbackResult{}, nil, errors.New(errBannedAccount)
|
|
}
|
|
legacy := ToLegacyUser(&user, "")
|
|
return OAuthCallbackResult{Status: "logged_in", User: &legacy}, nil, nil
|
|
}
|
|
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return OAuthCallbackResult{}, nil, err
|
|
}
|
|
|
|
if currentUserID != nil && *currentUserID > 0 {
|
|
var user model.User
|
|
if err := db.DB(ctx).Where("id = ?", *currentUserID).First(&user).Error; err != nil {
|
|
return OAuthCallbackResult{}, nil, err
|
|
}
|
|
if !user.IsActive {
|
|
return OAuthCallbackResult{}, nil, errors.New(errBannedAccount)
|
|
}
|
|
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
|
|
AuthSourceID: source.ID,
|
|
UserID: user.ID,
|
|
ExternalID: profile.Sub,
|
|
ExternalUsername: profile.Username,
|
|
Email: profile.Email,
|
|
}); err != nil {
|
|
return OAuthCallbackResult{}, nil, err
|
|
}
|
|
legacy := ToLegacyUser(&user, "")
|
|
return OAuthCallbackResult{Status: "linked", User: &legacy}, nil, nil
|
|
}
|
|
|
|
registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
|
if regErr != nil {
|
|
registrationEnabled = true
|
|
}
|
|
if !registrationEnabled {
|
|
pending := &PendingExternalAccount{
|
|
AuthSourceID: source.ID,
|
|
ExternalID: profile.Sub,
|
|
ExternalUsername: profile.Username,
|
|
DisplayName: profile.Name,
|
|
Email: profile.Email,
|
|
}
|
|
return OAuthCallbackResult{Status: "link_required"}, pending, nil
|
|
}
|
|
|
|
user, err := createUserFromOAuthProfile(ctx, source, profile)
|
|
if err != nil {
|
|
return OAuthCallbackResult{}, nil, err
|
|
}
|
|
legacy := ToLegacyUser(&user, "")
|
|
return OAuthCallbackResult{Status: "logged_in", User: &legacy}, nil, nil
|
|
}
|
|
|
|
func createUserFromOAuthProfile(ctx context.Context, source *model.AuthSource, profile *model.OAuthUserInfo) (model.User, error) {
|
|
username, err := uniqueLegacyUsername(ctx, profile.Username)
|
|
if err != nil {
|
|
return model.User{}, err
|
|
}
|
|
profile.Username = username
|
|
|
|
var user model.User
|
|
if err := user.CreateUser(ctx, db.DB(ctx), profile); err != nil {
|
|
return model.User{}, err
|
|
}
|
|
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
|
|
AuthSourceID: source.ID,
|
|
UserID: user.ID,
|
|
ExternalID: profile.Sub,
|
|
ExternalUsername: profile.Username,
|
|
Email: profile.Email,
|
|
}); err != nil {
|
|
return model.User{}, err
|
|
}
|
|
logger.InfoF(ctx, "[LoginAudit] successful legacy OAuth registration via source: %s, user: %s, ID: %d", source.Name, user.Username, user.ID)
|
|
return user, nil
|
|
}
|
|
|
|
func linkPendingExternalAccount(ctx context.Context, pending *PendingExternalAccount, input LinkExistingInput) (*model.User, error) {
|
|
if pending == nil || pending.AuthSourceID == 0 || pending.ExternalID == "" {
|
|
return nil, errors.New(errPendingOAuthExpired)
|
|
}
|
|
input.Username = strings.TrimSpace(input.Username)
|
|
if input.Username == "" || input.Password == "" {
|
|
return nil, errors.New(errInvalidParams)
|
|
}
|
|
|
|
var user model.User
|
|
if err := db.DB(ctx).Where("username = ? OR email = ?", input.Username, input.Username).First(&user).Error; err != nil {
|
|
return nil, errors.New(errUsernameOrPasswordWrong)
|
|
}
|
|
if !user.IsActive {
|
|
return nil, errors.New(errBannedAccount)
|
|
}
|
|
if !user.CheckPassword(input.Password) {
|
|
return nil, errors.New(errUsernameOrPasswordWrong)
|
|
}
|
|
|
|
if existing, err := model.FindExternalAccount(ctx, pending.AuthSourceID, pending.ExternalID); err == nil {
|
|
if existing.UserID != user.ID {
|
|
return nil, errors.New("该第三方账号已绑定其他用户")
|
|
}
|
|
return &user, nil
|
|
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return nil, err
|
|
}
|
|
|
|
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
|
|
AuthSourceID: pending.AuthSourceID,
|
|
UserID: user.ID,
|
|
ExternalID: pending.ExternalID,
|
|
ExternalUsername: pending.ExternalUsername,
|
|
Email: pending.Email,
|
|
}); err != nil {
|
|
return nil, err
|
|
}
|
|
return &user, nil
|
|
}
|
|
|
|
func currentUserFromLegacyToken(ctx context.Context, c *gin.Context) *model.User {
|
|
token := strings.TrimSpace(c.GetHeader(compat.OpenFlareTokenHeader()))
|
|
if token == "" {
|
|
return nil
|
|
}
|
|
tokenHash := model.HashToken(token)
|
|
var record model.AccessToken
|
|
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&record).Error; err != nil {
|
|
return nil
|
|
}
|
|
var user model.User
|
|
if err := db.DB(ctx).Where("id = ? AND is_active = ?", record.UserID, true).First(&user).Error; err != nil {
|
|
return nil
|
|
}
|
|
return &user
|
|
}
|
|
|
|
func uniqueLegacyUsername(ctx context.Context, base string) (string, error) {
|
|
base = strings.TrimSpace(base)
|
|
if base == "" {
|
|
base = "user"
|
|
}
|
|
candidate := base
|
|
for i := 0; i <= 1000; i++ {
|
|
if i > 0 {
|
|
candidate = fmt.Sprintf("%s-%d", base, i)
|
|
}
|
|
count, err := repository.CountUsersByUsername(ctx, candidate)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
if count == 0 {
|
|
return candidate, nil
|
|
}
|
|
}
|
|
return "", errors.New("无法生成唯一用户名")
|
|
}
|
|
|
|
func parseUint64(raw string) (uint64, error) {
|
|
var id uint64
|
|
_, err := fmt.Sscanf(raw, "%d", &id)
|
|
return id, err
|
|
}
|
|
|
|
func isOIDCLoginEnabled(ctx context.Context) bool {
|
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
|
|
return err != nil || enabled
|
|
}
|
|
|
|
// The following functions mirror oauth package internals for legacy GET callback support.
|
|
// They intentionally duplicate minimal logic to avoid modifying the core oauth module.
|
|
|
|
func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code, nonce, redirectURL string) (*model.OAuthUserInfo, error) {
|
|
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
token, err := authConfig.Exchange(ctx, code)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
userInfo := &model.OAuthUserInfo{Active: true}
|
|
if verifier != nil {
|
|
if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
|
|
return nil, verifyErr
|
|
}
|
|
}
|
|
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
|
|
userInfo.Username = userInfo.PreferredUsername
|
|
}
|
|
if userInfo.Username == "" && userInfo.Email != "" {
|
|
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
|
|
}
|
|
if userInfo.Username == "" && userInfo.Sub != "" {
|
|
userInfo.Username = userInfo.Sub
|
|
}
|
|
if userInfo.Name == "" {
|
|
userInfo.Name = userInfo.Username
|
|
}
|
|
return userInfo, nil
|
|
}
|
|
|
|
func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
|
|
userInfo.Username = strings.TrimSpace(userInfo.Username)
|
|
userInfo.Email = strings.TrimSpace(userInfo.Email)
|
|
userInfo.Name = strings.TrimSpace(userInfo.Name)
|
|
if userInfo.Username == "" {
|
|
return errors.New("无法从认证源获取用户名")
|
|
}
|
|
if userInfo.Name == "" {
|
|
userInfo.Name = userInfo.Username
|
|
}
|
|
if !userInfo.Active {
|
|
userInfo.Active = true
|
|
}
|
|
return nil
|
|
}
|