mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 07:36:37 +08:00
refactor(auth): modularize auth plugin with physical subpackages and decoupled services
This commit is contained in:
@@ -0,0 +1,49 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service implements domain business services and orchestration for the auth plugin.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// AuthRegistryImpl implements contracts.AuthRegistry.
|
||||
type AuthRegistryImpl struct {
|
||||
mu sync.RWMutex
|
||||
providers map[string]contracts.OAuthProvider
|
||||
}
|
||||
|
||||
// NewAuthRegistry creates a new AuthRegistryImpl.
|
||||
func NewAuthRegistry() *AuthRegistryImpl {
|
||||
return &AuthRegistryImpl{
|
||||
providers: make(map[string]contracts.OAuthProvider),
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterOAuthProvider registers an OAuthProvider by name.
|
||||
func (r *AuthRegistryImpl) RegisterOAuthProvider(name string, provider contracts.OAuthProvider) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.providers[name] = provider
|
||||
}
|
||||
|
||||
// GetOAuthProvider retrieves an OAuthProvider by name.
|
||||
func (r *AuthRegistryImpl) GetOAuthProvider(name string) (contracts.OAuthProvider, bool) {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
p, ok := r.providers[name]
|
||||
return p, ok
|
||||
}
|
||||
|
||||
// ListOAuthProviders lists all registered provider names.
|
||||
func (r *AuthRegistryImpl) ListOAuthProviders() []string {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
res := make([]string, 0, len(r.providers))
|
||||
for name := range r.providers {
|
||||
res = append(res, name)
|
||||
}
|
||||
return res
|
||||
}
|
||||
@@ -0,0 +1,278 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service implements domain business services and orchestration for the auth plugin.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/auth/consts"
|
||||
"Wavelet/plugins/domain/auth/dao"
|
||||
"Wavelet/plugins/domain/auth/model/entity"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
)
|
||||
|
||||
// HashToken computes SHA-256 hex digest of access token.
|
||||
func HashToken(token string) string {
|
||||
h := sha256.New()
|
||||
h.Write([]byte(token))
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
// UserIDExtractor extracts user ID from a request context.
|
||||
type UserIDExtractor func(ctx context.Context) (uint64, bool)
|
||||
|
||||
// AuthServiceImpl implements contracts.AuthService.
|
||||
type AuthServiceImpl struct {
|
||||
dao *dao.DAO
|
||||
requireAuthMiddleware any
|
||||
requireAdminMiddleware any
|
||||
disallowTokenMiddleware any
|
||||
userIDExtractor UserIDExtractor
|
||||
}
|
||||
|
||||
// NewAuthService creates a new AuthServiceImpl.
|
||||
func NewAuthService(
|
||||
d *dao.DAO,
|
||||
requireAuth any,
|
||||
requireAdmin any,
|
||||
disallowToken any,
|
||||
extractor UserIDExtractor,
|
||||
) *AuthServiceImpl {
|
||||
return &AuthServiceImpl{
|
||||
dao: d,
|
||||
requireAuthMiddleware: requireAuth,
|
||||
requireAdminMiddleware: requireAdmin,
|
||||
disallowTokenMiddleware: disallowToken,
|
||||
userIDExtractor: extractor,
|
||||
}
|
||||
}
|
||||
|
||||
// SetMiddlewareHandlers wires middleware handlers into AuthService after controller initialization.
|
||||
func (s *AuthServiceImpl) SetMiddlewareHandlers(requireAuth, requireAdmin, disallowToken any, extractor UserIDExtractor) {
|
||||
s.requireAuthMiddleware = requireAuth
|
||||
s.requireAdminMiddleware = requireAdmin
|
||||
s.disallowTokenMiddleware = disallowToken
|
||||
s.userIDExtractor = extractor
|
||||
}
|
||||
|
||||
// RequireAuthMiddleware returns the authentication check middleware.
|
||||
func (s *AuthServiceImpl) RequireAuthMiddleware() any {
|
||||
return s.requireAuthMiddleware
|
||||
}
|
||||
|
||||
// RequireAdminMiddleware returns the admin authorization middleware.
|
||||
func (s *AuthServiceImpl) RequireAdminMiddleware() any {
|
||||
return s.requireAdminMiddleware
|
||||
}
|
||||
|
||||
// DisallowTokenAuthMiddleware returns the token rejection middleware.
|
||||
func (s *AuthServiceImpl) DisallowTokenAuthMiddleware() any {
|
||||
return s.disallowTokenMiddleware
|
||||
}
|
||||
|
||||
// GetCurrentUser 从 context 中读取登录用户。
|
||||
func (s *AuthServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
|
||||
if v := ctx.Value(contracts.AuthUserObjKey); v != nil {
|
||||
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
|
||||
return u, nil
|
||||
}
|
||||
}
|
||||
|
||||
return nil, errors.New(consts.ErrUserNotInContext)
|
||||
}
|
||||
|
||||
// GetCurrentUserID 从请求登录态中读取用户 ID。
|
||||
func (s *AuthServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) {
|
||||
if s.userIDExtractor != nil {
|
||||
if userID, ok := s.userIDExtractor(ctx); ok {
|
||||
return userID, nil
|
||||
}
|
||||
}
|
||||
return 0, errors.New(consts.ErrUserNotInContext)
|
||||
}
|
||||
|
||||
// VerifyToken 验证访问令牌并返回对应的用户。
|
||||
func (s *AuthServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) {
|
||||
if token == "" {
|
||||
return nil, errors.New(consts.ErrEmptyToken)
|
||||
}
|
||||
|
||||
tokenHash := HashToken(token)
|
||||
tokenRecord, err := s.dao.GetCachedToken(ctx, tokenHash)
|
||||
if err != nil {
|
||||
tokenRecord, err = s.dao.GetAccessTokenByHash(ctx, tokenHash)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s.dao.SetCachedToken(ctx, tokenHash, tokenRecord)
|
||||
}
|
||||
|
||||
user, err := s.dao.GetCachedUser(ctx, tokenRecord.UserID)
|
||||
if err != nil || user == nil || !user.IsActive {
|
||||
user, err = s.dao.GetActiveUserByID(ctx, tokenRecord.UserID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s.dao.SetCachedUser(ctx, tokenRecord.UserID, user)
|
||||
}
|
||||
|
||||
if user.Username == consts.SystemUsername {
|
||||
return nil, errors.New(consts.ErrSystemUserTokenNotAllowed)
|
||||
}
|
||||
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// CreateSession establishes an authenticated session.
|
||||
func (s *AuthServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// RevokeUserSessions revokes active sessions and cached tokens for a user.
|
||||
func (s *AuthServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error {
|
||||
s.dao.InvalidateCachedUser(ctx, userID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// RevokeToken invalidates a cached token by its hash.
|
||||
func (s *AuthServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error {
|
||||
s.dao.InvalidateCachedToken(ctx, tokenHash)
|
||||
return nil
|
||||
}
|
||||
|
||||
// InvalidateCachedUser invalidates cached user profile.
|
||||
func (s *AuthServiceImpl) InvalidateCachedUser(ctx context.Context, userID uint64) {
|
||||
s.dao.InvalidateCachedUser(ctx, userID)
|
||||
}
|
||||
|
||||
// InvalidateCachedToken invalidates cached access token.
|
||||
func (s *AuthServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash string) {
|
||||
s.dao.InvalidateCachedToken(ctx, tokenHash)
|
||||
}
|
||||
|
||||
// ListAuthSources lists all configured authentication sources.
|
||||
func (s *AuthServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
|
||||
sources, err := s.dao.ListAllAuthSources(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
views := make([]contracts.AuthSourceViewDTO, len(sources))
|
||||
for i := range sources {
|
||||
views[i] = contracts.AuthSourceViewDTO{
|
||||
ID: sources[i].ID,
|
||||
Name: sources[i].Name,
|
||||
Type: sources[i].Type,
|
||||
DisplayName: sources[i].DisplayName,
|
||||
IsActive: sources[i].IsActive,
|
||||
IconURL: sources[i].IconURL,
|
||||
ClientSecretConfigured: sources[i].ClientSecret != "",
|
||||
}
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
// CreateAuthSource creates a new authentication source.
|
||||
func (s *AuthServiceImpl) CreateAuthSource(ctx context.Context, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
|
||||
model := entity.AuthSource{
|
||||
ID: source.ID,
|
||||
Name: source.Name,
|
||||
Type: source.Type,
|
||||
DisplayName: source.DisplayName,
|
||||
ClientID: source.ClientID,
|
||||
ClientSecret: source.ClientSecret,
|
||||
OpenIDDiscoveryURL: source.OpenIDDiscoveryURL,
|
||||
Scopes: source.Scopes,
|
||||
IconURL: source.IconURL,
|
||||
IsActive: source.IsActive,
|
||||
}
|
||||
|
||||
if err := model.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := s.dao.CreateAuthSource(ctx, &model); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
model.Sanitize()
|
||||
return toAuthSourceDTO(&model), nil
|
||||
}
|
||||
|
||||
// UpdateAuthSource updates an existing authentication source.
|
||||
func (s *AuthServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
|
||||
existing, err := s.dao.GetAuthSourceByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
existing.DisplayName = source.DisplayName
|
||||
existing.ClientID = source.ClientID
|
||||
if source.ClientSecret != "" {
|
||||
existing.ClientSecret = source.ClientSecret
|
||||
}
|
||||
existing.OpenIDDiscoveryURL = source.OpenIDDiscoveryURL
|
||||
existing.Scopes = source.Scopes
|
||||
existing.IconURL = source.IconURL
|
||||
|
||||
if err := existing.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := s.dao.SaveAuthSource(ctx, existing); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
existing.Sanitize()
|
||||
return toAuthSourceDTO(existing), nil
|
||||
}
|
||||
|
||||
// DeleteAuthSource deletes an authentication source.
|
||||
func (s *AuthServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
|
||||
existing, err := s.dao.GetAuthSourceByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return s.dao.DeleteAuthSource(ctx, existing)
|
||||
}
|
||||
|
||||
// ToggleAuthSource toggles active status of an authentication source.
|
||||
func (s *AuthServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
|
||||
existing, err := s.dao.GetAuthSourceByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
existing.IsActive = !existing.IsActive
|
||||
if err := s.dao.SaveAuthSource(ctx, existing); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
existing.Sanitize()
|
||||
return toAuthSourceDTO(existing), nil
|
||||
}
|
||||
|
||||
func toAuthSourceDTO(s *entity.AuthSource) *contracts.AuthSourceDTO {
|
||||
if s == nil {
|
||||
return nil
|
||||
}
|
||||
return &contracts.AuthSourceDTO{
|
||||
ID: s.ID,
|
||||
Name: s.Name,
|
||||
Type: s.Type,
|
||||
DisplayName: s.DisplayName,
|
||||
ClientID: s.ClientID,
|
||||
ClientSecret: s.ClientSecret,
|
||||
OpenIDDiscoveryURL: s.OpenIDDiscoveryURL,
|
||||
Scopes: s.Scopes,
|
||||
IconURL: s.IconURL,
|
||||
IsActive: s.IsActive,
|
||||
CreatedAt: s.CreatedAt,
|
||||
UpdatedAt: s.UpdatedAt,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service implements domain business services and orchestration for the auth plugin.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/auth/consts"
|
||||
"Wavelet/plugins/domain/auth/model/dto"
|
||||
"Wavelet/plugins/domain/auth/pow"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// CaptchaManager orchestrates challenge generation and solution validation.
|
||||
type CaptchaManager struct {
|
||||
secret []byte
|
||||
store pow.Store
|
||||
settingsMgr *CapSettingsManager
|
||||
}
|
||||
|
||||
// NewCaptchaManager creates a new CAPTCHA Manager.
|
||||
func NewCaptchaManager(secret []byte, store pow.Store, settingsMgr *CapSettingsManager) *CaptchaManager {
|
||||
return &CaptchaManager{
|
||||
secret: secret,
|
||||
store: store,
|
||||
settingsMgr: settingsMgr,
|
||||
}
|
||||
}
|
||||
|
||||
// SetSecret updates the shared secret used for PoW generation and validation.
|
||||
func (m *CaptchaManager) SetSecret(secret []byte) {
|
||||
m.secret = secret
|
||||
}
|
||||
|
||||
// Generate creates a challenge response.
|
||||
func (m *CaptchaManager) Generate(ctx context.Context, scope string) (*pow.ChallengeResponse, error) {
|
||||
settings, err := m.settingsMgr.Current(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
challengeConfig := pow.ChallengeConfig{
|
||||
Count: settings.ChallengeCount,
|
||||
Size: settings.ChallengeSize,
|
||||
Difficulty: settings.ChallengeDifficulty,
|
||||
Expires: settings.ChallengeTTL,
|
||||
}
|
||||
return pow.GenerateChallenge(m.secret, challengeConfig, scope)
|
||||
}
|
||||
|
||||
// Redeem verifies PoW solutions and returns a one-time redeem token.
|
||||
func (m *CaptchaManager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*dto.RedeemResponse, error) {
|
||||
sigHex := pow.JwtSigHex(token)
|
||||
if sigHex == "" {
|
||||
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrInvalidToken}, nil
|
||||
}
|
||||
|
||||
nonceKey := "cap:nonce:" + sigHex
|
||||
|
||||
payload, err := pow.VerifyChallengeSolutions(token, solutions, m.secret, scope)
|
||||
if err != nil {
|
||||
return &dto.RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors returned as response
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond
|
||||
if nonceTTL < time.Second {
|
||||
nonceTTL = time.Second
|
||||
}
|
||||
|
||||
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
|
||||
if err != nil {
|
||||
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrNonceStoreFailed}, err
|
||||
}
|
||||
if !set {
|
||||
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrAlreadyRedeemed}, nil
|
||||
}
|
||||
|
||||
settings, err := m.settingsMgr.Current(ctx)
|
||||
if err != nil {
|
||||
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrSettingsLoad}, err
|
||||
}
|
||||
|
||||
id := pow.RandomHex(consts.RedeemTokenIDLength)
|
||||
verToken := pow.RandomHex(consts.RedeemVerTokenLength)
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
tokenKey := "cap:token:" + id + ":" + verHashHex
|
||||
tokenExpires := time.Now().Add(settings.TokenTTL)
|
||||
storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope
|
||||
|
||||
if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil {
|
||||
return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrTokenStoreFailed}, err
|
||||
}
|
||||
|
||||
return &dto.RedeemResponse{
|
||||
Success: true,
|
||||
Token: id + ":" + verToken,
|
||||
Expires: tokenExpires.UnixNano() / int64(time.Millisecond),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// VerifyToken validates and consumes the redeem token (single-use).
|
||||
func (m *CaptchaManager) VerifyToken(ctx context.Context, token, expectedScope string) (bool, error) {
|
||||
if token == "" {
|
||||
return false, nil
|
||||
}
|
||||
parts := strings.Split(token, ":")
|
||||
if len(parts) != consts.TokenPartsCount {
|
||||
return false, nil
|
||||
}
|
||||
id := parts[0]
|
||||
verToken := parts[1]
|
||||
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
tokenKey := "cap:token:" + id + ":" + verHashHex
|
||||
|
||||
if m.store == nil {
|
||||
return false, nil
|
||||
}
|
||||
val, exists, err := m.store.GetAndDelete(ctx, tokenKey)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !exists {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
valParts := strings.Split(val, "|")
|
||||
if len(valParts) != consts.ValuePartsCount {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
expNano, err := strconv.ParseInt(valParts[0], 10, 64)
|
||||
if err != nil {
|
||||
return false, nil //nolint:nilerr // invalid format is failure
|
||||
}
|
||||
tokenScope := valParts[1]
|
||||
|
||||
if expectedScope != "" && tokenScope != expectedScope {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if time.Now().UnixNano() > expNano {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// CaptchaServiceImpl implements contracts.CaptchaService.
|
||||
type CaptchaServiceImpl struct {
|
||||
manager *CaptchaManager
|
||||
verifyMiddleware func(scope string) any
|
||||
challengeHandler any
|
||||
redeemHandler any
|
||||
}
|
||||
|
||||
// NewCaptchaService creates a new CaptchaServiceImpl.
|
||||
func NewCaptchaService(mgr *CaptchaManager, verifyMiddleware func(scope string) any, challengeHandler any, redeemHandler any) contracts.CaptchaService {
|
||||
return &CaptchaServiceImpl{
|
||||
manager: mgr,
|
||||
verifyMiddleware: verifyMiddleware,
|
||||
challengeHandler: challengeHandler,
|
||||
redeemHandler: redeemHandler,
|
||||
}
|
||||
}
|
||||
|
||||
// VerifyMiddleware returns the captcha verification middleware.
|
||||
func (s *CaptchaServiceImpl) VerifyMiddleware(scope string) any {
|
||||
if s.verifyMiddleware != nil {
|
||||
return s.verifyMiddleware(scope)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ChallengeHandler returns the challenge HTTP handler.
|
||||
func (s *CaptchaServiceImpl) ChallengeHandler() any {
|
||||
return s.challengeHandler
|
||||
}
|
||||
|
||||
// RedeemHandler returns the redeem HTTP handler.
|
||||
func (s *CaptchaServiceImpl) RedeemHandler() any {
|
||||
return s.redeemHandler
|
||||
}
|
||||
@@ -0,0 +1,119 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service implements domain business services and orchestration for the auth plugin.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/auth/consts"
|
||||
"Wavelet/plugins/domain/auth/dao"
|
||||
"Wavelet/plugins/domain/auth/model/do"
|
||||
"context"
|
||||
"errors"
|
||||
"sync/atomic"
|
||||
|
||||
"golang.org/x/sync/singleflight"
|
||||
)
|
||||
|
||||
var capRuntimeConfigKeys = []string{
|
||||
consts.ConfigKeyCapLoginEnabled,
|
||||
consts.ConfigKeyCapChallengeCount,
|
||||
consts.ConfigKeyCapChallengeSize,
|
||||
consts.ConfigKeyCapChallengeDifficulty,
|
||||
consts.ConfigKeyCapChallengeTTL,
|
||||
consts.ConfigKeyCapTokenTTL,
|
||||
}
|
||||
|
||||
var capRuntimeConfigKeySet = func() map[string]struct{} {
|
||||
set := make(map[string]struct{}, len(capRuntimeConfigKeys))
|
||||
for _, key := range capRuntimeConfigKeys {
|
||||
set[key] = struct{}{}
|
||||
}
|
||||
return set
|
||||
}()
|
||||
|
||||
// IsCapRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings.
|
||||
func IsCapRuntimeConfigKey(key string) bool {
|
||||
_, ok := capRuntimeConfigKeySet[key]
|
||||
return ok
|
||||
}
|
||||
|
||||
// CapSettingsManager manages dynamic CAPTCHA configuration cache.
|
||||
type CapSettingsManager struct {
|
||||
dao *dao.DAO
|
||||
snapshot atomic.Pointer[do.CapRuntimeSettings]
|
||||
loadGroup singleflight.Group
|
||||
}
|
||||
|
||||
// NewCapSettingsManager creates a new CapSettingsManager.
|
||||
func NewCapSettingsManager(d *dao.DAO) *CapSettingsManager {
|
||||
return &CapSettingsManager{
|
||||
dao: d,
|
||||
}
|
||||
}
|
||||
|
||||
// Invalidate drops the in-process CAPTCHA settings snapshot.
|
||||
func (m *CapSettingsManager) Invalidate() {
|
||||
m.snapshot.Store(nil)
|
||||
}
|
||||
|
||||
// Current returns the cached CAPTCHA runtime settings snapshot.
|
||||
func (m *CapSettingsManager) Current(ctx context.Context) (do.CapRuntimeSettings, error) {
|
||||
if snapshot := m.snapshot.Load(); snapshot != nil {
|
||||
return *snapshot, nil
|
||||
}
|
||||
|
||||
loaded, err, _ := m.loadGroup.Do("cap-runtime-settings", func() (any, error) {
|
||||
if snapshot := m.snapshot.Load(); snapshot != nil {
|
||||
return *snapshot, nil
|
||||
}
|
||||
|
||||
settings, loadErr := m.loadSettings(ctx)
|
||||
if loadErr != nil {
|
||||
return do.CapRuntimeSettings{}, loadErr
|
||||
}
|
||||
|
||||
m.snapshot.Store(&settings)
|
||||
return settings, nil
|
||||
})
|
||||
if err != nil {
|
||||
return do.CapRuntimeSettings{}, err
|
||||
}
|
||||
|
||||
settings, ok := loaded.(do.CapRuntimeSettings)
|
||||
if !ok {
|
||||
return do.CapRuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type")
|
||||
}
|
||||
return settings, nil
|
||||
}
|
||||
|
||||
// CapProtectionEnabled reports whether CAPTCHA verification is required for protected routes.
|
||||
func (m *CapSettingsManager) CapProtectionEnabled(ctx context.Context) bool {
|
||||
settings, err := m.Current(ctx)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return settings.LoginEnabled
|
||||
}
|
||||
|
||||
// InstallTestSnapshot installs a fixed snapshot for unit tests.
|
||||
func (m *CapSettingsManager) InstallTestSnapshot(settings do.CapRuntimeSettings) func() {
|
||||
snapshot := settings
|
||||
m.snapshot.Store(&snapshot)
|
||||
return m.Invalidate
|
||||
}
|
||||
|
||||
func (m *CapSettingsManager) loadSettings(ctx context.Context) (do.CapRuntimeSettings, error) {
|
||||
if m.dao == nil {
|
||||
return do.ParseCapRuntimeSettings(nil), nil
|
||||
}
|
||||
records, err := m.dao.ListSystemConfigsByKeys(ctx, capRuntimeConfigKeys)
|
||||
if err != nil {
|
||||
return do.CapRuntimeSettings{}, err
|
||||
}
|
||||
configs := make(map[string]string, len(records))
|
||||
for _, r := range records {
|
||||
configs[r.Key] = r.Value
|
||||
}
|
||||
return do.ParseCapRuntimeSettings(configs), nil
|
||||
}
|
||||
@@ -0,0 +1,440 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service implements domain business services and orchestration for the auth plugin.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/plugins/domain/auth/consts"
|
||||
"Wavelet/plugins/domain/auth/dao"
|
||||
"Wavelet/plugins/domain/auth/model/dto"
|
||||
"Wavelet/plugins/domain/auth/model/entity"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"golang.org/x/oauth2"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// OAuthService orchestrates OAuth/OIDC operations.
|
||||
type OAuthService struct {
|
||||
dao *dao.DAO
|
||||
providerCache *OIDCProviderCache
|
||||
sessionSvc *SessionService
|
||||
}
|
||||
|
||||
// NewOAuthService creates a new OAuthService.
|
||||
func NewOAuthService(d *dao.DAO, cache *OIDCProviderCache, sessSvc *SessionService) *OAuthService {
|
||||
return &OAuthService{
|
||||
dao: d,
|
||||
providerCache: cache,
|
||||
sessionSvc: sessSvc,
|
||||
}
|
||||
}
|
||||
|
||||
// IsOIDCLoginEnabled checks if OIDC login is globally enabled.
|
||||
func (s *OAuthService) IsOIDCLoginEnabled(ctx context.Context) bool {
|
||||
val, err := s.dao.GetSystemConfigValue(ctx, "oidc_login_enabled")
|
||||
if err != nil || val == "" {
|
||||
return true
|
||||
}
|
||||
b, err := strconv.ParseBool(val)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// ResolveAuthSource retrieves the specified or default active auth source.
|
||||
func (s *OAuthService) ResolveAuthSource(ctx context.Context, sourceName string) (*entity.AuthSource, error) {
|
||||
name := strings.TrimSpace(strings.ToLower(sourceName))
|
||||
if name == "" {
|
||||
sources, err := s.dao.ListActiveAuthSources(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(sources) == 0 {
|
||||
return nil, errors.New(consts.ErrNoActiveAuthSource)
|
||||
}
|
||||
src, err := s.dao.GetAuthSourceByName(ctx, sources[0].Name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return src, nil
|
||||
}
|
||||
src, err := s.dao.GetAuthSourceByName(ctx, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return src, nil
|
||||
}
|
||||
|
||||
// ActiveLoginSources returns all active login sources formatted for display.
|
||||
func (s *OAuthService) ActiveLoginSources(ctx context.Context) ([]dto.AuthSourceView, error) {
|
||||
if !s.IsOIDCLoginEnabled(ctx) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
dbSources, err := s.dao.ListActiveAuthSources(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sources := make([]dto.AuthSourceView, 0, len(dbSources))
|
||||
for _, source := range dbSources {
|
||||
sources = append(sources, dto.AuthSourceView{
|
||||
ID: source.ID,
|
||||
Name: source.Name,
|
||||
Type: source.Type,
|
||||
DisplayName: source.DisplayName,
|
||||
IsActive: source.IsActive,
|
||||
IconURL: source.IconURL,
|
||||
ClientSecretConfigured: source.ClientSecretConfigured,
|
||||
})
|
||||
}
|
||||
return sources, nil
|
||||
}
|
||||
|
||||
// GetFrontendLoginRedirectURL constructs the OAuth frontend redirect URL.
|
||||
func (s *OAuthService) GetFrontendLoginRedirectURL(ctx context.Context) (string, error) {
|
||||
val, err := s.dao.GetSystemConfigValue(ctx, "server_address")
|
||||
if err != nil || strings.TrimSpace(val) == "" {
|
||||
return "", errors.New(consts.ErrServerAddressMissing)
|
||||
}
|
||||
return strings.TrimRight(val, "/") + "/login", nil
|
||||
}
|
||||
|
||||
// ReserveOAuthStateSlot ensures that a session does not abuse OAuth state generation.
|
||||
func (s *OAuthService) ReserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
|
||||
if sessionHash == "" {
|
||||
return nil
|
||||
}
|
||||
if limiter := s.dao.Limiter(); limiter != nil {
|
||||
key := fmt.Sprintf(consts.OAuthStateLimitKeyFormat, sessionHash)
|
||||
res, err := limiter.Allow(ctx, key, contracts.Rate{
|
||||
Limit: consts.OAuthStateLimitMax,
|
||||
Period: consts.OAuthStateCacheKeyExpiration,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !res.Allowed {
|
||||
return errors.New(consts.ErrOAuthStateRateLimited)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
cache := s.dao.Cache()
|
||||
if cache == nil {
|
||||
return nil
|
||||
}
|
||||
key := fmt.Sprintf(consts.OAuthStateLimitKeyFormat, sessionHash)
|
||||
var count int
|
||||
_ = cache.Get(ctx, key, &count)
|
||||
count++
|
||||
_ = cache.Set(ctx, key, count, consts.OAuthStateCacheKeyExpiration)
|
||||
if count > consts.OAuthStateLimitMax {
|
||||
return errors.New(consts.ErrOAuthStateRateLimited)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// BuildOAuthConfig builds oauth2.Config and oidc.IDTokenVerifier.
|
||||
func (s *OAuthService) BuildOAuthConfig(ctx context.Context, source *entity.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
|
||||
if source == nil {
|
||||
return nil, nil, errors.New(consts.ErrAuthSourceRequired)
|
||||
}
|
||||
|
||||
if source.OpenIDDiscoveryURL == "" {
|
||||
return nil, nil, errors.New(consts.ErrDiscoveryURLRequired)
|
||||
}
|
||||
|
||||
issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/")
|
||||
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
|
||||
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
|
||||
|
||||
provider, err := s.providerCache.Get(ctx, issuer)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID})
|
||||
scopes := strings.Fields(source.Scopes)
|
||||
if len(scopes) == 0 {
|
||||
scopes = []string{oidc.ScopeOpenID, "profile", "email"}
|
||||
}
|
||||
if !containsScope(scopes, oidc.ScopeOpenID) {
|
||||
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
|
||||
}
|
||||
|
||||
return &oauth2.Config{
|
||||
ClientID: source.ClientID,
|
||||
ClientSecret: source.ClientSecret,
|
||||
RedirectURL: redirectURL,
|
||||
Scopes: scopes,
|
||||
Endpoint: provider.Endpoint(),
|
||||
}, verifier, nil
|
||||
}
|
||||
|
||||
func containsScope(scopes []string, scope string) bool {
|
||||
for _, item := range scopes {
|
||||
if item == scope {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// BuildAuthorizeURL generates the redirect authorize URL for the source and state.
|
||||
func (s *OAuthService) BuildAuthorizeURL(ctx context.Context, source *entity.AuthSource, state string) (string, error) {
|
||||
redirectURL, err := s.GetFrontendLoginRedirectURL(ctx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
authConfig, verifier, err := s.BuildOAuthConfig(ctx, source, redirectURL)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if verifier != nil {
|
||||
return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil
|
||||
}
|
||||
return authConfig.AuthCodeURL(state), nil
|
||||
}
|
||||
|
||||
// BuildOAuthUserInfo exchanges the auth code and retrieves user identity claims.
|
||||
func (s *OAuthService) BuildOAuthUserInfo(ctx context.Context, source *entity.AuthSource, code, nonce, redirectURL string) (*contracts.OAuthUserInfoDTO, error) {
|
||||
authConfig, verifier, err := s.BuildOAuthConfig(ctx, source, redirectURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
token, err := authConfig.Exchange(ctx, code)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
userInfo := &contracts.OAuthUserInfoDTO{Active: true}
|
||||
if verifier != nil {
|
||||
if verifyErr := s.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 (s *OAuthService) verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *contracts.OAuthUserInfoDTO) error {
|
||||
rawIDToken, ok := token.Extra("id_token").(string)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
|
||||
if verifyErr != nil {
|
||||
return fmt.Errorf(consts.ErrIDTokenVerifyFailedFormat, consts.ErrIDTokenVerifyFailed, verifyErr)
|
||||
}
|
||||
if nonce != "" && idToken.Nonce != nonce {
|
||||
return errors.New(consts.ErrNonceMismatch)
|
||||
}
|
||||
if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
|
||||
return claimsErr
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// NormalizeOAuthUserInfo sanitizes user claims.
|
||||
func (s *OAuthService) NormalizeOAuthUserInfo(userInfo *contracts.OAuthUserInfoDTO) error {
|
||||
userInfo.Username = strings.TrimSpace(userInfo.Username)
|
||||
userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
|
||||
userInfo.Email = strings.TrimSpace(userInfo.Email)
|
||||
userInfo.Name = strings.TrimSpace(userInfo.Name)
|
||||
userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL)
|
||||
|
||||
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.Username == "" {
|
||||
return errors.New(consts.ErrUsernameFromSourceFailed)
|
||||
}
|
||||
if userInfo.Name == "" {
|
||||
userInfo.Name = userInfo.Username
|
||||
}
|
||||
if !userInfo.Active {
|
||||
userInfo.Active = true
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UniqueUsername generates a unique username given a base candidate.
|
||||
func (s *OAuthService) UniqueUsername(ctx context.Context, base string) (string, error) {
|
||||
base = strings.TrimSpace(base)
|
||||
if base == "" {
|
||||
base = "user"
|
||||
}
|
||||
|
||||
existingUsernames, err := s.dao.ListSimilarUsernames(ctx, base)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
exists := make(map[string]bool, len(existingUsernames))
|
||||
for _, u := range existingUsernames {
|
||||
exists[strings.ToLower(u)] = true
|
||||
}
|
||||
|
||||
if !exists[strings.ToLower(base)] {
|
||||
return base, nil
|
||||
}
|
||||
|
||||
for i := 1; i <= 1000; i++ {
|
||||
candidate := fmt.Sprintf("%s-%d", base, i)
|
||||
if !exists[strings.ToLower(candidate)] {
|
||||
return candidate, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", errors.New(consts.ErrUsernameGenerateFailed)
|
||||
}
|
||||
|
||||
// BindExternalAccount binds an external identity to an existing user.
|
||||
func (s *OAuthService) BindExternalAccount(ctx context.Context, sourceID, userID uint64, userInfo *contracts.OAuthUserInfoDTO) error {
|
||||
user, err := s.dao.GetUserByID(ctx, userID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.dao.BindExternalAccount(ctx, &entity.ExternalAccount{
|
||||
AuthSourceID: sourceID,
|
||||
UserID: user.ID,
|
||||
ExternalID: userInfo.Sub,
|
||||
ExternalUsername: userInfo.Username,
|
||||
Email: userInfo.Email,
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = s.dao.TouchUserLastLogin(ctx, user.ID, user.LastLoginAt)
|
||||
return nil
|
||||
}
|
||||
|
||||
// AuthenticateOrRegisterUser finds existing binding or creates a new user.
|
||||
func (s *OAuthService) AuthenticateOrRegisterUser(ctx context.Context, source *entity.AuthSource, userInfo *contracts.OAuthUserInfoDTO) (*contracts.UserDTO, bool, error) {
|
||||
account, err := s.dao.FindExternalAccount(ctx, source.ID, userInfo.Sub)
|
||||
if err == nil {
|
||||
user, loadErr := s.dao.GetUserByID(ctx, account.UserID)
|
||||
if loadErr != nil {
|
||||
return nil, false, loadErr
|
||||
}
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = s.dao.TouchUserLastLogin(ctx, user.ID, user.LastLoginAt)
|
||||
return user, true, nil
|
||||
}
|
||||
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
// Not found -> check registration
|
||||
registrationEnabled := true
|
||||
val, cfgErr := s.dao.GetSystemConfigValue(ctx, "registration_enabled")
|
||||
if cfgErr == nil && val != "" {
|
||||
if b, err := strconv.ParseBool(val); err == nil {
|
||||
registrationEnabled = b
|
||||
}
|
||||
}
|
||||
|
||||
if !registrationEnabled {
|
||||
return nil, false, nil // registration disabled -> need bind
|
||||
}
|
||||
|
||||
username, uniqueErr := s.UniqueUsername(ctx, userInfo.Username)
|
||||
if uniqueErr != nil {
|
||||
return nil, false, uniqueErr
|
||||
}
|
||||
userInfo.Username = username
|
||||
|
||||
now := time.Now()
|
||||
user := contracts.UserDTO{
|
||||
ID: idgen.NextUint64ID(),
|
||||
Username: userInfo.Username,
|
||||
Nickname: userInfo.Name,
|
||||
Email: userInfo.Email,
|
||||
AvatarURL: userInfo.AvatarURL,
|
||||
IsActive: userInfo.Active,
|
||||
LastLoginAt: now,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
|
||||
if err := s.dao.InsertUser(ctx, &user); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
if err := s.dao.BindExternalAccount(ctx, &entity.ExternalAccount{
|
||||
AuthSourceID: source.ID,
|
||||
UserID: user.ID,
|
||||
ExternalID: userInfo.Sub,
|
||||
ExternalUsername: userInfo.Username,
|
||||
Email: userInfo.Email,
|
||||
}); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
return &user, true, nil
|
||||
}
|
||||
|
||||
// ListExternalAccounts returns sanitized external account bindings.
|
||||
func (s *OAuthService) ListExternalAccounts(ctx context.Context, userID uint64) ([]dto.ExternalAccountView, error) {
|
||||
accounts, err := s.dao.ListExternalAccountsByUserID(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]dto.ExternalAccountView, len(accounts))
|
||||
for i, acc := range accounts {
|
||||
source, _ := s.dao.GetAuthSourceByID(ctx, acc.AuthSourceID)
|
||||
sourceName, sourceType, sourceLabel := "", "", ""
|
||||
if source != nil {
|
||||
sourceName = source.Name
|
||||
sourceType = source.Type
|
||||
sourceLabel = source.DisplayName
|
||||
}
|
||||
views[i] = dto.ExternalAccountView{
|
||||
ID: acc.ID,
|
||||
AuthSourceID: acc.AuthSourceID,
|
||||
AuthSourceName: sourceName,
|
||||
AuthSourceType: sourceType,
|
||||
AuthSourceLabel: sourceLabel,
|
||||
ExternalUsername: acc.ExternalUsername,
|
||||
Email: acc.Email,
|
||||
CreatedAt: acc.CreatedAt.Format(time.RFC3339),
|
||||
}
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
// DeleteExternalAccount unbinds an external account.
|
||||
func (s *OAuthService) DeleteExternalAccount(ctx context.Context, id, userID uint64) error {
|
||||
return s.dao.UnbindExternalAccount(ctx, id, userID)
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service implements domain business services and orchestration for the auth plugin.
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"sync"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"golang.org/x/oauth2"
|
||||
"golang.org/x/sync/singleflight"
|
||||
)
|
||||
|
||||
// OIDCProviderCache 进程级 OIDC provider 缓存。
|
||||
type OIDCProviderCache struct {
|
||||
mu sync.RWMutex
|
||||
entries map[string]*oidc.Provider // key: normalized issuer URL
|
||||
sfGroup singleflight.Group
|
||||
}
|
||||
|
||||
// NewOIDCProviderCache creates a new OIDCProviderCache.
|
||||
func NewOIDCProviderCache() *OIDCProviderCache {
|
||||
return &OIDCProviderCache{
|
||||
entries: make(map[string]*oidc.Provider),
|
||||
}
|
||||
}
|
||||
|
||||
// discoveryContext 从请求 ctx 提取 HTTP 客户端,并绑定到不可取消的 Background ctx。
|
||||
func discoveryContext(ctx context.Context) context.Context {
|
||||
bg := context.Background()
|
||||
if client, ok := ctx.Value(oauth2.HTTPClient).(*http.Client); ok && client != nil {
|
||||
bg = oidc.ClientContext(bg, client)
|
||||
}
|
||||
return bg
|
||||
}
|
||||
|
||||
// Get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。
|
||||
func (c *OIDCProviderCache) Get(ctx context.Context, issuer string) (*oidc.Provider, error) {
|
||||
c.mu.RLock()
|
||||
if p, ok := c.entries[issuer]; ok {
|
||||
c.mu.RUnlock()
|
||||
return p, nil
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
discCtx := discoveryContext(ctx)
|
||||
v, err, _ := c.sfGroup.Do(issuer, func() (any, error) {
|
||||
c.mu.RLock()
|
||||
if p, ok := c.entries[issuer]; ok {
|
||||
c.mu.RUnlock()
|
||||
return p, nil
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
p, err := oidc.NewProvider(discCtx, issuer)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
c.entries[issuer] = p
|
||||
c.mu.Unlock()
|
||||
return p, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return v.(*oidc.Provider), nil //nolint:forcetypeassert
|
||||
}
|
||||
|
||||
// Invalidate 从缓存中移除指定 issuer 对应的 provider。
|
||||
func (c *OIDCProviderCache) Invalidate(issuer string) {
|
||||
c.mu.Lock()
|
||||
delete(c.entries, issuer)
|
||||
c.mu.Unlock()
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service implements domain business services and orchestration for the auth plugin.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/auth/dao"
|
||||
"Wavelet/plugins/domain/auth/pow"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Service aggregates all domain services for the auth plugin.
|
||||
type Service struct {
|
||||
DAO *dao.DAO
|
||||
Session *SessionService
|
||||
OAuth *OAuthService
|
||||
OIDCProviderCache *OIDCProviderCache
|
||||
CapSettings *CapSettingsManager
|
||||
CapManager *CaptchaManager
|
||||
AuthSvc *AuthServiceImpl
|
||||
AuthRegistry *AuthRegistryImpl
|
||||
}
|
||||
|
||||
// New creates a new Service container with all domain services wired up.
|
||||
func New(d *dao.DAO, sessionCfg SessionConfig, capSecret []byte) *Service {
|
||||
sessionSvc := NewSessionService(sessionCfg, d)
|
||||
oidcCache := NewOIDCProviderCache()
|
||||
oauthSvc := NewOAuthService(d, oidcCache, sessionSvc)
|
||||
capSettings := NewCapSettingsManager(d)
|
||||
|
||||
var capStore pow.Store
|
||||
if len(capSecret) > 0 {
|
||||
capStore = pow.NewMemoryStore(1 * time.Minute)
|
||||
}
|
||||
capMgr := NewCaptchaManager(capSecret, capStore, capSettings)
|
||||
|
||||
authSvc := NewAuthService(d, nil, nil, nil, nil)
|
||||
authRegistry := NewAuthRegistry()
|
||||
|
||||
return &Service{
|
||||
DAO: d,
|
||||
Session: sessionSvc,
|
||||
OAuth: oauthSvc,
|
||||
OIDCProviderCache: oidcCache,
|
||||
CapSettings: capSettings,
|
||||
CapManager: capMgr,
|
||||
AuthSvc: authSvc,
|
||||
AuthRegistry: authRegistry,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service implements domain business services and orchestration for the auth plugin.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/auth/consts"
|
||||
"Wavelet/plugins/domain/auth/dao"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/google/uuid"
|
||||
gsessions "github.com/gorilla/sessions"
|
||||
)
|
||||
|
||||
// SessionConfig defines session settings.
|
||||
type SessionConfig struct {
|
||||
SessionCookieName string `config:"session_cookie_name" env:"APP_SESSION_COOKIE_NAME" default:"wavelet_session"`
|
||||
SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"`
|
||||
SessionDomain string `config:"session_domain" env:"APP_SESSION_DOMAIN"`
|
||||
SessionAge int `config:"session_age" env:"APP_SESSION_AGE" default:"86400"`
|
||||
SessionHTTPOnly bool `config:"session_http_only" env:"APP_SESSION_HTTP_ONLY" default:"true"`
|
||||
SessionSecure bool `config:"session_secure" env:"APP_SESSION_SECURE"`
|
||||
}
|
||||
|
||||
// SessionService manages HTTP session operations, cookies, and tokens.
|
||||
type SessionService struct {
|
||||
mu sync.RWMutex
|
||||
config SessionConfig
|
||||
dao *dao.DAO
|
||||
}
|
||||
|
||||
// NewSessionService creates a new SessionService.
|
||||
func NewSessionService(cfg SessionConfig, d *dao.DAO) *SessionService {
|
||||
return &SessionService{
|
||||
config: cfg,
|
||||
dao: d,
|
||||
}
|
||||
}
|
||||
|
||||
// SetConfig updates the active session configuration.
|
||||
func (s *SessionService) SetConfig(cfg SessionConfig) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.config = cfg
|
||||
}
|
||||
|
||||
// Config returns the current session configuration.
|
||||
func (s *SessionService) Config() SessionConfig {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.config
|
||||
}
|
||||
|
||||
// GetSessionOptions 根据配置构建 Session 选项
|
||||
func (s *SessionService) GetSessionOptions(maxAge int) sessions.Options {
|
||||
cfg := s.Config()
|
||||
return sessions.Options{
|
||||
Path: "/",
|
||||
Domain: cfg.SessionDomain,
|
||||
MaxAge: maxAge,
|
||||
HttpOnly: cfg.SessionHTTPOnly,
|
||||
Secure: cfg.SessionSecure,
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
}
|
||||
}
|
||||
|
||||
// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie
|
||||
func (s *SessionService) StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
|
||||
headers := header["Set-Cookie"]
|
||||
if len(headers) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
newHeaders := make([]string, 0, len(headers))
|
||||
for _, h := range headers {
|
||||
if strings.HasPrefix(h, cookieName+"=") {
|
||||
parts := strings.Split(h, ";")
|
||||
newParts := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
trimmed := strings.TrimSpace(p)
|
||||
lower := strings.ToLower(trimmed)
|
||||
if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") {
|
||||
continue
|
||||
}
|
||||
newParts = append(newParts, p)
|
||||
}
|
||||
newHeaders = append(newHeaders, strings.Join(newParts, ";"))
|
||||
} else {
|
||||
newHeaders = append(newHeaders, h)
|
||||
}
|
||||
}
|
||||
header["Set-Cookie"] = newHeaders
|
||||
}
|
||||
|
||||
// EnsureSessionToken returns or generates the session unique token.
|
||||
func (s *SessionService) EnsureSessionToken(session sessions.Session) (string, bool) {
|
||||
token, ok := session.Get(consts.SessionTokenKey).(string)
|
||||
if !ok || token == "" {
|
||||
token = uuid.NewString()
|
||||
session.Set(consts.SessionTokenKey, token)
|
||||
return token, true
|
||||
}
|
||||
return token, false
|
||||
}
|
||||
|
||||
// HashSessionToken hashes the session token using SHA-256.
|
||||
func (s *SessionService) HashSessionToken(token string) string {
|
||||
h := sha256.New()
|
||||
h.Write([]byte(token))
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
// RotateSessionID forces session ID rotation to prevent session fixation attacks.
|
||||
func (s *SessionService) RotateSessionID(session sessions.Session) {
|
||||
if inner, ok := session.(interface{ Session() *gsessions.Session }); ok {
|
||||
if sess := inner.Session(); sess != nil {
|
||||
sess.ID = ""
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// CalculateSessionMaxAge dynamically calculates max age and whether it's a browser-session cookie.
|
||||
func (s *SessionService) CalculateSessionMaxAge(ctx context.Context) (int, bool) {
|
||||
cfg := s.Config()
|
||||
maxAge := cfg.SessionAge
|
||||
isSessionCookie := false
|
||||
|
||||
if s.dao != nil {
|
||||
val, err := s.dao.GetSystemConfigValue(ctx, "login_session_ttl_hours")
|
||||
if err == nil && val != "" {
|
||||
if ttlHours, err := strconv.Atoi(val); err == nil {
|
||||
switch {
|
||||
case ttlHours == -1:
|
||||
// 永不过期,设置为 10 年
|
||||
maxAge = 10 * 365 * 24 * 3600
|
||||
case ttlHours > 0:
|
||||
maxAge = ttlHours * 3600
|
||||
case ttlHours == 0:
|
||||
isSessionCookie = true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return maxAge, isSessionCookie
|
||||
}
|
||||
|
||||
// ApplyLoginSession writes the authenticated user into a freshly rotated session.
|
||||
func (s *SessionService) ApplyLoginSession(ctx context.Context, session sessions.Session, user *contracts.UserDTO, extras ...map[string]any) (bool, error) {
|
||||
session.Clear()
|
||||
s.RotateSessionID(session)
|
||||
|
||||
session.Set(consts.UserIDKey, strconv.FormatUint(user.ID, 10))
|
||||
session.Set(consts.UserNameKey, user.Username)
|
||||
if len(extras) > 0 {
|
||||
for key, value := range extras[0] {
|
||||
session.Set(key, value)
|
||||
}
|
||||
}
|
||||
|
||||
maxAge, isSessionCookie := s.CalculateSessionMaxAge(ctx)
|
||||
session.Options(s.GetSessionOptions(maxAge))
|
||||
|
||||
if err := session.Save(); err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
return isSessionCookie, nil
|
||||
}
|
||||
Reference in New Issue
Block a user