refactor(core): align with cordis spatiotemporal composability architecture

- Purify core micro-kernel by removing context hardcoded helpers and reverse dependencies
- Eliminate init() side effects in infra plugins with reversible lifecycle disposal
- Completely isolate plugins by removing cross-plugin imports and using core/contracts
- Introduce TaskService and RiskControlService contracts for unified cross-plugin APIs
- Regenerate Swagger documentation and update developer guide matrix
- Achieve 0 violations in check_cordis_architecture.sh and 100% test pass
This commit is contained in:
ryan
2026-08-28 15:05:31 +08:00
parent fc7fae7b0e
commit 299ac30ee4
150 changed files with 4328 additions and 2923 deletions
+2 -1
View File
@@ -7,9 +7,10 @@ import (
"context"
"encoding/json"
"github.com/gin-gonic/gin"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
)
// LogForAudit 将登录鉴权审计日志写入 Logger
@@ -11,7 +11,6 @@ import (
"strings"
"Wavelet/core/contracts"
db "Wavelet/plugins/infra/database"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
@@ -19,7 +18,7 @@ import (
func isOIDCLoginEnabled(ctx context.Context) bool {
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "oidc_login_enabled").Pluck("value", &val).Error; err != nil || val == "" {
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "oidc_login_enabled").Pluck("value", &val).Error; err != nil || val == "" {
return true
}
b, err := strconv.ParseBool(val)
@@ -78,7 +77,7 @@ func activeLoginSources(ctx context.Context) []AuthSourceView {
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
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) == "" {
if err := getDB(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(val, "/") + "/login", nil
+14 -156
View File
@@ -6,23 +6,15 @@ package auth
import (
"context"
"fmt"
"strconv"
"sync"
"time"
"Wavelet/core/contracts"
"Wavelet/pkg/cache/ram"
"Wavelet/pkg/util"
db "Wavelet/plugins/infra/cache"
)
const (
tokenCacheTTL = 5 * time.Minute
userCacheTTL = 5 * time.Minute
//nolint:gosec // This is a Redis Pub/Sub channel name, not a credential
oauthTokenInvalidationChannel = "oauth:token_invalidation"
oauthUserInvalidationChannel = "oauth:user_invalidation"
)
// CachedToken represents the minimal cached representation of an access token.
@@ -35,16 +27,6 @@ type CachedToken struct {
var (
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
tokenListenerCancel context.CancelFunc
tokenListenerDone chan struct{}
userListenerOnce sync.Once
userListenerCtx context.Context
userListenerCancel context.CancelFunc
userListenerDone chan struct{}
)
func tokenCacheKey(tokenHash string) string {
@@ -55,107 +37,16 @@ func userCacheKey(userID uint64) string {
return fmt.Sprintf("oauth:user:%d", userID)
}
func ensureTokenCacheListener() {
if db.Redis == nil {
return
}
tokenListenerOnce.Do(startTokenCacheInvalidationListener)
}
func startTokenCacheInvalidationListener() {
tokenListenerCtx, tokenListenerCancel = context.WithCancel(context.Background())
tokenListenerDone = make(chan struct{})
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
util.Go(func() {
listenerCtx := tokenListenerCtx
defer close(tokenListenerDone)
pubsub := redisClient.Subscribe(listenerCtx, oauthTokenInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
util.Go(func() {
<-listenerCtx.Done()
_ = pubsub.Close()
})
for msg := range pubsub.Channel() {
tokenHash := msg.Payload
if tokenHash == "" || tokenHash == "*" || tokenHash == "reset" {
tokenRAM.InvalidateAll()
} else {
tokenRAM.Invalidate(tokenHash)
}
}
})
}
func publishTokenRAMInvalidation(ctx context.Context, tokenHash string) {
if db.Redis == nil {
return
}
_ = db.Redis.Publish(ctx, oauthTokenInvalidationChannel, tokenHash).Err()
}
func ensureUserCacheListener() {
if db.Redis == nil {
return
}
userListenerOnce.Do(startUserCacheInvalidationListener)
}
func startUserCacheInvalidationListener() {
userListenerCtx, userListenerCancel = context.WithCancel(context.Background())
userListenerDone = make(chan struct{})
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
util.Go(func() {
listenerCtx := userListenerCtx
defer close(userListenerDone)
pubsub := redisClient.Subscribe(listenerCtx, oauthUserInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
util.Go(func() {
<-listenerCtx.Done()
_ = pubsub.Close()
})
for msg := range pubsub.Channel() {
userIDStr := msg.Payload
if userIDStr == "" || userIDStr == "*" || userIDStr == "reset" {
userRAM.InvalidateAll()
} else if userID, err := strconv.ParseUint(userIDStr, 10, 64); err == nil {
userRAM.Invalidate(userID)
}
}
})
}
func publishUserRAMInvalidation(ctx context.Context, userID uint64) {
if db.Redis == nil {
return
}
_ = db.Redis.Publish(ctx, oauthUserInvalidationChannel, strconv.FormatUint(userID, 10)).Err()
}
// GetCachedToken 获取缓存的 Token
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
ensureTokenCacheListener()
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
return val, nil
}
if db.Redis != nil {
if cache := getCache(ctx); cache != nil {
var token CachedToken
key := tokenCacheKey(tokenHash)
if err := db.GetJSON(ctx, key, &token); err == nil {
// Write back to local cache
if err := cache.Get(ctx, key, &token); err == nil {
tokenRAM.Set(tokenHash, &token)
return &token, nil
}
@@ -165,40 +56,32 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error)
// SetCachedToken 设置 Token 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
ensureTokenCacheListener()
tokenRAM.Set(tokenHash, token)
if db.Redis != nil {
if cache := getCache(ctx); cache != nil {
key := tokenCacheKey(tokenHash)
_ = db.SetJSON(ctx, key, token, tokenCacheTTL)
_ = cache.Set(ctx, key, token, tokenCacheTTL)
}
}
// InvalidateCachedToken 吊销/删除 token 缓存
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
ensureTokenCacheListener()
tokenRAM.Invalidate(tokenHash)
if db.Redis != nil {
if cache := getCache(ctx); cache != nil {
key := tokenCacheKey(tokenHash)
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
publishTokenRAMInvalidation(ctx, tokenHash)
_ = cache.Delete(ctx, key)
}
}
// GetCachedUser 获取缓存的 UserDTO
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
ensureUserCacheListener()
if val, ok := userRAM.GetIfPresent(userID); ok {
return val, nil
}
if db.Redis != nil {
if cache := getCache(ctx); cache != nil {
var u contracts.UserDTO
key := userCacheKey(userID)
if err := db.GetJSON(ctx, key, &u); err == nil {
// Write back to local cache
if err := cache.Get(ctx, key, &u); err == nil {
userRAM.Set(userID, &u)
return &u, nil
}
@@ -208,49 +91,24 @@ func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, erro
// SetCachedUser 设置 UserDTO 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
ensureUserCacheListener()
userRAM.Set(userID, u)
if db.Redis != nil {
if cache := getCache(ctx); cache != nil {
key := userCacheKey(userID)
_ = db.SetJSON(ctx, key, u, userCacheTTL)
_ = cache.Set(ctx, key, u, userCacheTTL)
}
}
// InvalidateCachedUser 吊销/失效 UserDTO 缓存
func InvalidateCachedUser(ctx context.Context, userID uint64) {
ensureUserCacheListener()
userRAM.Invalidate(userID)
if db.Redis != nil {
if cache := getCache(ctx); cache != nil {
key := userCacheKey(userID)
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
publishUserRAMInvalidation(ctx, userID)
_ = cache.Delete(ctx, key)
}
}
// StopAuthCacheListener stops both token and user Redis Pub/Sub subscription listeners and resets the sync.Once guards.
func StopAuthCacheListener() {
if tokenListenerCancel != nil {
tokenListenerCancel()
if tokenListenerDone != nil {
<-tokenListenerDone
}
tokenListenerCancel = nil
tokenListenerDone = nil
}
tokenListenerOnce = sync.Once{}
if userListenerCancel != nil {
userListenerCancel()
if userListenerDone != nil {
<-userListenerDone
}
userListenerCancel = nil
userListenerDone = nil
}
userListenerOnce = sync.Once{}
}
// StopAuthCacheListener compatibility stub for tests
func StopAuthCacheListener() {}
// ResetAuthRAMCacheForTest clears only the process-local RAM cache.
func ResetAuthRAMCacheForTest() {
+50 -29
View File
@@ -5,48 +5,69 @@ package auth_test
import (
"context"
"encoding/json"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth"
db "Wavelet/plugins/infra/cache"
)
func setupOauthCacheTest(t *testing.T) (*miniredis.Miniredis, func()) {
t.Helper()
type mockCacheService struct {
items map[string][]byte
}
miniRedis, err := miniredis.Run()
func newMockCacheService() *mockCacheService {
return &mockCacheService{items: make(map[string][]byte)}
}
func (m *mockCacheService) Get(ctx context.Context, key string, target any) error {
b, ok := m.items[key]
if !ok {
return contracts.ErrCacheMiss
}
return json.Unmarshal(b, target)
}
func (m *mockCacheService) Set(ctx context.Context, key string, value any, ttl time.Duration) error {
b, err := json.Marshal(value)
if err != nil {
t.Fatalf("failed to start miniredis: %v", err)
return err
}
m.items[key] = b
return nil
}
db.Redis = redis.NewClient(&redis.Options{
Addr: miniRedis.Addr(),
MaintNotificationsConfig: &maintnotifications.Config{
Mode: maintnotifications.ModeDisabled,
},
})
func (m *mockCacheService) Delete(ctx context.Context, key string) error {
delete(m.items, key)
return nil
}
auth.ResetAuthRAMCacheForTest()
func (m *mockCacheService) Invalidate(ctx context.Context, key string) error {
return m.Delete(ctx, key)
}
cleanup := func() {
auth.StopAuthCacheListener()
auth.ResetAuthRAMCacheForTest()
_ = db.Redis.Close()
miniRedis.Close()
db.Redis = nil
func (m *mockCacheService) GetOrSet(ctx context.Context, key string, target any, ttl time.Duration, loader func() (any, error)) error {
err := m.Get(ctx, key, target)
if err == nil {
return nil
}
return miniRedis, cleanup
val, err := loader()
if err != nil {
return err
}
if err := m.Set(ctx, key, val, ttl); err != nil {
return err
}
b, _ := json.Marshal(val)
return json.Unmarshal(b, target)
}
func TestTokenCache_GetSetInvalidate(t *testing.T) {
_, cleanup := setupOauthCacheTest(t)
defer cleanup()
ctx := context.Background()
ctx := core.NewContext(context.Background())
mockCache := newMockCacheService()
core.Provide[contracts.CacheService](ctx, mockCache)
tokenHash := "test-token-hash"
token := &auth.CachedToken{
@@ -84,9 +105,9 @@ func TestTokenCache_GetSetInvalidate(t *testing.T) {
}
func TestUserCache_GetSetInvalidate(t *testing.T) {
_, cleanup := setupOauthCacheTest(t)
defer cleanup()
ctx := context.Background()
ctx := core.NewContext(context.Background())
mockCache := newMockCacheService()
core.Provide[contracts.CacheService](ctx, mockCache)
userID := uint64(789)
user := &contracts.UserDTO{
+43 -32
View File
@@ -14,17 +14,16 @@ import (
"Wavelet/core/contracts"
"Wavelet/pkg/idgen"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
cachepkg "Wavelet/plugins/infra/cache"
db "Wavelet/plugins/infra/database"
"github.com/coreos/go-oidc/v3/oidc"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
"Wavelet/pkg/idgen"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
)
// GetLoginSources 获取可用登录源列表
@@ -78,9 +77,12 @@ func GetLoginURL(c *gin.Context) {
response.AbortInternal(c, err.Error())
return
}
if err := cachepkg.Redis.Set(ctx, cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
response.AbortInternal(c, err.Error())
return
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state)
if cache := getCache(ctx); cache != nil {
if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
@@ -107,18 +109,19 @@ func buildAuthorizeURL(ctx context.Context, source *AuthSource, state string) (s
}
func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
if cachepkg.Redis == nil || sessionHash == "" {
if sessionHash == "" {
return nil
}
key := cachepkg.PrefixedKey(fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash))
n, err := cachepkg.Redis.Incr(ctx, key).Result()
if err != nil {
return err
cache := getCache(ctx)
if cache == nil {
return nil
}
if n == 1 {
_ = cachepkg.Redis.Expire(ctx, key, OAuthStateCacheKeyExpiration).Err()
}
if n > oauthStateLimitMax {
key := fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash)
var count int
_ = cache.Get(ctx, key, &count)
count++
_ = cache.Set(ctx, key, count, OAuthStateCacheKeyExpiration)
if count > oauthStateLimitMax {
return errors.New(errOAuthStateRateLimited)
}
return nil
@@ -179,9 +182,12 @@ func Authorize(c *gin.Context) {
response.AbortInternal(c, err.Error())
return
}
if err := cachepkg.Redis.Set(ctx, cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
response.AbortInternal(c, err.Error())
return
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state)
if cache := getCache(ctx); cache != nil {
if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
@@ -201,13 +207,18 @@ func Callback(c *gin.Context) {
}
ctx := c.Request.Context()
stateKey := cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
payloadRaw, err := cachepkg.Redis.Get(ctx, stateKey).Result()
if err != nil {
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)
var payloadRaw string
cache := getCache(ctx)
if cache == nil {
response.AbortBadRequest(c, errInvalidState)
return
}
_ = cachepkg.Redis.Del(ctx, stateKey)
if err := cache.Get(ctx, stateKey, &payloadRaw); err != nil {
response.AbortBadRequest(c, errInvalidState)
return
}
_ = cache.Delete(ctx, stateKey)
payload, err := decodeOAuthStatePayload(payloadRaw)
if err != nil {
@@ -289,7 +300,7 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource,
return
}
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
@@ -304,7 +315,7 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource,
return
}
user.LastLoginAt = time.Now()
_ = db.DB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
_ = 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")))
}
@@ -314,7 +325,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource
account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch {
case err == nil:
if loadErr := db.DB(ctx).Table("w_users").Where("id = ?", account.UserID).First(&user).Error; loadErr != nil {
if loadErr := getDB(ctx).Table("w_users").Where("id = ?", account.UserID).First(&user).Error; loadErr != nil {
response.AbortInternal(c, loadErr.Error())
return
}
@@ -330,7 +341,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource
}
user.LastLoginAt = time.Now()
_ = db.DB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
_ = getDB(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
@@ -348,7 +359,7 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
}
var existingUsernames []string
if err := db.DB(ctx).Table("w_users").
if err := getDB(ctx).Table("w_users").
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
return "", err
@@ -376,7 +387,7 @@ 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 := db.DB(ctx).Table("w_system_configs").Where("key = ?", "registration_enabled").Pluck("value", &val).Error; err == nil && val != "" {
if err := getDB(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
}
@@ -407,7 +418,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSou
UpdatedAt: now,
}
if err := db.DB(ctx).Table("w_users").Create(&user).Error; err != nil {
if err := getDB(ctx).Table("w_users").Create(&user).Error; err != nil {
response.AbortInternal(c, err.Error())
return contracts.UserDTO{}, false
}
+5 -5
View File
@@ -9,12 +9,12 @@ import (
"encoding/hex"
"errors"
"github.com/gin-gonic/gin"
"Wavelet/core/contracts"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"Wavelet/pkg/util"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
)
func hashToken(token string) string {
@@ -38,7 +38,7 @@ func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *
UserID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
return nil, nil, err
}
tokenRecord = &CachedToken{
@@ -49,7 +49,7 @@ func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *
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 {
if err := getDB(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)
@@ -90,7 +90,7 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
user, err := GetCachedUser(ctx, userID)
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 {
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&dbUser).Error; err != nil {
return nil, err
}
user = &dbUser
+21
View File
@@ -76,6 +76,27 @@ func (p *Plugin) Manifest() core.Manifest {
// Apply registers the auth migrations, services, routes, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 0. Bind DBService & CacheService from Context
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
setDBService(db)
} else {
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
setDBService(db)
})
}
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
setCacheService(cache)
} else {
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
setCacheService(cache)
})
}
ctx.OnDispose(func() error {
setDBService(nil)
setCacheService(nil)
return nil
})
// 1. Register migrations
ctx.Migrations().Register("auth", authMigrations)
+17 -2
View File
@@ -19,9 +19,24 @@ import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/auth"
db "Wavelet/plugins/infra/database"
)
type mockDBService struct {
db *gorm.DB
}
func (m *mockDBService) GORM() *gorm.DB {
return m.db
}
func (m *mockDBService) DB(ctx context.Context) *gorm.DB {
return m.db.WithContext(ctx)
}
func (m *mockDBService) Named(_ string) *gorm.DB {
return m.db
}
type testUser struct {
ID uint64 `gorm:"primaryKey"`
Username string
@@ -60,7 +75,6 @@ func setupTestDB(t *testing.T) *gorm.DB {
&auth.ExternalAccount{},
))
db.SetDB(testDB)
return testDB
}
@@ -81,6 +95,7 @@ func (m *mockProvider) ExchangeCode(ctx context.Context, code string) (*contract
func TestAuthPluginUnit(t *testing.T) {
ctx := core.NewContext(context.Background())
testDB := setupTestDB(t)
core.Provide[contracts.DBService](ctx, &mockDBService{db: testDB})
p := auth.New()
assert.Equal(t, "auth", p.Name())
+58 -8
View File
@@ -5,14 +5,64 @@ package auth
import (
"context"
"sync"
db "Wavelet/plugins/infra/database"
"gorm.io/gorm"
"Wavelet/core"
"Wavelet/core/contracts"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
cacheMu sync.RWMutex
cacheSvc contracts.CacheService
)
func setDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
}
func setCacheService(s contracts.CacheService) {
cacheMu.Lock()
defer cacheMu.Unlock()
cacheSvc = s
}
func getDB(ctx context.Context) *gorm.DB {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
return s.DB(ctx)
}
}
dbMu.RLock()
s := dbSvc
dbMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
func getCache(ctx context.Context) contracts.CacheService {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
return s
}
}
cacheMu.RLock()
s := cacheSvc
cacheMu.RUnlock()
return s
}
// 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 {
if err := getDB(ctx).First(&src, id).Error; err != nil {
return nil, err
}
return &src, nil
@@ -21,7 +71,7 @@ func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
// 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 {
if err := getDB(ctx).Where("name = ?", name).First(&src).Error; err != nil {
return nil, err
}
return &src, nil
@@ -30,7 +80,7 @@ func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error)
// 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 {
if err := getDB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
return sources, nil
@@ -49,7 +99,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*AuthSource, e
// 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 {
if err := getDB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
@@ -57,13 +107,13 @@ func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID st
// BindExternalAccount 绑定外部账号
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
return db.DB(ctx).Create(account).Error
return getDB(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 {
if err := getDB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil {
return nil, err
}
return accounts, nil
@@ -71,5 +121,5 @@ func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]Externa
// 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
return getDB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
}
+12 -12
View File
@@ -8,10 +8,10 @@ import (
"errors"
"sync"
"github.com/gin-gonic/gin"
"Wavelet/core/contracts"
"Wavelet/pkg/util"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
)
type authServiceImpl struct{}
@@ -57,7 +57,7 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
UserID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
return nil, err
}
tokenRecord = &CachedToken{
@@ -71,7 +71,7 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
user, err := GetCachedUser(ctx, tokenRecord.UserID)
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 {
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil {
return nil, err
}
user = &dbUser
@@ -120,7 +120,7 @@ func (s *authServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash s
func (s *authServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
var sources []AuthSource
if err := db.DB(ctx).Order("id ASC").Find(&sources).Error; err != nil {
if err := getDB(ctx).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
@@ -157,7 +157,7 @@ func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts
return nil, err
}
if err := db.DB(ctx).Create(&model).Error; err != nil {
if err := getDB(ctx).Create(&model).Error; err != nil {
return nil, err
}
@@ -167,7 +167,7 @@ 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 := db.DB(ctx).First(&existing, id).Error; err != nil {
if err := getDB(ctx).First(&existing, id).Error; err != nil {
return nil, err
}
@@ -184,7 +184,7 @@ func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, sourc
return nil, err
}
if err := db.DB(ctx).Save(&existing).Error; err != nil {
if err := getDB(ctx).Save(&existing).Error; err != nil {
return nil, err
}
@@ -194,21 +194,21 @@ func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, sourc
func (s *authServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
var existing AuthSource
if err := db.DB(ctx).First(&existing, id).Error; err != nil {
if err := getDB(ctx).First(&existing, id).Error; err != nil {
return err
}
return db.DB(ctx).Delete(&existing).Error
return getDB(ctx).Delete(&existing).Error
}
func (s *authServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
var existing AuthSource
if err := db.DB(ctx).First(&existing, id).Error; err != nil {
if err := getDB(ctx).First(&existing, id).Error; err != nil {
return nil, err
}
existing.IsActive = !existing.IsActive
if err := db.DB(ctx).Save(&existing).Error; err != nil {
if err := getDB(ctx).Save(&existing).Error; err != nil {
return nil, err
}
+4 -4
View File
@@ -11,13 +11,13 @@ import (
"strconv"
"strings"
"Wavelet/core/contracts"
"Wavelet/pkg/config"
db "Wavelet/plugins/infra/database"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
gsessions "github.com/gorilla/sessions"
"Wavelet/core/contracts"
"Wavelet/pkg/config"
)
// GetSessionOptions 根据配置构建 Session 选项
@@ -118,7 +118,7 @@ func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDT
isSessionCookie := false
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 err := getDB(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: