refactor(auth): modularize auth plugin with physical subpackages and decoupled services

This commit is contained in:
ryan
2026-09-03 09:12:44 +08:00
parent 2124bce7ca
commit 4407589b62
51 changed files with 3859 additions and 2915 deletions
@@ -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
}