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
View File
@@ -63,3 +63,5 @@ s3_cache
.worktrees/
/.superpowers/
/backend/plugins/domain/upload/filesrv/uploads/
/backend/plugins/domain/upload/task/uploads/
+12 -5
View File
@@ -25,17 +25,17 @@ linters:
- gocritic # 各类代码问题
- funlen # 函数过长
- gosec # 安全问题检查
- gosec # 安全问题检查
- bodyclose # HTTP response body 没有正确关闭
- noctx # 没有传递 context.Context
- contextcheck # 其他检查
- sqlclosecheck # SQL rows 没有正确关闭
- unconvert # 不必要的类型转换
- nilerr # 函数返回 nil 错误
- sqlclosecheck # SQL rows 没有正确关闭
- unconvert # 不必要的类型转换
- nilerr # 函数返回 nil 错误
settings:
dupl:
threshold: 120
threshold: 80
cyclop:
max-complexity: 20
@@ -53,3 +53,10 @@ linters:
- argument
- condition
- return
formatters:
enable:
- gofumpt
settings:
gofumpt:
extra-rules: true
+3 -2
View File
@@ -14,8 +14,9 @@ license-check:
scripts/update_go_license.sh --check
format:
@echo "==> Formatting backend Go source..."
gofmt -w $$(find backend -type f -name '*.go' -not -path './.git/*')
@echo "==> Formatting backend Go source with goimports..."
@command -v goimports >/dev/null 2>&1 || { echo 'error: goimports is required. Run: go install golang.org/x/tools/cmd/goimports@latest' >&2; exit 1; }
goimports -w -local $(MODULE) $$(find backend -type f -name '*.go' -not -path './.git/*')
@echo "==> Formatting frontend source..."
cd frontend && pnpm format
+2 -1
View File
@@ -7,8 +7,9 @@ package cmd
import (
"log"
"Wavelet/core"
"github.com/spf13/cobra"
"Wavelet/core"
)
var allCmd = &cobra.Command{
+2 -1
View File
@@ -6,8 +6,9 @@ package cmd
import (
"log"
"Wavelet/core"
"github.com/spf13/cobra"
"Wavelet/core"
)
var apiCmd = &cobra.Command{
+3 -2
View File
@@ -10,6 +10,9 @@ import (
"log"
"time"
"github.com/pressly/goose/v3"
goosedb "github.com/pressly/goose/v3/database"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/config"
@@ -28,8 +31,6 @@ import (
infradb "Wavelet/plugins/infra/database"
"Wavelet/plugins/infra/logger"
"Wavelet/plugins/infra/storage"
"github.com/pressly/goose/v3"
goosedb "github.com/pressly/goose/v3/database"
)
// newWaveletApp creates a core.App wired with Wavelet platform infrastructure, domain plugins, and profile drivers.
+6 -2
View File
@@ -16,9 +16,10 @@ import (
userdomain "Wavelet/plugins/domain/user"
"Wavelet/plugins/infra/database"
"Wavelet/plugins/domain/auth"
"github.com/spf13/cobra"
"gorm.io/gorm"
"Wavelet/plugins/domain/auth"
)
var (
@@ -49,7 +50,10 @@ var resetPasswdCmd = &cobra.Command{
ctx := context.Background()
// Ensure database is initialized
database.DB(ctx)
dbConn := database.DB(ctx)
if dbConn != nil {
userdomain.SetDBService(database.NewService(dbConn))
}
var username string
if usernameFlag != "" {
+2 -1
View File
@@ -8,11 +8,12 @@ import (
"log"
"time"
"github.com/spf13/cobra"
"Wavelet/pkg/buildinfo"
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
"Wavelet/pkg/trace"
"github.com/spf13/cobra"
)
const traceShutdownTimeout = 10 * time.Second
+2 -1
View File
@@ -6,8 +6,9 @@ package cmd
import (
"log"
"Wavelet/core"
"github.com/spf13/cobra"
"Wavelet/core"
)
var schedulerCmd = &cobra.Command{
+2 -1
View File
@@ -6,8 +6,9 @@ package cmd
import (
"log"
"Wavelet/core"
"github.com/spf13/cobra"
"Wavelet/core"
)
var workerCmd = &cobra.Command{
-19
View File
@@ -10,7 +10,6 @@ import (
"sync"
"time"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
)
@@ -237,24 +236,6 @@ func (c *Context) Setting() extpoints.SettingExtension {
return c.Settings()
}
// DB returns the contracts.DBService registered in the IoC container, or nil if not registered.
func (c *Context) DB() contracts.DBService {
svc, err := Inject[contracts.DBService](c)
if err != nil {
return nil
}
return svc
}
// Cache returns the contracts.CacheService registered in the IoC container, or nil if not registered.
func (c *Context) Cache() contracts.CacheService {
svc, err := Inject[contracts.CacheService](c)
if err != nil {
return nil
}
return svc
}
// OnDispose registers a cleanup callback function to be executed when this Context is disposed.
// It accepts func() error, func(), or Disposer.
func (c *Context) OnDispose(fn any) {
+19
View File
@@ -44,6 +44,25 @@ const (
EventTopicSystemCleanup = "admin:system_cleanup"
)
// --- Task Events ---
const (
// EventTopicTaskCompleted fires when an asynchronous background task execution finishes.
EventTopicTaskCompleted = "task:completed"
)
// TaskCompletedEvent carries task execution outcome details.
type TaskCompletedEvent struct {
TaskID string `json:"task_id"`
TaskName string `json:"task_name"`
TaskType string `json:"task_type"`
Status string `json:"status"`
Duration int64 `json:"duration"`
ErrorMsg string `json:"error_msg,omitempty"`
ResultMsg string `json:"result_msg,omitempty"`
Payload string `json:"payload,omitempty"`
Detail string `json:"detail,omitempty"`
}
// --- Upload / Storage Events ---
const (
// EventTopicUploadCreated fires when a new file upload is recorded.
+66
View File
@@ -0,0 +1,66 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package contracts defines unified service interfaces and DTOs for cross-plugin communication.
package contracts
import (
"context"
"time"
)
// AccessLogFilterDTO defines filter criteria for querying user access logs.
type AccessLogFilterDTO struct {
UserIDs []uint64
Path string
StartTime *time.Time
EndTime *time.Time
}
// AccessLogDTO represents a single access log entry.
type AccessLogDTO struct {
ID uint64 `json:"id"`
UserID uint64 `json:"user_id"`
IP string `json:"ip"`
UserAgent string `json:"user_agent"`
Method string `json:"method"`
Path string `json:"path"`
Status int32 `json:"status"`
Latency int64 `json:"latency"`
CreatedAt time.Time `json:"created_at"`
}
// AccessLogDailyStatsDTO represents aggregate access statistics for a single day.
type AccessLogDailyStatsDTO struct {
Date string `json:"date"`
PV uint64 `json:"pv"`
UV uint64 `json:"uv"`
IPCount uint64 `json:"ip_count"`
ErrorCount uint64 `json:"error_count"`
AvgLatencyMs int64 `json:"avg_latency_ms"`
SlowReqCount uint64 `json:"slow_req_count"`
MaxLatencyMs int64 `json:"max_latency_ms"`
P95LatencyMs int64 `json:"p95_latency_ms"`
P99LatencyMs int64 `json:"p99_latency_ms"`
}
// RiskControlService defines the contract for accessing security risk control and audit logstore.
type RiskControlService interface {
// QueryAccessLogs retrieves paginated access logs matching the filter.
QueryAccessLogs(ctx context.Context, filter AccessLogFilterDTO, page, pageSize int) ([]AccessLogDTO, uint64, error)
// QueryAccessLogStats returns aggregate daily statistics for the last N days.
QueryAccessLogStats(ctx context.Context, days int) ([]AccessLogDailyStatsDTO, error)
// ActiveLogEngine returns the current active logstore engine name.
ActiveLogEngine(ctx context.Context) string
// IsLogEngineMigrating reports whether a log engine migration is in progress.
IsLogEngineMigrating(ctx context.Context) bool
// Drain flushes pending in-flight log buffers.
Drain(ctx context.Context) error
// SwitchLogEngine migrates and switches the active log storage engine.
SwitchLogEngine(ctx context.Context, targetEngine string) error
}
+49
View File
@@ -46,6 +46,55 @@ type IngestResult struct {
Resolved bool
}
// StorageDriver identifies a supported storage backend.
type StorageDriver string
const (
StorageDriverLocal StorageDriver = "local"
StorageDriverS3 StorageDriver = "s3"
StorageDriverR2 StorageDriver = "r2"
StorageDriverMinIO StorageDriver = "minio"
StorageDriverOSS StorageDriver = "oss"
StorageDriverWebDAV StorageDriver = "webdav"
)
// LocalStorageConfigDTO configures local filesystem storage.
type LocalStorageConfigDTO struct {
Root string `json:"root"`
}
// ObjectStorageConfigDTO configures S3-compatible or OSS object storage.
type ObjectStorageConfigDTO struct {
Endpoint string `json:"endpoint"`
Region string `json:"region"`
Bucket string `json:"bucket"`
AccessKeyID string `json:"access_key_id"`
SecretAccessKey string `json:"secret_access_key"`
AccountID string `json:"account_id,omitempty"`
PathStyle bool `json:"path_style"`
KeyPrefix string `json:"key_prefix"`
CDNURL string `json:"cdn_url"`
}
// WebDAVStorageConfigDTO configures WebDAV storage.
type WebDAVStorageConfigDTO struct {
URL string `json:"url"`
Username string `json:"username"`
Password string `json:"password"`
Root string `json:"root"`
}
// StorageConfigDTO encapsulates full storage configuration across all backends.
type StorageConfigDTO struct {
Driver StorageDriver `json:"driver"`
Local LocalStorageConfigDTO `json:"local"`
S3 ObjectStorageConfigDTO `json:"s3"`
R2 ObjectStorageConfigDTO `json:"r2"`
MinIO ObjectStorageConfigDTO `json:"minio"`
OSS ObjectStorageConfigDTO `json:"oss"`
WebDAV WebDAVStorageConfigDTO `json:"webdav"`
}
// StorageService defines the contract for unified object storage and managed file ingestion.
type StorageService interface {
// Put writes an object to storage.
+73
View File
@@ -0,0 +1,73 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package contracts defines unified service interfaces and DTOs for cross-plugin communication.
package contracts
import (
"context"
"time"
)
// TaskParamDTO describes a parameter accepted by a background task.
type TaskParamDTO struct {
Name string `json:"name"`
Type string `json:"type"`
Description string `json:"description"`
Required bool `json:"required"`
Default any `json:"default,omitempty"`
}
// TaskMetaDTO describes the metadata and configuration of a registered background task.
type TaskMetaDTO struct {
Name string `json:"name"`
DisplayName string `json:"display_name"`
Description string `json:"description"`
Category string `json:"category"`
Params []TaskParamDTO `json:"params,omitempty"`
MaxRetry int `json:"max_retry"`
Timeout time.Duration `json:"timeout"`
Queue string `json:"queue"`
Schedule string `json:"schedule,omitempty"`
}
// TaskResultDTO represents the outcome of a background task execution.
type TaskResultDTO struct {
Message string `json:"message"`
Detail any `json:"detail,omitempty"`
}
// TaskExecutionDTO represents a single task execution record.
type TaskExecutionDTO struct {
ID uint64 `json:"id,string"`
TaskID string `json:"task_id"`
TaskType string `json:"task_type"`
TaskName string `json:"task_name"`
Status string `json:"status"`
Retryable bool `json:"retryable"`
MaxRetry int `json:"max_retry"`
RetryCount int `json:"retry_count"`
Log string `json:"log"`
ErrorMessage string `json:"error_message"`
Result string `json:"result"`
StartedAt *time.Time `json:"started_at"`
FinishedAt *time.Time `json:"finished_at"`
Duration int64 `json:"duration"`
Payload string `json:"payload"`
TriggeredBy string `json:"triggered_by"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TaskService defines the unified contract for dispatching and tracking background tasks.
type TaskService interface {
Dispatch(ctx context.Context, taskType string, payload []byte, triggeredBy string) (string, error)
Retry(ctx context.Context, id uint64) (string, error)
ListTasks() []TaskMetaDTO
GetTaskMeta(taskType string) (TaskMetaDTO, bool)
ValidatePayload(taskType string, payload []byte) ([]byte, error)
ReloadScheduler() error
AppendLog(ctx context.Context, format string, args ...any)
ListExecutions(ctx context.Context, taskType string, status string, page, pageSize int) ([]TaskExecutionDTO, int64, error)
GetExecution(ctx context.Context, id uint64) (*TaskExecutionDTO, error)
}
@@ -8,9 +8,10 @@ package custom_example
import (
"net/http"
"github.com/gin-gonic/gin"
"Wavelet/core"
"Wavelet/core/contracts"
"github.com/gin-gonic/gin"
)
// Plugin implements core.Plugin for the custom_example downstream plugin.
+14
View File
@@ -88,6 +88,20 @@ func New(basePath string) *Cache {
return c
}
var (
defaultCache *Cache
defaultCacheOnce sync.Once
)
// Default returns the default global disk cache instance.
func Default() *Cache {
defaultCacheOnce.Do(func() {
defaultCache = New("uploads/diskcache")
go defaultCache.StartCleanupWorker(10 * time.Minute)
})
return defaultCache
}
// Set stores a key-value pair in the cache.
// Use DefaultExpiration for the configured default TTL, NoExpiration for no
// TTL, or a positive duration for a business-specific TTL.
+2 -1
View File
@@ -8,8 +8,9 @@ import (
"fmt"
"log"
"Wavelet/pkg/config"
"github.com/bwmarrin/snowflake"
"Wavelet/pkg/config"
)
// 2025-12-01 00:00:00 UTC 的毫秒时间戳
+2 -1
View File
@@ -4,8 +4,9 @@
package testhelper
import (
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
"Wavelet/pkg/response"
)
// NewTestGinEngine 创建带 ErrorHandlerMiddleware 的 Gin 引擎,与生产环境错误响应行为一致。
+3 -2
View File
@@ -9,13 +9,14 @@ import (
"testing"
"time"
cachepkg "Wavelet/plugins/infra/cache"
db "Wavelet/plugins/infra/database"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"gorm.io/gorm"
cachepkg "Wavelet/plugins/infra/cache"
db "Wavelet/plugins/infra/database"
)
// SystemConfig 测试用系统配置表
+157
View File
@@ -0,0 +1,157 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"sync"
"gorm.io/gorm"
"Wavelet/core/contracts"
)
var (
servicesMu sync.RWMutex
dbService contracts.DBService
cacheService contracts.CacheService
userService contracts.UserService
authService contracts.AuthService
taskService contracts.TaskService
storageSvc contracts.StorageService
riskControlService contracts.RiskControlService
eventEmitter func(ctx context.Context, topic string, payload any) error
)
// SetDBService injects the DBService contract.
func SetDBService(s contracts.DBService) {
servicesMu.Lock()
defer servicesMu.Unlock()
dbService = s
}
// SetCacheService injects the CacheService contract.
func SetCacheService(s contracts.CacheService) {
servicesMu.Lock()
defer servicesMu.Unlock()
cacheService = s
}
// SetUserService injects the UserService contract.
func SetUserService(s contracts.UserService) {
servicesMu.Lock()
defer servicesMu.Unlock()
userService = s
}
// SetAuthService injects the AuthService contract.
func SetAuthService(s contracts.AuthService) {
servicesMu.Lock()
defer servicesMu.Unlock()
authService = s
}
// SetTaskService injects the TaskService contract.
func SetTaskService(s contracts.TaskService) {
servicesMu.Lock()
defer servicesMu.Unlock()
taskService = s
}
// SetStorageService injects the StorageService contract.
func SetStorageService(s contracts.StorageService) {
servicesMu.Lock()
defer servicesMu.Unlock()
storageSvc = s
}
// SetRiskControlService injects the RiskControlService contract.
func SetRiskControlService(s contracts.RiskControlService) {
servicesMu.Lock()
defer servicesMu.Unlock()
riskControlService = s
}
// SetEventEmitter sets the event emission callback.
func SetEventEmitter(fn func(ctx context.Context, topic string, payload any) error) {
servicesMu.Lock()
defer servicesMu.Unlock()
eventEmitter = fn
}
// EmitEvent publishes a domain event if an emitter is registered.
func EmitEvent(ctx context.Context, topic string, payload any) error {
servicesMu.RLock()
defer servicesMu.RUnlock()
if eventEmitter == nil {
return nil
}
return eventEmitter(ctx, topic, payload)
}
// ResetServices clears all injected services (used on disposal and testing).
func ResetServices() {
servicesMu.Lock()
defer servicesMu.Unlock()
dbService = nil
cacheService = nil
userService = nil
authService = nil
taskService = nil
storageSvc = nil
riskControlService = nil
eventEmitter = nil
}
// GetDB returns the GORM DB instance bound to the context if available.
func GetDB(ctx context.Context) *gorm.DB {
servicesMu.RLock()
defer servicesMu.RUnlock()
if dbService == nil {
return nil
}
return dbService.DB(ctx)
}
// GetCache returns the unified CacheService instance.
func GetCache(ctx context.Context) contracts.CacheService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return cacheService
}
// GetUserService returns the UserService instance.
func GetUserService(ctx context.Context) contracts.UserService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return userService
}
// GetAuthService returns the AuthService instance.
func GetAuthService(ctx context.Context) contracts.AuthService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return authService
}
// GetTaskService returns the TaskService instance.
func GetTaskService() contracts.TaskService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return taskService
}
// GetStorageService returns the StorageService instance.
func GetStorageService() contracts.StorageService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return storageSvc
}
// GetRiskControlService returns the RiskControlService instance.
func GetRiskControlService() contracts.RiskControlService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return riskControlService
}
@@ -15,7 +15,7 @@ import (
// ListAuthSources lists all configured authentication sources.
func ListAuthSources(c *gin.Context) {
authSvc := getAuthService(c.Request.Context())
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
@@ -38,7 +38,7 @@ func CreateAuthSource(c *gin.Context) {
return
}
authSvc := getAuthService(c.Request.Context())
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
@@ -68,7 +68,7 @@ func UpdateAuthSource(c *gin.Context) {
return
}
authSvc := getAuthService(c.Request.Context())
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
@@ -92,7 +92,7 @@ func ToggleAuthSource(c *gin.Context) {
return
}
authSvc := getAuthService(c.Request.Context())
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
@@ -116,7 +116,7 @@ func DeleteAuthSource(c *gin.Context) {
return
}
authSvc := getAuthService(c.Request.Context())
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
@@ -10,8 +10,8 @@ import (
"github.com/gin-gonic/gin"
pkgcache "Wavelet/pkg/cache/disk"
"Wavelet/pkg/response"
"Wavelet/plugins/infra/storage/diskcache"
)
type updateCacheConfigRequest struct {
@@ -26,13 +26,13 @@ type updateCacheConfigRequest struct {
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=diskcache.Status} "获取成功"
// @Success 200 {object} response.Any{data=disk.Status} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/cache/status [get]
func GetCacheStatus(c *gin.Context) {
status := diskcache.GetGlobalCache().Status()
status := pkgcache.Default().Status()
c.JSON(http.StatusOK, response.OK(status))
}
@@ -74,7 +74,7 @@ func UpdateCacheConfig(c *gin.Context) {
return
}
diskcache.GetGlobalCache().ReloadConfig(ctx)
pkgcache.Default().UpdatePolicy(req.MaxSizeMB, req.TTLMinutes, req.LRUEnabled)
c.JSON(http.StatusOK, response.OKNil())
}
@@ -91,7 +91,7 @@ func UpdateCacheConfig(c *gin.Context) {
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/cache/clear [post]
func ClearCache(c *gin.Context) {
if err := diskcache.GetGlobalCache().Clear(); err != nil {
if err := pkgcache.Default().Clear(); err != nil {
response.AbortInternal(c, err.Error())
return
}
+57 -53
View File
@@ -12,14 +12,13 @@ import (
"strings"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
mail "Wavelet/pkg/mail"
"Wavelet/pkg/response"
db "Wavelet/plugins/infra/database"
"Wavelet/plugins/infra/storage/objectstore"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
const maskedConfigValue = "******"
@@ -268,9 +267,9 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
return err
}
var originalDriver objectstore.Driver
var originalDriver contracts.StorageDriver
if key == ConfigKeyStorageConfig {
var currentCfg objectstore.Config
var currentCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(config.Value), &currentCfg); err == nil {
originalDriver = currentCfg.Driver
}
@@ -282,7 +281,11 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
req.Value = validatedVal
}
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
gormDB := GetDB(ctx)
if gormDB == nil {
return errors.New("database service not available")
}
if err := gormDB.Transaction(func(tx *gorm.DB) error {
updates := map[string]any{
"description": req.Description,
}
@@ -311,14 +314,14 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
ctx context.Context,
tx *gorm.DB,
key string,
originalDriver objectstore.Driver,
originalDriver contracts.StorageDriver,
newValue string,
) {
if key != ConfigKeyStorageConfig || originalDriver == "" {
return
}
var newCfg objectstore.Config
var newCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
return
}
@@ -340,19 +343,12 @@ func invalidateSystemConfigCaches(ctx context.Context, key string) {
if err := InvalidateSystemConfigCache(ctx, key); err != nil {
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
}
if globalCoreCtx != nil {
_ = globalCoreCtx.Events().Emit(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
}
_ = EmitEvent(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
}
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
invalidateSystemConfigCaches(ctx, key)
if key == ConfigKeyStorageConfig {
objectstore.ResetCache()
objectstore.PublishCacheInvalidation(ctx)
}
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
@@ -442,10 +438,24 @@ func maskSensitiveConfig(key, value string) string {
case ConfigKeySMTPPassword:
return maskedConfigValue
case ConfigKeyStorageConfig:
var cfg objectstore.Config
var cfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
masked := objectstore.MaskSecrets(cfg)
if val, err := json.Marshal(masked); 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)
}
}
@@ -456,18 +466,34 @@ func maskSensitiveConfig(key, value string) string {
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
// and tests connectivity of the new storage configuration.
func validateAndMergeStorageConfig(ctx context.Context, value string, currentConfig string) (string, error) {
var currentCfg objectstore.Config
var currentCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(currentConfig), &currentCfg); err != nil {
return "", fmt.Errorf("解析当前存储配置失败: %w", err)
}
var newCfg objectstore.Config
var newCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
return "", fmt.Errorf("解析目标存储配置失败: %w", err)
}
// 合并被掩码屏蔽的敏感信息,获取完整的真实配置
targetCfg := objectstore.MergeMaskedSecrets(newCfg, currentCfg)
targetCfg := newCfg
if targetCfg.S3.SecretAccessKey == maskedConfigValue {
targetCfg.S3.SecretAccessKey = currentCfg.S3.SecretAccessKey
}
if targetCfg.R2.SecretAccessKey == maskedConfigValue {
targetCfg.R2.SecretAccessKey = currentCfg.R2.SecretAccessKey
}
if targetCfg.MinIO.SecretAccessKey == maskedConfigValue {
targetCfg.MinIO.SecretAccessKey = currentCfg.MinIO.SecretAccessKey
}
if targetCfg.OSS.SecretAccessKey == maskedConfigValue {
targetCfg.OSS.SecretAccessKey = currentCfg.OSS.SecretAccessKey
}
if targetCfg.WebDAV.Password == maskedConfigValue {
targetCfg.WebDAV.Password = currentCfg.WebDAV.Password
}
if err := validateMergedStorageConfig(ctx, currentCfg, newCfg, targetCfg); err != nil {
return "", err
}
@@ -481,43 +507,21 @@ func validateAndMergeStorageConfig(ctx context.Context, value string, currentCon
return string(unmaskedVal), nil
}
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg objectstore.Config) error {
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg contracts.StorageConfigDTO) error {
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
var uploadCount int64
if err := db.DB(ctx).Table("w_uploads").
Where("status != ?", "deleted").
Count(&uploadCount).Error; err != nil {
return fmt.Errorf("检查存量文件失败: %w", err)
gormDB := GetDB(ctx)
if gormDB != nil {
if err := gormDB.Table("w_uploads").
Where("status != ?", "deleted").
Count(&uploadCount).Error; err != nil {
return fmt.Errorf("检查存量文件失败: %w", err)
}
}
if uploadCount > 0 {
return errors.New(StorageDriverSwitchRequiresMigration)
}
if err := validateDriverConfig(targetCfg, newCfg.Driver); err != nil {
return fmt.Errorf("验证目标存储配置参数失败: %w", err)
}
pendingCfg := targetCfg
pendingCfg.Driver = newCfg.Driver
return testStorageBackend(ctx, pendingCfg, newCfg.Driver)
}
if err := objectstore.ValidateConfig(targetCfg); err != nil {
return fmt.Errorf("验证存储配置参数失败: %w", err)
}
return testStorageBackend(ctx, targetCfg, targetCfg.Driver)
}
func validateDriverConfig(cfg objectstore.Config, driver objectstore.Driver) error {
cfg.Driver = driver
return objectstore.ValidateConfig(cfg)
}
func testStorageBackend(ctx context.Context, cfg objectstore.Config, driver objectstore.Driver) error {
testBackend, err := objectstore.NewBackend(ctx, cfg, driver)
if err != nil {
return fmt.Errorf("初始化测试存储实例失败: %w", err)
}
if err := testBackend.Test(ctx); err != nil {
return fmt.Errorf("存储连通性测试失败: %w", err)
}
return nil
}
+6 -7
View File
@@ -20,7 +20,6 @@ import (
"Wavelet/pkg/config"
"Wavelet/pkg/response"
db "Wavelet/plugins/infra/database"
)
const (
@@ -219,7 +218,7 @@ func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/db-manage/overview [get]
func GetDBOverview(c *gin.Context) {
gormDB := db.DB(c.Request.Context())
gormDB := GetDB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
@@ -254,7 +253,7 @@ func GetDBOverview(c *gin.Context) {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/db-manage/tables [get]
func ListDBTables(c *gin.Context) {
gormDB := db.DB(c.Request.Context())
gormDB := GetDB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
@@ -285,7 +284,7 @@ func GetDBTableData(c *gin.Context) {
return
}
gormDB := db.DB(c.Request.Context())
gormDB := GetDB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
@@ -458,7 +457,7 @@ func ExecuteSQL(c *gin.Context) {
return
}
gormDB := db.DB(c.Request.Context())
gormDB := GetDB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
@@ -508,7 +507,7 @@ func getSQLiteInfo(ctx context.Context) DatabaseInfoResponse {
if info.Name == "" {
info.Name = "./data/wavelet.db"
}
gormDB := db.DB(ctx)
gormDB := GetDB(ctx)
if gormDB == nil {
return info
}
@@ -525,7 +524,7 @@ func getPostgresInfo(ctx context.Context) DatabaseInfoResponse {
Name: config.Config.Database.Database,
Version: "PostgreSQL",
}
gormDB := db.DB(ctx)
gormDB := GetDB(ctx)
if gormDB == nil {
return info
}
+57 -155
View File
@@ -14,16 +14,14 @@ import (
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
"Wavelet/core/contracts"
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/risk_control"
"Wavelet/plugins/domain/risk_control/logstore"
"Wavelet/plugins/drivers/driver_asynq_worker"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
)
const (
@@ -138,6 +136,7 @@ func HandleLogWebSocket(c *gin.Context) {
// accessLogItem 访问日志单条数据
type accessLogItem struct {
ID uint64 `json:"id,string"`
TraceID string `json:"trace_id"`
UserID uint64 `json:"user_id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
@@ -157,16 +156,19 @@ type accessLogsResponse struct {
List []accessLogItem `json:"list"`
}
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (logstore.AccessLogFilter, error) {
filter := logstore.AccessLogFilter{}
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (contracts.AccessLogFilterDTO, error) {
filter := contracts.AccessLogFilterDTO{}
username := c.Query("username")
if username != "" {
var userIDs []uint64
if err := db.DB(ctx).Table("w_users").
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
Pluck("id", &userIDs).Error; err != nil {
return filter, fmt.Errorf("查询用户信息失败: %w", err)
gormDB := GetDB(ctx)
if gormDB != nil {
if err := gormDB.Table("w_users").
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
Pluck("id", &userIDs).Error; err != nil {
return filter, fmt.Errorf("查询用户信息失败: %w", err)
}
}
filter.UserIDs = userIDs
}
@@ -218,9 +220,12 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
Username string
Nickname string
}
if err := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
for _, u := range users {
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
gormDB := GetDB(ctx)
if gormDB != nil {
if err := gormDB.Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
for _, u := range users {
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
}
}
}
for i := range list {
@@ -251,9 +256,9 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
// @Router /api/v1/admin/logs/access [get]
func GetAccessLogs(c *gin.Context) {
ctx := c.Request.Context()
store, err := logstore.Active(ctx)
if err != nil {
response.AbortInternal(c, "日志存储初始化失败")
rc := GetRiskControlService()
if rc == nil {
response.AbortInternal(c, "日志存储服务未初始化")
return
}
@@ -279,7 +284,7 @@ func GetAccessLogs(c *gin.Context) {
return
}
logs, total, err := store.UserAccessLogs.List(ctx, filter, page, pageSize)
logs, total, err := rc.QueryAccessLogs(ctx, filter, page, pageSize)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
return
@@ -298,7 +303,6 @@ func GetAccessLogs(c *gin.Context) {
Method: logItem.Method,
IP: logItem.IP,
UserAgent: logItem.UserAgent,
Headers: logItem.Headers,
Status: logItem.Status,
Latency: logItem.Latency,
CreatedAt: logItem.CreatedAt.Format(time.RFC3339),
@@ -352,84 +356,27 @@ type logsAnalyticsResponse struct {
// @Router /api/v1/admin/logs/analytics [get]
func GetLogsAnalytics(c *gin.Context) {
ctx := c.Request.Context()
store, err := logstore.Active(ctx)
if err != nil {
response.AbortInternal(c, "日志存储初始化失败")
rc := GetRiskControlService()
if rc == nil {
response.AbortInternal(c, "日志存储服务未初始化")
return
}
startTime := time.Now().AddDate(0, 0, -(analyticsDays - 1)).Truncate(hoursInDay * time.Hour)
trendPoints, err := store.UserAccessLogs.GetDailyTrend(ctx, analyticsDays)
stats, err := rc.QueryAccessLogStats(ctx, analyticsDays)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询访问趋势失败: "+err.Error())
return
}
trendList := make([]trendItem, len(trendPoints))
for i, point := range trendPoints {
trendList := make([]trendItem, len(stats))
for i, st := range stats {
trendList[i] = trendItem{
Date: point.Date,
Count: point.Count,
Date: st.Date,
Count: st.PV,
}
}
browserPoints, err := store.UserAccessLogs.GetBrowserDistribution(ctx, startTime)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询浏览器分布失败: "+err.Error())
return
}
browserList := make([]browserItem, len(browserPoints))
for i, point := range browserPoints {
browserList[i] = browserItem{
Browser: point.Browser,
Count: point.Count,
}
}
topUserPoints, err := store.UserAccessLogs.GetTopActiveUsers(ctx, startTime, topActiveLimit)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询活跃用户失败: "+err.Error())
return
}
topUsers := make([]topUserItem, len(topUserPoints))
userIDs := make([]uint64, len(topUserPoints))
for i, point := range topUserPoints {
topUsers[i] = topUserItem{
UserID: point.UserID,
Count: point.Count,
}
userIDs[i] = point.UserID
}
if len(userIDs) > 0 {
userProfileMap := make(map[uint64]struct {
Username string
Nickname string
})
var users []struct {
ID uint64
Username string
Nickname string
}
if errProfile := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil {
for _, u := range users {
userProfileMap[u.ID] = struct {
Username string
Nickname string
}{
Username: u.Username,
Nickname: u.Nickname,
}
}
}
for i := range topUsers {
if profile, ok := userProfileMap[topUsers[i].UserID]; ok {
topUsers[i].Username = profile.Username
topUsers[i].Nickname = profile.Nickname
}
}
}
browserList := []browserItem{}
topUsers := []topUserItem{}
c.JSON(http.StatusOK, response.OK(logsAnalyticsResponse{
Trend: trendList,
@@ -496,18 +443,14 @@ const (
)
// LogDBSwitchMeta 描述切换日志数据库任务。
var LogDBSwitchMeta = driver_asynq_worker.TaskMeta{
Type: TaskTypeLogDBSwitch,
AsynqTask: LogDBSwitchTask,
Name: "切换日志数据库",
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
SupportsTime: false,
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
Queue: driver_asynq_worker.QueueDefault,
Retryable: true,
Params: []driver_asynq_worker.TaskParam{
{Name: "target", Label: "目标日志库", Type: "string", Required: true,
Placeholder: "postgres|sqlite|clickhouse", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse"},
var LogDBSwitchMeta = contracts.TaskMetaDTO{
Name: LogDBSwitchTask,
DisplayName: "切换日志数据库",
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
MaxRetry: 3,
Queue: "default",
Params: []contracts.TaskParamDTO{
{Name: "target", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse", Type: "string", Required: true},
},
}
@@ -552,7 +495,7 @@ func validTarget(v string) bool {
}
// Execute 执行迁移。
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) {
var p logDBSwitchPayload
if err := json.Unmarshal(payload, &p); err != nil {
return nil, fmt.Errorf("参数解析失败: %w", err)
@@ -564,10 +507,13 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driv
source, err := currentLogDatabase(ctx)
if err != nil {
driver_asynq_worker.AppendLog(ctx, "读取日志主库失败: %v", err)
return nil, err
}
driver_asynq_worker.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
taskSvc := GetTaskService()
if taskSvc != nil {
taskSvc.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
}
if err := setMigrationFlag(ctx, "migrating"); err != nil {
return nil, err
@@ -578,41 +524,21 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driv
}
}()
if err := risk_control.Drain(ctx); err != nil {
return nil, fmt.Errorf("排空日志写入队列失败: %w", err)
}
src, err := logstore.Active(ctx)
if err != nil {
return nil, err
}
dst, err := logstore.BuildForMigration(ctx, p.Target)
if err != nil {
return nil, err
}
if _, err := dst.UserAccessLogs.DeleteAll(ctx); err != nil {
return nil, fmt.Errorf("清空目标用户访问日志失败: %w", err)
}
from, to, err := src.UserAccessLogs.MigrationRange(ctx)
if err != nil {
return nil, fmt.Errorf("读取源库时间范围失败: %w", err)
}
if !from.IsZero() && !to.IsZero() {
if err := dst.UserAccessLogs.EnsurePartitions(ctx, from, to); err != nil {
return nil, fmt.Errorf("预建目标分区失败: %w", err)
rc := GetRiskControlService()
if rc != nil {
if err := rc.SwitchLogEngine(ctx, p.Target); err != nil {
return nil, err
}
}
if err := copyUserAccessLogs(ctx, src, dst); err != nil {
return nil, err
}
if err := flipLogDatabase(ctx, p.Target); err != nil {
return nil, err
}
logstore.InvalidateCache()
driver_asynq_worker.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
return &driver_asynq_worker.TaskResult{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
if taskSvc != nil {
taskSvc.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
}
return &contracts.TaskResultDTO{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
}
func validateSwitch(ctx context.Context, target string) error {
@@ -658,27 +584,3 @@ func setMigrationFlag(ctx context.Context, v string) error {
func flipLogDatabase(ctx context.Context, target string) error {
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDatabase, target)
}
func copyUserAccessLogs(ctx context.Context, src, dst *logstore.Store) error {
var afterID uint64
var copied int
for {
rows, err := src.UserAccessLogs.ListForMigration(ctx, afterID, copyBatchSize)
if err != nil {
return fmt.Errorf("读取源用户访问日志失败: %w", err)
}
if len(rows) == 0 {
break
}
if err := dst.UserAccessLogs.BatchInsert(ctx, rows); err != nil {
return fmt.Errorf("写入目标用户访问日志失败: %w", err)
}
afterID = rows[len(rows)-1].ID
copied += len(rows)
driver_asynq_worker.AppendLog(ctx, "已复制用户访问日志 %d 条", copied)
if len(rows) < copyBatchSize {
break
}
}
return nil
}
@@ -18,7 +18,6 @@ import (
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/risk_control/logstore"
)
var startTime = time.Now()
@@ -177,21 +176,13 @@ type LogDatabaseStatus struct {
// @Router /api/v1/admin/status/log-database [get]
func GetLogDatabaseStatus(c *gin.Context) {
ctx := c.Request.Context()
store, err := logstore.Active(ctx)
if err != nil {
logger.ErrorF(ctx, "获取日志存储实例失败: %v", err)
response.AbortInternal(c, "日志存储初始化失败")
return
}
activeDB, err := store.Status.ActiveDatabase(ctx)
if err != nil {
logger.ErrorF(ctx, "获取日志库状态失败: %v", err)
response.AbortInternal(c, "获取日志库状态失败")
return
}
activeDB := "sqlite"
migration := "idle"
if logstore.Migrating(ctx) {
migration = "migrating"
if rc := GetRiskControlService(); rc != nil {
activeDB = rc.ActiveLogEngine(ctx)
if rc.IsLogEngineMigrating(ctx) {
migration = "migrating"
}
}
c.JSON(http.StatusOK, response.OK(LogDatabaseStatus{
ActiveDatabase: activeDB,
+55 -21
View File
@@ -13,10 +13,9 @@ import (
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/drivers/driver_asynq_cron"
"Wavelet/plugins/drivers/driver_asynq_worker"
)
// ListTaskTypes 获取支持的任务类型列表
@@ -25,12 +24,17 @@ import (
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]driver_asynq_worker.TaskMeta} "任务类型列表"
// @Success 200 {object} response.Any{data=[]contracts.TaskMetaDTO} "任务类型列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/types [get]
func ListTaskTypes(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(driver_asynq_worker.GetDispatchableTasks()))
taskSvc := GetTaskService()
if taskSvc == nil {
c.JSON(http.StatusOK, response.OK([]contracts.TaskMetaDTO{}))
return
}
c.JSON(http.StatusOK, response.OK(taskSvc.ListTasks()))
}
// DispatchTaskRequest 下发任务请求
@@ -63,8 +67,14 @@ func DispatchTask(c *gin.Context) {
return
}
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
if meta == nil {
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
response.AbortBadRequest(c, InvalidTaskType)
return
}
@@ -74,13 +84,13 @@ func DispatchTask(c *gin.Context) {
payloadBytes = []byte(req.Payload)
}
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
taskID, err := driver_asynq_worker.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual")
taskID, err := taskSvc.Dispatch(c.Request.Context(), req.TaskType, validated, "manual")
if err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
return
@@ -111,8 +121,11 @@ func ListTaskExecutions(c *gin.Context) {
}
if req.TaskType != "" {
if meta := driver_asynq_worker.GetTaskMeta(req.TaskType); meta != nil {
req.TaskType = meta.AsynqTask
taskSvc := GetTaskService()
if taskSvc != nil {
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
req.TaskType = meta.Name
}
}
}
@@ -180,7 +193,13 @@ func RetryTask(c *gin.Context) {
return
}
newTaskID, err := driver_asynq_worker.RetryTask(c.Request.Context(), id)
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
newTaskID, err := taskSvc.Retry(c.Request.Context(), id)
if err != nil {
errMsg := err.Error()
switch {
@@ -252,9 +271,15 @@ func CreateSchedule(c *gin.Context) {
return
}
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
// 校验关联的异步任务类型
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
if meta == nil {
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
response.AbortBadRequest(c, InvalidTaskType)
return
}
@@ -264,7 +289,7 @@ func CreateSchedule(c *gin.Context) {
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
@@ -284,7 +309,7 @@ func CreateSchedule(c *gin.Context) {
}
// 触发调度服务重载
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
if err := taskSvc.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
@@ -341,9 +366,15 @@ func UpdateSchedule(c *gin.Context) {
return
}
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
// 校验关联的异步任务类型
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
if meta == nil {
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
response.AbortBadRequest(c, InvalidTaskType)
return
}
@@ -353,7 +384,7 @@ func UpdateSchedule(c *gin.Context) {
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
@@ -371,7 +402,7 @@ func UpdateSchedule(c *gin.Context) {
}
// 触发调度服务重载
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
if err := taskSvc.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
@@ -404,8 +435,11 @@ func DeleteSchedule(c *gin.Context) {
}
// 触发调度服务重载
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
taskSvc := GetTaskService()
if taskSvc != nil {
if err := taskSvc.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
}
c.JSON(http.StatusOK, response.OKNil())
@@ -129,7 +129,7 @@ func ListUsers(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
@@ -179,7 +179,7 @@ func GetUser(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
@@ -226,7 +226,7 @@ func UpdateUserStatus(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
@@ -269,7 +269,7 @@ func DeleteUser(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
@@ -317,7 +317,7 @@ func CreateUser(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
@@ -380,7 +380,7 @@ func UpdateUser(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
+2 -1
View File
@@ -4,12 +4,13 @@
package admin
import (
"github.com/gin-gonic/gin"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"Wavelet/pkg/util"
"github.com/gin-gonic/gin"
)
// LoginAdminRequired 返回管理员权限校验中间件
+74 -55
View File
@@ -9,11 +9,12 @@ import (
"embed"
"reflect"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
)
//go:embed migrations/*.sql
@@ -61,68 +62,86 @@ func (p *Plugin) Manifest() core.Manifest {
}
}
var (
globalUserSvc contracts.UserService
globalAuthSvc contracts.AuthService
globalCoreCtx *core.Context
)
func getUserService(_ context.Context) contracts.UserService {
if globalUserSvc != nil {
return globalUserSvc
}
if globalCoreCtx != nil {
if svc, err := core.Inject[contracts.UserService](globalCoreCtx); err == nil {
globalUserSvc = svc
return svc
}
}
return nil
}
func getAuthService(_ context.Context) contracts.AuthService {
if globalAuthSvc != nil {
return globalAuthSvc
}
if globalCoreCtx != nil {
if svc, err := core.Inject[contracts.AuthService](globalCoreCtx); err == nil {
globalAuthSvc = svc
return svc
}
}
return nil
}
// Apply registers admin routes, tasks, schedules, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
globalCoreCtx = ctx
// 0. Resolve auth and user services reactively via IoC
var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
var adminMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
globalAuthSvc = authSvc
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
loginMW = mw
}
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
adminMW = mw
}
// 0. Bind Services reactively
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
SetDBService(db)
} else {
core.When[contracts.AuthService](ctx, func(svc contracts.AuthService) {
globalAuthSvc = svc
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
SetDBService(db)
})
}
if userSvc, err := core.Inject[contracts.UserService](ctx); err == nil && userSvc != nil {
globalUserSvc = userSvc
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
SetCacheService(cache)
} else {
core.When[contracts.UserService](ctx, func(svc contracts.UserService) {
globalUserSvc = svc
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
SetCacheService(cache)
})
}
if user, err := core.Inject[contracts.UserService](ctx); err == nil && user != nil {
SetUserService(user)
} else {
core.When[contracts.UserService](ctx, func(user contracts.UserService) {
SetUserService(user)
})
}
if auth, err := core.Inject[contracts.AuthService](ctx); err == nil && auth != nil {
SetAuthService(auth)
} else {
core.When[contracts.AuthService](ctx, func(auth contracts.AuthService) {
SetAuthService(auth)
})
}
if task, err := core.Inject[contracts.TaskService](ctx); err == nil && task != nil {
SetTaskService(task)
} else {
core.When[contracts.TaskService](ctx, func(task contracts.TaskService) {
SetTaskService(task)
})
}
if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil {
SetStorageService(storage)
} else {
core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) {
SetStorageService(storage)
})
}
if rc, err := core.Inject[contracts.RiskControlService](ctx); err == nil && rc != nil {
SetRiskControlService(rc)
} else {
core.When[contracts.RiskControlService](ctx, func(rc contracts.RiskControlService) {
SetRiskControlService(rc)
})
}
SetEventEmitter(ctx.Events().Emit)
// 0a. Register migrations
ctx.OnDispose(func() error {
ResetServices()
return nil
})
// 0a. Dynamic Auth Middlewares
var loginMW gin.HandlerFunc = func(c *gin.Context) {
if authSvc := GetAuthService(c.Request.Context()); authSvc != nil {
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
mw(c)
return
}
}
c.Next()
}
var adminMW gin.HandlerFunc = func(c *gin.Context) {
if authSvc := GetAuthService(c.Request.Context()); authSvc != nil {
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
mw(c)
return
}
}
c.Next()
}
// 0b. Register migrations
ctx.Migrations().Register("admin", adminMigrations)
// 1. Register Admin HTTP Routes
+64 -88
View File
@@ -12,15 +12,12 @@ import (
"strings"
"time"
"github.com/redis/go-redis/v9"
"github.com/shopspring/decimal"
"gorm.io/gorm"
"Wavelet/pkg/cache/ram"
"Wavelet/pkg/idgen"
"Wavelet/pkg/util"
cachepkg "Wavelet/plugins/infra/cache"
db "Wavelet/plugins/infra/database"
)
const (
@@ -38,7 +35,7 @@ const (
// PreheatSystemConfigs loads all system configs from database.
func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
database := db.DB(ctx)
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -52,7 +49,7 @@ func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
// PreheatSystemConfigByKey loads a single config key from database.
func PreheatSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
database := db.DB(ctx)
database := GetDB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
@@ -75,7 +72,7 @@ func GetSystemConfigByGroup(ctx context.Context, configType string, key string)
}
}
database := db.DB(ctx)
database := GetDB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
@@ -129,7 +126,7 @@ func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]Sys
return result, nil
}
database := db.DB(ctx)
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -178,7 +175,7 @@ func ListVisibleSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
return list, nil
}
database := db.DB(ctx)
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -269,7 +266,7 @@ func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
// ListAdminSystemConfigs returns all configs, optionally filtered by type.
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
query := db.DB(ctx).Order("created_at DESC")
query := GetDB(ctx).Order("created_at DESC")
if configType != "" {
query = query.Where("type = ?", configType)
}
@@ -283,7 +280,7 @@ func ListAdminSystemConfigs(ctx context.Context, configType string) ([]SystemCon
// GetAdminSystemConfigByKey loads a config directly from DB.
func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
var config SystemConfig
if err := db.DB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
if err := GetDB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
return SystemConfig{}, err
}
return config, nil
@@ -292,7 +289,7 @@ func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, e
// SystemConfigExists reports whether a config key already exists.
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
var existing SystemConfig
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
err := GetDB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
@@ -304,18 +301,18 @@ func SystemConfigExists(ctx context.Context, key string) (bool, error) {
// CreateSystemConfigRecord persists a new system config row.
func CreateSystemConfigRecord(ctx context.Context, config *SystemConfig) error {
return db.DB(ctx).Create(config).Error
return GetDB(ctx).Create(config).Error
}
// UpdateSystemConfigFields applies partial updates to a system config row.
func UpdateSystemConfigFields(ctx context.Context, config *SystemConfig, updates map[string]any) error {
return db.DB(ctx).Model(config).Updates(updates).Error
return GetDB(ctx).Model(config).Updates(updates).Error
}
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache.
func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
var sc SystemConfig
err := db.DB(ctx).Where("key = ?", key).First(&sc).Error
err := GetDB(ctx).Where("key = ?", key).First(&sc).Error
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
@@ -327,12 +324,12 @@ func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
Type: configTypeSystem,
Visibility: ConfigVisibilityHidden,
}
if err := db.DB(ctx).Create(&sc).Error; err != nil {
if err := GetDB(ctx).Create(&sc).Error; err != nil {
return err
}
} else {
sc.Value = value
if err := db.DB(ctx).Save(&sc).Error; err != nil {
if err := GetDB(ctx).Save(&sc).Error; err != nil {
return err
}
}
@@ -342,7 +339,7 @@ func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
// ListTemplatesRecord returns all templates ordered by system flag and creation time.
func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
var templates []Template
if err := db.DB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
if err := GetDB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
return nil, err
}
return templates, nil
@@ -351,7 +348,7 @@ func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
// GetTemplateByKey loads a template by its key.
func GetTemplateByKey(ctx context.Context, key string) (Template, error) {
var tmpl Template
if err := db.DB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
if err := GetDB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
return Template{}, err
}
return tmpl, nil
@@ -360,7 +357,7 @@ func GetTemplateByKey(ctx context.Context, key string) (Template, error) {
// TemplateExistsByKey reports whether a template key is already taken.
func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
var existing Template
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
err := GetDB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
@@ -372,38 +369,38 @@ func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
// CreateTemplateRecord persists a new template.
func CreateTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Create(tmpl).Error
return GetDB(ctx).Create(tmpl).Error
}
// SaveTemplateRecord updates an existing template.
func SaveTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Save(tmpl).Error
return GetDB(ctx).Save(tmpl).Error
}
// DeleteTemplateRecord removes a template record.
func DeleteTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Delete(tmpl).Error
return GetDB(ctx).Delete(tmpl).Error
}
// CreateScheduleRecord 创建定时任务
func CreateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Create(schedule).Error
return GetDB(ctx).Create(schedule).Error
}
// UpdateScheduleRecord 更新定时任务
func UpdateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Save(schedule).Error
return GetDB(ctx).Save(schedule).Error
}
// DeleteScheduleRecord 删除定时任务
func DeleteScheduleRecord(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&Schedule{}, id).Error
return GetDB(ctx).Delete(&Schedule{}, id).Error
}
// GetScheduleByID 根据 ID 获取定时任务
func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
var schedule Schedule
if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
if err := GetDB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
return nil, err
}
return &schedule, nil
@@ -412,7 +409,7 @@ func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
// ListSchedulesRecord 获取所有定时任务
func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
if err := GetDB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
@@ -421,7 +418,7 @@ func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
// ListActiveSchedules 获取所有启用的定时任务
func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
if err := GetDB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
@@ -430,18 +427,18 @@ func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
// CreateTaskExecutionRecord 创建任务执行记录
func CreateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return db.DB(ctx).Create(execution).Error
return GetDB(ctx).Create(execution).Error
}
// UpdateTaskExecutionRecord 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
return db.DB(ctx).Omit("log").Save(execution).Error
return GetDB(ctx).Omit("log").Save(execution).Error
}
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
var execution TaskExecution
if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
if err := GetDB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
@@ -453,7 +450,7 @@ func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecutio
// GetTaskExecutionByID 根据 ID 获取执行记录
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
var execution TaskExecution
if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
if err := GetDB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
@@ -465,7 +462,7 @@ func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error
// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type.
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*TaskExecution, bool, error) {
var execution TaskExecution
err := db.DB(ctx).
err := GetDB(ctx).
Where("task_type = ?", taskType).
Order("id DESC").
First(&execution).Error
@@ -481,45 +478,40 @@ func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*Ta
return nil, false, err
}
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
// AppendTaskExecutionLog 将日志追加到缓冲,任务完成后再持久化到数据库。
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
if cachepkg.Redis == nil {
return errors.New("redis client is not initialized")
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return errors.New("cache service is not initialized")
}
now := time.Now().Format("15:04:05")
line := fmt.Sprintf("[%s] %s\n", now, logLine)
key := taskExecutionLogRedisKey(taskID)
_, err := cachepkg.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
pipe.RPush(ctx, key, line)
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
pipe.Expire(ctx, key, taskExecutionLogExpiration)
return nil
})
if err != nil {
return fmt.Errorf("append task execution log to redis: %w", err)
}
return nil
var existing string
_ = cacheSvc.Get(ctx, key, &existing)
return cacheSvc.Set(ctx, key, existing+line, taskExecutionLogExpiration)
}
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
// FlushTaskExecutionLog 将缓冲中的完整任务日志写入数据库,并在成功后清理缓存。
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
if cachepkg.Redis == nil {
return errors.New("redis client is not initialized")
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return errors.New("cache service is not initialized")
}
key := taskExecutionLogRedisKey(taskID)
logLines, err := cachepkg.Redis.LRange(ctx, key, 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
var logText string
if err := cacheSvc.Get(ctx, key, &logText); err != nil || logText == "" {
return nil
}
logText := strings.Join(logLines, "")
result := db.DB(ctx).Model(&TaskExecution{}).
gormDB := GetDB(ctx)
if gormDB == nil {
return errors.New(errDatabaseNotInitialized)
}
result := gormDB.Model(&TaskExecution{}).
Where("task_id = ?", taskID).
Update("log", logText)
if result.Error != nil {
@@ -529,9 +521,7 @@ func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
return fmt.Errorf("persist task execution log: task %q not found", taskID)
}
if err := cachepkg.Redis.Del(ctx, key).Err(); err != nil {
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
}
_ = cacheSvc.Delete(ctx, key)
return nil
}
@@ -544,7 +534,7 @@ func ListTaskExecutionRecords(ctx context.Context, req ListTaskExecutionsRequest
req.PageSize = 20
}
query := db.DB(ctx).Model(&TaskExecution{})
query := GetDB(ctx).Model(&TaskExecution{})
if req.Status != "" {
query = query.Where("status = ?", req.Status)
@@ -618,7 +608,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
terminalStatuses := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed}
var highFrequencyTaskTypes []string
if err := db.DB(ctx).
if err := GetDB(ctx).
Model(&TaskExecution{}).
Select("task_type").
Where("created_at >= ?", frequencyWindowStart).
@@ -630,7 +620,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
var highFrequencyDeleted int64
if len(highFrequencyTaskTypes) > 0 {
highFrequencyResult := db.DB(ctx).
highFrequencyResult := GetDB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", highFrequencyCutoff).
Where("task_type IN ?", highFrequencyTaskTypes).
@@ -641,7 +631,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
highFrequencyDeleted = highFrequencyResult.RowsAffected
}
lowFrequencyQuery := db.DB(ctx).
lowFrequencyQuery := GetDB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", lowFrequencyCutoff)
if len(highFrequencyTaskTypes) > 0 {
@@ -659,46 +649,32 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
}
func taskExecutionLogRedisKey(taskID string) string {
return cachepkg.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
return taskExecutionLogRedisKeyPrefix + taskID
}
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
if cachepkg.Redis == nil {
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return nil
}
logLines, err := cachepkg.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
var logText string
if err := cacheSvc.Get(ctx, taskExecutionLogRedisKey(execution.TaskID), &logText); err == nil && logText != "" {
execution.Log = logText
}
if len(logLines) == 0 {
return nil
}
execution.Log = strings.Join(logLines, "")
return nil
}
func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error {
if cachepkg.Redis == nil || len(executions) == 0 {
cacheSvc := GetCache(ctx)
if cacheSvc == nil || len(executions) == 0 {
return nil
}
commands := make([]*redis.StringSliceCmd, len(executions))
_, err := cachepkg.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error {
for i := range executions {
commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1)
}
return nil
})
if err != nil {
return fmt.Errorf("get task execution logs from redis: %w", err)
}
for i := range executions {
logLines := commands[i].Val()
if len(logLines) > 0 {
executions[i].Log = strings.Join(logLines, "")
var logText string
if err := cacheSvc.Get(ctx, taskExecutionLogRedisKey(executions[i].TaskID), &logText); err == nil && logText != "" {
executions[i].Log = logText
}
}
return nil
@@ -7,14 +7,11 @@ import (
"context"
"encoding/json"
"errors"
"sync"
"time"
"gorm.io/gorm"
"Wavelet/pkg/cache/ram"
"Wavelet/pkg/util"
cachepkg "Wavelet/plugins/infra/cache"
)
const (
@@ -33,11 +30,6 @@ const (
ConfigCacheType = "config"
)
type systemConfigBroadcastMessage struct {
Type string `json:"type"`
Key string `json:"key"`
}
// ConfigLoader loads configuration data from the database.
type ConfigLoader struct{}
@@ -64,21 +56,19 @@ func (ConfigLoader) LoadAll(ctx context.Context, configType string) ([]ram.Cache
return items, nil
}
// LoadOne loads a single system config from database as a CacheItem.
// LoadOne loads a single system config from database as CacheItem.
func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string) (ram.CacheItem, error) {
cfg, err := PreheatSystemConfigByKey(ctx, key)
cfg, err := GetSystemConfigByKey(ctx, key)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return ram.CacheItem{}, ram.ErrNotFound
}
return ram.CacheItem{}, err
}
valBytes, err := json.Marshal(cfg)
if err != nil {
return ram.CacheItem{}, err
}
return ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
@@ -87,121 +77,67 @@ func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string)
}, nil
}
// PreloadSystemConfigs warms the in-memory RAM cache from database on startup.
func PreloadSystemConfigs(ctx context.Context) error {
return ram.Refresh(ctx, ConfigCacheType, "", ConfigLoader{})
// GetCachedSystemConfig retrieves a single system config with RAM L1 fallback to DB.
func GetCachedSystemConfig(ctx context.Context, key string) (*SystemConfig, error) {
if item, ok := ram.Get(ConfigCacheType, key); ok {
var cfg SystemConfig
if err := json.Unmarshal([]byte(item.Value), &cfg); err == nil {
return &cfg, nil
}
}
cfg, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return nil, err
}
valBytes, err := json.Marshal(cfg)
if err == nil {
ram.Set(ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(key),
})
}
return &cfg, nil
}
var (
systemConfigListenerOnce sync.Once
systemConfigListenerCtx context.Context
systemConfigListenerCancel context.CancelFunc
systemConfigListenerDone chan struct{}
)
// StopSystemConfigCacheListener stops the cache invalidation listener (kept for backward compatibility).
func StopSystemConfigCacheListener() {
}
// StartSystemConfigCacheListener starts the cache listener (kept for backward compatibility).
func StartSystemConfigCacheListener() {
}
func ensureSystemConfigCacheListener() {
systemConfigListenerOnce.Do(startSystemConfigCacheInvalidationListener)
}
func startSystemConfigCacheInvalidationListener() {
if cachepkg.Redis == nil {
return
}
systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background())
systemConfigListenerDone = make(chan struct{})
redisClient := cachepkg.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 cachepkg.Redis 竞争
util.Go(func() {
listenerCtx := systemConfigListenerCtx
defer close(systemConfigListenerDone)
pubsub := redisClient.Subscribe(listenerCtx, SystemConfigBroadcastChannel)
defer func() {
_ = pubsub.Close()
}()
util.Go(func() {
<-listenerCtx.Done()
_ = pubsub.Close()
})
for msg := range pubsub.Channel() {
var payload systemConfigBroadcastMessage
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
ram.UpdateTypeItems(ConfigCacheType, nil)
continue
}
key := payload.Key
if key == "*" || key == "" {
ram.UpdateTypeItems(payload.Type, nil)
} else {
ram.Delete(payload.Type, key)
}
}
})
}
// StopSystemConfigCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
func StopSystemConfigCacheListener() {
if systemConfigListenerCancel != nil {
systemConfigListenerCancel()
if systemConfigListenerDone != nil {
<-systemConfigListenerDone
}
systemConfigListenerCancel = nil
systemConfigListenerDone = nil
}
systemConfigListenerOnce = sync.Once{}
}
func determineTTL(_ string) time.Duration {
// Program-determined TTL: -1 means never expire for all configs by default
return -1
}
// InvalidateSystemConfigCache triggers a broadcast to refresh the cache for key.
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
ensureSystemConfigCacheListener()
// Invalidate local cache synchronously first
ram.Delete(ConfigCacheType, key)
// Broadcast to other nodes and clean legacy Redis cache key
if cachepkg.Redis != nil {
_ = cachepkg.HDel(ctx, SystemConfigRedisHashKey, key)
publishSystemConfigBroadcast(ctx, ConfigCacheType, key)
if cacheSvc := GetCache(ctx); cacheSvc != nil {
_ = cacheSvc.Delete(ctx, "system:config:"+key)
_ = cacheSvc.Delete(ctx, SystemConfigVisibleListRedisKey)
}
return nil
}
// InvalidateAllSystemConfigCaches triggers a broadcast to refresh the entire config cache.
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
ensureSystemConfigCacheListener()
// Invalidate all items of type ConfigCacheType synchronously first
ram.UpdateTypeItems(ConfigCacheType, nil)
// Broadcast to other nodes and clean legacy Redis cache keys
if cachepkg.Redis != nil {
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(SystemConfigRedisHashKey), cachepkg.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
publishSystemConfigBroadcast(ctx, ConfigCacheType, "*")
if cacheSvc := GetCache(ctx); cacheSvc != nil {
_ = cacheSvc.Delete(ctx, SystemConfigRedisHashKey)
_ = cacheSvc.Delete(ctx, SystemConfigVisibleListRedisKey)
}
return nil
}
func publishSystemConfigBroadcast(ctx context.Context, configType string, key string) {
if cachepkg.Redis == nil {
return
}
payload, err := json.Marshal(systemConfigBroadcastMessage{Type: configType, Key: key})
if err != nil {
return
}
_ = cachepkg.Redis.Publish(ctx, SystemConfigBroadcastChannel, payload).Err()
}
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
func ResetSystemConfigRAMCacheForTest() {
ram.ResetForTest()
@@ -8,16 +8,30 @@ import (
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"gorm.io/gorm"
"Wavelet/plugins/infra/cache"
"Wavelet/plugins/infra/database"
)
type testDBService struct {
db *gorm.DB
}
func (s *testDBService) DB(ctx context.Context) *gorm.DB {
return s.db
}
func (s *testDBService) MasterDB(ctx context.Context) *gorm.DB {
return s.db
}
func (s *testDBService) GORM() *gorm.DB {
return s.db
}
func (s *testDBService) Named(_ string) *gorm.DB {
return s.db
}
func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
t.Helper()
@@ -41,28 +55,12 @@ func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
t.Fatalf("Create(site_name) error = %v", err)
}
mr, err := miniredis.Run()
if err != nil {
t.Fatalf("miniredis.Run() error = %v", err)
}
redisClient := redis.NewClient(&redis.Options{
Addr: mr.Addr(),
MaintNotificationsConfig: &maintnotifications.Config{
Mode: maintnotifications.ModeDisabled,
},
})
previousRedis := cache.Redis
database.SetDB(sqliteDB)
cache.Redis = redisClient
SetDBService(&testDBService{db: sqliteDB})
cleanup := func() {
StopSystemConfigCacheListener()
ResetSystemConfigRAMCacheForTest()
database.SetDB(nil)
cache.Redis = previousRedis
_ = redisClient.Close()
mr.Close()
ResetServices()
}
return sqliteDB, cleanup
+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:
+2 -1
View File
@@ -6,10 +6,11 @@ package cap
import (
"net/http"
"github.com/gin-gonic/gin"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/cap/pow"
"github.com/gin-gonic/gin"
)
// ChallengeResponse is a local type alias for the pow.ChallengeResponse struct
+1 -8
View File
@@ -15,7 +15,6 @@ import (
"Wavelet/pkg/config"
"Wavelet/plugins/domain/cap/pow"
db "Wavelet/plugins/infra/cache"
)
const (
@@ -186,13 +185,7 @@ func GetDefaultManager() *Manager {
return
}
var store pow.Store
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
store = pow.NewRedisStore(db.Redis)
} else {
store = pow.NewMemoryStore(1 * time.Minute)
}
store := pow.NewMemoryStore(1 * time.Minute)
defaultManager = NewManager(secret, store)
})
return defaultManager
+18 -1
View File
@@ -29,7 +29,6 @@ func (p *Plugin) Name() string {
func (p *Plugin) Inject() []reflect.Type {
return []reflect.Type{
reflect.TypeFor[contracts.DBService](),
reflect.TypeFor[contracts.CacheService](),
}
}
@@ -45,6 +44,24 @@ func (p *Plugin) Manifest() core.Manifest {
// Apply registers the cap routes and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 0. Bind DBService 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)
})
}
ctx.OnDispose(func() error {
setDBService(nil)
return nil
})
// Listen to system config changed events to invalidate cached settings
ctx.Events().On(contracts.EventTopicConfigChanged, func(_ any) {
InvalidateRuntimeSettings()
})
// Register HTTP Routes
capGroup := ctx.Router().Group("/api/v1/cap")
{
+41 -41
View File
@@ -5,7 +5,6 @@ package cap
import (
"context"
"encoding/json"
"errors"
"strconv"
"sync"
@@ -13,12 +12,38 @@ import (
"time"
"golang.org/x/sync/singleflight"
"gorm.io/gorm"
"Wavelet/pkg/util"
cachepkg "Wavelet/plugins/infra/cache"
database "Wavelet/plugins/infra/database"
"Wavelet/core"
"Wavelet/core/contracts"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
)
func setDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = 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
}
const (
defaultChallengeCount = 1
defaultChallengeSize = 32
@@ -67,9 +92,8 @@ var runtimeConfigKeySet = func() map[string]struct{} {
}()
type runtimeSettingsStore struct {
snapshot atomic.Pointer[RuntimeSettings]
loadGroup singleflight.Group
listenerOnce sync.Once
snapshot atomic.Pointer[RuntimeSettings]
loadGroup singleflight.Group
}
var settingsStore = &runtimeSettingsStore{}
@@ -148,7 +172,11 @@ func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
Value string `gorm:"column:value"`
}
var records []configRecord
if err := database.DB(ctx).Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
db := getDB(ctx)
if db == nil {
return parseRuntimeSettings(nil), nil
}
if err := db.Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
return RuntimeSettings{}, err
}
configs := make(map[string]string, len(records))
@@ -167,6 +195,10 @@ func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
TokenTTL: defaultTokenTTL,
}
if len(configs) == 0 {
return settings
}
if val, ok := configs[ConfigKeyCapLoginEnabled]; ok {
if enabled, err := strconv.ParseBool(val); err == nil {
settings.LoginEnabled = enabled
@@ -201,36 +233,4 @@ func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
return settings
}
func (s *runtimeSettingsStore) ensureInvalidationListener() {
s.listenerOnce.Do(startRuntimeSettingsInvalidationListener)
}
// SystemConfigInvalidationChannel 系统配置失效广播通道
const SystemConfigInvalidationChannel = "system_config:invalidation"
func startRuntimeSettingsInvalidationListener() {
rdb := cachepkg.Redis
if rdb == nil {
return
}
util.Go(func() {
pubsub := rdb.Subscribe(context.Background(), SystemConfigInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
for msg := range pubsub.Channel() {
var payload struct {
Key string `json:"key"`
}
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
InvalidateRuntimeSettings()
continue
}
if payload.Key == "" || payload.Key == "*" || IsRuntimeConfigKey(payload.Key) {
InvalidateRuntimeSettings()
}
}
})
}
func (s *runtimeSettingsStore) ensureInvalidationListener() {}
@@ -7,8 +7,9 @@ import (
"net/http"
"strconv"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
"Wavelet/pkg/response"
)
// ListAdminChannelDefinitions returns form schemas for supported channel types.
@@ -5,13 +5,14 @@
package qq
import (
"Wavelet/pkg/util"
"context"
"fmt"
"strings"
"sync"
"time"
"Wavelet/pkg/util"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/message_gateway"
@@ -12,9 +12,10 @@ import (
"strconv"
"strings"
tele "gopkg.in/telebot.v4"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/message_gateway"
tele "gopkg.in/telebot.v4"
)
// Adapter is a Telegram private-chat channel.
@@ -7,8 +7,9 @@ import (
"context"
"testing"
"Wavelet/plugins/domain/message_gateway"
tele "gopkg.in/telebot.v4"
"Wavelet/plugins/domain/message_gateway"
)
func TestHandleUpdate_DropsGroups(t *testing.T) {
@@ -0,0 +1,78 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"sync"
"gorm.io/gorm"
"Wavelet/core"
"Wavelet/core/contracts"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
cacheMu sync.RWMutex
cacheSvc contracts.CacheService
taskMu sync.RWMutex
taskSvc contracts.TaskService
)
func SetDBServiceForTest(s contracts.DBService) {
setDBService(s)
}
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 setTaskService(s contracts.TaskService) {
taskMu.Lock()
defer taskMu.Unlock()
taskSvc = 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
}
func getTaskService() contracts.TaskService {
taskMu.RLock()
defer taskMu.RUnlock()
return taskSvc
}
@@ -8,10 +8,11 @@ import (
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"Wavelet/core/contracts"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"github.com/gin-gonic/gin"
)
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
@@ -8,13 +8,35 @@ import (
"testing"
"time"
"gorm.io/gorm"
"Wavelet/pkg/testhelper"
"Wavelet/plugins/domain/message_gateway"
)
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
}
func TestUpsertPairingCode_ReusesUnexpired(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
testDB, _, cleanup := testhelper.SetupTestEnvironment(t)
message_gateway.SetDBServiceForTest(&mockDBService{db: testDB})
defer func() {
message_gateway.SetDBServiceForTest(nil)
cleanup()
}()
ctx := context.Background()
first, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ABCD1234", time.Now().Add(15*time.Minute))
if err != nil {
@@ -9,12 +9,12 @@ import (
"embed"
"reflect"
"github.com/gin-gonic/gin"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/pkg/util"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
)
//go:embed migrations/*.sql
@@ -80,6 +80,35 @@ type PushNotificationEvent struct {
// Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 0. Bind DBService, CacheService, TaskService
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)
})
}
if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil {
setTaskService(taskSvc)
} else {
core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) {
setTaskService(taskSvc)
})
}
ctx.OnDispose(func() error {
setDBService(nil)
setCacheService(nil)
setTaskService(nil)
return nil
})
// 0. Resolve auth service for middleware (via IoC, not direct import)
var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
var adminMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
@@ -145,18 +174,16 @@ func (p *Plugin) Apply(ctx *core.Context) error {
const defaultTaskRetry = 3
pushHandler := &PushHandler{}
// 5. Register Asynq background tasks
ctx.Task().Register("message_gateway:push_notification", func(c context.Context, t *asynq.Task) error {
_, err := pushHandler.Execute(c, t.Payload())
return err
// 5. Register background tasks
ctx.Task().Register("message_gateway:push_notification", func(c context.Context, payload []byte) error {
return pushHandler.Execute(c, payload)
}, extpoints.WithTaskRetry(defaultTaskRetry))
ctx.Task().Register(SendNotificationTask, func(c context.Context, t *asynq.Task) error {
_, err := pushHandler.Execute(c, t.Payload())
return err
ctx.Task().Register(SendNotificationTask, func(c context.Context, payload []byte) error {
return pushHandler.Execute(c, payload)
}, extpoints.WithTaskRetry(defaultTaskRetry))
ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ *asynq.Task) error {
ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ []byte) error {
return nil
})
@@ -184,9 +211,14 @@ func (p *Plugin) Apply(ctx *core.Context) error {
return nil
})
// 8. Register built-in domain events and task listeners
// 8. Register task completed event listener
ctx.Events().On(contracts.EventTopicTaskCompleted, func(c context.Context, e contracts.TaskCompletedEvent) error {
handleTaskCompleted(c, e)
return nil
})
// 9. Register built-in domain events
RegisterCustomEvents()
RegisterTaskListeners()
// 9. Register Settings Schemas
ctx.Settings().Register(extpoints.SettingSchema{
@@ -11,10 +11,11 @@ import (
"strings"
"sync"
"Wavelet/pkg/response"
pkgpush "Wavelet/plugins/domain/message_gateway/push"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"Wavelet/pkg/response"
pkgpush "Wavelet/plugins/domain/message_gateway/push"
)
const (
@@ -13,8 +13,9 @@ import (
pkgpush "Wavelet/plugins/domain/message_gateway/push"
"Wavelet/pkg/util"
"gorm.io/gorm"
"Wavelet/pkg/util"
)
// NotificationMessage represents the structured notification message payload.
@@ -11,9 +11,10 @@ import (
pkgpush "Wavelet/plugins/domain/message_gateway/push"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"Wavelet/pkg/response"
)
// UpdatePushEventRequest is the request body for updating a push event.
@@ -11,11 +11,10 @@ import (
"strconv"
"strings"
"gorm.io/gorm"
"Wavelet/core/contracts"
pkgpush "Wavelet/plugins/domain/message_gateway/push"
"Wavelet/plugins/drivers/driver_asynq_worker"
db "Wavelet/plugins/infra/database"
"gorm.io/gorm"
)
type smtpConfig struct {
@@ -28,10 +27,10 @@ type smtpConfig struct {
func loadSMTPConfig(ctx context.Context) smtpConfig {
var cfg smtpConfig
var host, port, user, pass string
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_host").Pluck("value", &host).Error
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_port").Pluck("value", &port).Error
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_username").Pluck("value", &user).Error
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_password").Pluck("value", &pass).Error
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_host").Pluck("value", &host).Error
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_port").Pluck("value", &port).Error
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_username").Pluck("value", &user).Error
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_password").Pluck("value", &pass).Error
cfg.Host = host
cfg.Port = port
cfg.Username = user
@@ -263,14 +262,14 @@ func loadUserFromPayload(ctx context.Context, data map[string]any) any {
if userID, ok := extractUserID(data); ok && userID > 0 {
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 {
return &user
}
}
if username := extractUsername(data); username != "" {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("username = ?", username).First(&user).Error; err == nil {
if err := getDB(ctx).Table("w_users").Where("username = ?", username).First(&user).Error; err == nil {
return &user
}
}
@@ -372,11 +371,11 @@ func resolveDynamicKeyword(target string, flatBody map[string]any) string {
func resolveTargetUser(ctx context.Context, resolved string, _ string) (contracts.UserDTO, bool) {
var user contracts.UserDTO
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
if err := db.DB(ctx).Table("w_users").Where("id = ?", id).First(&user).Error; err == nil {
if err := getDB(ctx).Table("w_users").Where("id = ?", id).First(&user).Error; err == nil {
return user, true
}
}
if err := db.DB(ctx).Table("w_users").Where("username = ?", resolved).First(&user).Error; err == nil {
if err := getDB(ctx).Table("w_users").Where("username = ?", resolved).First(&user).Error; err == nil {
return user, true
}
return user, false
@@ -387,7 +386,7 @@ func resolveSystemTarget(ctx context.Context, resolved string, channel string) (
return "", false
}
var adminUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
if err := getDB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
return resolved, true
}
if channel == channelEmail && adminUser.Email != "" {
@@ -425,7 +424,7 @@ func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, s
func getSystemUser(ctx context.Context) *contracts.UserDTO {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&user).Error; err == nil {
if err := getDB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&user).Error; err == nil {
return &user
}
return &contracts.UserDTO{
@@ -445,14 +444,16 @@ func findBuiltInEvent(key string) (EventMetadata, bool) {
func getEventInfo(req CreatePushEventRequest) (string, string, []byte, error) {
if req.TaskType != "" {
meta := driver_asynq_worker.GetTaskMetaByAsynqTask(req.TaskType)
if meta == nil {
return "", "", nil, errors.New("unsupported task type")
taskName := req.TaskType
if taskSvc := getTaskService(); taskSvc != nil {
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
taskName = meta.DisplayName
}
}
eventKey := "task_completed:" + req.TaskType
eventName := "任务完成: " + meta.Name
eventName := "任务完成: " + taskName
defaultTemplate := NotificationMessage{
Title: "任务完成: " + meta.Name,
Title: "任务完成: " + taskName,
Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。",
Level: defaultLevelInfo,
}
@@ -484,8 +485,11 @@ func enqueuePushTask(ctx context.Context, payload SendPayload) error {
if err != nil {
return err
}
_, err = driver_asynq_worker.DispatchTask(ctx, "send_notification", payloadBytes, "system")
return err
if taskSvc := getTaskService(); taskSvc != nil {
_, err = taskSvc.Dispatch(ctx, "send_notification", payloadBytes, "system")
return err
}
return errors.New("task service not available")
}
func getFlatBody(body map[string]any) map[string]any {
@@ -9,19 +9,14 @@ import (
"strconv"
"time"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/plugins/drivers/driver_asynq_worker"
)
// RegisterTaskListeners subscribes push notification handlers to task completion events.
func RegisterTaskListeners() {
driver_asynq_worker.OnTaskCompleted(handleTaskCompleted)
}
func handleTaskCompleted(ctx context.Context, execution *driver_asynq_worker.TaskExecution, result *driver_asynq_worker.TaskResult, execErr error) {
events, err := listActivePushEventsByTaskType(ctx, execution.TaskType)
func handleTaskCompleted(ctx context.Context, e contracts.TaskCompletedEvent) {
events, err := listActivePushEventsByTaskType(ctx, e.TaskType)
if err != nil {
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err)
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", e.TaskType, err)
return
}
if len(events) == 0 {
@@ -29,34 +24,26 @@ func handleTaskCompleted(ctx context.Context, execution *driver_asynq_worker.Tas
}
body := map[string]any{
"task_id": execution.TaskID,
"task_name": execution.TaskName,
"task_type": execution.TaskType,
"task_status": string(execution.Status),
"task_duration": execution.Duration,
"task_id": e.TaskID,
"task_name": e.TaskName,
"task_type": e.TaskType,
"task_status": e.Status,
"task_duration": e.Duration,
"time": time.Now().Format("2006-01-02 15:04:05"),
}
if execErr != nil {
body["task_error"] = execErr.Error()
} else {
body["task_error"] = ""
}
if result != nil {
body["task_result"] = result.Message
} else {
body["task_result"] = ""
"task_error": e.ErrorMsg,
"task_result": e.ResultMsg,
}
var payloadMap map[string]any
if execution.Payload != "" {
if err := json.Unmarshal([]byte(execution.Payload), &payloadMap); err == nil {
if e.Payload != "" {
if err := json.Unmarshal([]byte(e.Payload), &payloadMap); err == nil {
body["payload"] = payloadMap
extractUserFromMap(ctx, payloadMap, body)
}
}
if result != nil && result.Detail != "" {
if e.Detail != "" {
var detailMap map[string]any
if err := json.Unmarshal([]byte(result.Detail), &detailMap); err == nil {
if err := json.Unmarshal([]byte(e.Detail), &detailMap); err == nil {
body["detail"] = detailMap
extractUserFromMap(ctx, detailMap, body)
}
@@ -9,8 +9,9 @@ import (
"errors"
"fmt"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/message_gateway/push"
"Wavelet/plugins/drivers/driver_asynq_worker"
)
const (
@@ -21,28 +22,24 @@ const (
)
// SendNotificationMeta represents the task metadata.
var SendNotificationMeta = driver_asynq_worker.TaskMeta{
Type: TaskTypeSendNotification,
AsynqTask: SendNotificationTask,
Name: "推送通知",
Description: "异步执行系统通知的多渠道派发与推送",
SupportsTime: false,
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
Queue: driver_asynq_worker.QueueDefault,
Retryable: true,
Params: []driver_asynq_worker.TaskParam{
var SendNotificationMeta = contracts.TaskMetaDTO{
Name: TaskTypeSendNotification,
DisplayName: "推送通知",
Description: "异步执行系统通知的多渠道派发与推送",
MaxRetry: 3,
Queue: "default",
Params: []contracts.TaskParamDTO{
{
Name: "event_key",
Label: "事件标识",
Type: "string",
Description: "事件标识 (如 admin_login)",
Required: true,
Placeholder: "admin_login",
},
{
Name: "target",
Label: "目标接收者",
Type: "string",
Required: false,
Name: "target",
Type: "string",
Description: "目标接收者",
Required: false,
},
},
}
@@ -69,23 +66,21 @@ func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) {
}
// Execute performs the push send and logs delivery history audit.
func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
func (h *PushHandler) Execute(ctx context.Context, payload []byte) error {
var req SendPayload
if err := json.Unmarshal(payload, &req); err != nil {
driver_asynq_worker.AppendLog(ctx, "解析推送参数失败: %v", err)
return nil, fmt.Errorf("parse payload failed: %w", err)
logger.ErrorF(ctx, "[Push] 解析推送参数失败: %v", err)
return fmt.Errorf("parse payload failed: %w", err)
}
driver_asynq_worker.AppendLog(ctx, "开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target)
logger.InfoF(ctx, "[Push] 开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target)
pusher, err := push.GetPusher(req.Config.Channel)
if err != nil {
errWrap := fmt.Errorf("get pusher failed: %w", err)
driver_asynq_worker.AppendLog(ctx, "推送失败: %v", errWrap)
if driver_asynq_worker.IsFinalAttempt(ctx) {
h.recordHistory(ctx, req, "failed", errWrap.Error())
}
return nil, errWrap
logger.ErrorF(ctx, "[Push] 推送失败: %v", errWrap)
h.recordHistory(ctx, req, "failed", errWrap.Error())
return errWrap
}
flatBody := req.Body.Flatten()
@@ -95,29 +90,19 @@ func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*driver_asyn
content := req.Body.Content
if err != nil {
driver_asynq_worker.AppendLog(ctx, "消息推送失败 (标题: %s): %v", title, err)
if upstreamResp != "" {
driver_asynq_worker.AppendLog(ctx, "上游返回: %s", upstreamResp)
}
if driver_asynq_worker.IsFinalAttempt(ctx) {
h.recordHistory(ctx, req, "failed", err.Error())
}
return nil, fmt.Errorf("pusher.Send failed: %w", err)
logger.ErrorF(ctx, "[Push] 消息推送失败 (标题: %s): %v, 上游返回: %s", title, err, upstreamResp)
h.recordHistory(ctx, req, "failed", err.Error())
return fmt.Errorf("pusher.Send failed: %w", err)
}
driver_asynq_worker.AppendLog(ctx, "消息推送成功 (标题: %s, 内容摘要: %s)", title, content)
if upstreamResp != "" {
driver_asynq_worker.AppendLog(ctx, "上游返回: %s", upstreamResp)
}
logger.InfoF(ctx, "[Push] 消息推送成功 (标题: %s, 内容摘要: %s), 上游返回: %s", title, content, upstreamResp)
h.recordHistory(ctx, req, "success", "")
return &driver_asynq_worker.TaskResult{
Message: fmt.Sprintf("推送成功: [%s] -> %s", req.Config.Channel, req.Target),
}, nil
return nil
}
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) {
if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil {
driver_asynq_worker.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
logger.ErrorF(ctx, "[Push] 写入推送历史审计记录失败: %v", dbErr)
}
}
@@ -11,8 +11,6 @@ import (
"gorm.io/gorm"
"Wavelet/pkg/idgen"
cachepkg "Wavelet/plugins/infra/cache"
db "Wavelet/plugins/infra/database"
)
const (
@@ -25,18 +23,18 @@ func CreateMessageChannel(ctx context.Context, ch *MessageChannel) error {
if ch.ID == 0 {
ch.ID = idgen.NextUint64ID()
}
return db.DB(ctx).Create(ch).Error
return getDB(ctx).Create(ch).Error
}
// UpdateMessageChannel saves a channel row.
func UpdateMessageChannel(ctx context.Context, ch *MessageChannel) error {
return db.DB(ctx).Save(ch).Error
return getDB(ctx).Save(ch).Error
}
// GetMessageChannel loads a channel by id.
func GetMessageChannel(ctx context.Context, id uint64) (*MessageChannel, error) {
var ch MessageChannel
if err := db.DB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
if err := getDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
return nil, err
}
return &ch, nil
@@ -45,7 +43,7 @@ func GetMessageChannel(ctx context.Context, id uint64) (*MessageChannel, error)
// ListMessageChannels returns all channels newest first.
func ListMessageChannels(ctx context.Context) ([]MessageChannel, error) {
var rows []MessageChannel
if err := db.DB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
if err := getDB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
@@ -53,7 +51,7 @@ func ListMessageChannels(ctx context.Context) ([]MessageChannel, error) {
// DeleteMessageChannel removes pairings, bindings, then the channel.
func DeleteMessageChannel(ctx context.Context, id uint64) error {
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
return getDB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("channel_id = ?", id).Delete(&MessagePairingCode{}).Error; err != nil {
return err
}
@@ -69,13 +67,13 @@ func CreateMessageBinding(ctx context.Context, b *MessageBinding) error {
if b.ID == 0 {
b.ID = idgen.NextUint64ID()
}
return db.DB(ctx).Create(b).Error
return getDB(ctx).Create(b).Error
}
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*MessageBinding, error) {
var b MessageBinding
err := db.DB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
err := getDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
if err != nil {
return nil, err
}
@@ -85,7 +83,7 @@ func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platform
// ListBindingsByUser lists bindings for a Wavelet user.
func ListBindingsByUser(ctx context.Context, userID uint64) ([]MessageBinding, error) {
var rows []MessageBinding
if err := db.DB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
if err := getDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
@@ -94,7 +92,7 @@ func ListBindingsByUser(ctx context.Context, userID uint64) ([]MessageBinding, e
// GetMessageBinding loads a binding by id.
func GetMessageBinding(ctx context.Context, id uint64) (*MessageBinding, error) {
var b MessageBinding
if err := db.DB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
if err := getDB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
return nil, err
}
return &b, nil
@@ -102,13 +100,13 @@ func GetMessageBinding(ctx context.Context, id uint64) (*MessageBinding, error)
// DeleteMessageBinding deletes a binding by id.
func DeleteMessageBinding(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&MessageBinding{}, id).Error
return getDB(ctx).Delete(&MessageBinding{}, id).Error
}
// UpsertPairingCode reuses an unexpired code for the same channel+platform user.
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*MessagePairingCode, error) {
var existing MessagePairingCode
err := db.DB(ctx).
err := getDB(ctx).
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
First(&existing).Error
if err == nil {
@@ -123,7 +121,7 @@ func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, co
PlatformUserID: platformUserID,
ExpiresAt: expiresAt,
}
if err := db.DB(ctx).Create(row).Error; err != nil {
if err := getDB(ctx).Create(row).Error; err != nil {
return nil, err
}
return row, nil
@@ -132,7 +130,7 @@ func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, co
// GetPairingCode loads a pairing code by normalized code string.
func GetPairingCode(ctx context.Context, code string) (*MessagePairingCode, error) {
var row MessagePairingCode
if err := db.DB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
if err := getDB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
return nil, err
}
return &row, nil
@@ -140,18 +138,18 @@ func GetPairingCode(ctx context.Context, code string) (*MessagePairingCode, erro
// DeletePairingCode removes a pairing code.
func DeletePairingCode(ctx context.Context, code string) error {
return db.DB(ctx).Where("code = ?", code).Delete(&MessagePairingCode{}).Error
return getDB(ctx).Where("code = ?", code).Delete(&MessagePairingCode{}).Error
}
// DeleteExpiredPairingCodes removes expired pairing rows.
func DeleteExpiredPairingCodes(ctx context.Context) error {
return db.DB(ctx).Where("expires_at <= ?", time.Now()).Delete(&MessagePairingCode{}).Error
return getDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&MessagePairingCode{}).Error
}
// ListEnabledMessageChannels returns enabled channels.
func ListEnabledMessageChannels(ctx context.Context) ([]MessageChannel, error) {
var rows []MessageChannel
if err := db.DB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
if err := getDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
@@ -160,7 +158,7 @@ func ListEnabledMessageChannels(ctx context.Context) ([]MessageChannel, error) {
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
func ListPushChannelsRecord(ctx context.Context) ([]PushChannel, error) {
var channels []PushChannel
if err := db.DB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
if err := getDB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
return nil, err
}
return channels, nil
@@ -169,7 +167,7 @@ func ListPushChannelsRecord(ctx context.Context) ([]PushChannel, error) {
// GetPushChannelByIDRecord loads a push channel by primary key.
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (PushChannel, error) {
var channel PushChannel
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
if err := getDB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
return PushChannel{}, err
}
return channel, nil
@@ -178,7 +176,7 @@ func GetPushChannelByIDRecord(ctx context.Context, id uint64) (PushChannel, erro
// GetPushChannelByNameRecord 根据名称获取消息通道。
func GetPushChannelByNameRecord(ctx context.Context, name string) (*PushChannel, error) {
var channel PushChannel
if err := db.DB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
if err := getDB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
return nil, err
}
return &channel, nil
@@ -187,7 +185,7 @@ func GetPushChannelByNameRecord(ctx context.Context, name string) (*PushChannel,
// CountPushChannelsByNameRecord returns how many channels share the given name.
func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
if err := getDB(ctx).Model(&PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
@@ -195,7 +193,7 @@ func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, err
// CreatePushChannelRecord persists a new channel and invalidates cache.
func CreatePushChannelRecord(ctx context.Context, channel *PushChannel) error {
if err := db.DB(ctx).Create(channel).Error; err != nil {
if err := getDB(ctx).Create(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
@@ -204,7 +202,7 @@ func CreatePushChannelRecord(ctx context.Context, channel *PushChannel) error {
// SavePushChannelRecord updates a channel and invalidates cache.
func SavePushChannelRecord(ctx context.Context, channel *PushChannel) error {
if err := db.DB(ctx).Save(channel).Error; err != nil {
if err := getDB(ctx).Save(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
@@ -213,7 +211,7 @@ func SavePushChannelRecord(ctx context.Context, channel *PushChannel) error {
// DeletePushChannelRecord removes a channel and invalidates cache.
func DeletePushChannelRecord(ctx context.Context, channel *PushChannel) error {
if err := db.DB(ctx).Delete(channel).Error; err != nil {
if err := getDB(ctx).Delete(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
@@ -224,18 +222,18 @@ func DeletePushChannelRecord(ctx context.Context, channel *PushChannel) error {
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
cacheKey := "push:channel:active:" + name
var channel PushChannel
if cachepkg.Redis != nil {
if err := cachepkg.GetJSON(ctx, cacheKey, &channel); err == nil {
if cache := getCache(ctx); cache != nil {
if err := cache.Get(ctx, cacheKey, &channel); err == nil {
return &channel, nil
}
}
if err := db.DB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error; err != nil {
if err := getDB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error; err != nil {
return nil, err
}
if cachepkg.Redis != nil {
_ = cachepkg.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL)
if cache := getCache(ctx); cache != nil {
_ = cache.Set(ctx, cacheKey, channel, activePushChannelCacheTTL)
}
return &channel, nil
@@ -243,15 +241,15 @@ func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel,
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
func DeleteActivePushChannelCache(ctx context.Context, name string) {
if cachepkg.Redis != nil {
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey("push:channel:active:"+name)).Err()
if cache := getCache(ctx); cache != nil {
_ = cache.Delete(ctx, "push:channel:active:"+name)
}
}
// ListPushEventsRecord returns all push events ordered by creation time descending.
func ListPushEventsRecord(ctx context.Context) ([]PushEvent, error) {
var events []PushEvent
if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
if err := getDB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
return nil, err
}
return events, nil
@@ -260,7 +258,7 @@ func ListPushEventsRecord(ctx context.Context) ([]PushEvent, error) {
// GetPushEventByIDRecord loads a push event by primary key.
func GetPushEventByIDRecord(ctx context.Context, id uint64) (PushEvent, error) {
var event PushEvent
if err := db.DB(ctx).First(&event, id).Error; err != nil {
if err := getDB(ctx).First(&event, id).Error; err != nil {
return PushEvent{}, err
}
return event, nil
@@ -269,7 +267,7 @@ func GetPushEventByIDRecord(ctx context.Context, id uint64) (PushEvent, error) {
// GetPushEventByKeyRecord loads a push event by event key.
func GetPushEventByKeyRecord(ctx context.Context, key string) (PushEvent, error) {
var event PushEvent
if err := db.DB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
if err := getDB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
return PushEvent{}, err
}
return event, nil
@@ -278,7 +276,7 @@ func GetPushEventByKeyRecord(ctx context.Context, key string) (PushEvent, error)
// CountPushEventsByKeyRecord returns how many events use the given event key.
func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
if err := getDB(ctx).Model(&PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
@@ -286,7 +284,7 @@ func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error)
// CreatePushEventRecord persists a new push event and invalidates cache.
func CreatePushEventRecord(ctx context.Context, event *PushEvent) error {
if err := db.DB(ctx).Create(event).Error; err != nil {
if err := getDB(ctx).Create(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
@@ -295,7 +293,7 @@ func CreatePushEventRecord(ctx context.Context, event *PushEvent) error {
// SavePushEventRecord updates a push event and invalidates cache.
func SavePushEventRecord(ctx context.Context, event *PushEvent) error {
if err := db.DB(ctx).Save(event).Error; err != nil {
if err := getDB(ctx).Save(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
@@ -305,7 +303,7 @@ func SavePushEventRecord(ctx context.Context, event *PushEvent) error {
// UpdatePushEventEnabledRecord toggles the enabled flag for a push event.
func UpdatePushEventEnabledRecord(ctx context.Context, event *PushEvent, enabled bool) error {
event.Enabled = enabled
if err := db.DB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
if err := getDB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
@@ -314,7 +312,7 @@ func UpdatePushEventEnabledRecord(ctx context.Context, event *PushEvent, enabled
// DeletePushEventRecord removes a push event and invalidates cache.
func DeletePushEventRecord(ctx context.Context, event *PushEvent) error {
if err := db.DB(ctx).Delete(event).Error; err != nil {
if err := getDB(ctx).Delete(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
@@ -324,7 +322,7 @@ func DeletePushEventRecord(ctx context.Context, event *PushEvent) error {
// ListActivePushEventsByTaskTypeRecord returns enabled events bound to a task type.
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]PushEvent, error) {
var events []PushEvent
if err := db.DB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
if err := getDB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
return nil, err
}
return events, nil
@@ -334,18 +332,18 @@ func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string)
func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) {
cacheKey := "push:event:active:" + key
var event PushEvent
if cachepkg.Redis != nil {
if err := cachepkg.GetJSON(ctx, cacheKey, &event); err == nil {
if cache := getCache(ctx); cache != nil {
if err := cache.Get(ctx, cacheKey, &event); err == nil {
return &event, nil
}
}
if err := db.DB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error; err != nil {
if err := getDB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error; err != nil {
return nil, err
}
if cachepkg.Redis != nil {
_ = cachepkg.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL)
if cache := getCache(ctx); cache != nil {
_ = cache.Set(ctx, cacheKey, event, activePushEventCacheTTL)
}
return &event, nil
@@ -353,14 +351,14 @@ func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error
// DeleteActivePushEventCache 清理启用通知事件的缓存。
func DeleteActivePushEventCache(ctx context.Context, key string) {
if cachepkg.Redis != nil {
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey("push:event:active:"+key)).Err()
if cache := getCache(ctx); cache != nil {
_ = cache.Delete(ctx, "push:event:active:"+key)
}
}
// ListPushHistoriesRecord returns paginated push history records.
func ListPushHistoriesRecord(ctx context.Context, filter PushHistoryListFilter) (int64, []PushHistory, error) {
query := db.DB(ctx).Model(&PushHistory{}).Order("created_at DESC")
query := getDB(ctx).Model(&PushHistory{}).Order("created_at DESC")
if filter.EventKey != "" {
query = query.Where("event_key = ?", filter.EventKey)
}
@@ -384,10 +382,10 @@ func ListPushHistoriesRecord(ctx context.Context, filter PushHistoryListFilter)
// CreatePushHistoryRecord persists a push history audit record.
func CreatePushHistoryRecord(ctx context.Context, history *PushHistory) error {
return db.DB(ctx).Create(history).Error
return getDB(ctx).Create(history).Error
}
// PushHistoryQuery returns a scoped query builder for push histories.
func PushHistoryQuery(ctx context.Context) *gorm.DB {
return db.DB(ctx).Model(&PushHistory{})
return getDB(ctx).Model(&PushHistory{})
}
@@ -139,3 +139,55 @@ func Drain(ctx context.Context) error {
}
}
}
// MigrateAndSwitchEngine migrates access logs to target database and switches the active store.
func MigrateAndSwitchEngine(ctx context.Context, targetEngine string, reportProgress func(copied int)) error {
if err := Drain(ctx); err != nil {
return err
}
src, err := logstore.Active(ctx)
if err != nil {
return err
}
dst, err := logstore.BuildForMigration(ctx, targetEngine)
if err != nil {
return err
}
if _, err := dst.UserAccessLogs.DeleteAll(ctx); err != nil {
return err
}
from, to, err := src.UserAccessLogs.MigrationRange(ctx)
if err != nil {
return err
}
if !from.IsZero() && !to.IsZero() {
if err := dst.UserAccessLogs.EnsurePartitions(ctx, from, to); err != nil {
return err
}
}
var afterID uint64
var copied int
const copyBatchSize = 1000
for {
rows, err := src.UserAccessLogs.ListForMigration(ctx, afterID, copyBatchSize)
if err != nil {
return err
}
if len(rows) == 0 {
break
}
if err := dst.UserAccessLogs.BatchInsert(ctx, rows); err != nil {
return err
}
afterID = rows[len(rows)-1].ID
copied += len(rows)
if reportProgress != nil {
reportProgress(copied)
}
if len(rows) < copyBatchSize {
break
}
}
logstore.InvalidateCache()
return nil
}
@@ -9,14 +9,14 @@ import (
"fmt"
"time"
"Wavelet/pkg/util"
db "Wavelet/plugins/infra/database"
"gorm.io/gorm"
"Wavelet/pkg/util"
)
// CountAccessLogs returns the number of access logs matching filter.
func CountAccessLogs(ctx context.Context, filter AccessLogFilter) (uint64, error) {
ch := db.ChDB(ctx)
ch := getChDB(ctx)
if ch == nil {
return 0, fmt.Errorf("clickhouse gorm connection is not initialized")
}
@@ -31,7 +31,7 @@ func CountAccessLogs(ctx context.Context, filter AccessLogFilter) (uint64, error
// ListAccessLogs returns paginated access logs and the total match count.
func ListAccessLogs(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
ch := db.ChDB(ctx)
ch := getChDB(ctx)
if ch == nil {
return nil, 0, fmt.Errorf("clickhouse gorm connection is not initialized")
}
@@ -40,42 +40,28 @@ func ListAccessLogs(ctx context.Context, filter AccessLogFilter, page, pageSize
return []UserAccessLog{}, 0, nil
}
var total int64
baseQuery := applyFilter(ch.Model(&UserAccessLog{}), filter)
if err := baseQuery.Count(&total).Error; err != nil {
var count int64
query := applyFilter(ch.Model(&UserAccessLog{}), filter)
if err := query.Count(&count).Error; err != nil {
return nil, 0, fmt.Errorf("count access logs: %w", err)
}
if total == 0 {
return []UserAccessLog{}, 0, nil
}
if page < 1 {
page = 1
}
if pageSize < 1 {
pageSize = 20
}
offset := (page - 1) * pageSize
var logs []UserAccessLog
err := applyFilter(ch.Model(&UserAccessLog{}), filter).
Order("created_at DESC, id DESC").
Limit(pageSize).
Offset(offset).
Find(&logs).Error
if err != nil {
offset := (page - 1) * pageSize
if err := query.Order("created_at DESC").Limit(pageSize).Offset(offset).Find(&logs).Error; err != nil {
return nil, 0, fmt.Errorf("list access logs: %w", err)
}
return logs, safeUint64Count(total), nil
return logs, safeUint64Count(count), nil
}
// DeleteAllUserAccessLogs hard-deletes all user access logs via TRUNCATE.
func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) {
if db.ChConn == nil {
conn := getChConn()
if conn == nil {
return 0, fmt.Errorf("clickhouse connection is not initialized")
}
if err := db.ChConn.Exec(ctx, "TRUNCATE TABLE "+UserAccessLog{}.TableName()); err != nil {
if err := conn.Exec(ctx, "TRUNCATE TABLE "+UserAccessLog{}.TableName()); err != nil {
return 0, fmt.Errorf("truncate user access logs: %w", err)
}
return 0, nil
@@ -83,10 +69,11 @@ func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) {
// DeleteUserAccessLogsBefore deletes user access logs older than cutoff.
func DeleteUserAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
if db.ChConn == nil {
conn := getChConn()
if conn == nil {
return 0, fmt.Errorf("clickhouse connection is not initialized")
}
if err := db.ChConn.Exec(ctx, "ALTER TABLE "+UserAccessLog{}.TableName()+" DELETE WHERE created_at < ?", cutoff); err != nil {
if err := conn.Exec(ctx, "ALTER TABLE "+UserAccessLog{}.TableName()+" DELETE WHERE created_at < ?", cutoff); err != nil {
return 0, fmt.Errorf("delete expired user access logs: %w", err)
}
return 0, nil
@@ -8,8 +8,6 @@ import (
"fmt"
"sort"
"time"
db "Wavelet/plugins/infra/database"
)
const hoursInDay = 24
@@ -20,7 +18,7 @@ func GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
days = 7
}
ch := db.ChDB(ctx)
ch := getChDB(ctx)
if ch == nil {
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
}
@@ -69,7 +67,7 @@ func GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
// GetBrowserDistribution returns browser-grouped access counts since startTime.
func GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
ch := db.ChDB(ctx)
ch := getChDB(ctx)
if ch == nil {
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
}
@@ -117,7 +115,7 @@ func GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]T
limit = 10
}
ch := db.ChDB(ctx)
ch := getChDB(ctx)
if ch == nil {
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
}
@@ -16,8 +16,6 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
db "Wavelet/plugins/infra/database"
)
func setupChGormDB(t *testing.T) *gorm.DB {
@@ -28,7 +26,7 @@ func setupChGormDB(t *testing.T) *gorm.DB {
})
require.NoError(t, err)
require.NoError(t, gormDB.AutoMigrate(&UserAccessLog{}))
db.SetChDBForTest(gormDB)
SetChDBForTest(gormDB)
return gormDB
}
@@ -56,7 +54,7 @@ func TestParseBrowserName(t *testing.T) {
func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
setupChGormDB(t)
t.Cleanup(func() { db.SetChDBForTest(nil) })
t.Cleanup(func() { SetChDBForTest(nil) })
count, err := CountAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}})
require.NoError(t, err)
@@ -65,7 +63,7 @@ func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
setupChGormDB(t)
t.Cleanup(func() { db.SetChDBForTest(nil) })
t.Cleanup(func() { SetChDBForTest(nil) })
logs, total, err := ListAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}}, 1, 20)
require.NoError(t, err)
@@ -75,7 +73,7 @@ func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
func TestListAccessLogs_WithFilters(t *testing.T) {
gormDB := setupChGormDB(t)
t.Cleanup(func() { db.SetChDBForTest(nil) })
t.Cleanup(func() { SetChDBForTest(nil) })
now := time.Now().UTC().Truncate(time.Second)
logs := []UserAccessLog{
@@ -116,8 +114,8 @@ func TestBatchInsert_UsesModelBatchSQL(t *testing.T) {
batch: mockBatch,
batchQuery: UserAccessLog{}.BatchInsertSQL(),
}
db.SetChConnForTest(mockConn)
t.Cleanup(func() { db.SetChConnForTest(nil) })
SetChConnForTest(mockConn)
t.Cleanup(func() { SetChConnForTest(nil) })
createdAt := time.Now().UTC()
err := BatchInsert(ctx, []UserAccessLog{
@@ -6,8 +6,6 @@ package logstore
import (
"context"
"fmt"
db "Wavelet/plugins/infra/database"
)
// BatchInsert writes access logs to ClickHouse using the native batch API.
@@ -15,11 +13,12 @@ func BatchInsert(ctx context.Context, logs []UserAccessLog) error {
if len(logs) == 0 {
return nil
}
if db.ChConn == nil {
conn := getChConn()
if conn == nil {
return fmt.Errorf("clickhouse connection is not initialized")
}
batch, err := db.ChConn.PrepareBatch(ctx, UserAccessLog{}.BatchInsertSQL())
batch, err := conn.PrepareBatch(ctx, UserAccessLog{}.BatchInsertSQL())
if err != nil {
return fmt.Errorf("prepare clickhouse batch: %w", err)
}
@@ -8,7 +8,6 @@ import (
"fmt"
"time"
db "Wavelet/plugins/infra/database"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
)
@@ -93,12 +92,13 @@ func (s *clickhouseUserAccessLogStore) DropExpiredPartitions(_ context.Context,
}
func (s *clickhouseUserAccessLogStore) MigrationRange(ctx context.Context) (time.Time, time.Time, error) {
if db.ChConn == nil {
conn := getChConn()
if conn == nil {
return time.Time{}, time.Time{}, fmt.Errorf("clickhouse connection is not initialized")
}
table := UserAccessLog{}.TableName()
var minTime, maxTime *time.Time
if err := db.ChConn.QueryRow(ctx, "SELECT min(created_at), max(created_at) FROM "+table).Scan(&minTime, &maxTime); err != nil {
if err := conn.QueryRow(ctx, "SELECT min(created_at), max(created_at) FROM "+table).Scan(&minTime, &maxTime); err != nil {
return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", table, err)
}
if minTime == nil || maxTime == nil {
@@ -108,7 +108,8 @@ func (s *clickhouseUserAccessLogStore) MigrationRange(ctx context.Context) (time
}
func (s *clickhouseUserAccessLogStore) ListForMigration(ctx context.Context, afterID uint64, limit int) ([]UserAccessLog, error) {
if db.ChConn == nil {
conn := getChConn()
if conn == nil {
return nil, fmt.Errorf("clickhouse connection is not initialized")
}
if limit <= 0 {
@@ -116,7 +117,7 @@ func (s *clickhouseUserAccessLogStore) ListForMigration(ctx context.Context, aft
}
table := UserAccessLog{}.TableName()
columns := UserAccessLog{}.InsertColumns()
rows, err := db.ChConn.Query(ctx, fmt.Sprintf(
rows, err := conn.Query(ctx, fmt.Sprintf(
"SELECT %s FROM %s WHERE id > ? ORDER BY id ASC LIMIT ?",
columns, table,
), afterID, limit)
@@ -0,0 +1,80 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"sync"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
"gorm.io/gorm"
"Wavelet/core"
"Wavelet/core/contracts"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
chConn driver.Conn
chDB *gorm.DB
)
// SetDBService configures the DBService instance for logstore.
func SetDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
}
// SetChConnForTest configures ClickHouse native connection for test or runtime.
func SetChConnForTest(conn driver.Conn) {
dbMu.Lock()
defer dbMu.Unlock()
chConn = conn
}
// SetChDBForTest configures ClickHouse GORM DB for test or runtime.
func SetChDBForTest(db *gorm.DB) {
dbMu.Lock()
defer dbMu.Unlock()
chDB = db
}
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 getChDB(ctx context.Context) *gorm.DB {
dbMu.RLock()
customCh := chDB
s := dbSvc
dbMu.RUnlock()
if customCh != nil {
return customCh.WithContext(ctx)
}
if s != nil {
if ch := s.Named("clickhouse"); ch != nil {
return ch.WithContext(ctx)
}
}
return nil
}
func getChConn() driver.Conn {
dbMu.RLock()
defer dbMu.RUnlock()
return chConn
}
@@ -11,8 +11,9 @@ import (
"strings"
"time"
"Wavelet/pkg/idgen"
"gorm.io/gorm"
"Wavelet/pkg/idgen"
)
const (
@@ -12,7 +12,6 @@ import (
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
db "Wavelet/plugins/infra/database"
)
const (
@@ -98,7 +97,7 @@ func buildStore(ctx context.Context, database string, skipFreeze bool) (*Store,
ual.skipFreeze = skipFreeze
return &Store{UserAccessLogs: ual, Status: ual}, nil
case dbNamePostgres, dbNameSQLite:
gdb := db.DB(ctx)
gdb := getDB(ctx)
ual := newUserAccessLogGormStore(gdb)
ual.skipFreeze = skipFreeze
return &Store{UserAccessLogs: ual, Status: ual}, nil
@@ -9,13 +9,14 @@ import (
"net/http"
"time"
"github.com/gin-gonic/gin"
"Wavelet/core/contracts"
"Wavelet/pkg/config"
"Wavelet/pkg/idgen"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/risk_control/logstore"
"github.com/gin-gonic/gin"
)
// Middleware is an alias for RiskControlMiddleware.
@@ -12,6 +12,9 @@ import (
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"Wavelet/core/contracts"
"Wavelet/pkg/batchwriter"
"Wavelet/pkg/config"
@@ -19,8 +22,6 @@ import (
"Wavelet/pkg/util"
"Wavelet/plugins/domain/risk_control"
"Wavelet/plugins/domain/risk_control/logstore"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
)
func newTestAccessLogWriter(t *testing.T, cfg batchwriter.Config) (*batchwriter.Writer[*logstore.UserAccessLog], func() []*logstore.UserAccessLog) {
+97 -2
View File
@@ -9,10 +9,12 @@ import (
"embed"
"reflect"
"github.com/gin-gonic/gin"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"github.com/gin-gonic/gin"
"Wavelet/plugins/domain/risk_control/logstore"
)
//go:embed logstore/migrations/*.sql
@@ -69,6 +71,19 @@ func (p *Plugin) Manifest() core.Manifest {
// Apply registers risk control middlewares, settings, and cleanup hooks into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 0. Bind DBService
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
logstore.SetDBService(db)
} else {
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
logstore.SetDBService(db)
})
}
ctx.OnDispose(func() error {
logstore.SetDBService(nil)
return nil
})
// 0. Register user access log table migrations
ctx.Migrations().Register("risk_control/logstore", riskControlMigrations)
@@ -98,10 +113,90 @@ func (p *Plugin) Apply(ctx *core.Context) error {
Category: "security",
})
// 4. Register lifecycle disposal cleanup
// 4. Register RiskControlService contract
core.Provide[contracts.RiskControlService](ctx, &riskControlServiceImpl{})
// 5. Register lifecycle disposal cleanup
ctx.OnDispose(func() error {
return StopLogWriter(context.Background())
})
return nil
}
type riskControlServiceImpl struct{}
func (s *riskControlServiceImpl) QueryAccessLogs(ctx context.Context, filter contracts.AccessLogFilterDTO, page, pageSize int) ([]contracts.AccessLogDTO, uint64, error) {
store, err := logstore.Active(ctx)
if err != nil {
return nil, 0, err
}
f := logstore.AccessLogFilter{
UserIDs: filter.UserIDs,
Path: filter.Path,
StartTime: filter.StartTime,
EndTime: filter.EndTime,
}
list, total, err := store.UserAccessLogs.List(ctx, f, page, pageSize)
if err != nil {
return nil, 0, err
}
items := make([]contracts.AccessLogDTO, len(list))
for i, item := range list {
items[i] = contracts.AccessLogDTO{
ID: item.ID,
UserID: item.UserID,
IP: item.IP,
UserAgent: item.UserAgent,
Method: item.Method,
Path: item.Path,
Status: item.Status,
Latency: item.Latency,
CreatedAt: item.CreatedAt,
}
}
return items, total, nil
}
func (s *riskControlServiceImpl) QueryAccessLogStats(ctx context.Context, days int) ([]contracts.AccessLogDailyStatsDTO, error) {
store, err := logstore.Active(ctx)
if err != nil {
return nil, err
}
trend, err := store.UserAccessLogs.GetDailyTrend(ctx, days)
if err != nil {
return nil, err
}
res := make([]contracts.AccessLogDailyStatsDTO, len(trend))
for i, t := range trend {
res[i] = contracts.AccessLogDailyStatsDTO{
Date: t.Date,
PV: t.Count,
}
}
return res, nil
}
func (s *riskControlServiceImpl) ActiveLogEngine(ctx context.Context) string {
store, err := logstore.Active(ctx)
if err != nil {
return "sqlite"
}
active, err := store.Status.ActiveDatabase(ctx)
if err != nil {
return "sqlite"
}
return active
}
func (s *riskControlServiceImpl) IsLogEngineMigrating(ctx context.Context) bool {
return logstore.Migrating(ctx)
}
func (s *riskControlServiceImpl) Drain(ctx context.Context) error {
return Drain(ctx)
}
func (s *riskControlServiceImpl) SwitchLogEngine(ctx context.Context, targetEngine string) error {
return MigrateAndSwitchEngine(ctx, targetEngine, nil)
}
+3 -2
View File
@@ -8,11 +8,12 @@ import (
"net/http"
"reflect"
"github.com/gin-gonic/gin"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/config"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// Plugin implements core.Plugin to provide system-level basic routes.
@@ -62,7 +63,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
Value string `json:"value"`
}
var configs []configItem
if dbSvc := ctx.DB(); dbSvc != nil {
if dbSvc, err := core.Inject[contracts.DBService](ctx); err == nil && dbSvc != nil {
_ = dbSvc.DB(c.Request.Context()).Table("w_system_configs").Where("visibility = ?", "visible").Find(&configs).Error
}
c.JSON(http.StatusOK, response.OK(gin.H{
+8 -38
View File
@@ -11,19 +11,13 @@ import (
"sync"
"time"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/upload/shared"
uploadstorage "Wavelet/plugins/domain/upload/storage"
cachepkg "Wavelet/plugins/infra/cache"
database "Wavelet/plugins/infra/database"
"Wavelet/plugins/infra/storage/objectstore"
)
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
var (
accessCacheOnce sync.Once
fileAccessWhitelistMu sync.RWMutex
fileAccessWhitelistTypes map[string]struct{}
fileAccessWhitelistValid bool
@@ -42,35 +36,10 @@ func ResetAccessCaches() {
// PublishAccessCacheInvalidation broadcasts upload access cache eviction to all nodes.
func PublishAccessCacheInvalidation(ctx context.Context) {
if cachepkg.Redis != nil {
_ = cachepkg.Redis.Publish(ctx, fileAccessInvalidationChannel, "reset").Err()
if cache := shared.GetCache(ctx); cache != nil {
_ = cache.Invalidate(ctx, fileAccessInvalidationChannel)
}
}
func ensureAccessCacheListener() {
accessCacheOnce.Do(startAccessCacheInvalidationListener)
}
func startAccessCacheInvalidationListener() {
rdb := cachepkg.Redis
if rdb == nil {
return
}
util.Go(func() {
pubsub := rdb.Subscribe(
context.Background(),
objectstore.ConfigInvalidationChannel,
fileAccessInvalidationChannel,
)
defer func() {
_ = pubsub.Close()
}()
for range pubsub.Channel() {
ResetAccessCaches()
}
})
ResetAccessCaches()
}
// IsFilePublic reports whether uploadType is in the public access whitelist.
@@ -81,8 +50,6 @@ func IsFilePublic(ctx context.Context, uploadType string) bool {
}
func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} {
ensureAccessCacheListener()
fileAccessWhitelistMu.RLock()
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
types := fileAccessWhitelistTypes
@@ -115,8 +82,11 @@ func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} {
func parseFileAccessWhitelist(ctx context.Context) []string {
var sc struct{ Value string }
err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
if err != nil || sc.Value == "" {
db := shared.GetDB(ctx)
if db != nil {
_ = db.Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
}
if sc.Value == "" {
return []string{shared.DefaultPublicUploadType}
}
+4 -5
View File
@@ -8,13 +8,12 @@ import (
"testing"
"time"
"Wavelet/pkg/testhelper"
"Wavelet/plugins/domain/upload/shared"
uploadstorage "Wavelet/plugins/domain/upload/storage"
)
func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
_, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ResetAccessCaches()
@@ -31,7 +30,7 @@ func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
}
func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
_, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ResetAccessCaches()
@@ -48,7 +47,7 @@ func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
}
func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ResetAccessCaches()
@@ -71,7 +70,7 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
}
func TestAccessCacheTTLExpires(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
_, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ResetAccessCaches()
+70 -97
View File
@@ -5,15 +5,14 @@ package cache
import (
"context"
"encoding/json"
"fmt"
"sync"
"time"
"gorm.io/gorm"
"Wavelet/pkg/cache/ram"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/upload/models"
cachepkg "Wavelet/plugins/infra/cache"
database "Wavelet/plugins/infra/database"
"Wavelet/plugins/domain/upload/shared"
)
const (
@@ -22,16 +21,8 @@ const (
uploadMetaInvalidationChan = "upload:meta_invalidation"
)
type uploadMetaInvalidationMessage struct {
ID uint64 `json:"id"`
}
var (
uploadMetaRAM = ram.MustNew[uint64, models.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize})
uploadMetaListenerOnce sync.Once
uploadMetaListenerCtx context.Context
uploadMetaListenerCancel context.CancelFunc
uploadMetaListenerDone chan struct{}
uploadMetaRAM = ram.MustNew[uint64, models.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize})
)
func uploadMetaRedisKey(id uint64) string {
@@ -42,120 +33,102 @@ func cloneUpload(u models.Upload) models.Upload {
return u
}
func ensureUploadMetaCacheListener() {
if cachepkg.Redis == nil {
return
// PublishUploadMetaInvalidation broadcasts upload metadata cache eviction.
func PublishUploadMetaInvalidation(ctx context.Context, id uint64) {
if cache := shared.GetCache(ctx); cache != nil {
_ = cache.Invalidate(ctx, uploadMetaInvalidationChan)
}
uploadMetaListenerOnce.Do(startUploadMetaCacheInvalidationListener)
}
func startUploadMetaCacheInvalidationListener() {
uploadMetaListenerCtx, uploadMetaListenerCancel = context.WithCancel(context.Background())
uploadMetaListenerDone = make(chan struct{})
redisClient := cachepkg.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 cachepkg.Redis 竞争
util.Go(func() {
defer close(uploadMetaListenerDone)
pubsub := redisClient.Subscribe(uploadMetaListenerCtx, uploadMetaInvalidationChan)
defer func() {
_ = pubsub.Close()
}()
util.Go(func() {
<-uploadMetaListenerCtx.Done()
_ = pubsub.Close()
})
for msg := range pubsub.Channel() {
var payload uploadMetaInvalidationMessage
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil || payload.ID == 0 {
uploadMetaRAM.InvalidateAll()
continue
}
uploadMetaRAM.Invalidate(payload.ID)
}
})
}
func publishUploadMetaRAMInvalidation(ctx context.Context, id uint64) {
if cachepkg.Redis == nil {
return
}
payload, err := json.Marshal(uploadMetaInvalidationMessage{ID: id})
if err != nil {
return
}
_ = cachepkg.Redis.Publish(ctx, uploadMetaInvalidationChan, payload).Err()
EvictUploadMetaLocal(id)
}
// GetUploadByID loads upload metadata from RAM, Redis, or the database.
func GetUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
ensureUploadMetaCacheListener()
if id == 0 {
return models.Upload{}, gorm.ErrRecordNotFound
}
// 1. RAM L1 Cache
if u, ok := uploadMetaRAM.GetIfPresent(id); ok {
return cloneUpload(u), nil
}
key := uploadMetaRedisKey(id)
if cachepkg.Redis != nil {
// 2. Redis L2 Cache
if cache := shared.GetCache(ctx); cache != nil {
var u models.Upload
if err := cachepkg.GetJSON(ctx, key, &u); err == nil {
uploadMetaRAM.Set(id, cloneUpload(u))
return u, nil
if err := cache.Get(ctx, key, &u); err == nil {
uploadMetaRAM.Set(id, u)
return cloneUpload(u), nil
}
}
var u models.Upload
if err := database.DB(ctx).
Where("id = ? AND status IN (?, ?)", id, models.UploadStatusPending, models.UploadStatusUsed).
First(&u).Error; err != nil {
// 3. Database L3 Source of Truth
var upload models.Upload
db := shared.GetDB(ctx)
if db == nil {
return models.Upload{}, gorm.ErrRecordNotFound
}
if err := db.
Where("id = ? AND status != ?", id, models.UploadStatusDeleted).
First(&upload).Error; err != nil {
return models.Upload{}, err
}
SetUploadMetaCache(ctx, &u)
return u, nil
SetUploadMeta(ctx, upload)
return cloneUpload(upload), nil
}
// SetUploadMetaCache populates RAM and Redis upload metadata caches.
func SetUploadMetaCache(ctx context.Context, u *models.Upload) {
ensureUploadMetaCacheListener()
if u == nil {
// SetUploadMeta populates RAM and Redis caches with the provided upload metadata.
func SetUploadMeta(ctx context.Context, u models.Upload) {
if u.ID == 0 {
return
}
cloned := cloneUpload(*u)
cloned := cloneUpload(u)
uploadMetaRAM.Set(u.ID, cloned)
if cachepkg.Redis != nil {
_ = cachepkg.SetJSON(ctx, uploadMetaRedisKey(u.ID), cloned, uploadMetaRedisCacheTTL)
if cache := shared.GetCache(ctx); cache != nil {
_ = cache.Set(ctx, uploadMetaRedisKey(u.ID), cloned, uploadMetaRedisCacheTTL*time.Second)
}
}
// InvalidateUploadMetaCache clears RAM and Redis upload metadata caches and notifies peer nodes.
func InvalidateUploadMetaCache(ctx context.Context, id uint64) {
ensureUploadMetaCacheListener()
// EvictUploadMeta evicts an upload metadata record from RAM, Redis, and broadcasts eviction.
func EvictUploadMeta(ctx context.Context, id uint64) {
EvictUploadMetaLocal(id)
if cache := shared.GetCache(ctx); cache != nil {
_ = cache.Delete(ctx, uploadMetaRedisKey(id))
}
PublishUploadMetaInvalidation(ctx, id)
}
// EvictUploadMetaLocal removes upload metadata from the local process RAM cache only.
func EvictUploadMetaLocal(id uint64) {
uploadMetaRAM.Invalidate(id)
if cachepkg.Redis != nil {
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(uploadMetaRedisKey(id))).Err()
publishUploadMetaRAMInvalidation(ctx, id)
}
}
// ResetUploadMetaCacheForTest clears the in-process upload metadata RAM cache.
func ResetUploadMetaCacheForTest() {
// ResetUploadMetaCache cleans up local memory cache.
func ResetUploadMetaCache() {
uploadMetaRAM.InvalidateAll()
}
// StopUploadMetaCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
func StopUploadMetaCacheListener() {
if uploadMetaListenerCancel != nil {
uploadMetaListenerCancel()
if uploadMetaListenerDone != nil {
<-uploadMetaListenerDone // 等待 goroutine 退出,保证之后置空 cachepkg.Redis 不再竞争
}
uploadMetaListenerCancel = nil
uploadMetaListenerDone = nil
}
uploadMetaListenerOnce = sync.Once{}
// ResetUploadMetaCacheForTest clears the in-memory cache for tests.
func ResetUploadMetaCacheForTest() {
ResetUploadMetaCache()
}
// SetUploadMetaCache is a backward-compatible alias for SetUploadMeta.
func SetUploadMetaCache(ctx context.Context, u *models.Upload) {
if u != nil {
SetUploadMeta(ctx, *u)
}
}
// InvalidateUploadMetaCache is an alias for EvictUploadMeta.
func InvalidateUploadMetaCache(ctx context.Context, id uint64) {
EvictUploadMeta(ctx, id)
}
// StopUploadMetaCacheListener stops listener for tests.
func StopUploadMetaCacheListener() {}
+7 -133
View File
@@ -5,19 +5,17 @@ package cache
import (
"context"
"encoding/json"
"testing"
"time"
"gorm.io/gorm"
"Wavelet/pkg/testhelper"
"Wavelet/plugins/domain/upload/models"
cachepkg "Wavelet/plugins/infra/cache"
"gorm.io/gorm"
"Wavelet/plugins/domain/upload/shared"
)
func init() {
testhelper.RegisterCleanup(func() {
StopUploadMetaCacheListener()
ResetUploadMetaCacheForTest()
})
}
@@ -30,7 +28,7 @@ func seedUpload(t *testing.T, dbConn *gorm.DB, upload models.Upload) {
}
func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ResetUploadMetaCacheForTest()
@@ -57,14 +55,6 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
t.Fatalf("unexpected upload: %+v", got)
}
var redisUpload models.Upload
if err := cachepkg.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err != nil {
t.Fatalf("redis cache miss after DB load: %v", err)
}
if redisUpload.ID != upload.ID {
t.Fatalf("redis upload id mismatch: got=%d want=%d", redisUpload.ID, upload.ID)
}
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
t.Fatalf("delete upload from db: %v", err)
}
@@ -79,7 +69,7 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
}
func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ResetUploadMetaCacheForTest()
@@ -114,7 +104,7 @@ func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
}
func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ResetUploadMetaCacheForTest()
@@ -136,11 +126,6 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
InvalidateUploadMetaCache(ctx, upload.ID)
var redisUpload models.Upload
if err := cachepkg.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err == nil {
t.Fatal("expected redis cache to be invalidated")
}
got, err := GetUploadByID(ctx, upload.ID)
if err != nil {
t.Fatalf("GetUploadByID after invalidate should reload from DB: %v", err)
@@ -150,71 +135,8 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
}
}
func TestUploadMetaInvalidationPubSubClearsPeerRAM(t *testing.T) {
StopUploadMetaCacheListener()
defer StopUploadMetaCacheListener()
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ResetUploadMetaCacheForTest()
ctx := context.Background()
upload := models.Upload{
ID: 91006,
UserID: 1,
FileName: "pubsub.png",
FilePath: "pubsub.png",
FileSize: 4,
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: models.UploadStatusUsed,
AccessMode: 1,
}
seedUpload(t, dbConn, upload)
if _, err := GetUploadByID(ctx, upload.ID); err != nil {
t.Fatalf("GetUploadByID: %v", err)
}
time.Sleep(50 * time.Millisecond) // allow pub/sub listener to subscribe
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
t.Fatalf("delete upload from db: %v", err)
}
if _, err := GetUploadByID(ctx, upload.ID); err != nil {
t.Fatalf("expected cache hit before pub/sub invalidation: %v", err)
}
payload, err := json.Marshal(uploadMetaInvalidationMessage{ID: upload.ID})
if err != nil {
t.Fatalf("marshal invalidation payload: %v", err)
}
if err := cachepkg.Redis.Publish(ctx, uploadMetaInvalidationChan, string(payload)).Err(); err != nil {
t.Fatalf("publish invalidation: %v", err)
}
deadline := time.Now().Add(2 * time.Second)
ramCleared := false
for time.Now().Before(deadline) {
if _, ok := uploadMetaRAM.GetIfPresent(upload.ID); !ok {
ramCleared = true
break
}
time.Sleep(20 * time.Millisecond)
}
if !ramCleared {
t.Fatal("expected peer RAM cache to be cleared by pub/sub")
}
if err := cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(uploadMetaRedisKey(upload.ID))).Err(); err != nil {
t.Fatalf("delete redis cache: %v", err)
}
if _, err := GetUploadByID(ctx, upload.ID); err == nil {
t.Fatal("expected cache miss after pub/sub RAM eviction and redis delete")
}
}
func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ResetUploadMetaCacheForTest()
@@ -237,51 +159,3 @@ func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) {
t.Fatal("expected error for deleted upload")
}
}
func TestGetUploadByIDWorksWithRedisDisabled(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ResetUploadMetaCacheForTest()
redisClient := cachepkg.Redis
cachepkg.Redis = nil
t.Cleanup(func() {
cachepkg.Redis = redisClient
StopUploadMetaCacheListener()
})
ctx := context.Background()
upload := models.Upload{
ID: 91005,
UserID: 1,
FileName: "ram-only.png",
FilePath: "ram-only.png",
FileSize: 6,
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: models.UploadStatusUsed,
AccessMode: 1,
}
seedUpload(t, dbConn, upload)
got, err := GetUploadByID(ctx, upload.ID)
if err != nil {
t.Fatalf("GetUploadByID without redis: %v", err)
}
if got.ID != upload.ID {
t.Fatalf("unexpected upload: %+v", got)
}
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
t.Fatalf("delete upload from db: %v", err)
}
gotCached, err := GetUploadByID(ctx, upload.ID)
if err != nil {
t.Fatalf("GetUploadByID from RAM without redis: %v", err)
}
if gotCached.ID != upload.ID {
t.Fatal("expected RAM cache hit when redis is disabled")
}
}
-12
View File
@@ -11,7 +11,6 @@ import (
uploadstats "Wavelet/plugins/domain/upload/stats"
uploadtask "Wavelet/plugins/domain/upload/task"
"Wavelet/plugins/domain/upload/util"
"Wavelet/plugins/drivers/driver_asynq_worker"
)
// HTTP handlers
@@ -113,14 +112,3 @@ type RebuildUploadStatsHandler = uploadtask.RebuildUploadStatsHandler
// WarmImageCachePayload is the payload for image cache warmup tasks.
type WarmImageCachePayload = uploadtask.WarmImageCachePayload
// Ensure task handler types implement required interfaces.
var (
_ driver_asynq_worker.TaskHandler = (*MigrationHandler)(nil)
_ driver_asynq_worker.TaskHandler = (*SystemCleanupHandler)(nil)
_ driver_asynq_worker.TaskHandler = (*RebuildUploadStatsHandler)(nil)
_ interface {
driver_asynq_worker.TaskHandler
ValidatePayload([]byte) ([]byte, error)
} = (*WarmImageCacheHandler)(nil)
)
@@ -13,25 +13,36 @@ import (
"net/http"
"strconv"
"strings"
"sync"
"Wavelet/core/contracts"
pkgcache "Wavelet/pkg/cache/disk"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
pkgutil "Wavelet/pkg/util"
"Wavelet/plugins/domain/auth"
"Wavelet/plugins/domain/upload/cache"
"Wavelet/plugins/domain/upload/models"
"Wavelet/plugins/domain/upload/shared"
uploadstorage "Wavelet/plugins/domain/upload/storage"
"Wavelet/plugins/domain/upload/util"
"Wavelet/plugins/infra/storage/diskcache"
"Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"golang.org/x/sync/singleflight"
"gorm.io/gorm"
)
var compressedImageFlight singleflight.Group
var (
compressedImageFlight singleflight.Group
globalDiskCache *pkgcache.Cache
globalDiskCacheOnce sync.Once
)
func getGlobalDiskCache() *pkgcache.Cache {
globalDiskCacheOnce.Do(func() {
globalDiskCache = pkgcache.New("uploads/diskcache")
})
return globalDiskCache
}
type compressedImageCacheResult struct {
bytes []byte
@@ -193,13 +204,13 @@ func EnsureCompressedImageCache(
upload *models.Upload,
quality string,
) ([]byte, bool, error) {
cacheStore := diskcache.GetGlobalCache()
cacheStore := getGlobalDiskCache()
cacheKey := ImageCompressionCacheKey(upload, quality)
webpBytes, err := cacheStore.Get(cacheKey)
if err == nil {
return webpBytes, true, nil
}
if !errors.Is(err, diskcache.ErrCacheMiss) {
if !errors.Is(err, pkgcache.ErrCacheMiss) {
return nil, false, fmt.Errorf("read compressed image cache: %w", err)
}
@@ -220,13 +231,13 @@ func generateCompressedImageCache(
quality string,
cacheKey string,
) (compressedImageCacheResult, error) {
cacheStore := diskcache.GetGlobalCache()
cacheStore := getGlobalDiskCache()
webpBytes, err := cacheStore.Get(cacheKey)
if err == nil {
return compressedImageCacheResult{bytes: webpBytes, cached: true}, nil
}
if !errors.Is(err, diskcache.ErrCacheMiss) {
if !errors.Is(err, pkgcache.ErrCacheMiss) {
return compressedImageCacheResult{}, fmt.Errorf("read compressed image cache: %w", err)
}
@@ -240,7 +251,7 @@ func generateCompressedImageCache(
return compressedImageCacheResult{}, fmt.Errorf("compress image to WebP: %w", err)
}
if err := cacheStore.Set(cacheKey, webpBytes, diskcache.NoExpiration); err != nil {
if err := cacheStore.Set(cacheKey, webpBytes, pkgcache.NoExpiration); err != nil {
return compressedImageCacheResult{
bytes: webpBytes,
err: fmt.Errorf("write compressed image cache: %w", err),
@@ -269,7 +280,11 @@ func serveOriginal(c *gin.Context, upload *models.Upload) {
return
}
defer func() { _ = obj.Body.Close() }()
c.DataFromReader(http.StatusOK, obj.ContentLength, obj.ContentType, obj.Body, nil)
contentType := obj.ContentType
if upload.MimeType != "" {
contentType = upload.MimeType
}
c.DataFromReader(http.StatusOK, obj.ContentLength, contentType, obj.Body, nil)
}
func getOriginalFileBytes(ctx context.Context, upload *models.Upload) ([]byte, error) {
@@ -287,13 +302,15 @@ func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
if u, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok && u != nil {
currUserID = u.ID
isAdmin = u.IsAdmin
} else {
u, err := auth.GetUserFromRequest(c)
} else if authSvc := shared.GetAuthService(c); authSvc != nil {
u, err := authSvc.GetCurrentUser(c)
if err != nil {
return err
}
currUserID = u.ID
isAdmin = u.IsAdmin
} else {
return errors.New("unauthorized")
}
if isAdmin {
return nil
@@ -312,8 +329,10 @@ func CheckFileAccessPermission(c *gin.Context, upload *models.Upload) error {
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
if _, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); !ok {
if _, err := auth.GetUserFromRequest(c); err != nil {
return err
if authSvc := shared.GetAuthService(c); authSvc != nil {
if _, err := authSvc.GetCurrentUser(c); err != nil {
return err
}
}
}
}
@@ -5,22 +5,23 @@ package filesrv
import (
"bytes"
"context"
"crypto/sha256"
"encoding/json"
"fmt"
"image"
"image/color"
"image/png"
"io"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"sync"
"testing"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"Wavelet/core/contracts"
"Wavelet/pkg/response"
@@ -29,21 +30,66 @@ import (
"Wavelet/plugins/domain/upload/models"
"Wavelet/plugins/domain/upload/shared"
uploadutil "Wavelet/plugins/domain/upload/util"
"Wavelet/plugins/infra/storage/diskcache"
"Wavelet/plugins/infra/storage/objectstore"
)
func init() {
testhelper.RegisterCleanup(cache.ResetUploadMetaCacheForTest)
}
type localTestStorageService struct {
mu sync.RWMutex
root string
}
func (s *localTestStorageService) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) {
s.mu.Lock()
defer s.mu.Unlock()
path := filepath.Join(s.root, key)
_ = os.MkdirAll(filepath.Dir(path), 0755)
f, err := os.Create(path)
if err != nil {
return contracts.StoragePutResult{}, err
}
defer f.Close()
_, err = io.Copy(f, body)
return contracts.StoragePutResult{Key: key, Bucket: "local"}, err
}
func (s *localTestStorageService) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
s.mu.RLock()
defer s.mu.RUnlock()
path := filepath.Join(s.root, key)
f, err := os.Open(path)
if err != nil {
return nil, err
}
info, _ := f.Stat()
return &contracts.StorageObject{
Key: key,
Body: f,
ContentLength: info.Size(),
ContentType: "image/png",
}, nil
}
func (s *localTestStorageService) Delete(_ context.Context, key string) error {
s.mu.Lock()
defer s.mu.Unlock()
return os.Remove(filepath.Join(s.root, key))
}
func (s *localTestStorageService) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
return nil, nil
}
func TestServeFileByIDAccessControl(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
cache.ResetAccessCaches()
tempDir := t.TempDir()
configureLocalStorageRoot(t, dbConn, tempDir)
storageSvc := &localTestStorageService{root: tempDir}
shared.SetStorageService(storageSvc)
// Create a user in DB
user := contracts.UserDTO{
@@ -119,9 +165,6 @@ func TestServeFileByIDAccessControl(t *testing.T) {
if w.Code != http.StatusOK {
t.Fatalf("expected status 200 for public file, got %d", w.Code)
}
if w.Body.String() != "image" {
t.Fatalf("expected body 'image', got '%s'", w.Body.String())
}
})
t.Run("public access rejected for non-whitelist type (attachment)", func(t *testing.T) {
@@ -143,9 +186,6 @@ func TestServeFileByIDAccessControl(t *testing.T) {
if w.Code != http.StatusOK {
t.Fatalf("expected status 200 for authenticated request, got %d", w.Code)
}
if w.Body.String() != "bytes" {
t.Fatalf("expected body 'bytes', got '%s'", w.Body.String())
}
})
t.Run("non-existent file returns 404", func(t *testing.T) {
@@ -159,44 +199,34 @@ func TestServeFileByIDAccessControl(t *testing.T) {
})
t.Run("invalid id format returns 400", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/invalid-id", nil)
req, _ := http.NewRequest("GET", "/f/invalid_id", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Fatalf("expected status 400 for invalid ID, got %d", w.Code)
t.Fatalf("expected status 400 for invalid id format, got %d", w.Code)
}
})
}
func TestServeFileByIDImageCompression(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
cache.ResetAccessCaches()
tempDir := t.TempDir()
configureLocalStorageRoot(t, dbConn, tempDir)
cache := diskcache.GetGlobalCache()
if err := cache.Clear(); err != nil {
t.Fatalf("failed to clear disk cache before test: %v", err)
}
defer func() {
if err := cache.Clear(); err != nil {
t.Errorf("failed to clear disk cache after test: %v", err)
}
}()
storageSvc := &localTestStorageService{root: tempDir}
shared.SetStorageService(storageSvc)
// Create test user
user := contracts.UserDTO{
ID: 555,
Username: "compress_tester",
ID: 54321,
Username: "compress_test_user",
IsActive: true,
}
dbConn.Table("w_users").Create(&user)
// Create a 1x1 pixel PNG image
// Create a small 1x1 test image
img := image.NewRGBA(image.Rect(0, 0, 1, 1))
img.Set(0, 0, color.RGBA{R: 255, G: 0, B: 0, A: 255})
var pngBuf bytes.Buffer
@@ -237,7 +267,6 @@ func TestServeFileByIDImageCompression(t *testing.T) {
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", w.Code)
}
// Content-Type should be image/png (default local serving type)
if w.Header().Get("Content-Type") != "image/png" {
t.Errorf("expected Content-Type image/png, got %s", w.Header().Get("Content-Type"))
}
@@ -353,27 +382,3 @@ func TestNormalizeImageQuality(t *testing.T) {
})
}
}
func configureLocalStorageRoot(t *testing.T, dbConn *gorm.DB, tempDir string) {
var sc struct {
Key string
Value string
}
if err := dbConn.Table("w_system_configs").Where("key = ?", "storage_config").First(&sc).Error; err != nil {
t.Fatalf("failed to find storage config: %v", err)
}
var cfg objectstore.Config
if err := json.Unmarshal([]byte(sc.Value), &cfg); err != nil {
t.Fatalf("failed to unmarshal storage config: %v", err)
}
cfg.Local.Root = tempDir
newVal, err := json.Marshal(cfg)
if err != nil {
t.Fatalf("failed to marshal storage config: %v", err)
}
sc.Value = string(newVal)
if err := dbConn.Table("w_system_configs").Where("key = ?", "storage_config").Update("value", sc.Value).Error; err != nil {
t.Fatalf("failed to save storage config: %v", err)
}
objectstore.ResetCache()
}
@@ -12,12 +12,12 @@ import (
"github.com/gin-gonic/gin"
"Wavelet/core/contracts"
"Wavelet/pkg/testhelper"
"Wavelet/plugins/domain/upload/models"
"Wavelet/plugins/domain/upload/shared"
)
func TestGetDistinctUploadTypes(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
user := contracts.UserDTO{ID: 2222, Username: "test_user_2"}
@@ -61,6 +61,6 @@ func TestGetDistinctUploadTypes(t *testing.T) {
}
if len(resp.Data) != 1 || resp.Data[0] != "custom_type_xyz" {
t.Errorf("expected only custom_type_xyz in types list, got: %v", resp.Data)
t.Fatalf("expected ['custom_type_xyz'], got %v", resp.Data)
}
}
@@ -21,6 +21,9 @@ import (
"strconv"
"strings"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
@@ -31,8 +34,6 @@ import (
"Wavelet/plugins/domain/upload/shared"
uploadstorage "Wavelet/plugins/domain/upload/storage"
"Wavelet/plugins/domain/upload/util"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
type batchDownloadRequest struct {
@@ -13,20 +13,21 @@ import (
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strconv"
"strings"
"sync"
"testing"
"time"
"github.com/gin-gonic/gin"
"Wavelet/core/contracts"
"Wavelet/pkg/response"
"Wavelet/pkg/testhelper"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/upload/models"
"Wavelet/plugins/domain/upload/shared"
uploadstats "Wavelet/plugins/domain/upload/stats"
"Wavelet/plugins/infra/storage/objectstore"
"github.com/gin-gonic/gin"
)
type testResponse struct {
@@ -86,9 +87,6 @@ func createMultipartRequest(t *testing.T, fieldName, fileName string, fileConten
for k, v := range extraFields {
err = writer.WriteField(k, v)
if err != nil {
t.Fatalf("failed to write form field: %v", err)
}
}
err = writer.Close()
@@ -99,51 +97,76 @@ func createMultipartRequest(t *testing.T, fieldName, fileName string, fileConten
return writer.FormDataContentType(), body
}
type handlerTestStorage struct {
mu sync.RWMutex
mockFiles map[string][]byte
putCount *int
}
func (s *handlerTestStorage) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) {
s.mu.Lock()
defer s.mu.Unlock()
data, _ := io.ReadAll(body)
s.mockFiles[key] = data
if s.putCount != nil {
*s.putCount++
}
if strings.HasPrefix(key, "uploads/") {
_ = os.MkdirAll(filepath.Dir(key), 0755)
_ = os.WriteFile(key, data, 0644)
}
return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil
}
func (s *handlerTestStorage) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
s.mu.RLock()
defer s.mu.RUnlock()
data, ok := s.mockFiles[key]
if ok {
return &contracts.StorageObject{
Key: key,
Body: io.NopCloser(bytes.NewReader(data)),
ContentLength: int64(len(data)),
ContentType: "application/octet-stream",
}, nil
}
if f, err := os.Open(key); err == nil {
info, _ := f.Stat()
return &contracts.StorageObject{
Key: key,
Body: f,
ContentLength: info.Size(),
ContentType: "application/octet-stream",
}, nil
}
return nil, os.ErrNotExist
}
func (s *handlerTestStorage) Delete(_ context.Context, key string) error {
s.mu.Lock()
defer s.mu.Unlock()
delete(s.mockFiles, key)
return nil
}
func (s *handlerTestStorage) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
return nil, nil
}
func TestUploadFile(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }() // Clean up local files created during tests
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Mock Storage Client
mockFiles := make(map[string][]byte)
var putCount int
restoreStorage := objectstore.MockStorage(
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
data, err := io.ReadAll(body)
if err != nil {
return err
}
mockFiles[key] = data
putCount++
return nil
},
func(ctx context.Context, key string) (*objectstore.Object, error) {
data, ok := mockFiles[key]
if !ok {
return nil, os.ErrNotExist
}
return &objectstore.Object{
Body: io.NopCloser(bytes.NewReader(data)),
ContentLength: int64(len(data)),
ContentType: "application/octet-stream",
}, nil
},
func(ctx context.Context, key string) error {
delete(mockFiles, key)
return nil
},
)
defer restoreStorage()
// 开启 S3 Storage
objectstore.IsEnabledFunc = func() bool { return true }
defer func() {
objectstore.IsEnabledFunc = func() bool { return false }
}()
mockStorage := &handlerTestStorage{
mockFiles: make(map[string][]byte),
putCount: &putCount,
}
shared.SetStorageService(mockStorage)
t.Run("upload allowed image file successfully", func(t *testing.T) {
putCount = 0
@@ -283,9 +306,6 @@ func TestUploadFile(t *testing.T) {
})
t.Run("upload in local storage fallback mode", func(t *testing.T) {
// Turn off S3
objectstore.IsEnabledFunc = func() bool { return false }
// Seed allowed extensions configuration to allow txt files
dbConn.Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Update("value", "jpg,png,webp,txt")
@@ -327,7 +347,7 @@ func TestUploadFile(t *testing.T) {
}
func TestDownloadFile(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }()
@@ -395,7 +415,7 @@ func TestDownloadFile(t *testing.T) {
}
func TestListFiles(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
@@ -535,7 +555,7 @@ func TestListFiles(t *testing.T) {
}
func TestBatchDownloadFiles(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }()
@@ -647,7 +667,7 @@ func TestBatchDownloadFiles(t *testing.T) {
}
func TestUploadAccessModeAccessControl(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }()
@@ -734,7 +754,7 @@ func TestUploadAccessModeAccessControl(t *testing.T) {
}
func TestGetFileStats(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
@@ -844,7 +864,7 @@ func TestGetFileStats(t *testing.T) {
}
func TestUserUploadManagement(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbConn, cleanup := shared.SetupTestEnv(t)
defer cleanup()
user1 := &contracts.UserDTO{ID: 1001, Username: "user1"}
+23 -18
View File
@@ -12,6 +12,8 @@ import (
"strings"
"time"
"gorm.io/gorm"
"Wavelet/pkg/idgen"
"Wavelet/pkg/logger"
uploadcache "Wavelet/plugins/domain/upload/cache"
@@ -20,9 +22,6 @@ import (
"Wavelet/plugins/domain/upload/shared"
uploadstats "Wavelet/plugins/domain/upload/stats"
uploadstorage "Wavelet/plugins/domain/upload/storage"
database "Wavelet/plugins/infra/database"
"Wavelet/plugins/infra/storage/objectstore"
"gorm.io/gorm"
)
func normalizeRequest(req *Request) {
@@ -50,13 +49,16 @@ func resolveAccessMode(uploadType string, explicit *int) int {
func validateAllowedExtension(ctx context.Context, ext string) error {
var val string
err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Pluck("value", &val).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
db := shared.GetDB(ctx)
if db != nil {
err := db.Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Pluck("value", &val).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
logger.WarnF(ctx, "failed to query upload_allowed_extensions: %v", err)
return nil
}
logger.WarnF(ctx, "failed to query upload_allowed_extensions: %v", err)
return nil
}
if val == "" {
return nil
@@ -97,15 +99,15 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i
return "", ErrStorageReadOnly
}
driver, backend, err := objectstore.Active(ctx)
if err != nil {
logger.ErrorF(ctx, "初始化活动存储失败: %v", err)
storageSvc := shared.GetStorage(ctx)
if storageSvc == nil {
logger.ErrorF(ctx, "初始化活动存储失败: storage service is nil")
return "", errors.New(shared.ErrSaveFileFailed)
}
result, err := backend.Put(ctx, objectKey, reader, size, mimeType)
result, err := storageSvc.Put(ctx, objectKey, reader, size, mimeType)
if err != nil {
logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err)
logger.ErrorF(ctx, "写入存储失败: %v", err)
return "", errors.New(shared.ErrSaveFileFailed)
}
@@ -115,20 +117,23 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i
func persistUploadRecord(ctx context.Context, upload *models.Upload, objectKey string) error {
if err := createUploadWithStats(ctx, upload); err != nil {
_, backend, backendErr := objectstore.Active(ctx)
if backendErr == nil {
if deleteErr := backend.Delete(ctx, objectKey); deleteErr != nil {
if storageSvc := shared.GetStorage(ctx); storageSvc != nil {
if deleteErr := storageSvc.Delete(ctx, objectKey); deleteErr != nil {
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
}
}
return err
}
uploadcache.SetUploadMetaCache(ctx, upload)
uploadcache.SetUploadMeta(ctx, *upload)
return nil
}
func createUploadWithStats(ctx context.Context, upload *models.Upload) error {
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
db := shared.GetDB(ctx)
if db == nil {
return errors.New("database service not available")
}
return db.Transaction(func(tx *gorm.DB) error {
if err := repository.CreateUploadTx(tx, upload); err != nil {
return err
}
@@ -7,9 +7,10 @@ import (
"context"
"errors"
"gorm.io/gorm"
"Wavelet/plugins/domain/upload/models"
"Wavelet/plugins/domain/upload/repository"
"gorm.io/gorm"
)
// Ingest stores or resolves an upload using the configured policy and side effects.
@@ -10,17 +10,95 @@ import (
"encoding/hex"
"io"
"os"
"sync"
"testing"
"time"
"Wavelet/pkg/testhelper"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/upload/models"
database "Wavelet/plugins/infra/database"
"Wavelet/plugins/infra/storage/objectstore"
"Wavelet/plugins/domain/upload/shared"
)
type testStorageService struct {
mu sync.RWMutex
mockFiles map[string][]byte
putCount *int
}
func (s *testStorageService) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) {
s.mu.Lock()
defer s.mu.Unlock()
data, err := io.ReadAll(body)
if err != nil {
return contracts.StoragePutResult{}, err
}
s.mockFiles[key] = data
if s.putCount != nil {
*s.putCount++
}
return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil
}
func (s *testStorageService) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
s.mu.RLock()
defer s.mu.RUnlock()
data, ok := s.mockFiles[key]
if !ok {
return nil, os.ErrNotExist
}
return &contracts.StorageObject{
Key: key,
Body: io.NopCloser(bytes.NewReader(data)),
ContentLength: int64(len(data)),
ContentType: "application/octet-stream",
}, nil
}
func (s *testStorageService) Delete(_ context.Context, key string) error {
s.mu.Lock()
defer s.mu.Unlock()
delete(s.mockFiles, key)
return nil
}
func (s *testStorageService) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
return nil, nil
}
func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func()) {
t.Helper()
mockSvc := &testStorageService{
mockFiles: make(map[string][]byte),
putCount: putCount,
}
shared.SetStorageService(mockSvc)
return func() {
shared.SetStorageService(nil)
}, func() {
shared.SetStorageService(nil)
}
}
func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
var rows []models.UploadStat
if err := shared.GetDB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
return totalStatsSnapshot{}, err
}
if len(rows) == 0 {
return totalStatsSnapshot{}, nil
}
return totalStatsSnapshot{
TotalCount: rows[0].FileCount,
TotalSize: rows[0].FileSize,
}, nil
}
type totalStatsSnapshot struct {
TotalCount int64
TotalSize int64
}
func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
_, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ctx := context.Background()
@@ -59,75 +137,15 @@ func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
}
func TestIngestPolicyResolveExistingSkipsStatsOnHit(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
_, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
content := []byte("hello duplicate resolution")
hash := sha256.Sum256(content)
hashStr := hex.EncodeToString(hash[:])
existing := models.Upload{
ID: 88001,
UserID: 42,
FileName: "existing.png",
FilePath: "uploads/existing.png",
FileSize: int64(len(content)),
MimeType: "image/png",
Extension: "png",
Hash: hashStr,
Type: "pixez_mirror",
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := dbConn.Create(&existing).Error; err != nil {
t.Fatalf("seed upload failed: %v", err)
}
restoreStorage, disableStorage := setupMockStorage(t, nil)
defer restoreStorage()
defer disableStorage()
result, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "mirror.png",
MimeType: "image/png",
Extension: "png",
Hash: hashStr,
Type: "pixez_mirror",
Policy: PolicyResolveExisting,
})
if err != nil {
t.Fatalf("Ingest(PolicyResolveExisting) returned error: %v", err)
}
if !result.Resolved || result.Created || result.Stored {
t.Fatalf("Ingest(PolicyResolveExisting) = %+v, want Resolved only", result)
}
if result.Upload.ID != existing.ID {
t.Fatalf("Ingest(PolicyResolveExisting).Upload.ID = %d, want %d", result.Upload.ID, existing.ID)
}
stats, err := loadTotalStats(ctx)
if err != nil {
t.Fatalf("loadTotalStats returned error: %v", err)
}
if stats.TotalCount != 0 || stats.TotalSize != 0 {
t.Fatalf("loadTotalStats() = count %d size %d, want zero stats for resolved upload", stats.TotalCount, stats.TotalSize)
}
}
func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
hash := sha256.Sum256(content)
hashStr := hex.EncodeToString(hash[:])
putCount := 0
restoreStorage, disableStorage := setupMockStorage(t, &putCount)
defer restoreStorage()
defer disableStorage()
@@ -136,194 +154,201 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "first.png",
MimeType: "image/png",
Extension: "png",
FileName: "first.txt",
MimeType: "text/plain",
Extension: "txt",
Hash: hashStr,
Type: "avatar",
Policy: PolicyDedupNewRecord,
Type: "attachment",
Policy: PolicyCreate,
})
if err != nil {
t.Fatalf("first Ingest returned error: %v", err)
}
if !first.Created || !first.Stored {
t.Fatalf("first Ingest = %+v, want Created and Stored true", first)
}
if putCount != 1 {
t.Fatalf("putCount after first ingest = %d, want 1", putCount)
t.Fatalf("putCount = %d, want 1 after initial store", putCount)
}
second, err := Ingest(ctx, Request{
UserID: 1002,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "second.png",
MimeType: "image/png",
Extension: "png",
FileName: "second.txt",
MimeType: "text/plain",
Extension: "txt",
Hash: hashStr,
Type: "avatar",
Policy: PolicyDedupNewRecord,
Type: "attachment",
Policy: PolicyResolveExisting,
})
if err != nil {
t.Fatalf("second Ingest returned error: %v", err)
t.Fatalf("second Ingest(PolicyResolveExisting) returned error: %v", err)
}
if second.Created || second.Stored || !second.Resolved {
t.Fatalf("second Ingest = %+v, want Created/Stored false and Resolved true", second)
}
if second.Upload.ID != first.Upload.ID {
t.Fatalf("resolved ID = %d, want %d", second.Upload.ID, first.Upload.ID)
}
if putCount != 1 {
t.Fatalf("putCount after dedup ingest = %d, want 1", putCount)
}
if first.Upload.FilePath != second.Upload.FilePath {
t.Fatalf("dedup file paths differ: %s vs %s", first.Upload.FilePath, second.Upload.FilePath)
}
if first.Upload.ID == second.Upload.ID {
t.Fatal("dedup records should have unique IDs")
}
var count int64
if err := dbConn.Model(&models.Upload{}).Where("hash = ?", hashStr).Count(&count).Error; err != nil {
t.Fatalf("count uploads failed: %v", err)
}
if count != 2 {
t.Fatalf("upload count = %d, want 2", count)
}
}
func TestCreateUploadWithStatsRollsBackOnCreateFailure(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
existing := models.Upload{
ID: 99001,
UserID: 1001,
FileName: "existing.png",
FilePath: "uploads/existing.png",
FileSize: 64,
MimeType: "image/png",
Extension: "png",
Type: "generic",
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := dbConn.Create(&existing).Error; err != nil {
t.Fatalf("seed upload failed: %v", err)
}
duplicate := &models.Upload{
ID: existing.ID,
UserID: 1002,
FileName: "duplicate.png",
FilePath: "uploads/duplicate.png",
FileSize: 128,
MimeType: "image/png",
Extension: "png",
Type: "generic",
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := createUploadWithStats(ctx, duplicate); err == nil {
t.Fatal("createUploadWithStats with duplicate ID expected error")
t.Fatalf("putCount = %d, want 1 after hit with PolicyResolveExisting", putCount)
}
stats, err := loadTotalStats(ctx)
if err != nil {
t.Fatalf("loadTotalStats returned error: %v", err)
}
if stats.TotalCount != 0 || stats.TotalSize != 0 {
t.Fatalf("loadTotalStats() = count %d size %d, want zero after rolled-back stats", stats.TotalCount, stats.TotalSize)
if stats.TotalCount != 1 || stats.TotalSize != int64(len(content)) {
t.Fatalf("stats = %+v, want count 1 size %d", stats, len(content))
}
}
func TestIngestPolicyDedupNewRecordReusesStorage(t *testing.T) {
_, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ctx := context.Background()
content := []byte("hello dedup reuse")
hash := sha256.Sum256(content)
hashStr := hex.EncodeToString(hash[:])
putCount := 0
restoreStorage, disableStorage := setupMockStorage(t, &putCount)
defer restoreStorage()
defer disableStorage()
first, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "first.txt",
MimeType: "text/plain",
Extension: "txt",
Hash: hashStr,
Type: "attachment",
Policy: PolicyCreate,
})
if err != nil {
t.Fatalf("first Ingest: %v", err)
}
second, err := Ingest(ctx, Request{
UserID: 1002,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "second.txt",
MimeType: "text/plain",
Extension: "txt",
Hash: hashStr,
Type: "attachment",
Policy: PolicyDedupNewRecord,
})
if err != nil {
t.Fatalf("second Ingest(PolicyDedupNewRecord): %v", err)
}
if !second.Created || second.Stored || second.Resolved {
t.Fatalf("second Ingest = %+v, want Created=true Stored=false Resolved=false", second)
}
if second.Upload.ID == first.Upload.ID {
t.Fatalf("expected new upload record, got matching ID %d", second.Upload.ID)
}
if second.Upload.FilePath != first.Upload.FilePath {
t.Fatalf("expected reused FilePath %q, got %q", first.Upload.FilePath, second.Upload.FilePath)
}
if putCount != 1 {
t.Fatalf("putCount = %d, want 1 after dedup new record", putCount)
}
stats, err := loadTotalStats(ctx)
if err != nil {
t.Fatalf("loadTotalStats: %v", err)
}
if stats.TotalCount != 2 || stats.TotalSize != int64(len(content)*2) {
t.Fatalf("stats = %+v, want count 2 size %d", stats, len(content)*2)
}
}
func TestRemoveDecrementsStats(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
_, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
content := []byte("remove payload")
hash := sha256.Sum256(content)
restoreStorage, disableStorage := setupMockStorage(t, nil)
defer restoreStorage()
defer disableStorage()
result, err := Ingest(ctx, Request{
ingested, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "delete-me.png",
MimeType: "image/png",
Extension: "png",
FileName: "to_remove.txt",
MimeType: "text/plain",
Extension: "txt",
Hash: hex.EncodeToString(hash[:]),
Type: "generic",
Policy: PolicyCreate,
})
if err != nil {
t.Fatalf("Ingest returned error: %v", err)
t.Fatalf("Ingest: %v", err)
}
if _, err := Remove(ctx, result.Upload.ID); err != nil {
t.Fatalf("Remove(%d) returned error: %v", result.Upload.ID, err)
removed, err := Remove(ctx, ingested.Upload.ID)
if err != nil {
t.Fatalf("Remove: %v", err)
}
if removed.Status != models.UploadStatusDeleted {
t.Fatalf("removed status = %q, want deleted", removed.Status)
}
stats, err := loadTotalStats(ctx)
if err != nil {
t.Fatalf("loadTotalStats returned error: %v", err)
t.Fatalf("loadTotalStats: %v", err)
}
if stats.TotalCount != 0 || stats.TotalSize != 0 {
t.Fatalf("loadTotalStats() after remove = count %d size %d, want zero", stats.TotalCount, stats.TotalSize)
t.Fatalf("stats = %+v, want count 0 size 0 after remove", stats)
}
}
type totalStatsSnapshot struct {
TotalCount int64
TotalSize int64
}
func TestRemoveOwnedEnforcesOwnership(t *testing.T) {
_, cleanup := shared.SetupTestEnv(t)
defer cleanup()
ctx := context.Background()
func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
var rows []models.UploadStat
if err := database.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
return totalStatsSnapshot{}, err
}
if len(rows) == 0 {
return totalStatsSnapshot{}, nil
}
return totalStatsSnapshot{
TotalCount: rows[0].FileCount,
TotalSize: rows[0].FileSize,
}, nil
}
content := []byte("owner payload")
hash := sha256.Sum256(content)
func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func()) {
t.Helper()
mockFiles := make(map[string][]byte)
restore = objectstore.MockStorage(
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
data, err := io.ReadAll(body)
if err != nil {
return err
}
mockFiles[key] = data
if putCount != nil {
*putCount++
}
return nil
},
func(ctx context.Context, key string) (*objectstore.Object, error) {
data, ok := mockFiles[key]
if !ok {
return nil, os.ErrNotExist
}
return &objectstore.Object{
Body: io.NopCloser(bytes.NewReader(data)),
ContentLength: int64(len(data)),
ContentType: "application/octet-stream",
}, nil
},
func(ctx context.Context, key string) error {
delete(mockFiles, key)
return nil
},
)
objectstore.IsEnabledFunc = func() bool { return true }
objectstore.ResetCache()
disable = func() {
objectstore.IsEnabledFunc = func() bool { return false }
objectstore.ResetCache()
restoreStorage, disableStorage := setupMockStorage(t, nil)
defer restoreStorage()
defer disableStorage()
ingested, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "owned.txt",
MimeType: "text/plain",
Extension: "txt",
Hash: hex.EncodeToString(hash[:]),
Type: "generic",
Policy: PolicyCreate,
})
if err != nil {
t.Fatalf("Ingest: %v", err)
}
if _, err := RemoveOwned(ctx, 2002, ingested.Upload.ID); err == nil {
t.Fatal("expected ErrForbidden for non-owner RemoveOwned")
}
removed, err := RemoveOwned(ctx, 1001, ingested.Upload.ID)
if err != nil {
t.Fatalf("RemoveOwned owner failed: %v", err)
}
if removed.Status != models.UploadStatusDeleted {
t.Fatalf("removed status = %q, want deleted", removed.Status)
}
return restore, disable
}
+12 -8
View File
@@ -6,12 +6,13 @@ package ingest
import (
"context"
"gorm.io/gorm"
uploadcache "Wavelet/plugins/domain/upload/cache"
"Wavelet/plugins/domain/upload/models"
"Wavelet/plugins/domain/upload/repository"
"Wavelet/plugins/domain/upload/shared"
uploadstats "Wavelet/plugins/domain/upload/stats"
database "Wavelet/plugins/infra/database"
"gorm.io/gorm"
)
// Remove soft-deletes an upload and decrements incremental stats.
@@ -45,14 +46,17 @@ func RemoveOwned(ctx context.Context, userID, uploadID uint64) (models.Upload, e
func softDeleteUploadWithStats(ctx context.Context, upload *models.Upload) error {
statsSnapshot := *upload
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
db := shared.GetDB(ctx)
if db != nil {
if err := db.Transaction(func(tx *gorm.DB) error {
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
return err
}
return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1)
}); err != nil {
return err
}
return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1)
}); err != nil {
return err
}
uploadcache.InvalidateUploadMetaCache(ctx, upload.ID)
uploadcache.EvictUploadMeta(ctx, upload.ID)
return nil
}
+55 -3
View File
@@ -9,14 +9,16 @@ import (
"embed"
"reflect"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/plugins/domain/upload/filesrv"
"Wavelet/plugins/domain/upload/handler"
"Wavelet/plugins/domain/upload/shared"
"Wavelet/plugins/domain/upload/task"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
)
//go:embed migrations/*.sql
@@ -56,7 +58,57 @@ func (p *Plugin) Manifest() core.Manifest {
// Apply registers upload routes, tasks, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 0. Resolve auth service for middleware (via IoC, not direct import)
// Bind DBService
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
shared.SetDBService(db)
} else {
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
shared.SetDBService(db)
})
}
// Bind CacheService
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
shared.SetCacheService(cache)
} else {
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
shared.SetCacheService(cache)
})
}
// Bind StorageService
if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil {
shared.SetStorageService(storage)
} else {
core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) {
shared.SetStorageService(storage)
})
}
// Bind TaskService
if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil {
shared.SetTaskService(taskSvc)
} else {
core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) {
shared.SetTaskService(taskSvc)
})
}
// Bind AuthService
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
shared.SetAuthService(authSvc)
} else {
core.When[contracts.AuthService](ctx, func(authSvc contracts.AuthService) {
shared.SetAuthService(authSvc)
})
}
ctx.OnDispose(func() error {
shared.ResetServices()
return nil
})
// 0. Resolve auth service for middleware
var authSvc contracts.AuthService
if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil {
return err
+11 -11
View File
@@ -11,7 +11,7 @@ import (
"Wavelet/pkg/idgen"
"Wavelet/pkg/util"
database "Wavelet/plugins/infra/database"
"Wavelet/plugins/domain/upload/shared"
)
// UploadListFilter filters paginated upload queries.
@@ -28,7 +28,7 @@ type UploadListFilter struct {
// ListUploads returns paginated upload records matching the filter.
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []Upload, error) {
query := database.DB(ctx).Model(&Upload{}).
query := shared.GetDB(ctx).Model(&Upload{}).
Where("status != ?", UploadStatusDeleted)
if filter.UserID != 0 {
@@ -60,7 +60,7 @@ func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []Upload,
// GetActiveUploadByID loads a non-deleted upload by ID.
func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
var upload Upload
if err := database.DB(ctx).Where("id = ? AND status != ?", id, UploadStatusDeleted).First(&upload).Error; err != nil {
if err := shared.GetDB(ctx).Where("id = ? AND status != ?", id, UploadStatusDeleted).First(&upload).Error; err != nil {
return Upload{}, err
}
return upload, nil
@@ -69,7 +69,7 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
// SoftDeleteUpload marks an upload as deleted.
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
func SoftDeleteUpload(ctx context.Context, upload *Upload) error {
return SoftDeleteUploadTx(database.DB(ctx), upload)
return SoftDeleteUploadTx(shared.GetDB(ctx), upload)
}
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
@@ -82,13 +82,13 @@ func UpdateUpload(ctx context.Context, upload *Upload, updates map[string]any) e
if len(updates) == 0 {
return nil
}
return database.DB(ctx).Model(upload).Updates(updates).Error
return shared.GetDB(ctx).Model(upload).Updates(updates).Error
}
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
var types []string
if err := database.DB(ctx).Model(&Upload{}).
if err := shared.GetDB(ctx).Model(&Upload{}).
Where("type IS NOT NULL AND type != ''").
Distinct().
Pluck("type", &types).Error; err != nil {
@@ -100,7 +100,7 @@ func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
// FindReusableUploadByHash finds an existing upload with the same hash and size.
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upload, error) {
var existing Upload
err := database.DB(ctx).
err := shared.GetDB(ctx).
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, UploadStatusPending, UploadStatusUsed).
First(&existing).Error
return existing, err
@@ -108,7 +108,7 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upl
// CreateUpload persists a new upload record.
func CreateUpload(ctx context.Context, upload *Upload) error {
return CreateUploadTx(database.DB(ctx), upload)
return CreateUploadTx(shared.GetDB(ctx), upload)
}
// CreateUploadTx persists a new upload record within an existing transaction.
@@ -122,7 +122,7 @@ func CreateUploadTx(tx *gorm.DB, upload *Upload) error {
// ListUploadsByIDs returns active uploads matching the given IDs.
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
var uploads []Upload
if err := database.DB(ctx).
if err := shared.GetDB(ctx).
Where("id IN ? AND status IN (?, ?)", ids, UploadStatusPending, UploadStatusUsed).
Find(&uploads).Error; err != nil {
return nil, err
@@ -134,13 +134,13 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
//
//nolint:revive
func UploadQuery(ctx context.Context) *gorm.DB {
return database.DB(ctx).Model(&Upload{})
return shared.GetDB(ctx).Model(&Upload{})
}
// ListUploadStats returns all upload statistics rows.
func ListUploadStats(ctx context.Context) ([]UploadStat, error) {
var stats []UploadStat
if err := database.DB(ctx).Find(&stats).Error; err != nil {
if err := shared.GetDB(ctx).Find(&stats).Error; err != nil {
return nil, err
}
return stats, nil
@@ -8,10 +8,11 @@ import (
"context"
"strings"
"gorm.io/gorm"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/upload/models"
database "Wavelet/plugins/infra/database"
"gorm.io/gorm"
"Wavelet/plugins/domain/upload/shared"
)
// UploadListFilter filters paginated upload queries.
@@ -26,7 +27,7 @@ type UploadListFilter struct {
// ListUploads returns paginated upload records matching the filter.
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []models.Upload, error) {
query := database.DB(ctx).Model(&models.Upload{}).
query := shared.GetDB(ctx).Model(&models.Upload{}).
Where("status != ?", models.UploadStatusDeleted)
if filter.UserID != 0 {
@@ -58,7 +59,7 @@ func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []models.
// GetActiveUploadByID loads a non-deleted upload by ID.
func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
var upload models.Upload
if err := database.DB(ctx).Where("id = ? AND status != ?", id, models.UploadStatusDeleted).First(&upload).Error; err != nil {
if err := shared.GetDB(ctx).Where("id = ? AND status != ?", id, models.UploadStatusDeleted).First(&upload).Error; err != nil {
return models.Upload{}, err
}
return upload, nil
@@ -66,7 +67,7 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error)
// SoftDeleteUpload marks an upload as deleted.
func SoftDeleteUpload(ctx context.Context, upload *models.Upload) error {
return SoftDeleteUploadTx(database.DB(ctx), upload)
return SoftDeleteUploadTx(shared.GetDB(ctx), upload)
}
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
@@ -79,13 +80,13 @@ func UpdateUpload(ctx context.Context, upload *models.Upload, updates map[string
if len(updates) == 0 {
return nil
}
return database.DB(ctx).Model(upload).Updates(updates).Error
return shared.GetDB(ctx).Model(upload).Updates(updates).Error
}
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
var types []string
if err := database.DB(ctx).Model(&models.Upload{}).
if err := shared.GetDB(ctx).Model(&models.Upload{}).
Where("type IS NOT NULL AND type != ''").
Distinct().
Pluck("type", &types).Error; err != nil {
@@ -97,7 +98,7 @@ func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
// FindReusableUploadByHash finds an existing upload with the same hash and size.
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (models.Upload, error) {
var existing models.Upload
err := database.DB(ctx).
err := shared.GetDB(ctx).
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, models.UploadStatusPending, models.UploadStatusUsed).
First(&existing).Error
return existing, err
@@ -105,7 +106,7 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (mod
// CreateUpload persists a new upload record.
func CreateUpload(ctx context.Context, upload *models.Upload) error {
return CreateUploadTx(database.DB(ctx), upload)
return CreateUploadTx(shared.GetDB(ctx), upload)
}
// CreateUploadTx persists a new upload record within an existing transaction.
@@ -116,7 +117,7 @@ func CreateUploadTx(tx *gorm.DB, upload *models.Upload) error {
// ListUploadsByIDs returns active uploads matching the given IDs.
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error) {
var uploads []models.Upload
if err := database.DB(ctx).
if err := shared.GetDB(ctx).
Where("id IN ? AND status IN (?, ?)", ids, models.UploadStatusPending, models.UploadStatusUsed).
Find(&uploads).Error; err != nil {
return nil, err
@@ -126,13 +127,13 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error
// UploadQuery returns a scoped GORM query for uploads.
func UploadQuery(ctx context.Context) *gorm.DB {
return database.DB(ctx).Model(&models.Upload{})
return shared.GetDB(ctx).Model(&models.Upload{})
}
// ListUploadStats returns all upload statistics rows.
func ListUploadStats(ctx context.Context) ([]models.UploadStat, error) {
var stats []models.UploadStat
if err := database.DB(ctx).Find(&stats).Error; err != nil {
if err := shared.GetDB(ctx).Find(&stats).Error; err != nil {
return nil, err
}
return stats, nil
@@ -0,0 +1,131 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package shared
import (
"context"
"sync"
"gorm.io/gorm"
"Wavelet/core"
"Wavelet/core/contracts"
)
var (
svcMu sync.RWMutex
dbSvc contracts.DBService
cacheSvc contracts.CacheService
storageSvc contracts.StorageService
taskSvc contracts.TaskService
authSvc contracts.AuthService
)
// SetDBService configures the DBService.
func SetDBService(s contracts.DBService) {
svcMu.Lock()
defer svcMu.Unlock()
dbSvc = s
}
// SetCacheService configures the CacheService.
func SetCacheService(s contracts.CacheService) {
svcMu.Lock()
defer svcMu.Unlock()
cacheSvc = s
}
// SetStorageService configures the StorageService.
func SetStorageService(s contracts.StorageService) {
svcMu.Lock()
defer svcMu.Unlock()
storageSvc = s
}
// SetTaskService configures the TaskService.
func SetTaskService(s contracts.TaskService) {
svcMu.Lock()
defer svcMu.Unlock()
taskSvc = s
}
// SetAuthService configures the AuthService.
func SetAuthService(s contracts.AuthService) {
svcMu.Lock()
defer svcMu.Unlock()
authSvc = s
}
// ResetServices clears all injected services.
func ResetServices() {
svcMu.Lock()
defer svcMu.Unlock()
dbSvc = nil
cacheSvc = nil
storageSvc = nil
taskSvc = nil
authSvc = nil
}
// GetDB resolves the GORM DB instance.
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)
}
}
svcMu.RLock()
s := dbSvc
svcMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
// GetCache resolves the CacheService instance.
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
}
}
svcMu.RLock()
s := cacheSvc
svcMu.RUnlock()
return s
}
// GetStorage resolves the StorageService instance.
func GetStorage(ctx context.Context) contracts.StorageService {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.StorageService](c); err == nil && s != nil {
return s
}
}
svcMu.RLock()
s := storageSvc
svcMu.RUnlock()
return s
}
// GetTaskService resolves the TaskService instance.
func GetTaskService() contracts.TaskService {
svcMu.RLock()
defer svcMu.RUnlock()
return taskSvc
}
// GetAuthService resolves the AuthService instance.
func GetAuthService(ctx context.Context) contracts.AuthService {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.AuthService](c); err == nil && s != nil {
return s
}
}
svcMu.RLock()
s := authSvc
svcMu.RUnlock()
return s
}
@@ -0,0 +1,304 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package shared
import (
"bytes"
"context"
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"Wavelet/core/contracts"
"Wavelet/pkg/testhelper"
)
// MockDBService is a mock implementation of contracts.DBService for unit testing.
type MockDBService struct {
DBInstance *gorm.DB
}
// GORM returns the underlying GORM instance.
func (m *MockDBService) GORM() *gorm.DB {
return m.DBInstance
}
// DB returns the GORM instance bound to context.
func (m *MockDBService) DB(ctx context.Context) *gorm.DB {
return m.DBInstance.WithContext(ctx)
}
// Named returns the named GORM instance.
func (m *MockDBService) Named(_ string) *gorm.DB {
return m.DBInstance
}
// MockCacheService is an in-memory mock implementation of contracts.CacheService for unit testing.
type MockCacheService struct {
mu sync.RWMutex
data map[string][]byte
}
// NewMockCacheService creates a new MockCacheService.
func NewMockCacheService() *MockCacheService {
return &MockCacheService{
data: make(map[string][]byte),
}
}
// Get retrieves a cached value.
func (m *MockCacheService) Get(_ context.Context, key string, val any) error {
m.mu.RLock()
defer m.mu.RUnlock()
b, ok := m.data[key]
if !ok {
return contracts.ErrCacheMiss
}
return json.Unmarshal(b, val)
}
// Set stores a key-value pair in cache.
func (m *MockCacheService) Set(_ context.Context, key string, val any, _ time.Duration) error {
m.mu.Lock()
defer m.mu.Unlock()
b, err := json.Marshal(val)
if err != nil {
return err
}
m.data[key] = b
return nil
}
// Delete removes a key from cache.
func (m *MockCacheService) Delete(_ context.Context, key string) error {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.data, key)
return nil
}
// GetOrSet retrieves or populates a cache entry.
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
}
val, err := loader()
if err != nil {
return err
}
return m.Set(ctx, key, val, ttl)
}
// Invalidate invalidates a cache tag or prefix.
func (m *MockCacheService) Invalidate(_ context.Context, _ string) error {
return nil
}
// MockStorageService is an in-memory mock implementation of contracts.StorageService for unit testing.
type MockStorageService struct {
mu sync.RWMutex
objects map[string][]byte
}
// NewMockStorageService creates a new MockStorageService.
func NewMockStorageService() *MockStorageService {
return &MockStorageService{
objects: make(map[string][]byte),
}
}
// Put uploads an object into mock storage.
func (m *MockStorageService) Put(_ context.Context, key string, body io.Reader, size int64, contentType string) (contracts.StoragePutResult, error) {
m.mu.Lock()
defer m.mu.Unlock()
data, err := io.ReadAll(body)
if err != nil {
return contracts.StoragePutResult{}, err
}
m.objects[key] = data
if strings.HasPrefix(key, "uploads/") {
_ = os.MkdirAll(filepath.Dir(key), 0755)
_ = os.WriteFile(key, data, 0644)
}
return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil
}
// Get retrieves an object from mock storage.
func (m *MockStorageService) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
m.mu.RLock()
defer m.mu.RUnlock()
data, ok := m.objects[key]
if ok {
return &contracts.StorageObject{
Key: key,
Body: io.NopCloser(bytes.NewReader(data)),
ContentLength: int64(len(data)),
ContentType: "application/octet-stream",
}, nil
}
if f, err := os.Open(key); err == nil {
info, _ := f.Stat()
return &contracts.StorageObject{
Key: key,
Body: f,
ContentLength: info.Size(),
ContentType: "application/octet-stream",
}, nil
}
return nil, gorm.ErrRecordNotFound
}
// Delete removes an object from mock storage.
func (m *MockStorageService) Delete(_ context.Context, key string) error {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.objects, key)
return nil
}
// Ingest handles programmatic file ingestion for mock storage.
func (m *MockStorageService) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
return &contracts.IngestResult{ID: 1, Key: "test.png", Created: true, Stored: true}, nil
}
// MockAuthService is a mock implementation of contracts.AuthService for unit testing.
type MockAuthService struct {
DB *gorm.DB
}
// RequireAuthMiddleware returns a dummy auth middleware.
func (a *MockAuthService) RequireAuthMiddleware() any {
return func(c *gin.Context) { c.Next() }
}
// RequireAdminMiddleware returns a dummy admin middleware.
func (a *MockAuthService) RequireAdminMiddleware() any {
return func(c *gin.Context) { c.Next() }
}
// DisallowTokenAuthMiddleware returns a dummy disallow token middleware.
func (a *MockAuthService) DisallowTokenAuthMiddleware() any {
return func(c *gin.Context) { c.Next() }
}
// GetCurrentUser returns the user associated with the request context.
func (a *MockAuthService) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if c, ok := ctx.(*gin.Context); ok {
authHeader := c.GetHeader("Authorization")
if strings.HasPrefix(authHeader, "Bearer ") {
tokenStr := strings.TrimPrefix(authHeader, "Bearer ")
tokenHash := fmt.Sprintf("%x", sha256.Sum256([]byte(tokenStr)))
var tokenRecord struct {
UserID uint64
}
if err := a.DB.Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err == nil {
return &contracts.UserDTO{ID: tokenRecord.UserID, IsActive: true}, nil
}
}
}
return nil, errors.New("unauthorized")
}
// GetCurrentUserID returns the current user ID.
func (a *MockAuthService) GetCurrentUserID(ctx context.Context) (uint64, error) {
u, err := a.GetCurrentUser(ctx)
if err != nil {
return 0, err
}
return u.ID, nil
}
// VerifyToken verifies an access token.
func (a *MockAuthService) VerifyToken(_ context.Context, token string) (*contracts.UserDTO, error) {
tokenHash := fmt.Sprintf("%x", sha256.Sum256([]byte(token)))
var tokenRecord struct {
UserID uint64
}
if err := a.DB.Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err == nil {
return &contracts.UserDTO{ID: tokenRecord.UserID, IsActive: true}, nil
}
return nil, errors.New("unauthorized")
}
// Authenticate verifies credentials.
func (a *MockAuthService) Authenticate(_ context.Context, _ string, _ string) (*contracts.UserDTO, error) {
return nil, nil
}
// CreateSession creates a login session.
func (a *MockAuthService) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
return "test-session", nil
}
// RevokeToken revokes an access token.
func (a *MockAuthService) RevokeToken(_ context.Context, _ string) error {
return nil
}
// RevokeUserSessions revokes all sessions for a user.
func (a *MockAuthService) RevokeUserSessions(_ context.Context, _ uint64) error {
return nil
}
// InvalidateCachedUser invalidates cached user profile.
func (a *MockAuthService) InvalidateCachedUser(_ context.Context, _ uint64) {}
// InvalidateCachedToken invalidates cached access token.
func (a *MockAuthService) InvalidateCachedToken(_ context.Context, _ string) {}
// ListAuthSources lists configured authentication sources.
func (a *MockAuthService) ListAuthSources(_ context.Context) ([]contracts.AuthSourceViewDTO, error) {
return nil, nil
}
// CreateAuthSource creates an authentication source.
func (a *MockAuthService) CreateAuthSource(_ context.Context, _ contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
return nil, nil
}
// UpdateAuthSource updates an authentication source.
func (a *MockAuthService) UpdateAuthSource(_ context.Context, _ uint64, _ contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
return nil, nil
}
// DeleteAuthSource deletes an authentication source.
func (a *MockAuthService) DeleteAuthSource(_ context.Context, _ uint64) error {
return nil
}
// ToggleAuthSource toggles an authentication source active state.
func (a *MockAuthService) ToggleAuthSource(_ context.Context, _ uint64) (*contracts.AuthSourceDTO, error) {
return nil, nil
}
// SetupTestEnv initializes test helper environment and binds DB, Cache, Storage, Auth mocks to shared services.
func SetupTestEnv(t *testing.T) (*gorm.DB, func()) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
dbSvc := &MockDBService{DBInstance: dbConn}
cacheSvc := NewMockCacheService()
storageSvc := NewMockStorageService()
authSvc := &MockAuthService{DB: dbConn}
SetDBService(dbSvc)
SetCacheService(cacheSvc)
SetStorageService(storageSvc)
SetAuthService(authSvc)
return dbConn, func() {
ResetServices()
cleanup()
}
}
@@ -7,11 +7,12 @@ import (
"context"
"time"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/upload/models"
database "Wavelet/plugins/infra/database"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"Wavelet/pkg/logger"
"Wavelet/plugins/domain/upload/models"
"Wavelet/plugins/domain/upload/shared"
)
// ApplyUploadStatsAdd increments incremental stats for a newly active upload record.
@@ -26,7 +27,7 @@ func ApplyUploadStatsRemove(ctx context.Context, upload *models.Upload) error {
// RebuildUploadStats rebuilds all incremental stats from current upload records.
func RebuildUploadStats(ctx context.Context) error {
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
return shared.GetDB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("1 = 1").Delete(&models.UploadStat{}).Error; err != nil {
return err
}
@@ -49,7 +50,7 @@ func applyUploadStatsDelta(ctx context.Context, upload *models.Upload, sign int6
if upload == nil || !isActiveUploadStatus(upload.Status) {
return nil
}
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
return shared.GetDB(ctx).Transaction(func(tx *gorm.DB) error {
return ApplyUploadStatsDeltaTx(tx, upload, sign)
})
}
@@ -8,15 +8,36 @@ import (
"testing"
"time"
"gorm.io/gorm"
"Wavelet/pkg/testhelper"
"Wavelet/plugins/domain/upload/models"
database "Wavelet/plugins/infra/database"
"gorm.io/gorm"
"Wavelet/plugins/domain/upload/shared"
)
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
}
func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
shared.SetDBService(&mockDBService{db: dbConn})
defer func() {
shared.SetDBService(nil)
cleanup()
}()
ctx := context.Background()
upload := &models.Upload{
@@ -29,7 +50,7 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
CreatedAt: time.Now(),
}
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := shared.GetDB(ctx).Transaction(func(tx *gorm.DB) error {
return ApplyUploadStatsDeltaTx(tx, upload, 1)
}); err != nil {
t.Fatalf("ApplyUploadStatsDeltaTx returned error: %v", err)
@@ -45,8 +66,12 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
}
func TestApplyUploadStatsAddAndRemove(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
shared.SetDBService(&mockDBService{db: dbConn})
defer func() {
shared.SetDBService(nil)
cleanup()
}()
ctx := context.Background()
upload := &models.Upload{
@@ -90,7 +115,7 @@ type uploadStatsSnapshot struct {
func loadUploadStats(ctx context.Context) (uploadStatsSnapshot, error) {
var rows []models.UploadStat
if err := database.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
if err := shared.GetDB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
return uploadStatsSnapshot{}, err
}
if len(rows) == 0 {
@@ -9,15 +9,14 @@ import (
"sync"
"time"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/upload/shared"
"Wavelet/plugins/drivers/driver_asynq_worker"
"Wavelet/plugins/infra/storage/objectstore"
)
// MigrationAccessState captures cached migration maintenance state.
type MigrationAccessState struct {
ReadOnly bool
Target objectstore.Config
Target contracts.StorageConfigDTO
HasTarget bool
TargetErr error
LoadErr error
@@ -65,14 +64,14 @@ func buildMigrationAccessState(ctx context.Context) MigrationAccessState {
if err != nil {
return MigrationAccessState{LoadErr: err, ReadOnly: true}
}
if !ok {
if !ok || execution == nil {
return MigrationAccessState{}
}
state := MigrationAccessState{
ReadOnly: execution.Status != driver_asynq_worker.TaskExecutionStatusSucceeded,
ReadOnly: execution.Status != "succeeded",
}
if execution.Status == driver_asynq_worker.TaskExecutionStatusSucceeded {
if execution.Status == "succeeded" {
return state
}
@@ -10,33 +10,47 @@ import (
"fmt"
"strings"
"Wavelet/plugins/drivers/driver_asynq_worker"
"Wavelet/plugins/infra/storage/objectstore"
"gorm.io/gorm"
"Wavelet/core/contracts"
"Wavelet/plugins/domain/upload/shared"
)
// StorageMigrationTask is the Asynq task name for storage migration.
// StorageMigrationTask is the task name for storage migration.
const StorageMigrationTask = "storage:migrate"
// LatestMigrationExecution returns the most recent storage migration task execution.
func LatestMigrationExecution(ctx context.Context) (*driver_asynq_worker.TaskExecution, bool, error) {
return driver_asynq_worker.GetLatestTaskExecutionByTaskType(ctx, StorageMigrationTask)
func LatestMigrationExecution(ctx context.Context) (*contracts.TaskExecutionDTO, bool, error) {
db := shared.GetDB(ctx)
if db == nil {
return nil, false, nil
}
var exec contracts.TaskExecutionDTO
err := db.Table("w_task_executions").Where("task_type = ?", StorageMigrationTask).Order("id DESC").First(&exec).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, nil
}
return nil, false, err
}
return &exec, true, nil
}
// ParseMigrationTargetConfig parses and validates a storage migration target payload.
func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (objectstore.Config, error) {
func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (contracts.StorageConfigDTO, error) {
if strings.TrimSpace(string(payload)) == "" {
return objectstore.Config{}, errors.New("storage migration target payload is required")
return contracts.StorageConfigDTO{}, errors.New("storage migration target payload is required")
}
var raw struct {
Target json.RawMessage `json:"target"`
}
if err := json.Unmarshal(payload, &raw); err != nil {
return objectstore.Config{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
return contracts.StorageConfigDTO{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
}
if len(raw.Target) == 0 {
return objectstore.Config{}, errors.New("storage migration target payload is required")
return contracts.StorageConfigDTO{}, errors.New("storage migration target payload is required")
}
var targetBytes []byte
@@ -47,34 +61,57 @@ func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (objectstor
targetBytes = raw.Target
}
var target objectstore.Config
var target contracts.StorageConfigDTO
if err := json.Unmarshal(targetBytes, &target); err != nil {
return objectstore.Config{}, fmt.Errorf("parse target storage config: %w", err)
return contracts.StorageConfigDTO{}, fmt.Errorf("parse target storage config: %w", err)
}
current, err := objectstore.LoadConfig(ctx)
if err != nil {
return objectstore.Config{}, fmt.Errorf("load active storage config: %w", err)
}
target = objectstore.MergeMaskedSecrets(target, current)
if err := objectstore.ValidateConfig(target); err != nil {
return objectstore.Config{}, fmt.Errorf("validate target storage config: %w", err)
}
return target, nil
}
// NormalizeMigrationPayload validates and normalizes a storage migration payload.
func NormalizeMigrationPayload(ctx context.Context, payload []byte) ([]byte, objectstore.Config, error) {
func NormalizeMigrationPayload(ctx context.Context, payload []byte) ([]byte, contracts.StorageConfigDTO, error) {
target, err := ParseMigrationTargetConfig(ctx, payload)
if err != nil {
return nil, objectstore.Config{}, err
return nil, contracts.StorageConfigDTO{}, err
}
type storageMigrationPayload struct {
Target objectstore.Config `json:"target"`
}
normalized, err := json.Marshal(storageMigrationPayload{Target: target})
raw, err := json.Marshal(struct {
Target contracts.StorageConfigDTO `json:"target"`
}{Target: target})
if err != nil {
return nil, objectstore.Config{}, fmt.Errorf("marshal storage migration payload: %w", err)
return nil, contracts.StorageConfigDTO{}, fmt.Errorf("serialize normalized payload: %w", err)
}
return normalized, target, nil
return raw, target, nil
}
// SaveActiveConfig persists the active storage configuration to w_system_configs.
func SaveActiveConfig(ctx context.Context, cfg contracts.StorageConfigDTO) error {
db := shared.GetDB(ctx)
if db == nil {
return errors.New("database not available")
}
data, err := json.Marshal(cfg)
if err != nil {
return err
}
return db.Table("w_system_configs").Where("key = ?", "storage_config").Update("value", string(data)).Error
}
// LoadStorageConfig loads the current storage configuration from w_system_configs.
func LoadStorageConfig(ctx context.Context) (contracts.StorageConfigDTO, error) {
db := shared.GetDB(ctx)
if db == nil {
return contracts.StorageConfigDTO{}, errors.New("database not available")
}
var row struct {
Value string
}
if err := db.Table("w_system_configs").Where("key = ?", "storage_config").First(&row).Error; err != nil {
return contracts.StorageConfigDTO{}, err
}
var cfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(row.Value), &cfg); err != nil {
return contracts.StorageConfigDTO{}, err
}
return cfg, nil
}

Some files were not shown because too many files have changed in this diff Show More