mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 14:06:36 +08:00
refactor(core): decouple gin from pkg/util and reduce code duplication
This commit is contained in:
@@ -19,27 +19,26 @@ type MigrationEntry struct {
|
||||
// MigrationExtension defines the interface for registering and querying plugin migrations.
|
||||
type MigrationExtension interface {
|
||||
Register(pluginID string, fsys fs.FS, dir ...string)
|
||||
Unregister(pluginID string) bool
|
||||
Entries() []MigrationEntry
|
||||
Get(pluginID string) (MigrationEntry, bool)
|
||||
Unregister(pluginID string) bool
|
||||
}
|
||||
|
||||
// MigrationRegistry collects and stores migration entries from plugins.
|
||||
// MigrationRegistry implements MigrationExtension.
|
||||
type MigrationRegistry struct {
|
||||
mu sync.RWMutex
|
||||
entries []MigrationEntry
|
||||
lookup map[string]MigrationEntry
|
||||
}
|
||||
|
||||
// NewMigrationRegistry creates a new migration registry.
|
||||
// NewMigrationRegistry creates a new MigrationRegistry.
|
||||
func NewMigrationRegistry() *MigrationRegistry {
|
||||
return &MigrationRegistry{
|
||||
lookup: make(map[string]MigrationEntry),
|
||||
}
|
||||
}
|
||||
|
||||
// Register adds a migration entry for a plugin.
|
||||
// If dir is not specified, it defaults to "migrations".
|
||||
// Register registers an embedded migration filesystem for a plugin.
|
||||
func (m *MigrationRegistry) Register(pluginID string, fsys fs.FS, dir ...string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
@@ -72,21 +71,9 @@ func (m *MigrationRegistry) Register(pluginID string, fsys fs.FS, dir ...string)
|
||||
|
||||
// Unregister removes a registered migration entry by plugin ID.
|
||||
func (m *MigrationRegistry) Unregister(pluginID string) bool {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if _, exists := m.lookup[pluginID]; !exists {
|
||||
return false
|
||||
}
|
||||
|
||||
delete(m.lookup, pluginID)
|
||||
for i, e := range m.entries {
|
||||
if e.PluginID == pluginID {
|
||||
m.entries = append(m.entries[:i], m.entries[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
return true
|
||||
return unregisterEntry(&m.mu, m.lookup, &m.entries, pluginID, func(e MigrationEntry) bool {
|
||||
return e.PluginID == pluginID
|
||||
})
|
||||
}
|
||||
|
||||
// Entries returns a copy of all registered migration entries in registration order.
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package extpoints
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"sync"
|
||||
)
|
||||
|
||||
func unregisterEntry[T any](mu *sync.RWMutex, lookup map[string]T, list *[]T, key string, matches func(T) bool) bool {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
|
||||
if _, exists := lookup[key]; !exists {
|
||||
return false
|
||||
}
|
||||
|
||||
delete(lookup, key)
|
||||
*list = slices.DeleteFunc(*list, matches)
|
||||
return true
|
||||
}
|
||||
@@ -88,21 +88,9 @@ func (s *ScheduleRegistry) RegisterCron(spec, taskType string, payload any, opts
|
||||
|
||||
// Unregister removes a registered schedule definition by its task type.
|
||||
func (s *ScheduleRegistry) Unregister(taskType string) bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
if _, exists := s.lookup[taskType]; !exists {
|
||||
return false
|
||||
}
|
||||
|
||||
delete(s.lookup, taskType)
|
||||
for i, item := range s.schedules {
|
||||
if item.TaskType == taskType {
|
||||
s.schedules = append(s.schedules[:i], s.schedules[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
return true
|
||||
return unregisterEntry(&s.mu, s.lookup, &s.schedules, taskType, func(item ScheduleDefinition) bool {
|
||||
return item.TaskType == taskType
|
||||
})
|
||||
}
|
||||
|
||||
// Schedules returns a copy of all registered ScheduleDefinitions.
|
||||
|
||||
@@ -65,21 +65,9 @@ func (s *SettingRegistry) Register(schema SettingSchema) {
|
||||
|
||||
// Unregister removes a registered SettingSchema by its key.
|
||||
func (s *SettingRegistry) Unregister(key string) bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
if _, exists := s.lookup[key]; !exists {
|
||||
return false
|
||||
}
|
||||
|
||||
delete(s.lookup, key)
|
||||
for i, item := range s.schemas {
|
||||
if item.Key == key {
|
||||
s.schemas = append(s.schemas[:i], s.schemas[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
return true
|
||||
return unregisterEntry(&s.mu, s.lookup, &s.schemas, key, func(item SettingSchema) bool {
|
||||
return item.Key == key
|
||||
})
|
||||
}
|
||||
|
||||
// Schemas returns a copy of all registered SettingSchemas.
|
||||
|
||||
@@ -107,21 +107,9 @@ func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption)
|
||||
|
||||
// Unregister removes a registered task definition by its pattern.
|
||||
func (t *TaskRegistry) Unregister(pattern string) bool {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
|
||||
if _, exists := t.lookup[pattern]; !exists {
|
||||
return false
|
||||
}
|
||||
|
||||
delete(t.lookup, pattern)
|
||||
for i, item := range t.tasks {
|
||||
if item.Pattern == pattern {
|
||||
t.tasks = append(t.tasks[:i], t.tasks[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
return true
|
||||
return unregisterEntry(&t.mu, t.lookup, &t.tasks, pattern, func(item TaskDefinition) bool {
|
||||
return item.Pattern == pattern
|
||||
})
|
||||
}
|
||||
|
||||
// Tasks returns a copy of all registered TaskDefinitions.
|
||||
|
||||
Vendored
+4
-1
@@ -5,6 +5,7 @@
|
||||
package disk
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/util"
|
||||
"container/list"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
@@ -100,7 +101,9 @@ var (
|
||||
func Default() *Cache {
|
||||
defaultCacheOnce.Do(func() {
|
||||
defaultCache = New("uploads/diskcache")
|
||||
go defaultCache.StartCleanupWorker(defaultCleanupInterval)
|
||||
util.Go(func() {
|
||||
defaultCache.StartCleanupWorker(defaultCleanupInterval)
|
||||
})
|
||||
})
|
||||
return defaultCache
|
||||
}
|
||||
|
||||
Vendored
+11
-24
@@ -5,15 +5,12 @@ package disk
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestDiskCacheBasic(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_basic"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
testDir := t.TempDir()
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
@@ -55,9 +52,7 @@ func TestDiskCacheBasic(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDiskCacheTTL(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_ttl"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
testDir := t.TempDir()
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
@@ -91,9 +86,7 @@ func TestDiskCacheTTL(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDiskCacheExpirationPolicies(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_expiration_policies"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
testDir := t.TempDir()
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
@@ -102,14 +95,14 @@ func TestDiskCacheExpirationPolicies(t *testing.T) {
|
||||
if err := c.Set("default", []byte("default"), DefaultExpiration); err != nil {
|
||||
t.Fatalf("Set(default, DefaultExpiration) returned error: %v", err)
|
||||
}
|
||||
if err := c.Set("custom", []byte("custom"), 100*time.Millisecond); err != nil {
|
||||
t.Fatalf("Set(custom, 100ms) returned error: %v", err)
|
||||
if err := c.Set("custom", []byte("custom"), 150*time.Millisecond); err != nil {
|
||||
t.Fatalf("Set(custom, 150ms) returned error: %v", err)
|
||||
}
|
||||
if err := c.Set("permanent", []byte("permanent"), NoExpiration); err != nil {
|
||||
t.Fatalf("Set(permanent, NoExpiration) returned error: %v", err)
|
||||
}
|
||||
|
||||
time.Sleep(75 * time.Millisecond)
|
||||
time.Sleep(80 * time.Millisecond)
|
||||
|
||||
if _, err := c.Get("default"); err != ErrCacheMiss {
|
||||
t.Errorf("Get(default) error = %v, want ErrCacheMiss", err)
|
||||
@@ -121,7 +114,7 @@ func TestDiskCacheExpirationPolicies(t *testing.T) {
|
||||
t.Errorf("Get(permanent) returned error: %v", err)
|
||||
}
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
if _, err := c.Get("custom"); err != ErrCacheMiss {
|
||||
t.Errorf("Get(custom) error = %v, want ErrCacheMiss", err)
|
||||
@@ -132,9 +125,7 @@ func TestDiskCacheExpirationPolicies(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDiskCacheNoExpirationSurvivesReload(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_no_expiration_reload"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
testDir := t.TempDir()
|
||||
|
||||
c := New(testDir)
|
||||
if err := c.Set("permanent", []byte("value"), NoExpiration); err != nil {
|
||||
@@ -154,9 +145,7 @@ func TestDiskCacheNoExpirationSurvivesReload(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDiskCacheLRUEviction(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_lru"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
testDir := t.TempDir()
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
@@ -186,10 +175,8 @@ func TestDiskCacheLRUEviction(t *testing.T) {
|
||||
t.Errorf("k2 should exist: %v", err)
|
||||
}
|
||||
|
||||
// Write item 3: 8 + 2 = 10 bytes -> total size would be 30, exceeding 20.
|
||||
// This should evict the oldest item. Since k1 was accessed, but then k2 was accessed,
|
||||
// wait, let's access k1 again to make it the most recently used, so k2 becomes oldest!
|
||||
_, _ = c.Get("k1") // k1 is now MRU, k2 is LRU
|
||||
// Access k1 again to make it MRU, k2 becomes LRU
|
||||
_, _ = c.Get("k1")
|
||||
|
||||
err = c.Set("k3", []byte("v3"), DefaultExpiration)
|
||||
if err != nil {
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
// Package ginutil provides helper utilities for Gin web framework contexts.
|
||||
package ginutil
|
||||
|
||||
import "github.com/gin-gonic/gin"
|
||||
|
||||
@@ -438,31 +438,38 @@ func maskSensitiveConfig(key, value string) string {
|
||||
case ConfigKeySMTPPassword:
|
||||
return maskedConfigValue
|
||||
case ConfigKeyStorageConfig:
|
||||
var cfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
|
||||
if cfg.S3.SecretAccessKey != "" {
|
||||
cfg.S3.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.R2.SecretAccessKey != "" {
|
||||
cfg.R2.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.MinIO.SecretAccessKey != "" {
|
||||
cfg.MinIO.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.OSS.SecretAccessKey != "" {
|
||||
cfg.OSS.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.WebDAV.Password != "" {
|
||||
cfg.WebDAV.Password = maskedConfigValue
|
||||
}
|
||||
if val, err := json.Marshal(cfg); err == nil {
|
||||
return string(val)
|
||||
}
|
||||
}
|
||||
return maskStorageConfig(value)
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func maskStorageConfig(value string) string {
|
||||
var cfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(value), &cfg); err != nil {
|
||||
return value
|
||||
}
|
||||
if cfg.S3.SecretAccessKey != "" {
|
||||
cfg.S3.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.R2.SecretAccessKey != "" {
|
||||
cfg.R2.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.MinIO.SecretAccessKey != "" {
|
||||
cfg.MinIO.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.OSS.SecretAccessKey != "" {
|
||||
cfg.OSS.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.WebDAV.Password != "" {
|
||||
cfg.WebDAV.Password = maskedConfigValue
|
||||
}
|
||||
val, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
return value
|
||||
}
|
||||
return string(val)
|
||||
}
|
||||
|
||||
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
|
||||
// and tests connectivity of the new storage configuration.
|
||||
func validateAndMergeStorageConfig(ctx context.Context, value, currentConfig string) (string, error) {
|
||||
|
||||
@@ -5,9 +5,9 @@ package admin
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
@@ -262,7 +262,7 @@ func DeleteUser(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
if currUser == nil {
|
||||
response.AbortUnauthorized(c, AdminRequired)
|
||||
return
|
||||
@@ -373,7 +373,7 @@ func UpdateUser(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
if currUser == nil {
|
||||
response.AbortUnauthorized(c, AdminRequired)
|
||||
return
|
||||
|
||||
@@ -5,10 +5,10 @@ package admin
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/trace"
|
||||
"Wavelet/pkg/util"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
@@ -19,15 +19,15 @@ func LoginAdminRequired() gin.HandlerFunc {
|
||||
ctx, span := trace.Start(c.Request.Context(), "LoginAdminRequired")
|
||||
defer span.End()
|
||||
|
||||
user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
if user == nil {
|
||||
response.AbortNotFound(c, AdminRequired)
|
||||
return
|
||||
}
|
||||
|
||||
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
|
||||
if tokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
|
||||
tokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
|
||||
if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
|
||||
tokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
|
||||
if !tokenAdmin {
|
||||
response.AbortNotFound(c, TokenAdminRequired)
|
||||
return
|
||||
|
||||
@@ -5,6 +5,7 @@ package auth
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
@@ -438,7 +439,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSou
|
||||
|
||||
// UserInfo 获取当前登录用户信息
|
||||
func UserInfo(c *gin.Context) {
|
||||
user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
session := sessions.Default(c)
|
||||
needChange := session.Get("need_change_password") == true
|
||||
|
||||
|
||||
@@ -5,9 +5,9 @@ package auth
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/trace"
|
||||
"Wavelet/pkg/util"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
@@ -25,43 +25,45 @@ func hashToken(token string) string {
|
||||
func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *CachedToken, error) {
|
||||
tokenHash := hashToken(tokenStr)
|
||||
tokenRecord, err := GetCachedToken(ctx, tokenHash)
|
||||
if err == nil {
|
||||
user, err := GetCachedUser(ctx, tokenRecord.UserID)
|
||||
if err == nil && user != nil && user.IsActive {
|
||||
return user, tokenRecord, nil
|
||||
if err != nil || tokenRecord == nil {
|
||||
var tokenRow struct {
|
||||
ID uint64
|
||||
UserID uint64
|
||||
IsAdmin bool
|
||||
}
|
||||
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
tokenRecord = &CachedToken{
|
||||
ID: tokenRow.ID,
|
||||
UserID: tokenRow.UserID,
|
||||
IsAdmin: tokenRow.IsAdmin,
|
||||
}
|
||||
SetCachedToken(ctx, tokenHash, tokenRecord)
|
||||
}
|
||||
|
||||
var tokenRow struct {
|
||||
ID uint64
|
||||
UserID uint64
|
||||
IsAdmin bool
|
||||
user, err := GetCachedUser(ctx, tokenRecord.UserID)
|
||||
if err != nil || user == nil || !user.IsActive {
|
||||
var userRow contracts.UserDTO
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&userRow).Error; err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
user = &userRow
|
||||
SetCachedUser(ctx, tokenRecord.UserID, user)
|
||||
}
|
||||
if err := getDB(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 := 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)
|
||||
return &userRow, tokenRecord, nil
|
||||
return user, tokenRecord, nil
|
||||
}
|
||||
|
||||
// GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error
|
||||
// GetUserFromRequest 从请求中获取当前用户(优先 Access Token,其次 Session)
|
||||
func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
|
||||
ctx := c.Request.Context()
|
||||
var tokenStr string
|
||||
|
||||
// Check token in headers
|
||||
tokenStr := c.GetHeader("X-Access-Token")
|
||||
if tokenStr == "" {
|
||||
tokenFromQuery := c.Query("token")
|
||||
if tokenFromQuery != "" {
|
||||
tokenStr = tokenFromQuery
|
||||
} else {
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
|
||||
tokenStr = authHeader[7:]
|
||||
@@ -74,8 +76,8 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
|
||||
if user.Username == SystemUsername {
|
||||
return nil, errors.New("system user is not allowed to login")
|
||||
}
|
||||
util.SetToContext(c, contracts.AuthTokenAuthKey, true)
|
||||
util.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
|
||||
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, true)
|
||||
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin)
|
||||
return user, nil
|
||||
}
|
||||
}
|
||||
@@ -96,8 +98,8 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
|
||||
SetCachedUser(ctx, userID, user)
|
||||
}
|
||||
|
||||
util.SetToContext(c, contracts.AuthTokenAuthKey, false)
|
||||
util.SetToContext(c, contracts.AuthTokenAdminKey, false)
|
||||
ginutil.SetToContext(c, contracts.AuthTokenAuthKey, false)
|
||||
ginutil.SetToContext(c, contracts.AuthTokenAdminKey, false)
|
||||
|
||||
if user.Username == "system" {
|
||||
return nil, errors.New("system user is not allowed to login")
|
||||
@@ -119,7 +121,7 @@ func LoginRequired() gin.HandlerFunc {
|
||||
}
|
||||
|
||||
LogForAudit(c.Request.Context(), user, c)
|
||||
util.SetToContext(c, contracts.AuthUserObjKey, user)
|
||||
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -136,8 +138,8 @@ func AdminRequired() gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
isTokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey)
|
||||
isTokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
|
||||
isTokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey)
|
||||
isTokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
|
||||
|
||||
// 如果是通过 Token 鉴权,要求该 Token 具备管理员权限或者用户本身为管理员
|
||||
if isTokenAuth && !isTokenAdmin && !user.IsAdmin {
|
||||
@@ -152,7 +154,7 @@ func AdminRequired() gin.HandlerFunc {
|
||||
}
|
||||
|
||||
LogForAudit(c.Request.Context(), user, c)
|
||||
util.SetToContext(c, contracts.AuthUserObjKey, user)
|
||||
ginutil.SetToContext(c, contracts.AuthUserObjKey, user)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -165,7 +167,7 @@ func LoginAdminRequired() gin.HandlerFunc {
|
||||
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
|
||||
func DisallowTokenAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if tokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
|
||||
if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
|
||||
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -5,7 +5,7 @@ package auth
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
@@ -29,7 +29,7 @@ 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 := util.GetFromContext[*contracts.UserDTO](ginCtx, contracts.AuthUserObjKey); ok && u != nil {
|
||||
if u, ok := ginutil.GetFromContext[*contracts.UserDTO](ginCtx, contracts.AuthUserObjKey); ok && u != nil {
|
||||
return u, nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -40,6 +40,23 @@ func ListAdminChannels(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(rows))
|
||||
}
|
||||
|
||||
func parseAdminChannelID(c *gin.Context) (uint64, bool) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid channel id")
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
}
|
||||
|
||||
func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
||||
if err.Error() == errChannelNotFound {
|
||||
response.AbortNotFound(c, err.Error())
|
||||
return
|
||||
}
|
||||
fallback(c, err.Error())
|
||||
}
|
||||
|
||||
// CreateAdminChannel creates a messaging channel.
|
||||
// @Summary Create message gateway channel
|
||||
// @Description Creates a Telegram or QQ channel with encrypted credentials
|
||||
@@ -52,17 +69,7 @@ func ListAdminChannels(c *gin.Context) {
|
||||
// @Failure 400 {object} response.Any
|
||||
// @Router /api/v1/admin/message-gateway/channels [post]
|
||||
func CreateAdminChannel(c *gin.Context) {
|
||||
var req CreateChannelRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
dto, err := createChannel(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(dto))
|
||||
handleJSONRequest(c, createChannel)
|
||||
}
|
||||
|
||||
// UpdateAdminChannel patches a messaging channel. Empty secrets keep the previous values.
|
||||
@@ -79,26 +86,9 @@ func CreateAdminChannel(c *gin.Context) {
|
||||
// @Failure 404 {object} response.Any
|
||||
// @Router /api/v1/admin/message-gateway/channels/{id} [patch]
|
||||
func UpdateAdminChannel(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid channel id")
|
||||
return
|
||||
}
|
||||
var req UpdateChannelRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
dto, err := updateChannel(c.Request.Context(), id, req)
|
||||
if err != nil {
|
||||
if err.Error() == errChannelNotFound {
|
||||
response.AbortNotFound(c, err.Error())
|
||||
return
|
||||
}
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(dto))
|
||||
handleEntityUpdate(c, parseAdminChannelID, updateChannel, func(c *gin.Context, err error) {
|
||||
handleAdminChannelError(c, err, response.AbortBadRequest)
|
||||
})
|
||||
}
|
||||
|
||||
// DeleteAdminChannel removes a channel and its bindings/pairing codes.
|
||||
@@ -112,17 +102,12 @@ func UpdateAdminChannel(c *gin.Context) {
|
||||
// @Failure 404 {object} response.Any
|
||||
// @Router /api/v1/admin/message-gateway/channels/{id} [delete]
|
||||
func DeleteAdminChannel(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid channel id")
|
||||
id, ok := parseAdminChannelID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := deleteChannel(c.Request.Context(), id); err != nil {
|
||||
if err.Error() == errChannelNotFound {
|
||||
response.AbortNotFound(c, err.Error())
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, err.Error())
|
||||
handleAdminChannelError(c, err, response.AbortInternal)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
@@ -140,17 +125,12 @@ func DeleteAdminChannel(c *gin.Context) {
|
||||
// @Failure 404 {object} response.Any
|
||||
// @Router /api/v1/admin/message-gateway/channels/{id}/test [post]
|
||||
func TestAdminChannel(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid channel id")
|
||||
id, ok := parseAdminChannelID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := probeChannel(c.Request.Context(), id); err != nil {
|
||||
if err.Error() == errChannelNotFound {
|
||||
response.AbortNotFound(c, err.Error())
|
||||
return
|
||||
}
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
handleAdminChannelError(c, err, response.AbortBadRequest)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) {
|
||||
var req Req
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
res, err := handler(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(res))
|
||||
}
|
||||
|
||||
func handleEntityUpdate[Req any, Res any](
|
||||
c *gin.Context,
|
||||
parseID func(*gin.Context) (uint64, bool),
|
||||
updater func(ctx context.Context, id uint64, req Req) (Res, error),
|
||||
onErr func(*gin.Context, error),
|
||||
) {
|
||||
id, ok := parseID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var req Req
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
dto, err := updater(c.Request.Context(), id, req)
|
||||
if err != nil {
|
||||
onErr(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(dto))
|
||||
}
|
||||
@@ -5,8 +5,8 @@ package message_gateway
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
@@ -15,7 +15,7 @@ import (
|
||||
)
|
||||
|
||||
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
|
||||
return util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
return ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
}
|
||||
|
||||
// ListChannels lists enabled channels a user can bind.
|
||||
|
||||
@@ -214,20 +214,26 @@ type CreatePushChannelRequest struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
func parsePushChannelID(c *gin.Context) (uint64, bool) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid channel id")
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
}
|
||||
|
||||
func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "channel not found")
|
||||
return
|
||||
}
|
||||
fallback(c, err.Error())
|
||||
}
|
||||
|
||||
// CreatePushChannel creates a push channel.
|
||||
func CreatePushChannel(c *gin.Context) {
|
||||
var req CreatePushChannelRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
channel, err := createPushChannel(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(channel))
|
||||
handleJSONRequest(c, createPushChannel)
|
||||
}
|
||||
|
||||
// UpdatePushChannelRequest is the update channel request payload.
|
||||
@@ -242,44 +248,20 @@ type UpdatePushChannelRequest struct {
|
||||
|
||||
// UpdatePushChannel updates a push channel.
|
||||
func UpdatePushChannel(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid channel id")
|
||||
return
|
||||
}
|
||||
|
||||
var req UpdatePushChannelRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
channel, err := updatePushChannel(c.Request.Context(), id, req)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "channel not found")
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(channel))
|
||||
handleEntityUpdate(c, parsePushChannelID, updatePushChannel, func(c *gin.Context, err error) {
|
||||
handlePushChannelNotFoundError(c, err, response.AbortInternal)
|
||||
})
|
||||
}
|
||||
|
||||
// DeletePushChannel deletes a push channel.
|
||||
func DeletePushChannel(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid channel id")
|
||||
id, ok := parsePushChannelID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
if err := deletePushChannel(c.Request.Context(), id); err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "channel not found")
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, err.Error())
|
||||
handlePushChannelNotFoundError(c, err, response.AbortInternal)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
|
||||
@@ -56,36 +56,37 @@ func ListBuiltInPushEvents(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(GetBuiltInEvents()))
|
||||
}
|
||||
|
||||
func parsePushEventID(c *gin.Context) (uint64, bool) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid event id")
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
}
|
||||
|
||||
func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "notification event not found")
|
||||
return
|
||||
}
|
||||
fallback(c, err.Error())
|
||||
}
|
||||
|
||||
// CreatePushEvent creates a new push event configuration.
|
||||
func CreatePushEvent(c *gin.Context) {
|
||||
var req CreatePushEventRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
event, err := createPushEvent(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(event))
|
||||
handleJSONRequest(c, createPushEvent)
|
||||
}
|
||||
|
||||
// DeletePushEvent deletes a push event configuration by ID.
|
||||
func DeletePushEvent(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid event id")
|
||||
id, ok := parsePushEventID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
if err := deletePushEvent(c.Request.Context(), id); err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "notification event not found")
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, err.Error())
|
||||
handlePushEventNotFoundError(c, err, response.AbortInternal)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
@@ -93,9 +94,8 @@ func DeletePushEvent(c *gin.Context) {
|
||||
|
||||
// UpdatePushEvent updates an existing push event.
|
||||
func UpdatePushEvent(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid event id")
|
||||
id, ok := parsePushEventID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -106,11 +106,7 @@ func UpdatePushEvent(c *gin.Context) {
|
||||
}
|
||||
|
||||
if err := updatePushEvent(c.Request.Context(), id, req); err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "notification event not found")
|
||||
return
|
||||
}
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
handlePushEventNotFoundError(c, err, response.AbortBadRequest)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
@@ -118,19 +114,14 @@ func UpdatePushEvent(c *gin.Context) {
|
||||
|
||||
// TogglePushEvent toggles the enabled state of a push event.
|
||||
func TogglePushEvent(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid event id")
|
||||
id, ok := parsePushEventID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
enabled, err := togglePushEvent(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "notification event not found")
|
||||
return
|
||||
}
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
handlePushEventNotFoundError(c, err, response.AbortBadRequest)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(enabled))
|
||||
|
||||
@@ -217,25 +217,31 @@ func DeletePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取)。
|
||||
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
|
||||
cacheKey := "push:channel:active:" + name
|
||||
var channel PushChannel
|
||||
func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Duration, query func(db *gorm.DB, dest *T) error) (*T, error) {
|
||||
var val T
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
if err := cache.Get(ctx, cacheKey, &channel); err == nil {
|
||||
return &channel, nil
|
||||
if err := cache.Get(ctx, cacheKey, &val); err == nil {
|
||||
return &val, nil
|
||||
}
|
||||
}
|
||||
|
||||
if err := getDB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error; err != nil {
|
||||
db := getDB(ctx)
|
||||
if err := query(db, &val); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
_ = cache.Set(ctx, cacheKey, channel, activePushChannelCacheTTL)
|
||||
_ = cache.Set(ctx, cacheKey, val, ttl)
|
||||
}
|
||||
|
||||
return &channel, nil
|
||||
return &val, nil
|
||||
}
|
||||
|
||||
// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取)。
|
||||
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
|
||||
return getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *PushChannel) error {
|
||||
return db.Where("name = ? AND enabled = ?", name, true).First(dest).Error
|
||||
})
|
||||
}
|
||||
|
||||
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
|
||||
@@ -329,23 +335,9 @@ func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string)
|
||||
|
||||
// GetActivePushEventByKey 获取启用的通知事件 (优先从 Redis 缓存获取)。
|
||||
func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) {
|
||||
cacheKey := "push:event:active:" + key
|
||||
var event PushEvent
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
if err := cache.Get(ctx, cacheKey, &event); err == nil {
|
||||
return &event, nil
|
||||
}
|
||||
}
|
||||
|
||||
if err := getDB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
_ = cache.Set(ctx, cacheKey, event, activePushEventCacheTTL)
|
||||
}
|
||||
|
||||
return &event, nil
|
||||
return getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *PushEvent) error {
|
||||
return db.Where("event_key = ? AND enabled = ?", key, true).First(dest).Error
|
||||
})
|
||||
}
|
||||
|
||||
// DeleteActivePushEventCache 清理启用通知事件的缓存。
|
||||
|
||||
@@ -7,9 +7,9 @@ package risk_control
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/risk_control/logstore"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
@@ -42,7 +42,7 @@ func RiskControlMiddleware() gin.HandlerFunc {
|
||||
c.Next()
|
||||
|
||||
// 3. 后置身份检查:仅记录通过认证的请求
|
||||
userObj, exists := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
userObj, exists := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
if !exists || userObj == nil {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -7,8 +7,8 @@ import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/batchwriter"
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/risk_control"
|
||||
"Wavelet/plugins/domain/risk_control/logstore"
|
||||
"context"
|
||||
@@ -99,7 +99,7 @@ func TestRiskControlMiddleware(t *testing.T) {
|
||||
r := gin.New()
|
||||
r.Use(func(c *gin.Context) {
|
||||
user := &contracts.UserDTO{ID: 12345}
|
||||
util.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, user)
|
||||
ginutil.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, user)
|
||||
c.Next()
|
||||
})
|
||||
r.Use(risk_control.RiskControlMiddleware())
|
||||
|
||||
@@ -23,8 +23,7 @@ import (
|
||||
"sync"
|
||||
|
||||
pkgcache "Wavelet/pkg/cache/disk"
|
||||
|
||||
pkgutil "Wavelet/pkg/util"
|
||||
"Wavelet/pkg/ginutil"
|
||||
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
|
||||
@@ -301,7 +300,7 @@ func getOriginalFileBytes(ctx context.Context, upload *models.Upload) ([]byte, e
|
||||
func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
|
||||
var currUserID uint64
|
||||
var isAdmin bool
|
||||
if u, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok && u != nil {
|
||||
if u, ok := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok && u != nil {
|
||||
currUserID = u.ID
|
||||
isAdmin = u.IsAdmin
|
||||
} else if authSvc := shared.GetAuthService(c); authSvc != nil {
|
||||
@@ -329,14 +328,19 @@ func CheckFileAccessPermission(c *gin.Context, upload *models.Upload) error {
|
||||
return checkPrivateFileOwner(c, upload.UserID)
|
||||
}
|
||||
|
||||
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
|
||||
if _, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); !ok {
|
||||
if authSvc := shared.GetAuthService(c); authSvc != nil {
|
||||
if _, err := authSvc.GetCurrentUser(c); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
if cache.IsFilePublic(c.Request.Context(), upload.Type) {
|
||||
return nil
|
||||
}
|
||||
return nil
|
||||
|
||||
if _, ok := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
authSvc := shared.GetAuthService(c)
|
||||
if authSvc == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
_, err := authSvc.GetCurrentUser(c)
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -5,8 +5,8 @@ package handler
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/ingest"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/repository"
|
||||
@@ -168,7 +168,7 @@ type listMyFilesResponse struct {
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Router /api/v1/upload/my [get]
|
||||
func ListMyFiles(c *gin.Context) {
|
||||
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
ctx := c.Request.Context()
|
||||
|
||||
var req listMyFilesRequest
|
||||
@@ -215,7 +215,7 @@ func ListMyFiles(c *gin.Context) {
|
||||
// @Failure 404 {object} response.Any "文件不存在"
|
||||
// @Router /api/v1/upload/{id} [delete]
|
||||
func DeleteMyFile(c *gin.Context) {
|
||||
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
ctx := c.Request.Context()
|
||||
if uploadstorage.ReadOnly(ctx) {
|
||||
response.AbortConflict(c, shared.ErrStorageReadOnly)
|
||||
@@ -262,7 +262,7 @@ type updateMyFileRequest struct {
|
||||
// @Failure 404 {object} response.Any "文件不存在"
|
||||
// @Router /api/v1/upload/{id} [put]
|
||||
func UpdateMyFile(c *gin.Context) {
|
||||
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
ctx := c.Request.Context()
|
||||
if uploadstorage.ReadOnly(ctx) {
|
||||
response.AbortConflict(c, shared.ErrStorageReadOnly)
|
||||
|
||||
@@ -29,11 +29,11 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"Wavelet/pkg/ginutil"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
pkgutil "Wavelet/pkg/util"
|
||||
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
)
|
||||
|
||||
@@ -64,7 +64,7 @@ func UploadFile(c *gin.Context) {
|
||||
|
||||
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize)
|
||||
|
||||
currUser, _ := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
currUser, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
ctx := c.Request.Context()
|
||||
|
||||
header, err := c.FormFile("file")
|
||||
|
||||
@@ -5,8 +5,8 @@ package handler
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/ginutil"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
"archive/zip"
|
||||
@@ -42,7 +42,7 @@ func setupTestRouter(authUser *contracts.UserDTO) *gin.Engine {
|
||||
|
||||
authMiddleware := func(c *gin.Context) {
|
||||
if authUser != nil {
|
||||
util.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, authUser)
|
||||
ginutil.SetToContext[*contracts.UserDTO](c, contracts.AuthUserObjKey, authUser)
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
|
||||
@@ -139,34 +139,7 @@ func (p *Plugin) Start(_ context.Context) error {
|
||||
}
|
||||
|
||||
if p.scheduler == nil {
|
||||
opts := p.schedulerOpts
|
||||
if opts == nil {
|
||||
opts = &asynq.SchedulerOpts{
|
||||
Location: p.location,
|
||||
}
|
||||
} else if opts.Location == nil && p.location != nil {
|
||||
opts.Location = p.location
|
||||
}
|
||||
|
||||
opt := p.redisOpt
|
||||
if opt == nil {
|
||||
if RedisOpt != nil {
|
||||
opt = RedisOpt
|
||||
} else {
|
||||
redisCfg := config.Config.Redis
|
||||
addr := "127.0.0.1:6379"
|
||||
if len(redisCfg.Addrs) > 0 && redisCfg.Addrs[0] != "" {
|
||||
addr = redisCfg.Addrs[0]
|
||||
}
|
||||
opt = asynq.RedisClientOpt{
|
||||
Addr: addr,
|
||||
Username: redisCfg.Username,
|
||||
Password: redisCfg.Password,
|
||||
DB: redisCfg.DB,
|
||||
}
|
||||
}
|
||||
}
|
||||
p.scheduler = asynq.NewScheduler(opt, opts)
|
||||
p.scheduler = p.initScheduler()
|
||||
}
|
||||
|
||||
if p.coreCtx != nil && p.coreCtx.Schedules() != nil {
|
||||
@@ -268,3 +241,37 @@ func buildAsynqOptions(opts map[string]any) []asynq.Option {
|
||||
|
||||
return res
|
||||
}
|
||||
|
||||
func (p *Plugin) initScheduler() *asynq.Scheduler {
|
||||
opts := p.schedulerOpts
|
||||
if opts == nil {
|
||||
opts = &asynq.SchedulerOpts{
|
||||
Location: p.location,
|
||||
}
|
||||
} else if opts.Location == nil && p.location != nil {
|
||||
opts.Location = p.location
|
||||
}
|
||||
|
||||
opt := p.resolveRedisOpt()
|
||||
return asynq.NewScheduler(opt, opts)
|
||||
}
|
||||
|
||||
func (p *Plugin) resolveRedisOpt() asynq.RedisConnOpt {
|
||||
if p.redisOpt != nil {
|
||||
return p.redisOpt
|
||||
}
|
||||
if RedisOpt != nil {
|
||||
return RedisOpt
|
||||
}
|
||||
redisCfg := config.Config.Redis
|
||||
addr := "127.0.0.1:6379"
|
||||
if len(redisCfg.Addrs) > 0 && redisCfg.Addrs[0] != "" {
|
||||
addr = redisCfg.Addrs[0]
|
||||
}
|
||||
return asynq.RedisClientOpt{
|
||||
Addr: addr,
|
||||
Username: redisCfg.Username,
|
||||
Password: redisCfg.Password,
|
||||
DB: redisCfg.DB,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -111,21 +111,7 @@ func ReloadScheduler() error {
|
||||
// 4. 遍历并注册任务
|
||||
taskSvc := getTaskService()
|
||||
for _, s := range schedules {
|
||||
taskName := s.TaskType
|
||||
maxRetry := 3
|
||||
queue := "default"
|
||||
|
||||
if taskSvc != nil {
|
||||
if meta, ok := taskSvc.GetTaskMeta(s.TaskType); ok {
|
||||
taskName = meta.Name
|
||||
if meta.MaxRetry > 0 {
|
||||
maxRetry = meta.MaxRetry
|
||||
}
|
||||
if meta.Queue != "" {
|
||||
queue = meta.Queue
|
||||
}
|
||||
}
|
||||
}
|
||||
taskName, maxRetry, queue := resolveTaskScheduleMeta(taskSvc, s.TaskType)
|
||||
|
||||
// 构造 Asynq 载荷
|
||||
t := asynq.NewTask(taskName, []byte(s.Payload))
|
||||
@@ -160,3 +146,24 @@ func waitForStop(done, signals <-chan struct{}) bool {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
func resolveTaskScheduleMeta(taskSvc contracts.TaskService, taskType string) (name string, maxRetry int, queue string) {
|
||||
name = taskType
|
||||
maxRetry = 3
|
||||
queue = "default"
|
||||
if taskSvc == nil {
|
||||
return
|
||||
}
|
||||
meta, ok := taskSvc.GetTaskMeta(taskType)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
name = meta.Name
|
||||
if meta.MaxRetry > 0 {
|
||||
maxRetry = meta.MaxRetry
|
||||
}
|
||||
if meta.Queue != "" {
|
||||
queue = meta.Queue
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -366,27 +366,8 @@ func (s *taskServiceImpl) ListExecutions(ctx context.Context, taskType, status s
|
||||
return nil, 0, err
|
||||
}
|
||||
res := make([]contracts.TaskExecutionDTO, 0, len(rows))
|
||||
for _, r := range rows {
|
||||
res = append(res, contracts.TaskExecutionDTO{
|
||||
ID: r.ID,
|
||||
TaskID: r.TaskID,
|
||||
TaskType: r.TaskType,
|
||||
TaskName: r.TaskName,
|
||||
Status: string(r.Status),
|
||||
Retryable: r.Retryable,
|
||||
MaxRetry: r.MaxRetry,
|
||||
RetryCount: r.RetryCount,
|
||||
Log: r.Log,
|
||||
ErrorMessage: r.ErrorMessage,
|
||||
Result: r.Result,
|
||||
StartedAt: r.StartedAt,
|
||||
FinishedAt: r.FinishedAt,
|
||||
Duration: r.Duration,
|
||||
Payload: r.Payload,
|
||||
TriggeredBy: r.TriggeredBy,
|
||||
CreatedAt: r.CreatedAt,
|
||||
UpdatedAt: r.UpdatedAt,
|
||||
})
|
||||
for i := range rows {
|
||||
res = append(res, toTaskExecutionDTO(&rows[i]))
|
||||
}
|
||||
return res, total, nil
|
||||
}
|
||||
@@ -416,7 +397,12 @@ func (s *taskServiceImpl) GetExecution(ctx context.Context, id uint64) (*contrac
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &contracts.TaskExecutionDTO{
|
||||
dto := toTaskExecutionDTO(exec)
|
||||
return &dto, nil
|
||||
}
|
||||
|
||||
func toTaskExecutionDTO(exec *TaskExecution) contracts.TaskExecutionDTO {
|
||||
return contracts.TaskExecutionDTO{
|
||||
ID: exec.ID,
|
||||
TaskID: exec.TaskID,
|
||||
TaskType: exec.TaskType,
|
||||
@@ -435,5 +421,5 @@ func (s *taskServiceImpl) GetExecution(ctx context.Context, id uint64) (*contrac
|
||||
TriggeredBy: exec.TriggeredBy,
|
||||
CreatedAt: exec.CreatedAt,
|
||||
UpdatedAt: exec.UpdatedAt,
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -78,7 +79,7 @@ func (s *inprocScheduler) registerJob(ctx context.Context, def extpoints.Schedul
|
||||
spec := def.Spec
|
||||
taskType := def.TaskType
|
||||
|
||||
fields := len(cronFields(spec))
|
||||
fields := len(strings.Fields(spec))
|
||||
cronSpec := spec
|
||||
if fields == standardCronFields {
|
||||
cronSpec = "0 " + spec
|
||||
@@ -141,22 +142,3 @@ func invokeHandler(ctx context.Context, handler any, payload []byte) error {
|
||||
return fmt.Errorf("unsupported handler type: %T", handler)
|
||||
}
|
||||
}
|
||||
|
||||
func cronFields(s string) []string {
|
||||
var fields []string
|
||||
var current []rune
|
||||
for _, r := range s {
|
||||
if r == ' ' || r == '\t' {
|
||||
if len(current) > 0 {
|
||||
fields = append(fields, string(current))
|
||||
current = nil
|
||||
}
|
||||
} else {
|
||||
current = append(current, r)
|
||||
}
|
||||
}
|
||||
if len(current) > 0 {
|
||||
fields = append(fields, string(current))
|
||||
}
|
||||
return fields
|
||||
}
|
||||
|
||||
@@ -101,7 +101,7 @@ func (q *InprocQueue) Start(ctx context.Context) {
|
||||
q.wg.Add(1)
|
||||
util.Go(func() {
|
||||
defer q.wg.Done()
|
||||
q.workerLoop()
|
||||
q.workerLoop(ctx)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -128,21 +128,23 @@ func (q *InprocQueue) Stop(ctx context.Context) error {
|
||||
}
|
||||
}
|
||||
|
||||
func (q *InprocQueue) workerLoop() {
|
||||
func (q *InprocQueue) workerLoop(ctx context.Context) {
|
||||
for {
|
||||
select {
|
||||
case <-q.stopCh:
|
||||
return
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case msg, ok := <-q.queue:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
q.executeTask(msg)
|
||||
q.executeTask(ctx, msg)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (q *InprocQueue) executeTask(msg TaskMessage) {
|
||||
func (q *InprocQueue) executeTask(ctx context.Context, msg TaskMessage) {
|
||||
if q.taskReg == nil {
|
||||
return
|
||||
}
|
||||
@@ -157,7 +159,7 @@ func (q *InprocQueue) executeTask(msg TaskMessage) {
|
||||
timeout = 5 * time.Minute
|
||||
}
|
||||
|
||||
taskCtx, cancel := context.WithTimeout(q.baseCtx, timeout)
|
||||
taskCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
err := invokeHandler(taskCtx, td.Handler, msg.Payload)
|
||||
@@ -165,7 +167,13 @@ func (q *InprocQueue) executeTask(msg TaskMessage) {
|
||||
msg.RetryLeft--
|
||||
// Retry with backoff
|
||||
util.Go(func() {
|
||||
time.Sleep(defaultRetryBackoff)
|
||||
select {
|
||||
case <-time.After(defaultRetryBackoff):
|
||||
case <-q.stopCh:
|
||||
return
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
if q.running.Load() {
|
||||
select {
|
||||
case q.queue <- msg:
|
||||
|
||||
@@ -106,12 +106,12 @@ else
|
||||
log_pass "backend/pkg/ 零插件依赖"
|
||||
fi
|
||||
|
||||
# 3.2 pkg/util/ 严禁导入 ORM / Session 框架
|
||||
UTIL_FRAMEWORK_IMPORTS=$(rg -n '"gorm.io/gorm"|"github.com/gorilla/sessions"' \
|
||||
# 3.2 pkg/util/ 严禁导入 Gin / ORM / Session 框架
|
||||
UTIL_FRAMEWORK_IMPORTS=$(rg -n '"gorm.io/gorm"|"github.com/gorilla/sessions"|"github.com/gin-gonic/gin"' \
|
||||
"${BACKEND_DIR}/pkg/util/" --glob '*.go' -g '!*_test.go' || true)
|
||||
|
||||
if [ -n "${UTIL_FRAMEWORK_IMPORTS}" ]; then
|
||||
log_fail "backend/pkg/util/ 必须保持纯粹,禁止导入 gorm、sessions 等数据库/会话框架包:"
|
||||
log_fail "backend/pkg/util/ 必须保持纯粹,禁止导入 gin、gorm、sessions 等 Web/数据库/会话框架包:"
|
||||
echo "${UTIL_FRAMEWORK_IMPORTS}" >&2
|
||||
else
|
||||
log_pass "backend/pkg/util/ 保持纯净无状态"
|
||||
|
||||
Reference in New Issue
Block a user