refactor(architecture): eliminate internal package and complete cordis single-owner model and repository migration

- Physically purged all legacy internal/ packages, centralized pkg/model/ and pkg/repository/
- Migrated domain models and database repositories into self-contained owner plugins (user, auth, message_gateway, admin, upload, risk_control)
- Decoupled cross-plugin interactions via pure core/contracts and typed EventBus
- Ensured 100% test coverage pass, zero data races (-race clean), and 0 lint issues in make code-check
This commit is contained in:
ryan
2026-08-28 08:40:43 +08:00
parent 1f348fd425
commit fb6a3edb89
323 changed files with 8222 additions and 17693 deletions
+2 -2
View File
@@ -8,13 +8,13 @@ import (
"context"
"encoding/json"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
)
// LogForAudit 将登录鉴权审计日志写入 Logger
func LogForAudit(ctx context.Context, user *model.User, c *gin.Context) {
func LogForAudit(ctx context.Context, user *contracts.UserDTO, c *gin.Context) {
if user == nil || c == nil {
return
}
+34 -50
View File
@@ -7,44 +7,58 @@ import (
"context"
"errors"
"fmt"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
)
func isOIDCLoginEnabled(ctx context.Context) bool {
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "oidc_login_enabled").Pluck("value", &val).Error; err != nil || val == "" {
return true
}
b, err := strconv.ParseBool(val)
if err != nil {
return true
}
return enabled
return b
}
func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) {
func resolveAuthSource(ctx context.Context, sourceName string) (*AuthSource, error) {
name := strings.TrimSpace(strings.ToLower(sourceName))
if name == "" {
sources, err := repository.GetActiveAuthSourcesCached(ctx)
sources, err := GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil, err
}
if len(sources) == 0 {
return nil, errors.New(errNoActiveAuthSource)
}
return repository.GetAuthSourceByNameCached(ctx, sources[0].Name)
src, err := GetAuthSourceByNameCached(ctx, sources[0].Name)
if err != nil {
return nil, err
}
return src, nil
}
return repository.GetAuthSourceByNameCached(ctx, name)
src, err := GetAuthSourceByNameCached(ctx, name)
if err != nil {
return nil, err
}
return src, nil
}
func activeLoginSources(ctx context.Context) []AuthSourceView {
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
if err == nil && !enabled {
if !isOIDCLoginEnabled(ctx) {
return nil
}
dbSources, err := repository.GetActiveAuthSourcesCached(ctx)
dbSources, err := GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil
}
@@ -64,14 +78,14 @@ func activeLoginSources(ctx context.Context) []AuthSourceView {
}
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
if err != nil || strings.TrimSpace(sc.Value) == "" {
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || strings.TrimSpace(val) == "" {
return "", errors.New(errServerAddressMissing)
}
return strings.TrimRight(sc.Value, "/") + "/login", nil
return strings.TrimRight(val, "/") + "/login", nil
}
func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
func buildOAuthConfig(ctx context.Context, source *AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
if source == nil {
return nil, nil, errors.New(errAuthSourceRequired)
}
@@ -116,37 +130,7 @@ func containsScope(scopes []string, scope string) bool {
return false
}
func uniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = "user"
}
existingUsernames, err := repository.ListUsernamesMatchingBase(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(errUsernameGenerateFailed)
}
func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) {
func buildOAuthUserInfo(ctx context.Context, source *AuthSource, code string, nonce string, redirectURL string) (*contracts.OAuthUserInfoDTO, error) {
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return nil, err
@@ -157,7 +141,7 @@ func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code stri
return nil, err
}
userInfo := &model.OAuthUserInfo{Active: true}
userInfo := &contracts.OAuthUserInfoDTO{Active: true}
if verifier != nil {
if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
return nil, verifyErr
@@ -180,7 +164,7 @@ func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code stri
return userInfo, nil
}
func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *model.OAuthUserInfo) error {
func 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
@@ -198,7 +182,7 @@ func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *o
return nil
}
func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
func normalizeOAuthUserInfo(userInfo *contracts.OAuthUserInfoDTO) error {
userInfo.Username = strings.TrimSpace(userInfo.Username)
userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
userInfo.Email = strings.TrimSpace(userInfo.Email)
@@ -226,7 +210,7 @@ func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
return nil
}
func buildCallbackResult(user *model.User, status string) OAuthCallbackResult {
func buildCallbackResult(user *contracts.UserDTO, status string) OAuthCallbackResult {
result := OAuthCallbackResult{Status: status}
if user != nil {
info := BuildBasicUserInfo(user, false)
+22 -15
View File
@@ -10,9 +10,9 @@ import (
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/util"
)
@@ -25,9 +25,16 @@ const (
oauthUserInvalidationChannel = "oauth:user_invalidation"
)
// CachedToken represents the minimal cached representation of an access token.
type CachedToken struct {
ID uint64 `json:"id"`
UserID uint64 `json:"user_id"`
IsAdmin bool `json:"is_admin"`
}
var (
tokenRAM = ram.MustNew[string, *model.AccessToken](ram.Options{MaximumSize: 2048})
userRAM = ram.MustNew[uint64, *model.User](ram.Options{MaximumSize: 2048})
tokenRAM = ram.MustNew[string, *CachedToken](ram.Options{MaximumSize: 2048})
userRAM = ram.MustNew[uint64, *contracts.UserDTO](ram.Options{MaximumSize: 2048})
tokenListenerOnce sync.Once
tokenListenerCtx context.Context
@@ -136,8 +143,8 @@ func publishUserRAMInvalidation(ctx context.Context, userID uint64) {
_ = db.Redis.Publish(ctx, oauthUserInvalidationChannel, strconv.FormatUint(userID, 10)).Err()
}
// GetCachedToken 获取缓存的 AccessToken
func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken, error) {
// GetCachedToken 获取缓存的 Token
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
ensureTokenCacheListener()
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
@@ -145,7 +152,7 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken,
}
if db.Redis != nil {
var token model.AccessToken
var token CachedToken
key := tokenCacheKey(tokenHash)
if err := db.GetJSON(ctx, key, &token); err == nil {
// Write back to local cache
@@ -156,8 +163,8 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken,
return nil, fmt.Errorf("cache miss")
}
// SetCachedToken 设置 AccessToken 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *model.AccessToken) {
// SetCachedToken 设置 Token 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
ensureTokenCacheListener()
tokenRAM.Set(tokenHash, token)
@@ -179,8 +186,8 @@ func InvalidateCachedToken(ctx context.Context, tokenHash string) {
}
}
// GetCachedUser 获取缓存的 User
func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
// GetCachedUser 获取缓存的 UserDTO
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
ensureUserCacheListener()
if val, ok := userRAM.GetIfPresent(userID); ok {
@@ -188,7 +195,7 @@ func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
}
if db.Redis != nil {
var u model.User
var u contracts.UserDTO
key := userCacheKey(userID)
if err := db.GetJSON(ctx, key, &u); err == nil {
// Write back to local cache
@@ -199,8 +206,8 @@ func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
return nil, fmt.Errorf("cache miss")
}
// SetCachedUser 设置 User 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *model.User) {
// SetCachedUser 设置 UserDTO 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
ensureUserCacheListener()
userRAM.Set(userID, u)
@@ -210,7 +217,7 @@ func SetCachedUser(ctx context.Context, userID uint64, u *model.User) {
}
}
// InvalidateCachedUser 吊销/失效 User 缓存
// InvalidateCachedUser 吊销/失效 UserDTO 缓存
func InvalidateCachedUser(ctx context.Context, userID uint64) {
ensureUserCacheListener()
+8 -9
View File
@@ -11,8 +11,8 @@ import (
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
)
@@ -49,11 +49,10 @@ func TestTokenCache_GetSetInvalidate(t *testing.T) {
ctx := context.Background()
tokenHash := "test-token-hash"
token := &model.AccessToken{
ID: 123,
UserID: 456,
TokenHash: tokenHash,
Name: "test-token",
token := &auth.CachedToken{
ID: 123,
UserID: 456,
IsAdmin: true,
}
// 1. Get from empty cache -> miss
@@ -70,7 +69,7 @@ func TestTokenCache_GetSetInvalidate(t *testing.T) {
if err != nil {
t.Fatalf("GetCachedToken() failed: %v", err)
}
if cached.ID != token.ID || cached.UserID != token.UserID {
if cached.ID != token.ID || cached.UserID != token.UserID || cached.IsAdmin != token.IsAdmin {
t.Fatalf("expected cached token %+v, got %+v", token, cached)
}
@@ -90,7 +89,7 @@ func TestUserCache_GetSetInvalidate(t *testing.T) {
ctx := context.Background()
userID := uint64(789)
user := &model.User{
user := &contracts.UserDTO{
ID: userID,
Username: "testuser",
Email: "test@example.com",
+83 -38
View File
@@ -13,13 +13,16 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/listener"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/shared"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/coreos/go-oidc/v3/oidc"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
@@ -91,7 +94,7 @@ func GetLoginURL(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state string) (string, error) {
func buildAuthorizeURL(ctx context.Context, source *AuthSource, state string) (string, error) {
redirectURL, err := getFrontendLoginRedirectURL(ctx)
if err != nil {
return "", err
@@ -282,18 +285,18 @@ func Callback(c *gin.Context) {
handleCallbackLogin(ctx, c, source, userInfo)
}
func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, shared.UnAuthorized)
return
}
user, err := repository.GetUserByID(ctx, userID)
if err != nil {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{
if err := BindExternalAccount(ctx, &ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
@@ -304,22 +307,20 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthS
return
}
user.LastLoginAt = time.Now()
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
_ = db.DB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
}
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
var user model.User
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
var user contracts.UserDTO
account, err := repository.FindExternalAccount(ctx, source.ID, userInfo.Sub)
account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch {
case err == nil:
loaded, loadErr := repository.GetUserByID(ctx, account.UserID)
if loadErr != nil {
if loadErr := db.DB(ctx).Table("w_users").Where("id = ?", account.UserID).First(&user).Error; 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 {
@@ -332,7 +333,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
}
user.LastLoginAt = time.Now()
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
_ = db.DB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
if err := SetLoginSession(ctx, c, &user); err != nil {
response.AbortInternal(c, err.Error())
return
@@ -340,37 +341,81 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
SetCachedUser(ctx, user.ID, &user)
logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP())
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in")))
}
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
if regErr != nil {
registrationEnabled = false
func uniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = "user"
}
var existingUsernames []string
if err := db.DB(ctx).Table("w_users").
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
Pluck("username", &existingUsernames).Error; 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(errUsernameGenerateFailed)
}
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) (contracts.UserDTO, bool) {
registrationEnabled := true
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "registration_enabled").Pluck("value", &val).Error; err == nil && val != "" {
if b, err := strconv.ParseBool(val); err == nil {
registrationEnabled = b
}
}
if !registrationEnabled {
c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
return model.User{}, false
return contracts.UserDTO{}, false
}
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
if uniqueErr != nil {
response.AbortInternal(c, uniqueErr.Error())
return model.User{}, false
return contracts.UserDTO{}, false
}
userInfo.Username = username
var user model.User
if err := repository.CreateUserFromOAuth(ctx, &user, userInfo); err != nil {
response.AbortInternal(c, err.Error())
return model.User{}, false
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 := repository.BindExternalAccount(ctx, &model.ExternalAccount{
if err := db.DB(ctx).Table("w_users").Create(&user).Error; err != nil {
response.AbortInternal(c, err.Error())
return contracts.UserDTO{}, false
}
if err := BindExternalAccount(ctx, &ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
@@ -378,7 +423,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return model.User{}, false
return contracts.UserDTO{}, false
}
logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
@@ -387,7 +432,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A
// UserInfo 获取当前登录用户信息
func UserInfo(c *gin.Context) {
user, _ := GetFromContext[*model.User](c, UserObjKey)
user, _ := GetFromContext[*contracts.UserDTO](c, UserObjKey)
session := sessions.Default(c)
needChange := session.Get("need_change_password") == true
@@ -420,7 +465,7 @@ func Logout(c *gin.Context) {
// ListExternalAccounts 获取当前用户的外部帐号绑定列表
func ListExternalAccounts(c *gin.Context) {
userID := GetUserIDFromContext(c)
accounts, err := repository.ListExternalAccountsByUserID(c.Request.Context(), userID)
accounts, err := ListExternalAccountsByUserID(c.Request.Context(), userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -441,7 +486,7 @@ func DeleteExternalAccount(c *gin.Context) {
response.AbortBadRequest(c, errInvalidExternalAccountBindingID)
return
}
if err := repository.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil {
if err := UnbindExternalAccount(c.Request.Context(), id, userID); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
+43 -36
View File
@@ -6,14 +6,15 @@ package auth
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/shared"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
@@ -33,32 +34,47 @@ func SetToContext[T any](c *gin.Context, key string, value T) {
c.Set(key, value)
}
func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.AccessToken, error) {
tokenHash := model.HashToken(tokenStr)
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *CachedToken, error) {
tokenHash := hashToken(tokenStr)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, nil, err
if err == nil {
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err == nil && user != nil && user.IsActive {
return user, tokenRecord, nil
}
tokenRecord = &dbToken
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || !user.IsActive {
dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, nil, err
}
user = &dbUser
SetCachedUser(ctx, tokenRecord.UserID, user)
var tokenRow struct {
ID uint64
UserID uint64
IsAdmin bool
}
return user, tokenRecord, nil
if err := db.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
return nil, nil, err
}
tokenRecord = &CachedToken{
ID: tokenRow.ID,
UserID: tokenRow.UserID,
IsAdmin: tokenRow.IsAdmin,
}
SetCachedToken(ctx, tokenHash, tokenRecord)
var userRow contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRow.UserID, true).First(&userRow).Error; err != nil {
return nil, nil, err
}
SetCachedUser(ctx, userRow.ID, &userRow)
return &userRow, tokenRecord, nil
}
// GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error
func GetUserFromRequest(c *gin.Context) (*model.User, error) {
func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
ctx := c.Request.Context()
// Check token in headers
@@ -89,24 +105,15 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) {
}
user, err := GetCachedUser(ctx, userID)
if err != nil || !user.IsActive {
dbUser, loadErr := repository.GetActiveUserByID(ctx, userID)
if loadErr != nil {
return nil, loadErr
if err != nil || user == nil || !user.IsActive {
var dbUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&dbUser).Error; err != nil {
return nil, err
}
user = &dbUser
SetCachedUser(ctx, userID, user)
}
// 密码哈希校验:当用户存在本地密码时,要求 Session 中的密码哈希必须与当前数据库中一致
if user.Password != "" {
session := sessions.Default(c)
sessionHash, _ := session.Get(PasswordHashKey).(string)
if sessionHash != user.Password {
return nil, errors.New("session expired due to password change")
}
}
SetToContext(c, TokenAuthKey, false)
SetToContext(c, TokenAdminKey, false)
+3 -3
View File
@@ -12,7 +12,7 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/core/contracts"
)
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
@@ -162,8 +162,8 @@ type BasicUserInfo struct {
Location string `json:"location"`
}
// BuildBasicUserInfo 将 User 模型转换为 BasicUserInfo
func BuildBasicUserInfo(user *model.User, needChange bool) BasicUserInfo {
// BuildBasicUserInfo 将 UserDTO 转换为 BasicUserInfo
func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo {
if user == nil {
return BasicUserInfo{}
}
+37 -10
View File
@@ -5,8 +5,11 @@ package auth_test
import (
"context"
"crypto/sha256"
"encoding/hex"
"path/filepath"
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
@@ -15,11 +18,35 @@ import (
"github.com/Rain-kl/Wavelet/core"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
)
type testUser struct {
ID uint64 `gorm:"primaryKey"`
Username string
IsActive bool
LastLoginAt time.Time
}
func (testUser) TableName() string { return "w_users" }
type testAccessToken struct {
ID uint64 `gorm:"primaryKey"`
UserID uint64
TokenHash string
Name string
IsAdmin bool
}
func (testAccessToken) TableName() string { return "w_access_tokens" }
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func setupTestDB(t *testing.T) *gorm.DB {
t.Helper()
dbPath := filepath.Join(t.TempDir(), "auth_test.db")
@@ -27,10 +54,10 @@ func setupTestDB(t *testing.T) *gorm.DB {
require.NoError(t, err)
require.NoError(t, testDB.AutoMigrate(
&model.User{},
&model.AccessToken{},
&model.AuthSource{},
&model.ExternalAccount{},
&testUser{},
&testAccessToken{},
&auth.AuthSource{},
&auth.ExternalAccount{},
))
db.SetDB(testDB)
@@ -75,7 +102,7 @@ func TestAuthPluginUnit(t *testing.T) {
assert.Equal(t, "custom", prov.Name())
// Test User Token Verification with dummy token
user := model.User{
user := testUser{
ID: 101,
Username: "token_user",
IsActive: true,
@@ -83,8 +110,8 @@ func TestAuthPluginUnit(t *testing.T) {
require.NoError(t, testDB.Create(&user).Error)
tokenStr := "test-secret-token-123456"
tokenHash := model.HashToken(tokenStr)
tokenRecord := model.AccessToken{
tokenHash := hashToken(tokenStr)
tokenRecord := testAccessToken{
ID: 201,
UserID: user.ID,
TokenHash: tokenHash,
@@ -106,7 +133,7 @@ func TestAuthPluginUnit(t *testing.T) {
require.NoError(t, authSvc.RevokeUserSessions(context.Background(), user.ID))
// GetCurrentUser from context
userCtx := context.WithValue(context.Background(), "user_obj", userDTO)
userCtx := context.WithValue(context.Background(), auth.UserObjKey, userDTO)
current, err := authSvc.GetCurrentUser(userCtx)
require.NoError(t, err)
assert.Equal(t, user.ID, current.ID)
+76
View File
@@ -0,0 +1,76 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
)
// GetAuthSourceByID 根据 ID 获取认证源
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
var src AuthSource
if err := db.DB(ctx).First(&src, id).Error; err != nil {
return nil, err
}
return &src, nil
}
// GetAuthSourceByName 根据名称获取认证源
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
var src AuthSource
if err := db.DB(ctx).Where("name = ?", name).First(&src).Error; err != nil {
return nil, err
}
return &src, nil
}
// ListActiveAuthSources 获取所有启用的认证源
func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
var sources []AuthSource
if err := db.DB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
return sources, nil
}
// GetActiveAuthSourcesCached 获取所有启用的认证源(带缓存或直接查询)
func GetActiveAuthSourcesCached(ctx context.Context) ([]AuthSource, error) {
return ListActiveAuthSources(ctx)
}
// GetAuthSourceByNameCached 根据名称获取认证源(带缓存或直接查询)
func GetAuthSourceByNameCached(ctx context.Context, name string) (*AuthSource, error) {
return GetAuthSourceByName(ctx, name)
}
// FindExternalAccount 查询指定认证源的外部账号绑定
func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*ExternalAccount, error) {
var account ExternalAccount
if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
// BindExternalAccount 绑定外部账号
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
return db.DB(ctx).Create(account).Error
}
// ListExternalAccountsByUserID 获取用户绑定的所有外部账号
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccount, error) {
var accounts []ExternalAccount
if err := db.DB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil {
return nil, err
}
return accounts, nil
}
// UnbindExternalAccount 解绑外部账号
func UnbindExternalAccount(ctx context.Context, id uint64, userID uint64) error {
return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
}
+19 -38
View File
@@ -9,34 +9,10 @@ import (
"sync"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/gin-gonic/gin"
)
func toUserDTO(u *model.User) *contracts.UserDTO {
if u == nil {
return nil
}
return &contracts.UserDTO{
ID: u.ID,
Username: u.Username,
Nickname: u.Nickname,
Email: u.Email,
AvatarURL: u.AvatarURL,
IsActive: u.IsActive,
IsAdmin: u.IsAdmin,
Bio: u.Bio,
Phone: u.Phone,
Gender: u.Gender,
Website: u.Website,
Location: u.Location,
LastLoginAt: u.LastLoginAt,
CreatedAt: u.CreatedAt,
UpdatedAt: u.UpdatedAt,
}
}
type authServiceImpl struct{}
func newAuthService() contracts.AuthService {
@@ -53,15 +29,12 @@ func (s *authServiceImpl) RequireAdminMiddleware() any {
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if ginCtx, ok := ctx.(*gin.Context); ok {
if u, ok := GetFromContext[*model.User](ginCtx, UserObjKey); ok && u != nil {
return toUserDTO(u), nil
if u, ok := GetFromContext[*contracts.UserDTO](ginCtx, UserObjKey); ok && u != nil {
return u, nil
}
}
if v := ctx.Value(UserObjKey); v != nil {
if u, ok := v.(*model.User); ok && u != nil {
return toUserDTO(u), nil
}
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil
}
@@ -75,21 +48,29 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
return nil, errors.New("auth: empty token")
}
tokenHash := model.HashToken(token)
tokenHash := hashToken(token)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
var tokenRow struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
return nil, err
}
tokenRecord = &dbToken
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.IsActive {
dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
if err != nil || user == nil || !user.IsActive {
var dbUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil {
return nil, err
}
user = &dbUser
@@ -100,7 +81,7 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
return nil, errors.New("auth: system user token not allowed")
}
return toUserDTO(user), nil
return user, nil
}
func (s *authServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
+21 -16
View File
@@ -8,11 +8,12 @@ import (
"crypto/sha256"
"encoding/hex"
"net/http"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/config"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
@@ -66,7 +67,10 @@ func GetUserIDFromSession(s sessions.Session) uint64 {
}
// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
func GetUserIDFromContext(c *gin.Context) uint64 {
func GetUserIDFromContext(c *gin.Context) (uid uint64) {
defer func() {
_ = recover()
}()
session := sessions.Default(c)
return GetUserIDFromSession(session)
}
@@ -96,14 +100,13 @@ func rotateSessionID(s sessions.Session) {
}
// SetLoginSession writes the authenticated user into a freshly rotated session.
func SetLoginSession(ctx context.Context, c *gin.Context, user *model.User, extras ...map[string]any) error {
func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDTO, extras ...map[string]any) error {
session := sessions.Default(c)
session.Clear()
rotateSessionID(session)
session.Set(UserIDKey, user.ID)
session.Set(UserNameKey, user.Username)
session.Set(PasswordHashKey, user.Password)
if len(extras) > 0 {
for key, value := range extras[0] {
session.Set(key, value)
@@ -114,16 +117,18 @@ func SetLoginSession(ctx context.Context, c *gin.Context, user *model.User, extr
maxAge := config.Config.App.SessionAge
isSessionCookie := false
ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
if err == nil {
switch {
case ttlHours == -1:
// 永不过期,设置为 10 年
maxAge = 10 * 365 * 24 * 3600
case ttlHours > 0:
maxAge = ttlHours * 3600
case ttlHours == 0:
isSessionCookie = true
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "login_session_ttl_hours").Pluck("value", &val).Error; 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
}
}
}
session.Options(GetSessionOptions(maxAge))