mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-03 15:06:36 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user