mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-05 15:26:36 +08:00
refactor(plugins): restructure admin and message_gateway into standard layered sub-packages
This commit is contained in:
@@ -16,8 +16,8 @@ import (
|
||||
)
|
||||
|
||||
func isOIDCLoginEnabled(ctx context.Context) bool {
|
||||
var val string
|
||||
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "oidc_login_enabled").Pluck("value", &val).Error; err != nil || val == "" {
|
||||
val, err := GetSystemConfigValue(ctx, "oidc_login_enabled")
|
||||
if err != nil || val == "" {
|
||||
return true
|
||||
}
|
||||
b, err := strconv.ParseBool(val)
|
||||
@@ -75,8 +75,8 @@ func activeLoginSources(ctx context.Context) []AuthSourceView {
|
||||
}
|
||||
|
||||
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
|
||||
var val string
|
||||
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || strings.TrimSpace(val) == "" {
|
||||
val, err := GetSystemConfigValue(ctx, "server_address")
|
||||
if err != nil || strings.TrimSpace(val) == "" {
|
||||
return "", errors.New(errServerAddressMissing)
|
||||
}
|
||||
return strings.TrimRight(val, "/") + "/login", nil
|
||||
|
||||
@@ -36,3 +36,19 @@ const (
|
||||
errBannedAccount = "账号已被封禁"
|
||||
errUnAuthorized = "未登录"
|
||||
)
|
||||
|
||||
// Service 层与鉴权中间件内部错误文案(保持与重构前逐字一致)
|
||||
const (
|
||||
errUserNotInContext = "auth: user not found in context"
|
||||
errEmptyToken = "auth: empty token" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errSystemUserTokenNotAllowed = "auth: system user token not allowed" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errUnauthorizedInternal = "unauthorized"
|
||||
errSystemUserLoginNotAllowed = "system user is not allowed to login"
|
||||
)
|
||||
|
||||
// OAuth 回调会话校验错误文案(保持与重构前逐字一致)
|
||||
const (
|
||||
errInvalidSessionContext = "invalid session context"
|
||||
errSessionMismatchForOAuth = "session mismatch for oauth state"
|
||||
errUserContextMismatch = "user context mismatch for oauth binding"
|
||||
)
|
||||
|
||||
@@ -9,7 +9,6 @@ import (
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
@@ -235,17 +234,17 @@ func Callback(c *gin.Context) {
|
||||
|
||||
token, ok := session.Get(SessionTokenKey).(string)
|
||||
if !ok || token == "" {
|
||||
response.AbortBadRequest(c, "invalid session context")
|
||||
response.AbortBadRequest(c, errInvalidSessionContext)
|
||||
return
|
||||
}
|
||||
|
||||
if hashSessionToken(token) != payload.SessionHash {
|
||||
response.AbortBadRequest(c, "session mismatch for oauth state")
|
||||
response.AbortBadRequest(c, errSessionMismatchForOAuth)
|
||||
return
|
||||
}
|
||||
|
||||
if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID {
|
||||
response.AbortBadRequest(c, "user context mismatch for oauth binding")
|
||||
response.AbortBadRequest(c, errUserContextMismatch)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -298,8 +297,8 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource,
|
||||
response.AbortUnauthorized(c, errUnAuthorized)
|
||||
return
|
||||
}
|
||||
var user contracts.UserDTO
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
|
||||
user, err := GetUserByID(ctx, userID)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
@@ -314,41 +313,43 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource,
|
||||
return
|
||||
}
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = getDB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
|
||||
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
|
||||
_ = TouchUserLastLogin(ctx, user.ID, user.LastLoginAt)
|
||||
c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "bound")))
|
||||
}
|
||||
|
||||
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
|
||||
var user contracts.UserDTO
|
||||
var user *contracts.UserDTO
|
||||
|
||||
account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub)
|
||||
switch {
|
||||
case err == nil:
|
||||
if loadErr := getDB(ctx).Table("w_users").Where("id = ?", account.UserID).First(&user).Error; loadErr != nil {
|
||||
loaded, loadErr := GetUserByID(ctx, account.UserID)
|
||||
if loadErr != nil {
|
||||
response.AbortInternal(c, loadErr.Error())
|
||||
return
|
||||
}
|
||||
user = loaded
|
||||
case errors.Is(err, gorm.ErrRecordNotFound):
|
||||
newUser, ok := handleCallbackRegister(ctx, c, source, userInfo)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
user = newUser
|
||||
user = &newUser
|
||||
default:
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = getDB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
|
||||
if err := SetLoginSession(ctx, c, &user); err != nil {
|
||||
_ = TouchUserLastLogin(ctx, user.ID, user.LastLoginAt)
|
||||
if err := SetLoginSession(ctx, c, user); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
SetCachedUser(ctx, user.ID, &user)
|
||||
SetCachedUser(ctx, user.ID, user)
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in")))
|
||||
c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "logged_in")))
|
||||
}
|
||||
|
||||
func uniqueUsername(ctx context.Context, base string) (string, error) {
|
||||
@@ -357,10 +358,8 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
|
||||
base = "user"
|
||||
}
|
||||
|
||||
var existingUsernames []string
|
||||
if err := getDB(ctx).Table("w_users").
|
||||
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
|
||||
Pluck("username", &existingUsernames).Error; err != nil {
|
||||
existingUsernames, err := ListSimilarUsernames(ctx, base)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
@@ -385,8 +384,8 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
|
||||
|
||||
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) (contracts.UserDTO, bool) {
|
||||
registrationEnabled := true
|
||||
var val string
|
||||
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "registration_enabled").Pluck("value", &val).Error; err == nil && val != "" {
|
||||
val, cfgErr := GetSystemConfigValue(ctx, "registration_enabled")
|
||||
if cfgErr == nil && val != "" {
|
||||
if b, err := strconv.ParseBool(val); err == nil {
|
||||
registrationEnabled = b
|
||||
}
|
||||
@@ -417,7 +416,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSou
|
||||
UpdatedAt: now,
|
||||
}
|
||||
|
||||
if err := getDB(ctx).Table("w_users").Create(&user).Error; err != nil {
|
||||
if err := InsertUser(ctx, &user); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return contracts.UserDTO{}, false
|
||||
}
|
||||
|
||||
@@ -22,33 +22,35 @@ func hashToken(token string) string {
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
// currentUserIDFromRequestContext 是接入层向 Service 层暴露的登录态桥接。
|
||||
//
|
||||
// Session 读取必须依赖 *gin.Context,而 Service 层禁止 import gin,
|
||||
// 因此该类型断言收敛在本(接入层)文件中。ok 为 false 表示 ctx 不是 *gin.Context。
|
||||
func currentUserIDFromRequestContext(ctx context.Context) (uint64, bool) {
|
||||
ginCtx, ok := ctx.(*gin.Context)
|
||||
if !ok {
|
||||
return 0, false
|
||||
}
|
||||
return GetUserIDFromContext(ginCtx), true
|
||||
}
|
||||
|
||||
func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *CachedToken, error) {
|
||||
tokenHash := hashToken(tokenStr)
|
||||
tokenRecord, err := GetCachedToken(ctx, tokenHash)
|
||||
if err != nil || tokenRecord == nil {
|
||||
var tokenRow struct {
|
||||
ID uint64
|
||||
UserID uint64
|
||||
IsAdmin bool
|
||||
}
|
||||
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
|
||||
tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
tokenRecord = &CachedToken{
|
||||
ID: tokenRow.ID,
|
||||
UserID: tokenRow.UserID,
|
||||
IsAdmin: tokenRow.IsAdmin,
|
||||
}
|
||||
SetCachedToken(ctx, tokenHash, tokenRecord)
|
||||
}
|
||||
|
||||
user, err := GetCachedUser(ctx, tokenRecord.UserID)
|
||||
if err != nil || user == nil || !user.IsActive {
|
||||
var userRow contracts.UserDTO
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&userRow).Error; err != nil {
|
||||
user, err = GetActiveUserByID(ctx, tokenRecord.UserID)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
user = &userRow
|
||||
SetCachedUser(ctx, tokenRecord.UserID, user)
|
||||
}
|
||||
|
||||
@@ -74,7 +76,7 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
|
||||
if tokenStr != "" {
|
||||
if user, tokenRecord, err := getUserByToken(ctx, tokenStr); err == nil {
|
||||
if user.Username == SystemUsername {
|
||||
return nil, errors.New("system user is not allowed to login")
|
||||
return nil, errors.New(errSystemUserLoginNotAllowed)
|
||||
}
|
||||
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, true)
|
||||
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
|
||||
@@ -85,16 +87,15 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
|
||||
// 降级使用 Session 鉴权
|
||||
userID := GetUserIDFromContext(c)
|
||||
if userID <= 0 {
|
||||
return nil, errors.New("unauthorized")
|
||||
return nil, errors.New(errUnauthorizedInternal)
|
||||
}
|
||||
|
||||
user, err := GetCachedUser(ctx, userID)
|
||||
if err != nil || user == nil || !user.IsActive {
|
||||
var dbUser contracts.UserDTO
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&dbUser).Error; err != nil {
|
||||
user, err = GetActiveUserByID(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
user = &dbUser
|
||||
SetCachedUser(ctx, userID, user)
|
||||
}
|
||||
|
||||
@@ -102,7 +103,7 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
|
||||
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, false)
|
||||
|
||||
if user.Username == "system" {
|
||||
return nil, errors.New("system user is not allowed to login")
|
||||
return nil, errors.New(errSystemUserLoginNotAllowed)
|
||||
}
|
||||
|
||||
return user, nil
|
||||
|
||||
@@ -6,8 +6,10 @@ package auth
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/util"
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
@@ -58,6 +60,80 @@ func getCache(ctx context.Context) contracts.CacheService {
|
||||
return s
|
||||
}
|
||||
|
||||
// GetAccessTokenByHash 按令牌哈希读取访问令牌记录(仅取鉴权所需字段)
|
||||
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*CachedToken, error) {
|
||||
var row struct {
|
||||
ID uint64
|
||||
UserID uint64
|
||||
IsAdmin bool
|
||||
}
|
||||
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&row).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &CachedToken{
|
||||
ID: row.ID,
|
||||
UserID: row.UserID,
|
||||
IsAdmin: row.IsAdmin,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetActiveUserByID 读取仍处于启用状态的用户
|
||||
func GetActiveUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
|
||||
var user contracts.UserDTO
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&user).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
// GetUserByID 按 ID 读取用户(不限制启用状态)
|
||||
func GetUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
|
||||
var user contracts.UserDTO
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
// InsertUser 新建用户记录
|
||||
func InsertUser(ctx context.Context, user *contracts.UserDTO) error {
|
||||
return getDB(ctx).Table("w_users").Create(user).Error
|
||||
}
|
||||
|
||||
// TouchUserLastLogin 刷新用户最后登录时间
|
||||
func TouchUserLastLogin(ctx context.Context, userID uint64, at time.Time) error {
|
||||
return getDB(ctx).Table("w_users").Where("id = ?", userID).Update("last_login_at", at).Error
|
||||
}
|
||||
|
||||
// ListSimilarUsernames 查询与基础用户名相同或带 `-序号` 后缀的用户名(用于用户名去重)
|
||||
func ListSimilarUsernames(ctx context.Context, base string) ([]string, error) {
|
||||
var existingUsernames []string
|
||||
if err := getDB(ctx).Table("w_users").
|
||||
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
|
||||
Pluck("username", &existingUsernames).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return existingUsernames, nil
|
||||
}
|
||||
|
||||
// GetSystemConfigValue 读取系统配置项原始值
|
||||
func GetSystemConfigValue(ctx context.Context, key string) (string, error) {
|
||||
var val string
|
||||
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", key).Pluck("value", &val).Error; err != nil {
|
||||
return "", err
|
||||
}
|
||||
return val, nil
|
||||
}
|
||||
|
||||
// ListAllAuthSources 获取全部认证源(含未启用),按 ID 升序
|
||||
func ListAllAuthSources(ctx context.Context) ([]AuthSource, error) {
|
||||
var sources []AuthSource
|
||||
if err := getDB(ctx).Order("id ASC").Find(&sources).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sources, nil
|
||||
}
|
||||
|
||||
// GetAuthSourceByID 根据 ID 获取认证源
|
||||
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
|
||||
var src AuthSource
|
||||
@@ -85,6 +161,21 @@ func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
|
||||
return sources, nil
|
||||
}
|
||||
|
||||
// CreateAuthSourceRecord 新建认证源记录
|
||||
func CreateAuthSourceRecord(ctx context.Context, source *AuthSource) error {
|
||||
return getDB(ctx).Create(source).Error
|
||||
}
|
||||
|
||||
// SaveAuthSourceRecord 全量保存认证源记录
|
||||
func SaveAuthSourceRecord(ctx context.Context, source *AuthSource) error {
|
||||
return getDB(ctx).Save(source).Error
|
||||
}
|
||||
|
||||
// DeleteAuthSourceRecord 删除认证源记录
|
||||
func DeleteAuthSourceRecord(ctx context.Context, source *AuthSource) error {
|
||||
return getDB(ctx).Delete(source).Error
|
||||
}
|
||||
|
||||
// GetActiveAuthSourcesCached 获取所有启用的认证源(带缓存或直接查询)
|
||||
func GetActiveAuthSourcesCached(ctx context.Context) ([]AuthSource, error) {
|
||||
return ListActiveAuthSources(ctx)
|
||||
|
||||
@@ -5,12 +5,9 @@ package auth
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type authServiceImpl struct{}
|
||||
@@ -27,58 +24,48 @@ func (s *authServiceImpl) RequireAdminMiddleware() any {
|
||||
return AdminRequired()
|
||||
}
|
||||
|
||||
// GetCurrentUser 从 context 中读取登录用户。
|
||||
//
|
||||
// 中间件通过 gin 的 c.Set(contracts.AuthUserObjKey, user) 写入登录态;
|
||||
// *gin.Context 自身实现了 context.Context,且其 Value(key) 对 string 类型 key
|
||||
// 等价于 c.Get(key)(未命中时再回落到 Request.Context().Value),
|
||||
// 因此这里无需感知 gin 即可读取同一份登录态。
|
||||
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
|
||||
if ginCtx, ok := ctx.(*gin.Context); ok {
|
||||
if u, ok := ginutil.GetFromContext[*contracts.UserDTO](ginCtx, contracts.AuthUserObjKey); ok && u != nil {
|
||||
return u, nil
|
||||
}
|
||||
}
|
||||
|
||||
if v := ctx.Value(contracts.AuthUserObjKey); v != nil {
|
||||
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
|
||||
return u, nil
|
||||
}
|
||||
}
|
||||
|
||||
return nil, errors.New("auth: user not found in context")
|
||||
return nil, errors.New(errUserNotInContext)
|
||||
}
|
||||
|
||||
func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) {
|
||||
if token == "" {
|
||||
return nil, errors.New("auth: empty token")
|
||||
return nil, errors.New(errEmptyToken)
|
||||
}
|
||||
|
||||
tokenHash := hashToken(token)
|
||||
tokenRecord, err := GetCachedToken(ctx, tokenHash)
|
||||
if err != nil {
|
||||
var tokenRow struct {
|
||||
ID uint64
|
||||
UserID uint64
|
||||
IsAdmin bool
|
||||
}
|
||||
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
|
||||
tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tokenRecord = &CachedToken{
|
||||
ID: tokenRow.ID,
|
||||
UserID: tokenRow.UserID,
|
||||
IsAdmin: tokenRow.IsAdmin,
|
||||
}
|
||||
SetCachedToken(ctx, tokenHash, tokenRecord)
|
||||
}
|
||||
|
||||
user, err := GetCachedUser(ctx, tokenRecord.UserID)
|
||||
if err != nil || user == nil || !user.IsActive {
|
||||
var dbUser contracts.UserDTO
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil {
|
||||
user, err = GetActiveUserByID(ctx, tokenRecord.UserID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
user = &dbUser
|
||||
SetCachedUser(ctx, tokenRecord.UserID, user)
|
||||
}
|
||||
|
||||
if user.Username == SystemUsername {
|
||||
return nil, errors.New("auth: system user token not allowed")
|
||||
return nil, errors.New(errSystemUserTokenNotAllowed)
|
||||
}
|
||||
|
||||
return user, nil
|
||||
@@ -93,11 +80,16 @@ func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64)
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetCurrentUserID 从请求登录态中读取用户 ID。
|
||||
//
|
||||
// Session 读取依赖 gin,属于接入层职责,因此这里通过接入层桥接函数
|
||||
// currentUserIDFromRequestContext(见 middleware.go)取值,Service 层本身不感知 gin。
|
||||
func (s *authServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) {
|
||||
if ginCtx, ok := ctx.(*gin.Context); ok {
|
||||
return GetUserIDFromContext(ginCtx), nil
|
||||
userID, ok := currentUserIDFromRequestContext(ctx)
|
||||
if !ok {
|
||||
return 0, errors.New(errUserNotInContext)
|
||||
}
|
||||
return 0, errors.New("auth: user not found in context")
|
||||
return userID, nil
|
||||
}
|
||||
|
||||
func (s *authServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error {
|
||||
@@ -118,8 +110,8 @@ func (s *authServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash s
|
||||
}
|
||||
|
||||
func (s *authServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
|
||||
var sources []AuthSource
|
||||
if err := getDB(ctx).Order("id ASC").Find(&sources).Error; err != nil {
|
||||
sources, err := ListAllAuthSources(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -156,7 +148,7 @@ func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := getDB(ctx).Create(&model).Error; err != nil {
|
||||
if err := CreateAuthSourceRecord(ctx, &model); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -165,8 +157,8 @@ func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts
|
||||
}
|
||||
|
||||
func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
|
||||
var existing AuthSource
|
||||
if err := getDB(ctx).First(&existing, id).Error; err != nil {
|
||||
existing, err := GetAuthSourceByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -183,36 +175,36 @@ func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, sourc
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := getDB(ctx).Save(&existing).Error; err != nil {
|
||||
if err := SaveAuthSourceRecord(ctx, existing); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
existing.Sanitize()
|
||||
return toAuthSourceDTO(&existing), nil
|
||||
return toAuthSourceDTO(existing), nil
|
||||
}
|
||||
|
||||
func (s *authServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
|
||||
var existing AuthSource
|
||||
if err := getDB(ctx).First(&existing, id).Error; err != nil {
|
||||
existing, err := GetAuthSourceByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return getDB(ctx).Delete(&existing).Error
|
||||
return DeleteAuthSourceRecord(ctx, existing)
|
||||
}
|
||||
|
||||
func (s *authServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
|
||||
var existing AuthSource
|
||||
if err := getDB(ctx).First(&existing, id).Error; err != nil {
|
||||
existing, err := GetAuthSourceByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
existing.IsActive = !existing.IsActive
|
||||
if err := getDB(ctx).Save(&existing).Error; err != nil {
|
||||
if err := SaveAuthSourceRecord(ctx, existing); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
existing.Sanitize()
|
||||
return toAuthSourceDTO(&existing), nil
|
||||
return toAuthSourceDTO(existing), nil
|
||||
}
|
||||
|
||||
func toAuthSourceDTO(s *AuthSource) *contracts.AuthSourceDTO {
|
||||
|
||||
@@ -0,0 +1,367 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package auth_test
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/auth"
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
testSessionCookieName = "auth-test-session"
|
||||
errUserNotInContext = "auth: user not found in context"
|
||||
)
|
||||
|
||||
// newTestAuthService 装配仅注册 auth 插件的 core.Context,并返回其对外契约实现。
|
||||
func newTestAuthService(t *testing.T, db *gorm.DB) contracts.AuthService {
|
||||
t.Helper()
|
||||
|
||||
ctx := core.NewContext(context.Background())
|
||||
if db != nil {
|
||||
core.Provide[contracts.DBService](ctx, &mockDBService{db: db})
|
||||
core.Provide[contracts.CacheService](ctx, newMockCacheService())
|
||||
}
|
||||
require.NoError(t, auth.New().Apply(ctx))
|
||||
|
||||
svc, err := core.Inject[contracts.AuthService](ctx)
|
||||
require.NoError(t, err)
|
||||
auth.ResetAuthRAMCacheForTest()
|
||||
|
||||
return svc
|
||||
}
|
||||
|
||||
// newSessionEngine 构造一个带 Session 中间件的 gin 引擎,用于走通真实登录态链路。
|
||||
//
|
||||
// response.Abort* 只把错误挂载到 gin 错误链,状态码由全局错误中间件渲染,
|
||||
// 因此这里必须同时装配 response.ErrorHandlerMiddleware()。
|
||||
func newSessionEngine() *gin.Engine {
|
||||
engine := gin.New()
|
||||
engine.Use(response.ErrorHandlerMiddleware())
|
||||
engine.Use(sessions.Sessions(testSessionCookieName, cookie.NewStore([]byte("test-secret"))))
|
||||
return engine
|
||||
}
|
||||
|
||||
func TestGetCurrentUserFromGinContext(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svc := newTestAuthService(t, nil)
|
||||
user := &contracts.UserDTO{ID: 4242, Username: "ctx_user", IsActive: true}
|
||||
|
||||
t.Run("gin 上下文已由中间件写入用户时返回该用户", func(t *testing.T) {
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/user-info", nil)
|
||||
c.Set(contracts.AuthUserObjKey, user)
|
||||
|
||||
got, err := svc.GetCurrentUser(c)
|
||||
require.NoError(t, err)
|
||||
assert.Same(t, user, got)
|
||||
})
|
||||
|
||||
t.Run("开启 ContextWithFallback 时可从请求 context 回落读取", func(t *testing.T) {
|
||||
// 说明:本项目引擎默认不开启 ContextWithFallback,此时 (*gin.Context).Value
|
||||
// 等价于 c.Get,与改造前 ginutil.GetFromContext 的读取路径完全一致;
|
||||
// 开启回落后还能额外读到写入 Request.Context() 的登录态。
|
||||
reqCtx := context.WithValue(context.Background(), contracts.AuthUserObjKey, user) //nolint:staticcheck // 模拟写入请求 context 的登录态
|
||||
engine := gin.New()
|
||||
engine.ContextWithFallback = true
|
||||
var (
|
||||
gotUser *contracts.UserDTO
|
||||
gotErr error
|
||||
)
|
||||
engine.GET("/probe", func(c *gin.Context) {
|
||||
gotUser, gotErr = svc.GetCurrentUser(c)
|
||||
c.Status(http.StatusNoContent)
|
||||
})
|
||||
|
||||
engine.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/probe", nil).WithContext(reqCtx))
|
||||
require.NoError(t, gotErr)
|
||||
assert.Same(t, user, gotUser)
|
||||
})
|
||||
|
||||
t.Run("未登录时报错且文案不变", func(t *testing.T) {
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/user-info", nil)
|
||||
|
||||
got, err := svc.GetCurrentUser(c)
|
||||
require.Error(t, err)
|
||||
assert.Nil(t, got)
|
||||
assert.Equal(t, errUserNotInContext, err.Error())
|
||||
})
|
||||
|
||||
t.Run("非 gin 的普通 context 仍按 Value 取值", func(t *testing.T) {
|
||||
got, err := svc.GetCurrentUser(context.WithValue(context.Background(), contracts.AuthUserObjKey, user)) //nolint:staticcheck // 与中间件写入的 key 语义一致
|
||||
require.NoError(t, err)
|
||||
assert.Same(t, user, got)
|
||||
|
||||
_, err = svc.GetCurrentUser(context.Background())
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, errUserNotInContext, err.Error())
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetCurrentUserID(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svc := newTestAuthService(t, nil)
|
||||
|
||||
t.Run("gin Session 中的用户 ID 可正常读取", func(t *testing.T) {
|
||||
engine := newSessionEngine()
|
||||
var (
|
||||
gotUID uint64
|
||||
gotErr error
|
||||
)
|
||||
engine.GET("/probe", func(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
session.Set(auth.UserIDKey, uint64(777))
|
||||
require.NoError(t, session.Save())
|
||||
|
||||
gotUID, gotErr = svc.GetCurrentUserID(c)
|
||||
c.Status(http.StatusNoContent)
|
||||
})
|
||||
|
||||
engine.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/probe", nil))
|
||||
require.NoError(t, gotErr)
|
||||
assert.Equal(t, uint64(777), gotUID)
|
||||
})
|
||||
|
||||
t.Run("gin 上下文存在但 Session 无用户时返回 0 且不报错", func(t *testing.T) {
|
||||
engine := newSessionEngine()
|
||||
var (
|
||||
gotUID uint64
|
||||
gotErr error
|
||||
)
|
||||
engine.GET("/probe", func(c *gin.Context) {
|
||||
gotUID, gotErr = svc.GetCurrentUserID(c)
|
||||
c.Status(http.StatusNoContent)
|
||||
})
|
||||
|
||||
engine.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/probe", nil))
|
||||
require.NoError(t, gotErr)
|
||||
assert.Equal(t, uint64(0), gotUID)
|
||||
})
|
||||
|
||||
t.Run("非 gin context 报错且文案不变", func(t *testing.T) {
|
||||
// 即使普通 context 中已写入用户对象,该方法的 Session 语义也保持不变。
|
||||
uid, err := svc.GetCurrentUserID(
|
||||
context.WithValue(context.Background(), contracts.AuthUserObjKey, &contracts.UserDTO{ID: 1}), //nolint:staticcheck // 同上
|
||||
)
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, uint64(0), uid)
|
||||
assert.Equal(t, errUserNotInContext, err.Error())
|
||||
})
|
||||
}
|
||||
|
||||
func TestLoginRequiredMiddlewarePopulatesServiceContext(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
db := setupTestDB(t)
|
||||
require.NoError(t, db.Create(&testUser{ID: 9001, Username: "session_user", IsActive: true}).Error)
|
||||
require.NoError(t, db.Create(&testUser{ID: 9002, Username: "token_user", IsActive: true}).Error)
|
||||
|
||||
tokenStr := "integration-secret-token"
|
||||
require.NoError(t, db.Create(&testAccessToken{
|
||||
ID: 9101,
|
||||
UserID: 9002,
|
||||
TokenHash: hashToken(tokenStr),
|
||||
Name: "integration",
|
||||
IsAdmin: false,
|
||||
}).Error)
|
||||
|
||||
svc := newTestAuthService(t, db)
|
||||
|
||||
t.Run("Session 鉴权链路上 GetCurrentUser 与 GetCurrentUserID 一致", func(t *testing.T) {
|
||||
engine := newSessionEngine()
|
||||
engine.Use(func(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
session.Set(auth.UserIDKey, uint64(9001))
|
||||
require.NoError(t, session.Save())
|
||||
c.Next()
|
||||
})
|
||||
|
||||
var (
|
||||
gotUser *contracts.UserDTO
|
||||
userErr error
|
||||
gotUID uint64
|
||||
uidErr error
|
||||
)
|
||||
engine.GET("/protected", auth.LoginRequired(), func(c *gin.Context) {
|
||||
gotUser, userErr = svc.GetCurrentUser(c)
|
||||
gotUID, uidErr = svc.GetCurrentUserID(c)
|
||||
c.Status(http.StatusNoContent)
|
||||
})
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
engine.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/protected", nil))
|
||||
require.Equal(t, http.StatusNoContent, recorder.Code)
|
||||
|
||||
require.NoError(t, userErr)
|
||||
require.NotNil(t, gotUser)
|
||||
assert.Equal(t, uint64(9001), gotUser.ID)
|
||||
assert.Equal(t, "session_user", gotUser.Username)
|
||||
|
||||
require.NoError(t, uidErr)
|
||||
assert.Equal(t, uint64(9001), gotUID)
|
||||
})
|
||||
|
||||
t.Run("Access Token 鉴权链路上 GetCurrentUser 可用", func(t *testing.T) {
|
||||
engine := newSessionEngine()
|
||||
var (
|
||||
gotUser *contracts.UserDTO
|
||||
userErr error
|
||||
)
|
||||
engine.GET("/protected", auth.LoginRequired(), func(c *gin.Context) {
|
||||
gotUser, userErr = svc.GetCurrentUser(c)
|
||||
c.Status(http.StatusNoContent)
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set("Authorization", "Bearer "+tokenStr)
|
||||
recorder := httptest.NewRecorder()
|
||||
engine.ServeHTTP(recorder, req)
|
||||
require.Equal(t, http.StatusNoContent, recorder.Code)
|
||||
|
||||
require.NoError(t, userErr)
|
||||
require.NotNil(t, gotUser)
|
||||
assert.Equal(t, uint64(9002), gotUser.ID)
|
||||
assert.Equal(t, "token_user", gotUser.Username)
|
||||
})
|
||||
|
||||
t.Run("未登录请求被中间件拒绝", func(t *testing.T) {
|
||||
engine := newSessionEngine()
|
||||
engine.GET("/protected", auth.LoginRequired(), func(c *gin.Context) {
|
||||
c.Status(http.StatusNoContent)
|
||||
})
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
engine.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/protected", nil))
|
||||
assert.Equal(t, http.StatusUnauthorized, recorder.Code)
|
||||
})
|
||||
}
|
||||
|
||||
// legacyGetCurrentUser 逐字复刻改造前 Service 层的取值实现
|
||||
// (*gin.Context 类型断言 + ginutil.GetFromContext + ctx.Value 回落),
|
||||
// 用于与新实现做 differential 等价性校验。
|
||||
func legacyGetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
|
||||
if ginCtx, ok := ctx.(*gin.Context); ok {
|
||||
if u, ok := ginutil.GetFromContext[*contracts.UserDTO](ginCtx, contracts.AuthUserObjKey); ok && u != nil {
|
||||
return u, nil
|
||||
}
|
||||
}
|
||||
|
||||
if v := ctx.Value(contracts.AuthUserObjKey); v != nil {
|
||||
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
|
||||
return u, nil
|
||||
}
|
||||
}
|
||||
|
||||
return nil, errors.New(errUserNotInContext)
|
||||
}
|
||||
|
||||
// legacyGetCurrentUserID 逐字复刻改造前 Service 层基于 gin Session 的实现。
|
||||
func legacyGetCurrentUserID(ctx context.Context) (uint64, error) {
|
||||
if ginCtx, ok := ctx.(*gin.Context); ok {
|
||||
return auth.GetUserIDFromContext(ginCtx), nil
|
||||
}
|
||||
|
||||
return 0, errors.New(errUserNotInContext)
|
||||
}
|
||||
|
||||
func errText(err error) string {
|
||||
if err == nil {
|
||||
return ""
|
||||
}
|
||||
return err.Error()
|
||||
}
|
||||
|
||||
// assertLoginStateParity 断言新实现与改造前实现在同一 ctx 上返回完全一致的结果与错误文案。
|
||||
func assertLoginStateParity(t *testing.T, svc contracts.AuthService, ctx context.Context) {
|
||||
t.Helper()
|
||||
|
||||
wantUser, wantUserErr := legacyGetCurrentUser(ctx)
|
||||
gotUser, gotUserErr := svc.GetCurrentUser(ctx)
|
||||
if (wantUser == nil) != (gotUser == nil) {
|
||||
t.Fatalf("GetCurrentUser nil-ness mismatch: want %v, got %v", wantUser, gotUser)
|
||||
}
|
||||
if wantUser != nil {
|
||||
assert.Same(t, wantUser, gotUser)
|
||||
}
|
||||
assert.Equal(t, errText(wantUserErr), errText(gotUserErr))
|
||||
|
||||
wantUID, wantUIDErr := legacyGetCurrentUserID(ctx)
|
||||
gotUID, gotUIDErr := svc.GetCurrentUserID(ctx)
|
||||
assert.Equal(t, wantUID, gotUID)
|
||||
assert.Equal(t, errText(wantUIDErr), errText(gotUIDErr))
|
||||
}
|
||||
|
||||
func TestLoginStateContextParityWithLegacyImplementation(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svc := newTestAuthService(t, nil)
|
||||
user := &contracts.UserDTO{ID: 5150, Username: "parity_user", IsActive: true}
|
||||
|
||||
t.Run("gin 上下文各分支", func(t *testing.T) {
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/user-info", nil)
|
||||
assertLoginStateParity(t, svc, c)
|
||||
|
||||
c.Set(contracts.AuthUserObjKey, user)
|
||||
assertLoginStateParity(t, svc, c)
|
||||
|
||||
c.Set(contracts.AuthUserObjKey, "not-a-user-dto")
|
||||
assertLoginStateParity(t, svc, c)
|
||||
|
||||
var typedNil *contracts.UserDTO
|
||||
c.Set(contracts.AuthUserObjKey, typedNil)
|
||||
assertLoginStateParity(t, svc, c)
|
||||
})
|
||||
|
||||
t.Run("普通 context 各分支", func(t *testing.T) {
|
||||
assertLoginStateParity(t, svc, context.Background())
|
||||
assertLoginStateParity(t, svc, context.WithValue(context.Background(), contracts.AuthUserObjKey, user))
|
||||
assertLoginStateParity(t, svc, context.WithValue(context.Background(), contracts.AuthUserObjKey, "nope"))
|
||||
})
|
||||
|
||||
t.Run("Session 登录态各分支", func(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
userID any
|
||||
}{
|
||||
{name: "无用户", userID: nil},
|
||||
{name: "uint64 用户 ID", userID: uint64(3301)},
|
||||
{name: "float64 用户 ID", userID: float64(3302)},
|
||||
{name: "string 用户 ID", userID: "3303"},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
engine := newSessionEngine()
|
||||
engine.GET("/probe", func(c *gin.Context) {
|
||||
if tc.userID != nil {
|
||||
session := sessions.Default(c)
|
||||
session.Set(auth.UserIDKey, tc.userID)
|
||||
require.NoError(t, session.Save())
|
||||
}
|
||||
assertLoginStateParity(t, svc, c)
|
||||
})
|
||||
|
||||
engine.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/probe", nil))
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -116,8 +116,8 @@ func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDT
|
||||
maxAge := config.Config.App.SessionAge
|
||||
isSessionCookie := false
|
||||
|
||||
var val string
|
||||
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "login_session_ttl_hours").Pluck("value", &val).Error; err == nil && val != "" {
|
||||
val, err := GetSystemConfigValue(ctx, "login_session_ttl_hours")
|
||||
if err == nil && val != "" {
|
||||
if ttlHours, err := strconv.Atoi(val); err == nil {
|
||||
switch {
|
||||
case ttlHours == -1:
|
||||
|
||||
Reference in New Issue
Block a user