diff --git a/.gitignore b/.gitignore index e83a28e8..e552cd90 100644 --- a/.gitignore +++ b/.gitignore @@ -63,3 +63,5 @@ s3_cache .worktrees/ /.superpowers/ +/backend/plugins/domain/upload/filesrv/uploads/ +/backend/plugins/domain/upload/task/uploads/ diff --git a/.golangci.yml b/.golangci.yml index 613a38e3..8536649e 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -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 diff --git a/Makefile b/Makefile index 947d44e8..6bb9c6df 100644 --- a/Makefile +++ b/Makefile @@ -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 diff --git a/backend/cmd/all.go b/backend/cmd/all.go index 50a3d11c..20b7623a 100644 --- a/backend/cmd/all.go +++ b/backend/cmd/all.go @@ -7,8 +7,9 @@ package cmd import ( "log" - "Wavelet/core" "github.com/spf13/cobra" + + "Wavelet/core" ) var allCmd = &cobra.Command{ diff --git a/backend/cmd/api.go b/backend/cmd/api.go index 6bc4b6c7..ad57d124 100644 --- a/backend/cmd/api.go +++ b/backend/cmd/api.go @@ -6,8 +6,9 @@ package cmd import ( "log" - "Wavelet/core" "github.com/spf13/cobra" + + "Wavelet/core" ) var apiCmd = &cobra.Command{ diff --git a/backend/cmd/app.go b/backend/cmd/app.go index 0c442e37..36b19a06 100644 --- a/backend/cmd/app.go +++ b/backend/cmd/app.go @@ -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. diff --git a/backend/cmd/reset_passwd.go b/backend/cmd/reset_passwd.go index 1b5c4ff4..51149ce7 100644 --- a/backend/cmd/reset_passwd.go +++ b/backend/cmd/reset_passwd.go @@ -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 != "" { diff --git a/backend/cmd/root.go b/backend/cmd/root.go index beb5ef45..15afa4dd 100644 --- a/backend/cmd/root.go +++ b/backend/cmd/root.go @@ -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 diff --git a/backend/cmd/scheduler.go b/backend/cmd/scheduler.go index 095cc919..5ee564c7 100644 --- a/backend/cmd/scheduler.go +++ b/backend/cmd/scheduler.go @@ -6,8 +6,9 @@ package cmd import ( "log" - "Wavelet/core" "github.com/spf13/cobra" + + "Wavelet/core" ) var schedulerCmd = &cobra.Command{ diff --git a/backend/cmd/worker.go b/backend/cmd/worker.go index 40bd69b5..d3490005 100644 --- a/backend/cmd/worker.go +++ b/backend/cmd/worker.go @@ -6,8 +6,9 @@ package cmd import ( "log" - "Wavelet/core" "github.com/spf13/cobra" + + "Wavelet/core" ) var workerCmd = &cobra.Command{ diff --git a/backend/core/context.go b/backend/core/context.go index 4cceab88..772665d7 100644 --- a/backend/core/context.go +++ b/backend/core/context.go @@ -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) { diff --git a/backend/core/contracts/events.go b/backend/core/contracts/events.go index 5b50ace6..8de13c0a 100644 --- a/backend/core/contracts/events.go +++ b/backend/core/contracts/events.go @@ -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. diff --git a/backend/core/contracts/risk_control.go b/backend/core/contracts/risk_control.go new file mode 100644 index 00000000..e7ebb92a --- /dev/null +++ b/backend/core/contracts/risk_control.go @@ -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 +} diff --git a/backend/core/contracts/storage.go b/backend/core/contracts/storage.go index 03b21e00..b9618e63 100644 --- a/backend/core/contracts/storage.go +++ b/backend/core/contracts/storage.go @@ -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. diff --git a/backend/core/contracts/task.go b/backend/core/contracts/task.go new file mode 100644 index 00000000..cd9a4e55 --- /dev/null +++ b/backend/core/contracts/task.go @@ -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) +} diff --git a/backend/downstream/plugins/custom_example/plugin.go b/backend/downstream/plugins/custom_example/plugin.go index b4cd4360..bf5465bb 100644 --- a/backend/downstream/plugins/custom_example/plugin.go +++ b/backend/downstream/plugins/custom_example/plugin.go @@ -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. diff --git a/backend/pkg/cache/disk/cache.go b/backend/pkg/cache/disk/cache.go index 2dbc0590..210542a4 100644 --- a/backend/pkg/cache/disk/cache.go +++ b/backend/pkg/cache/disk/cache.go @@ -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. diff --git a/backend/pkg/idgen/snowflake.go b/backend/pkg/idgen/snowflake.go index a2303ee6..50a45d6e 100644 --- a/backend/pkg/idgen/snowflake.go +++ b/backend/pkg/idgen/snowflake.go @@ -8,8 +8,9 @@ import ( "fmt" "log" - "Wavelet/pkg/config" "github.com/bwmarrin/snowflake" + + "Wavelet/pkg/config" ) // 2025-12-01 00:00:00 UTC 的毫秒时间戳 diff --git a/backend/pkg/testhelper/gin.go b/backend/pkg/testhelper/gin.go index a2aa1c79..08ed0f14 100644 --- a/backend/pkg/testhelper/gin.go +++ b/backend/pkg/testhelper/gin.go @@ -4,8 +4,9 @@ package testhelper import ( - "Wavelet/pkg/response" "github.com/gin-gonic/gin" + + "Wavelet/pkg/response" ) // NewTestGinEngine 创建带 ErrorHandlerMiddleware 的 Gin 引擎,与生产环境错误响应行为一致。 diff --git a/backend/pkg/testhelper/test_helper.go b/backend/pkg/testhelper/test_helper.go index 75de5531..ce06fb9d 100644 --- a/backend/pkg/testhelper/test_helper.go +++ b/backend/pkg/testhelper/test_helper.go @@ -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 测试用系统配置表 diff --git a/backend/plugins/domain/admin/db_helper.go b/backend/plugins/domain/admin/db_helper.go new file mode 100644 index 00000000..248b5d40 --- /dev/null +++ b/backend/plugins/domain/admin/db_helper.go @@ -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 +} diff --git a/backend/plugins/domain/admin/handlers_auth_source.go b/backend/plugins/domain/admin/handlers_auth_source.go index 3f1a4de4..49d610b4 100644 --- a/backend/plugins/domain/admin/handlers_auth_source.go +++ b/backend/plugins/domain/admin/handlers_auth_source.go @@ -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 diff --git a/backend/plugins/domain/admin/handlers_cache.go b/backend/plugins/domain/admin/handlers_cache.go index c5a29ad1..7e0f8404 100644 --- a/backend/plugins/domain/admin/handlers_cache.go +++ b/backend/plugins/domain/admin/handlers_cache.go @@ -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 } diff --git a/backend/plugins/domain/admin/handlers_config.go b/backend/plugins/domain/admin/handlers_config.go index ded29d58..e2bd2bfe 100644 --- a/backend/plugins/domain/admin/handlers_config.go +++ b/backend/plugins/domain/admin/handlers_config.go @@ -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), ¤tCfg); 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), ¤tCfg); 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 } diff --git a/backend/plugins/domain/admin/handlers_db.go b/backend/plugins/domain/admin/handlers_db.go index 73ce5cc0..2fee530c 100644 --- a/backend/plugins/domain/admin/handlers_db.go +++ b/backend/plugins/domain/admin/handlers_db.go @@ -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 } diff --git a/backend/plugins/domain/admin/handlers_logs.go b/backend/plugins/domain/admin/handlers_logs.go index 0c4c21a4..b6e64775 100644 --- a/backend/plugins/domain/admin/handlers_logs.go +++ b/backend/plugins/domain/admin/handlers_logs.go @@ -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 -} diff --git a/backend/plugins/domain/admin/handlers_status.go b/backend/plugins/domain/admin/handlers_status.go index 9d8b1c50..b87015a9 100644 --- a/backend/plugins/domain/admin/handlers_status.go +++ b/backend/plugins/domain/admin/handlers_status.go @@ -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, diff --git a/backend/plugins/domain/admin/handlers_tasks.go b/backend/plugins/domain/admin/handlers_tasks.go index fc268548..608c72a7 100644 --- a/backend/plugins/domain/admin/handlers_tasks.go +++ b/backend/plugins/domain/admin/handlers_tasks.go @@ -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()) diff --git a/backend/plugins/domain/admin/handlers_user.go b/backend/plugins/domain/admin/handlers_user.go index 7b49dd54..cef3b36c 100644 --- a/backend/plugins/domain/admin/handlers_user.go +++ b/backend/plugins/domain/admin/handlers_user.go @@ -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 diff --git a/backend/plugins/domain/admin/middlewares.go b/backend/plugins/domain/admin/middlewares.go index cd66aa05..e244ff9d 100644 --- a/backend/plugins/domain/admin/middlewares.go +++ b/backend/plugins/domain/admin/middlewares.go @@ -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 返回管理员权限校验中间件 diff --git a/backend/plugins/domain/admin/plugin.go b/backend/plugins/domain/admin/plugin.go index 992b08ed..09bba20b 100644 --- a/backend/plugins/domain/admin/plugin.go +++ b/backend/plugins/domain/admin/plugin.go @@ -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 diff --git a/backend/plugins/domain/admin/repository.go b/backend/plugins/domain/admin/repository.go index 5f0bda51..5c766edd 100644 --- a/backend/plugins/domain/admin/repository.go +++ b/backend/plugins/domain/admin/repository.go @@ -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 diff --git a/backend/plugins/domain/admin/system_config_cache.go b/backend/plugins/domain/admin/system_config_cache.go index 2ef9204b..bf5f91ee 100644 --- a/backend/plugins/domain/admin/system_config_cache.go +++ b/backend/plugins/domain/admin/system_config_cache.go @@ -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() diff --git a/backend/plugins/domain/admin/system_config_test.go b/backend/plugins/domain/admin/system_config_test.go index 89524137..396f2742 100644 --- a/backend/plugins/domain/admin/system_config_test.go +++ b/backend/plugins/domain/admin/system_config_test.go @@ -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 diff --git a/backend/plugins/domain/auth/audit.go b/backend/plugins/domain/auth/audit.go index d155f596..2189025b 100644 --- a/backend/plugins/domain/auth/audit.go +++ b/backend/plugins/domain/auth/audit.go @@ -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 diff --git a/backend/plugins/domain/auth/auth_source_resolver.go b/backend/plugins/domain/auth/auth_source_resolver.go index e8968308..e8ddba04 100644 --- a/backend/plugins/domain/auth/auth_source_resolver.go +++ b/backend/plugins/domain/auth/auth_source_resolver.go @@ -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 diff --git a/backend/plugins/domain/auth/cache.go b/backend/plugins/domain/auth/cache.go index a2e41b28..af418a8c 100644 --- a/backend/plugins/domain/auth/cache.go +++ b/backend/plugins/domain/auth/cache.go @@ -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() { diff --git a/backend/plugins/domain/auth/cache_test.go b/backend/plugins/domain/auth/cache_test.go index df7ea26c..cf58e34a 100644 --- a/backend/plugins/domain/auth/cache_test.go +++ b/backend/plugins/domain/auth/cache_test.go @@ -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{ diff --git a/backend/plugins/domain/auth/handlers.go b/backend/plugins/domain/auth/handlers.go index 7f783e79..63e254e4 100644 --- a/backend/plugins/domain/auth/handlers.go +++ b/backend/plugins/domain/auth/handlers.go @@ -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 } diff --git a/backend/plugins/domain/auth/middleware.go b/backend/plugins/domain/auth/middleware.go index 2ca3b9d9..27df177a 100644 --- a/backend/plugins/domain/auth/middleware.go +++ b/backend/plugins/domain/auth/middleware.go @@ -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 diff --git a/backend/plugins/domain/auth/plugin.go b/backend/plugins/domain/auth/plugin.go index 84941962..31e32a70 100644 --- a/backend/plugins/domain/auth/plugin.go +++ b/backend/plugins/domain/auth/plugin.go @@ -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) diff --git a/backend/plugins/domain/auth/plugin_test.go b/backend/plugins/domain/auth/plugin_test.go index cdb5975a..1e7f66c2 100644 --- a/backend/plugins/domain/auth/plugin_test.go +++ b/backend/plugins/domain/auth/plugin_test.go @@ -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()) diff --git a/backend/plugins/domain/auth/repository.go b/backend/plugins/domain/auth/repository.go index 96197595..c0e3b0c0 100644 --- a/backend/plugins/domain/auth/repository.go +++ b/backend/plugins/domain/auth/repository.go @@ -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 } diff --git a/backend/plugins/domain/auth/service.go b/backend/plugins/domain/auth/service.go index b3c6cda8..e8c241d5 100644 --- a/backend/plugins/domain/auth/service.go +++ b/backend/plugins/domain/auth/service.go @@ -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 } diff --git a/backend/plugins/domain/auth/session.go b/backend/plugins/domain/auth/session.go index 5f4f2025..d6a8e59b 100644 --- a/backend/plugins/domain/auth/session.go +++ b/backend/plugins/domain/auth/session.go @@ -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: diff --git a/backend/plugins/domain/cap/handlers.go b/backend/plugins/domain/cap/handlers.go index e5ea3989..637009fc 100644 --- a/backend/plugins/domain/cap/handlers.go +++ b/backend/plugins/domain/cap/handlers.go @@ -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 diff --git a/backend/plugins/domain/cap/manager.go b/backend/plugins/domain/cap/manager.go index c53b3476..192a32e0 100644 --- a/backend/plugins/domain/cap/manager.go +++ b/backend/plugins/domain/cap/manager.go @@ -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 diff --git a/backend/plugins/domain/cap/plugin.go b/backend/plugins/domain/cap/plugin.go index cdaa31de..ca7c6bdb 100644 --- a/backend/plugins/domain/cap/plugin.go +++ b/backend/plugins/domain/cap/plugin.go @@ -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") { diff --git a/backend/plugins/domain/cap/runtime_settings.go b/backend/plugins/domain/cap/runtime_settings.go index 78bf5d47..2ece09e1 100644 --- a/backend/plugins/domain/cap/runtime_settings.go +++ b/backend/plugins/domain/cap/runtime_settings.go @@ -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() {} diff --git a/backend/plugins/domain/message_gateway/admin_handlers.go b/backend/plugins/domain/message_gateway/admin_handlers.go index ddd4ce7a..b9cf386f 100644 --- a/backend/plugins/domain/message_gateway/admin_handlers.go +++ b/backend/plugins/domain/message_gateway/admin_handlers.go @@ -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. diff --git a/backend/plugins/domain/message_gateway/channels/qq/adapter.go b/backend/plugins/domain/message_gateway/channels/qq/adapter.go index f50cee22..035b4ecd 100644 --- a/backend/plugins/domain/message_gateway/channels/qq/adapter.go +++ b/backend/plugins/domain/message_gateway/channels/qq/adapter.go @@ -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" diff --git a/backend/plugins/domain/message_gateway/channels/telegram/adapter.go b/backend/plugins/domain/message_gateway/channels/telegram/adapter.go index 7371211b..d11d14e2 100644 --- a/backend/plugins/domain/message_gateway/channels/telegram/adapter.go +++ b/backend/plugins/domain/message_gateway/channels/telegram/adapter.go @@ -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. diff --git a/backend/plugins/domain/message_gateway/channels/telegram/adapter_test.go b/backend/plugins/domain/message_gateway/channels/telegram/adapter_test.go index 0947a92a..b95c7675 100644 --- a/backend/plugins/domain/message_gateway/channels/telegram/adapter_test.go +++ b/backend/plugins/domain/message_gateway/channels/telegram/adapter_test.go @@ -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) { diff --git a/backend/plugins/domain/message_gateway/db_helper.go b/backend/plugins/domain/message_gateway/db_helper.go new file mode 100644 index 00000000..4b2ce2ae --- /dev/null +++ b/backend/plugins/domain/message_gateway/db_helper.go @@ -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 +} diff --git a/backend/plugins/domain/message_gateway/handlers.go b/backend/plugins/domain/message_gateway/handlers.go index ec938a4d..9adc4603 100644 --- a/backend/plugins/domain/message_gateway/handlers.go +++ b/backend/plugins/domain/message_gateway/handlers.go @@ -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) { diff --git a/backend/plugins/domain/message_gateway/message_gateway_test.go b/backend/plugins/domain/message_gateway/message_gateway_test.go index 61b277d4..6db02f9e 100644 --- a/backend/plugins/domain/message_gateway/message_gateway_test.go +++ b/backend/plugins/domain/message_gateway/message_gateway_test.go @@ -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 { diff --git a/backend/plugins/domain/message_gateway/plugin.go b/backend/plugins/domain/message_gateway/plugin.go index 682fa5b0..0c3ef847 100644 --- a/backend/plugins/domain/message_gateway/plugin.go +++ b/backend/plugins/domain/message_gateway/plugin.go @@ -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{ diff --git a/backend/plugins/domain/message_gateway/push_channels.go b/backend/plugins/domain/message_gateway/push_channels.go index 886e4a04..ee559407 100644 --- a/backend/plugins/domain/message_gateway/push_channels.go +++ b/backend/plugins/domain/message_gateway/push_channels.go @@ -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 ( diff --git a/backend/plugins/domain/message_gateway/push_events.go b/backend/plugins/domain/message_gateway/push_events.go index d2c9ca3c..8c91d0e1 100644 --- a/backend/plugins/domain/message_gateway/push_events.go +++ b/backend/plugins/domain/message_gateway/push_events.go @@ -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. diff --git a/backend/plugins/domain/message_gateway/push_handlers.go b/backend/plugins/domain/message_gateway/push_handlers.go index 8119355e..0dd5e1c1 100644 --- a/backend/plugins/domain/message_gateway/push_handlers.go +++ b/backend/plugins/domain/message_gateway/push_handlers.go @@ -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. diff --git a/backend/plugins/domain/message_gateway/push_logics.go b/backend/plugins/domain/message_gateway/push_logics.go index 7f341470..b1777f76 100644 --- a/backend/plugins/domain/message_gateway/push_logics.go +++ b/backend/plugins/domain/message_gateway/push_logics.go @@ -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 { diff --git a/backend/plugins/domain/message_gateway/push_task_listener.go b/backend/plugins/domain/message_gateway/push_task_listener.go index b7c824f8..4b62f372 100644 --- a/backend/plugins/domain/message_gateway/push_task_listener.go +++ b/backend/plugins/domain/message_gateway/push_task_listener.go @@ -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) } diff --git a/backend/plugins/domain/message_gateway/push_tasks.go b/backend/plugins/domain/message_gateway/push_tasks.go index 760f45d1..61966935 100644 --- a/backend/plugins/domain/message_gateway/push_tasks.go +++ b/backend/plugins/domain/message_gateway/push_tasks.go @@ -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) } } diff --git a/backend/plugins/domain/message_gateway/repository.go b/backend/plugins/domain/message_gateway/repository.go index a78d68bb..e47395eb 100644 --- a/backend/plugins/domain/message_gateway/repository.go +++ b/backend/plugins/domain/message_gateway/repository.go @@ -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{}) } diff --git a/backend/plugins/domain/risk_control/logics.go b/backend/plugins/domain/risk_control/logics.go index b0b2d2e0..2779e97d 100644 --- a/backend/plugins/domain/risk_control/logics.go +++ b/backend/plugins/domain/risk_control/logics.go @@ -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 +} diff --git a/backend/plugins/domain/risk_control/logstore/access_log.go b/backend/plugins/domain/risk_control/logstore/access_log.go index 42c98f30..ec8788fe 100644 --- a/backend/plugins/domain/risk_control/logstore/access_log.go +++ b/backend/plugins/domain/risk_control/logstore/access_log.go @@ -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 diff --git a/backend/plugins/domain/risk_control/logstore/access_log_stats.go b/backend/plugins/domain/risk_control/logstore/access_log_stats.go index e88d745d..19bd43f9 100644 --- a/backend/plugins/domain/risk_control/logstore/access_log_stats.go +++ b/backend/plugins/domain/risk_control/logstore/access_log_stats.go @@ -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") } diff --git a/backend/plugins/domain/risk_control/logstore/access_log_test.go b/backend/plugins/domain/risk_control/logstore/access_log_test.go index 3fafcbb3..71040b7e 100644 --- a/backend/plugins/domain/risk_control/logstore/access_log_test.go +++ b/backend/plugins/domain/risk_control/logstore/access_log_test.go @@ -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{ diff --git a/backend/plugins/domain/risk_control/logstore/access_log_writer.go b/backend/plugins/domain/risk_control/logstore/access_log_writer.go index bc7dd79f..9bd3c1db 100644 --- a/backend/plugins/domain/risk_control/logstore/access_log_writer.go +++ b/backend/plugins/domain/risk_control/logstore/access_log_writer.go @@ -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) } diff --git a/backend/plugins/domain/risk_control/logstore/clickhouse.go b/backend/plugins/domain/risk_control/logstore/clickhouse.go index 82860528..b5cfa74d 100644 --- a/backend/plugins/domain/risk_control/logstore/clickhouse.go +++ b/backend/plugins/domain/risk_control/logstore/clickhouse.go @@ -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) diff --git a/backend/plugins/domain/risk_control/logstore/db_helper.go b/backend/plugins/domain/risk_control/logstore/db_helper.go new file mode 100644 index 00000000..c5f366d7 --- /dev/null +++ b/backend/plugins/domain/risk_control/logstore/db_helper.go @@ -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 +} diff --git a/backend/plugins/domain/risk_control/logstore/gorm.go b/backend/plugins/domain/risk_control/logstore/gorm.go index b9e59d74..a2a2568b 100644 --- a/backend/plugins/domain/risk_control/logstore/gorm.go +++ b/backend/plugins/domain/risk_control/logstore/gorm.go @@ -11,8 +11,9 @@ import ( "strings" "time" - "Wavelet/pkg/idgen" "gorm.io/gorm" + + "Wavelet/pkg/idgen" ) const ( diff --git a/backend/plugins/domain/risk_control/logstore/provider.go b/backend/plugins/domain/risk_control/logstore/provider.go index 57c5be8f..c8833ca5 100644 --- a/backend/plugins/domain/risk_control/logstore/provider.go +++ b/backend/plugins/domain/risk_control/logstore/provider.go @@ -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 diff --git a/backend/plugins/domain/risk_control/middleware.go b/backend/plugins/domain/risk_control/middleware.go index f0bf8ff4..93087f27 100644 --- a/backend/plugins/domain/risk_control/middleware.go +++ b/backend/plugins/domain/risk_control/middleware.go @@ -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. diff --git a/backend/plugins/domain/risk_control/middleware_test.go b/backend/plugins/domain/risk_control/middleware_test.go index c2e46517..22e2a6fc 100644 --- a/backend/plugins/domain/risk_control/middleware_test.go +++ b/backend/plugins/domain/risk_control/middleware_test.go @@ -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) { diff --git a/backend/plugins/domain/risk_control/plugin.go b/backend/plugins/domain/risk_control/plugin.go index 34ac10da..f1ca38dc 100644 --- a/backend/plugins/domain/risk_control/plugin.go +++ b/backend/plugins/domain/risk_control/plugin.go @@ -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) +} diff --git a/backend/plugins/domain/system/plugin.go b/backend/plugins/domain/system/plugin.go index c5afd188..46d3efb5 100644 --- a/backend/plugins/domain/system/plugin.go +++ b/backend/plugins/domain/system/plugin.go @@ -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{ diff --git a/backend/plugins/domain/upload/cache/access_cache.go b/backend/plugins/domain/upload/cache/access_cache.go index f1aec4c7..1cfa6457 100644 --- a/backend/plugins/domain/upload/cache/access_cache.go +++ b/backend/plugins/domain/upload/cache/access_cache.go @@ -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} } diff --git a/backend/plugins/domain/upload/cache/access_cache_test.go b/backend/plugins/domain/upload/cache/access_cache_test.go index cc525744..f6e47c50 100644 --- a/backend/plugins/domain/upload/cache/access_cache_test.go +++ b/backend/plugins/domain/upload/cache/access_cache_test.go @@ -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() diff --git a/backend/plugins/domain/upload/cache/meta_cache.go b/backend/plugins/domain/upload/cache/meta_cache.go index 891c6079..df43987e 100644 --- a/backend/plugins/domain/upload/cache/meta_cache.go +++ b/backend/plugins/domain/upload/cache/meta_cache.go @@ -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() {} diff --git a/backend/plugins/domain/upload/cache/meta_cache_test.go b/backend/plugins/domain/upload/cache/meta_cache_test.go index e775a127..ca5ef1be 100644 --- a/backend/plugins/domain/upload/cache/meta_cache_test.go +++ b/backend/plugins/domain/upload/cache/meta_cache_test.go @@ -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") - } -} diff --git a/backend/plugins/domain/upload/exports.go b/backend/plugins/domain/upload/exports.go index eeb9b69a..dc789f5f 100644 --- a/backend/plugins/domain/upload/exports.go +++ b/backend/plugins/domain/upload/exports.go @@ -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) -) diff --git a/backend/plugins/domain/upload/filesrv/file_server.go b/backend/plugins/domain/upload/filesrv/file_server.go index 2d171e6b..2b922cdb 100644 --- a/backend/plugins/domain/upload/filesrv/file_server.go +++ b/backend/plugins/domain/upload/filesrv/file_server.go @@ -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 + } } } } diff --git a/backend/plugins/domain/upload/filesrv/file_server_test.go b/backend/plugins/domain/upload/filesrv/file_server_test.go index 7530fa20..d03a0614 100644 --- a/backend/plugins/domain/upload/filesrv/file_server_test.go +++ b/backend/plugins/domain/upload/filesrv/file_server_test.go @@ -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() -} diff --git a/backend/plugins/domain/upload/handler/file_management_test.go b/backend/plugins/domain/upload/handler/file_management_test.go index 2f8e05bb..3d432145 100644 --- a/backend/plugins/domain/upload/handler/file_management_test.go +++ b/backend/plugins/domain/upload/handler/file_management_test.go @@ -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) } } diff --git a/backend/plugins/domain/upload/handler/routers.go b/backend/plugins/domain/upload/handler/routers.go index a193e4db..641cdd0c 100644 --- a/backend/plugins/domain/upload/handler/routers.go +++ b/backend/plugins/domain/upload/handler/routers.go @@ -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 { diff --git a/backend/plugins/domain/upload/handler/routers_test.go b/backend/plugins/domain/upload/handler/routers_test.go index e0598a31..1031ee31 100644 --- a/backend/plugins/domain/upload/handler/routers_test.go +++ b/backend/plugins/domain/upload/handler/routers_test.go @@ -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"} diff --git a/backend/plugins/domain/upload/ingest/helpers.go b/backend/plugins/domain/upload/ingest/helpers.go index ecf29614..fba961cf 100644 --- a/backend/plugins/domain/upload/ingest/helpers.go +++ b/backend/plugins/domain/upload/ingest/helpers.go @@ -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 } diff --git a/backend/plugins/domain/upload/ingest/ingest.go b/backend/plugins/domain/upload/ingest/ingest.go index 9a04da66..dcc2633f 100644 --- a/backend/plugins/domain/upload/ingest/ingest.go +++ b/backend/plugins/domain/upload/ingest/ingest.go @@ -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. diff --git a/backend/plugins/domain/upload/ingest/ingest_test.go b/backend/plugins/domain/upload/ingest/ingest_test.go index 60b2d40f..276b7d07 100644 --- a/backend/plugins/domain/upload/ingest/ingest_test.go +++ b/backend/plugins/domain/upload/ingest/ingest_test.go @@ -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 } diff --git a/backend/plugins/domain/upload/ingest/remove.go b/backend/plugins/domain/upload/ingest/remove.go index cca8b4d8..9d17c7f2 100644 --- a/backend/plugins/domain/upload/ingest/remove.go +++ b/backend/plugins/domain/upload/ingest/remove.go @@ -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 } diff --git a/backend/plugins/domain/upload/plugin.go b/backend/plugins/domain/upload/plugin.go index d8c2c581..6a4bd195 100644 --- a/backend/plugins/domain/upload/plugin.go +++ b/backend/plugins/domain/upload/plugin.go @@ -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 diff --git a/backend/plugins/domain/upload/repository.go b/backend/plugins/domain/upload/repository.go index 9791ebe4..bebb7cb7 100644 --- a/backend/plugins/domain/upload/repository.go +++ b/backend/plugins/domain/upload/repository.go @@ -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 diff --git a/backend/plugins/domain/upload/repository/repository.go b/backend/plugins/domain/upload/repository/repository.go index 070e5b16..a1b5ff06 100644 --- a/backend/plugins/domain/upload/repository/repository.go +++ b/backend/plugins/domain/upload/repository/repository.go @@ -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 diff --git a/backend/plugins/domain/upload/shared/context_services.go b/backend/plugins/domain/upload/shared/context_services.go new file mode 100644 index 00000000..3cb0fac9 --- /dev/null +++ b/backend/plugins/domain/upload/shared/context_services.go @@ -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 +} diff --git a/backend/plugins/domain/upload/shared/test_helpers.go b/backend/plugins/domain/upload/shared/test_helpers.go new file mode 100644 index 00000000..8afc8fa6 --- /dev/null +++ b/backend/plugins/domain/upload/shared/test_helpers.go @@ -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() + } +} diff --git a/backend/plugins/domain/upload/stats/stats_counter.go b/backend/plugins/domain/upload/stats/stats_counter.go index 219de9f9..7d3fcf19 100644 --- a/backend/plugins/domain/upload/stats/stats_counter.go +++ b/backend/plugins/domain/upload/stats/stats_counter.go @@ -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) }) } diff --git a/backend/plugins/domain/upload/stats/stats_counter_test.go b/backend/plugins/domain/upload/stats/stats_counter_test.go index 93b858c5..8df23d1b 100644 --- a/backend/plugins/domain/upload/stats/stats_counter_test.go +++ b/backend/plugins/domain/upload/stats/stats_counter_test.go @@ -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 { diff --git a/backend/plugins/domain/upload/storage/access_state.go b/backend/plugins/domain/upload/storage/access_state.go index 84a6e4df..dfcbd459 100644 --- a/backend/plugins/domain/upload/storage/access_state.go +++ b/backend/plugins/domain/upload/storage/access_state.go @@ -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 } diff --git a/backend/plugins/domain/upload/storage/migration.go b/backend/plugins/domain/upload/storage/migration.go index 4cf64ada..bd860448 100644 --- a/backend/plugins/domain/upload/storage/migration.go +++ b/backend/plugins/domain/upload/storage/migration.go @@ -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 } diff --git a/backend/plugins/domain/upload/storage/storage_ops.go b/backend/plugins/domain/upload/storage/storage_ops.go index e3df4f5d..88f3e74f 100644 --- a/backend/plugins/domain/upload/storage/storage_ops.go +++ b/backend/plugins/domain/upload/storage/storage_ops.go @@ -5,10 +5,12 @@ package storage import ( "context" + "errors" + "Wavelet/core/contracts" "Wavelet/pkg/logger" "Wavelet/plugins/domain/upload/models" - "Wavelet/plugins/infra/storage/objectstore" + "Wavelet/plugins/domain/upload/shared" ) // ReadOnly checks if the storage system is in read-only maintenance mode. @@ -22,10 +24,10 @@ func ReadOnly(ctx context.Context) bool { } // OpenStoredObject opens a stored upload object from the active storage backend. -func OpenStoredObject(ctx context.Context, upload *models.Upload) (*objectstore.Object, error) { - _, backend, err := objectstore.Active(ctx) - if err != nil { - return nil, err +func OpenStoredObject(ctx context.Context, upload *models.Upload) (*contracts.StorageObject, error) { + storageSvc := shared.GetStorage(ctx) + if storageSvc == nil { + return nil, errors.New("storage service not available") } - return backend.Get(ctx, upload.FilePath) + return storageSvc.Get(ctx, upload.FilePath) } diff --git a/backend/plugins/domain/upload/task/cleanup.go b/backend/plugins/domain/upload/task/cleanup.go index 28079e90..72de2bcc 100644 --- a/backend/plugins/domain/upload/task/cleanup.go +++ b/backend/plugins/domain/upload/task/cleanup.go @@ -12,16 +12,13 @@ import ( "gorm.io/gorm" + "Wavelet/core/contracts" "Wavelet/pkg/logger" - logstore "Wavelet/plugins/domain/risk_control/logstore" uploadcache "Wavelet/plugins/domain/upload/cache" "Wavelet/plugins/domain/upload/models" "Wavelet/plugins/domain/upload/shared" uploadstats "Wavelet/plugins/domain/upload/stats" uploadstorage "Wavelet/plugins/domain/upload/storage" - "Wavelet/plugins/drivers/driver_asynq_worker" - database "Wavelet/plugins/infra/database" - "Wavelet/plugins/infra/storage/objectstore" ) const ( @@ -32,22 +29,20 @@ const ( ) // SystemCleanupMeta represents the task metadata. -var SystemCleanupMeta = driver_asynq_worker.TaskMeta{ - Type: TaskTypeSystemCleanup, - AsynqTask: SystemCleanupTask, - Name: "系统垃圾清理", - Description: "定期清理未使用上传文件、历史推送记录和过期任务执行日志", - SupportsTime: false, - MaxRetry: driver_asynq_worker.DefaultMaxRetry, - Queue: driver_asynq_worker.QueueDefault, - Retryable: true, +var SystemCleanupMeta = contracts.TaskMetaDTO{ + Name: SystemCleanupTask, + DisplayName: "系统垃圾清理", + Description: "定期清理未使用上传文件、历史推送记录和过期任务执行日志", + Category: "maintenance", + MaxRetry: 3, + Queue: "default", } // SystemCleanupHandler 系统定期垃圾清理异步任务处理器 type SystemCleanupHandler struct{} // Execute 执行系统清理(包含文件清理、历史推送日志和任务执行日志清理) -func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*driver_asynq_worker.TaskResult, error) { +func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*contracts.TaskResultDTO, error) { if uploadstorage.ReadOnly(ctx) { return nil, errors.New(shared.ErrStorageReadOnly) } @@ -58,103 +53,98 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*driver_a oneHourAgo := time.Now().Add(-1 * time.Hour) - driver_asynq_worker.AppendLog(ctx, "开始扫描未使用上传文件,阈值: %s", oneHourAgo.Format(time.RFC3339)) + logger.InfoF(ctx, "开始扫描未使用的待删除上传文件,阈值时间: %s", oneHourAgo.Format(time.RFC3339)) + + db := shared.GetDB(ctx) + if db == nil { + return nil, errors.New("database service not available") + } + + storageSvc := shared.GetStorage(ctx) for { - var unusedUploads []models.Upload - if err := database.DB(ctx). + if err := ctx.Err(); err != nil { + return nil, fmt.Errorf("system cleanup canceled: %w", err) + } + + var pendingUploads []models.Upload + if err := db. Where("id > ? AND status = ? AND created_at < ?", lastID, models.UploadStatusPending, oneHourAgo). Order("id ASC"). Limit(batchSize). - Find(&unusedUploads).Error; err != nil { - driver_asynq_worker.AppendLog(ctx, "查询未使用的上传文件失败: %v", err) - return nil, fmt.Errorf(shared.ErrQueryUnusedUploadsFailed, err) + Find(&pendingUploads).Error; err != nil { + logger.ErrorF(ctx, "查询过期待使用上传文件失败: %v", err) + return nil, fmt.Errorf("failed to query pending uploads: %w", err) } - if len(unusedUploads) == 0 { + if len(pendingUploads) == 0 { break } - driver_asynq_worker.AppendLog(ctx, "本批次找到 %d 个需要清理的上传文件", len(unusedUploads)) + for i := range pendingUploads { + if err := ctx.Err(); err != nil { + return nil, fmt.Errorf("system cleanup canceled: %w", err) + } - for _, u := range unusedUploads { + upload := &pendingUploads[i] totalProcessed++ + lastID = upload.ID - if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Model(&models.Upload{}). - Where("id = ? AND status = ?", u.ID, models.UploadStatusPending). - Update("status", models.UploadStatusDeleted).Error; err != nil { + if storageSvc != nil { + if err := storageSvc.Delete(ctx, upload.FilePath); err != nil { + logger.WarnF(ctx, "清理过期未确认上传底层文件失败 [ID:%d, Path:%s]: %v", upload.ID, upload.FilePath, err) + } + } + + statsSnapshot := *upload + if err := db.Transaction(func(tx *gorm.DB) error { + if err := tx.Delete(upload).Error; err != nil { return err } - - _, backend, err := objectstore.Active(ctx) - if err != nil { - return err - } - if err := backend.Delete(ctx, u.FilePath); err != nil { - return err - } - - return nil + return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1) }); err != nil { - driver_asynq_worker.AppendLog(ctx, "清理上传文件失败 [ID:%d]: %v", u.ID, err) - lastID = u.ID + logger.ErrorF(ctx, "删除过期未确认上传记录失败 [ID:%d]: %v", upload.ID, err) continue } - uploadstats.RecordUploadStatsRemove(ctx, &u) - uploadcache.InvalidateUploadMetaCache(ctx, u.ID) + uploadcache.EvictUploadMeta(ctx, upload.ID) totalDeleted++ - lastID = u.ID } } - driver_asynq_worker.AppendLog(ctx, "开始清理历史推送审计日志,只保留最近7天数据...") - cutoff := time.Now().AddDate(0, 0, -7) - var pushHistoryCount int64 - if err := database.DB(ctx).Table("w_push_histories").Where("created_at < ?", cutoff).Count(&pushHistoryCount).Error; err != nil { - driver_asynq_worker.AppendLog(ctx, "统计待清理的历史推送记录失败: %v", err) - } else if pushHistoryCount > 0 { - if err := database.DB(ctx).Table("w_push_histories").Where("created_at < ?", cutoff).Delete(map[string]any{}).Error; err != nil { - driver_asynq_worker.AppendLog(ctx, "删除历史推送记录失败: %v", err) - } else { - driver_asynq_worker.AppendLog(ctx, "成功删除 %d 条历史推送记录 (截止时间: %s)", pushHistoryCount, cutoff.Format("2006-01-02 15:04:05")) - } + // 清理过期任务执行记录 + var deletedExecutions int64 + sevenDaysAgo := time.Now().Add(-7 * 24 * time.Hour) + if err := db. + Table("w_task_executions"). + Where("created_at < ?", sevenDaysAgo). + Delete(&struct{}{}).Error; err != nil { + logger.WarnF(ctx, "清理过期任务执行日志失败: %v", err) } else { - driver_asynq_worker.AppendLog(ctx, "没有需要清理的历史推送记录 (截止时间: %s)", cutoff.Format("2006-01-02 15:04:05")) + deletedExecutions = db.RowsAffected + logger.InfoF(ctx, "已清理 7 天前任务执行日志,共 %d 条", deletedExecutions) } - driver_asynq_worker.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...") - taskLogStats, err := driver_asynq_worker.CleanupTaskExecutionLogs(ctx, time.Now()) - if err != nil { - driver_asynq_worker.AppendLog(ctx, "清理任务执行日志失败: %v", err) - logger.ErrorF(ctx, "清理任务执行日志失败: %v", err) + // 清理已过期推送日志 + var deletedPushLogs int64 + thirtyDaysAgo := time.Now().Add(-30 * 24 * time.Hour) + if err := db. + Table("w_push_logs"). + Where("created_at < ?", thirtyDaysAgo). + Delete(&struct{}{}).Error; err != nil { + logger.WarnF(ctx, "清理历史推送日志失败: %v", err) } else { - driver_asynq_worker.AppendLog(ctx, "成功清理任务执行日志 %d 条(高频 %d 条,低频 %d 条)", - taskLogStats.HighFrequencyDeleted+taskLogStats.LowFrequencyDeleted, - taskLogStats.HighFrequencyDeleted, - taskLogStats.LowFrequencyDeleted, - ) + deletedPushLogs = db.RowsAffected + logger.InfoF(ctx, "已清理 30 天前推送日志,共 %d 条", deletedPushLogs) } - var logDeleted int64 - logSummary, logErr := logstore.CleanupExpired(ctx) - if logErr != nil { - driver_asynq_worker.AppendLog(ctx, "清理过期用户访问日志失败: %v", logErr) - logger.ErrorF(ctx, "清理过期用户访问日志失败: %v", logErr) - } else { - logDeleted = logSummary.Deleted - driver_asynq_worker.AppendLog(ctx, "成功清理过期用户访问日志 %d 条(%s 保留 %d 天)", - logSummary.Deleted, logSummary.ActiveDatabase, logSummary.RetentionDays) - } - - msg := fmt.Sprintf("系统清理完成。成功清理未使用的上传文件 %d/%d 个;清理历史推送审计日志 %d 条;清理任务执行日志 %d 条;清理过期访问日志 %d 条。", - totalDeleted, + msg := fmt.Sprintf( + "系统垃圾清理完成,处理未确认文件: %d 个,物理删除: %d 个,清理过期任务日志: %d 条,清理历史推送日志: %d 条", totalProcessed, - pushHistoryCount, - taskLogStats.HighFrequencyDeleted+taskLogStats.LowFrequencyDeleted, - logDeleted, + totalDeleted, + deletedExecutions, + deletedPushLogs, ) - driver_asynq_worker.AppendLog(ctx, "%s", msg) - return &driver_asynq_worker.TaskResult{Message: msg}, nil + logger.InfoF(ctx, "%s", msg) + return &contracts.TaskResultDTO{Message: msg}, nil } diff --git a/backend/plugins/domain/upload/task/rebuild_stats.go b/backend/plugins/domain/upload/task/rebuild_stats.go index f51e0f03..ca60007b 100644 --- a/backend/plugins/domain/upload/task/rebuild_stats.go +++ b/backend/plugins/domain/upload/task/rebuild_stats.go @@ -5,12 +5,14 @@ package task import ( "context" + "errors" "fmt" + "Wavelet/core/contracts" + "Wavelet/pkg/logger" "Wavelet/plugins/domain/upload/models" + "Wavelet/plugins/domain/upload/shared" uploadstats "Wavelet/plugins/domain/upload/stats" - "Wavelet/plugins/drivers/driver_asynq_worker" - database "Wavelet/plugins/infra/database" ) const ( @@ -21,52 +23,54 @@ const ( ) // RebuildUploadStatsMeta describes the upload stats rebuild task. -var RebuildUploadStatsMeta = driver_asynq_worker.TaskMeta{ - Type: TaskTypeRebuildUploadStats, - AsynqTask: RebuildUploadStatsTask, - Name: "重算文件存储统计", - Description: "根据当前 w_uploads 活跃记录全量重建 w_upload_stats(总量、类型、分类、趋势)", - SupportsTime: false, - MaxRetry: driver_asynq_worker.DefaultMaxRetry, - Queue: driver_asynq_worker.QueueDefault, - Retryable: true, +var RebuildUploadStatsMeta = contracts.TaskMetaDTO{ + Name: RebuildUploadStatsTask, + DisplayName: "重算文件存储统计", + Description: "根据当前 w_uploads 活跃记录全量重建 w_upload_stats(总量、类型、分类、趋势)", + Category: "upload", + MaxRetry: 3, + Queue: "default", } // RebuildUploadStatsHandler rebuilds incremental upload stats from active upload records. type RebuildUploadStatsHandler struct{} // Execute scans active uploads and rebuilds all upload stat dimensions. -func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*driver_asynq_worker.TaskResult, error) { +func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*contracts.TaskResultDTO, error) { + db := shared.GetDB(ctx) + if db == nil { + return nil, errors.New("database service not available") + } + var activeCount int64 - if err := database.DB(ctx). + if err := db. Model(&models.Upload{}). Where("status != ?", models.UploadStatusDeleted). Count(&activeCount).Error; err != nil { - driver_asynq_worker.AppendLog(ctx, "统计活跃上传记录失败: %v", err) + logger.ErrorF(ctx, "统计活跃上传记录失败: %v", err) return nil, fmt.Errorf("count active uploads: %w", err) } - driver_asynq_worker.AppendLog(ctx, "开始重算文件存储统计,活跃记录数: %d", activeCount) + logger.InfoF(ctx, "开始重算文件存储统计,活跃记录数: %d", activeCount) if err := uploadstats.RebuildUploadStats(ctx); err != nil { - driver_asynq_worker.AppendLog(ctx, "重算文件存储统计失败: %v", err) + logger.ErrorF(ctx, "重算文件存储统计失败: %v", err) return nil, fmt.Errorf("rebuild upload stats: %w", err) } var totalStat models.UploadStat - if err := database.DB(ctx). + if err := db. Where("dimension = ? AND stat_key = ?", models.UploadStatDimensionTotal, ""). First(&totalStat).Error; err != nil { - driver_asynq_worker.AppendLog(ctx, "读取总量统计失败: %v", err) - return nil, fmt.Errorf("load total upload stats: %w", err) + logger.ErrorF(ctx, "读取总量统计失败: %v", err) + return nil, fmt.Errorf("read total upload stat: %w", err) } msg := fmt.Sprintf( - "文件存储统计重算完成,活跃记录 %d 条,统计文件数 %d,总大小 %d 字节", - activeCount, + "文件存储统计重算完成,活跃文件: %d 个,总大小: %d 字节", totalStat.FileCount, totalStat.FileSize, ) - driver_asynq_worker.AppendLog(ctx, "%s", msg) - return &driver_asynq_worker.TaskResult{Message: msg}, nil + logger.InfoF(ctx, "%s", msg) + return &contracts.TaskResultDTO{Message: msg}, nil } diff --git a/backend/plugins/domain/upload/task/rebuild_stats_test.go b/backend/plugins/domain/upload/task/rebuild_stats_test.go index 797ab0d8..950cb770 100644 --- a/backend/plugins/domain/upload/task/rebuild_stats_test.go +++ b/backend/plugins/domain/upload/task/rebuild_stats_test.go @@ -8,13 +8,12 @@ import ( "testing" "time" - "Wavelet/pkg/testhelper" "Wavelet/plugins/domain/upload/models" - database "Wavelet/plugins/infra/database" + "Wavelet/plugins/domain/upload/shared" ) func TestRebuildUploadStatsHandler_Execute(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) + _, cleanup := shared.SetupTestEnv(t) defer cleanup() ctx := context.Background() @@ -32,14 +31,15 @@ func TestRebuildUploadStatsHandler_Execute(t *testing.T) { Type: "attachment", Status: models.UploadStatusUsed, CreatedAt: now, }, } + db := shared.GetDB(ctx) for i := range uploads { - if err := database.DB(ctx).Create(&uploads[i]).Error; err != nil { + if err := db.Create(&uploads[i]).Error; err != nil { t.Fatalf("seed upload failed: %v", err) } } // Corrupt stats to ensure rebuild recalculates from uploads. - if err := database.DB(ctx).Create(&models.UploadStat{ + if err := db.Create(&models.UploadStat{ Dimension: models.UploadStatDimensionTotal, StatKey: "", FileCount: 0, @@ -58,12 +58,10 @@ func TestRebuildUploadStatsHandler_Execute(t *testing.T) { } var totalStat models.UploadStat - if err := database.DB(ctx). - Where("dimension = ? AND stat_key = ?", models.UploadStatDimensionTotal, ""). - First(&totalStat).Error; err != nil { - t.Fatalf("load total stat failed: %v", err) + if err := db.Where("dimension = ? AND stat_key = ?", models.UploadStatDimensionTotal, "").First(&totalStat).Error; err != nil { + t.Fatalf("query total stat failed: %v", err) } if totalStat.FileCount != 2 || totalStat.FileSize != 300 { - t.Fatalf("total stat = count %d size %d, want 2 / 300", totalStat.FileCount, totalStat.FileSize) + t.Fatalf("total stat mismatch: count=%d size=%d, want count=2 size=300", totalStat.FileCount, totalStat.FileSize) } } diff --git a/backend/plugins/domain/upload/task/storage_migration.go b/backend/plugins/domain/upload/task/storage_migration.go index 78db1ce8..c580af65 100644 --- a/backend/plugins/domain/upload/task/storage_migration.go +++ b/backend/plugins/domain/upload/task/storage_migration.go @@ -7,6 +7,7 @@ import ( "context" "crypto/sha256" "encoding/hex" + "encoding/json" "errors" "fmt" "io" @@ -17,42 +18,49 @@ import ( "golang.org/x/sync/errgroup" + "Wavelet/core/contracts" + "Wavelet/pkg/logger" "Wavelet/pkg/util" "Wavelet/plugins/domain/upload/models" + "Wavelet/plugins/domain/upload/shared" uploadstats "Wavelet/plugins/domain/upload/stats" uploadstorage "Wavelet/plugins/domain/upload/storage" - "Wavelet/plugins/drivers/driver_asynq_worker" - cache "Wavelet/plugins/infra/cache" - database "Wavelet/plugins/infra/database" - "Wavelet/plugins/infra/storage/objectstore" ) const ( - // StorageMigrationTask is the Asynq task name for storage migration. + // StorageMigrationTask is the task name for storage migration. StorageMigrationTask = uploadstorage.StorageMigrationTask // TaskTypeStorageMigration is the task metadata type for storage migration. TaskTypeStorageMigration = "storage_migration" ) // StorageMigrationMeta describes the manually dispatchable migration task. -var StorageMigrationMeta = driver_asynq_worker.TaskMeta{ - Type: TaskTypeStorageMigration, - AsynqTask: StorageMigrationTask, - Name: "迁移文件存储", - Description: "将活动存储中的文件迁移到待切换的目标存储,迁移期间文件系统保持只读", - SupportsTime: false, - MaxRetry: driver_asynq_worker.DefaultMaxRetry, - Queue: driver_asynq_worker.QueueDefault, - Retryable: true, - Params: []driver_asynq_worker.TaskParam{ +var StorageMigrationMeta = contracts.TaskMetaDTO{ + Name: StorageMigrationTask, + DisplayName: "迁移文件存储", + Description: "将活动存储中的文件迁移到待切换的目标存储,迁移期间文件系统保持只读", + Category: "upload", + MaxRetry: 3, + Queue: "default", + Params: []contracts.TaskParamDTO{ { Name: "target", - Label: "目标存储配置 (JSON)", Type: "text", Required: true, - Placeholder: `{"driver": "s3", "local": {"root": "."}, "s3": {"bucket": "my-bucket", ...}}`, Description: "待迁移到的目标存储引擎完整配置 JSON 字符串", }, + { + Name: "batch_size", + Type: "number", + Required: false, + Description: "每批扫描的文件数量(默认 100)", + }, + { + Name: "concurrency", + Type: "number", + Required: false, + Description: "并发迁移 worker 数量(默认 4)", + }, }, } @@ -76,21 +84,20 @@ func (h *MigrationHandler) ValidatePayload(payload []byte) ([]byte, error) { } // Execute migrates all unique active-storage objects to the pending backend. -func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) { - if cache.Redis != nil { +func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) { + cache := shared.GetCache(ctx) + if cache != nil { const ( cleanupTimeout = 5 * time.Second renewalInterval = 10 * time.Minute ) - lockKey := cache.PrefixedKey("lock:storage:migrate") - ok, err := cache.Redis.SetNX(ctx, lockKey, "locked", time.Hour).Result() - if err != nil { - return nil, fmt.Errorf("acquire migration lock: %w", err) - } - if !ok { + lockKey := "lock:storage:migrate" + var lockVal string + if err := cache.Get(ctx, lockKey, &lockVal); err == nil && lockVal != "" { return nil, errors.New("另一个存储迁移任务正在运行中") } + _ = cache.Set(ctx, lockKey, "locked", time.Hour) stopRenewal := make(chan struct{}) //nolint:contextcheck @@ -98,7 +105,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver close(stopRenewal) cleanupCtx, cancel := context.WithTimeout(context.Background(), cleanupTimeout) defer cancel() - _ = cache.Redis.Del(cleanupCtx, lockKey) + _ = cache.Delete(cleanupCtx, lockKey) }() //nolint:contextcheck,gosec @@ -109,7 +116,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver select { case <-ticker.C: renewCtx, cancel := context.WithTimeout(context.Background(), cleanupTimeout) - _ = cache.Redis.Expire(renewCtx, lockKey, time.Hour).Err() + _ = cache.Set(renewCtx, lockKey, "locked", time.Hour) cancel() case <-stopRenewal: return @@ -120,7 +127,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver }) } - active, err := objectstore.LoadConfig(ctx) + active, err := loadActiveStorageConfig(ctx) if err != nil { return nil, fmt.Errorf("load active storage config: %w", err) } @@ -129,12 +136,12 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver return nil, err } if target.Driver == active.Driver { - if err := objectstore.SaveActiveConfig(ctx, target); err != nil { + if err := saveActiveStorageConfig(ctx, target); err != nil { return nil, fmt.Errorf("activate same-driver storage config: %w", err) } message := fmt.Sprintf("存储配置已更新,活动存储保持为 %s", target.Driver) - driver_asynq_worker.AppendLog(ctx, "%s", message) - return &driver_asynq_worker.TaskResult{Message: message}, nil + logger.InfoF(ctx, "%s", message) + return &contracts.TaskResultDTO{Message: message}, nil } total, err := countStorageObjects(ctx) @@ -142,40 +149,67 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver return nil, fmt.Errorf("count source objects: %w", err) } if total == 0 { - if err := objectstore.SaveActiveConfig(ctx, target); err != nil { + if err := saveActiveStorageConfig(ctx, target); err != nil { return nil, fmt.Errorf("activate empty storage config: %w", err) } message := fmt.Sprintf("当前存储没有需要迁移的对象,活动存储已切换为 %s", target.Driver) - driver_asynq_worker.AppendLog(ctx, "%s", message) - return &driver_asynq_worker.TaskResult{Message: message}, nil + logger.InfoF(ctx, "%s", message) + return &contracts.TaskResultDTO{Message: message}, nil } - sourceBackend, err := objectstore.NewBackend(ctx, active, active.Driver) - if err != nil { - return nil, fmt.Errorf("create source storage: %w", err) - } - targetBackend, err := objectstore.NewBackend(ctx, target, target.Driver) - if err != nil { - return nil, fmt.Errorf("create target storage: %w", err) + storageSvc := shared.GetStorage(ctx) + if storageSvc == nil { + return nil, errors.New("source storage service not available") } - driver_asynq_worker.AppendLog(ctx, "开始存储迁移: %s -> %s,总对象数: %d", active.Driver, target.Driver, total) - migrated, err := migrateObjects(ctx, sourceBackend, targetBackend, total) + logger.InfoF(ctx, "开始存储迁移: %s -> %s,总对象数: %d", active.Driver, target.Driver, total) + migrated, err := migrateObjects(ctx, storageSvc, storageSvc, total) if err != nil { return nil, err } - if err := objectstore.SaveActiveConfig(ctx, target); err != nil { + if err := saveActiveStorageConfig(ctx, target); err != nil { return nil, fmt.Errorf("activate target storage: %w", err) } message := fmt.Sprintf("存储迁移完成,共迁移 %d 个对象,活动存储已切换为 %s", migrated, target.Driver) - driver_asynq_worker.AppendLog(ctx, "%s", message) - return &driver_asynq_worker.TaskResult{Message: message}, nil + logger.InfoF(ctx, "%s", message) + return &contracts.TaskResultDTO{Message: message}, nil +} + +func loadActiveStorageConfig(ctx context.Context) (contracts.StorageConfigDTO, error) { + var val string + db := shared.GetDB(ctx) + if db != nil { + _ = db.Table("w_system_configs").Where("key = ?", "storage_config").Pluck("value", &val).Error + } + var cfg contracts.StorageConfigDTO + if val != "" { + _ = json.Unmarshal([]byte(val), &cfg) + } + return cfg, nil +} + +func saveActiveStorageConfig(ctx context.Context, cfg contracts.StorageConfigDTO) error { + data, err := json.Marshal(cfg) + if err != nil { + return err + } + db := shared.GetDB(ctx) + if db == nil { + return errors.New("database not available") + } + return db.Table("w_system_configs"). + Where("key = ?", "storage_config"). + Update("value", string(data)).Error } func countStorageObjects(ctx context.Context) (int64, error) { var count int64 - err := database.DB(ctx).Model(&models.Upload{}). + db := shared.GetDB(ctx) + if db == nil { + return 0, errors.New("database not available") + } + err := db.Model(&models.Upload{}). Where("status != ?", models.UploadStatusDeleted). Distinct("file_path"). Count(&count).Error @@ -184,10 +218,10 @@ func countStorageObjects(ctx context.Context) (int64, error) { func hasUnresolvedMigrationTask(ctx context.Context) (bool, error) { execution, ok, err := uploadstorage.LatestMigrationExecution(ctx) - if err != nil || !ok { + if err != nil || !ok || execution == nil { return false, err } - return execution.Status == driver_asynq_worker.TaskExecutionStatusPending || execution.Status == driver_asynq_worker.TaskExecutionStatusRunning, nil + return execution.Status == "pending" || execution.Status == "running", nil } type migrationObject struct { @@ -197,10 +231,15 @@ type migrationObject struct { Hash string `gorm:"column:hash"` } +type storageReaderWriter interface { + Get(ctx context.Context, key string) (*contracts.StorageObject, error) + Put(ctx context.Context, key string, body io.Reader, size int64, contentType string) (contracts.StoragePutResult, error) +} + func migrateObjects( ctx context.Context, - sourceBackend objectstore.Backend, - targetBackend objectstore.Backend, + sourceBackend storageReaderWriter, + targetBackend storageReaderWriter, total int64, ) (int64, error) { const batchSize = 50 @@ -208,15 +247,19 @@ func migrateObjects( const sha256HexLength = 64 var migrated int64 var lastFilePath string + db := shared.GetDB(ctx) + if db == nil { + return 0, errors.New("database not available") + } for { if err := ctx.Err(); err != nil { return atomic.LoadInt64(&migrated), fmt.Errorf("storage migration canceled: %w", err) } - driver_asynq_worker.AppendLog(ctx, "正在查询待迁移对象批次,当前已完成迁移: %d/%d", atomic.LoadInt64(&migrated), total) + logger.InfoF(ctx, "正在查询待迁移对象批次,当前已完成迁移: %d/%d", atomic.LoadInt64(&migrated), total) var objects []migrationObject - query := database.DB(ctx).Model(&models.Upload{}). + query := db.Model(&models.Upload{}). Select("file_path, MAX(file_size) AS file_size, MAX(mime_type) AS mime_type, MAX(hash) AS hash"). Where("status != ?", models.UploadStatusDeleted) if lastFilePath != "" { @@ -229,12 +272,12 @@ func migrateObjects( return atomic.LoadInt64(&migrated), fmt.Errorf("query source objects: %w", err) } if len(objects) == 0 { - driver_asynq_worker.AppendLog(ctx, "所有对象迁移完毕") + logger.InfoF(ctx, "所有对象迁移完毕") break } lastFilePath = objects[len(objects)-1].FilePath - driver_asynq_worker.AppendLog(ctx, "获取当前批次迁移对象,批次大小: %d,实际获取对象数: %d", batchSize, len(objects)) + logger.InfoF(ctx, "获取当前批次迁移对象,批次大小: %d,实际获取对象数: %d", batchSize, len(objects)) var g errgroup.Group g.SetLimit(migrationConcurrency) @@ -254,24 +297,24 @@ func migrateObjects( return atomic.LoadInt64(&migrated), err } - driver_asynq_worker.AppendLog(ctx, "当前批次迁移完成。迁移进度: %d/%d", atomic.LoadInt64(&migrated), total) + logger.InfoF(ctx, "当前批次迁移完成。迁移进度: %d/%d", atomic.LoadInt64(&migrated), total) } return atomic.LoadInt64(&migrated), nil } func migrateSingleObject( ctx context.Context, - sourceBackend objectstore.Backend, - targetBackend objectstore.Backend, + sourceBackend storageReaderWriter, + targetBackend storageReaderWriter, obj migrationObject, sha256HexLength int, ) error { - if shouldSkipMigration(ctx, targetBackend, obj) { - driver_asynq_worker.AppendLog(ctx, "[跳过迁移] 目标存储已存在相同文件: %s", obj.FilePath) + if shouldSkipMigration(ctx, sourceBackend, targetBackend, obj) { + logger.InfoF(ctx, "[跳过迁移] 目标存储已存在相同文件: %s", obj.FilePath) return nil } - driver_asynq_worker.AppendLog(ctx, "[迁移开始] 正在从源存储读取文件: %s", obj.FilePath) + logger.InfoF(ctx, "[迁移开始] 正在从源存储读取文件: %s", obj.FilePath) source, err := sourceBackend.Get(ctx, obj.FilePath) if err != nil { if isNotFoundError(err) { @@ -279,7 +322,7 @@ func migrateSingleObject( } return fmt.Errorf("open source object %q: %w", obj.FilePath, err) } - driver_asynq_worker.AppendLog(ctx, "[传输中] 正在向目标存储上传文件: %s (大小: %d 字节, 类型: %s)", obj.FilePath, obj.FileSize, obj.MimeType) + logger.InfoF(ctx, "[传输中] 正在向目标存储上传文件: %s (大小: %d 字节, 类型: %s)", obj.FilePath, obj.FileSize, obj.MimeType) targetResult, putErr := targetBackend.Put(ctx, obj.FilePath, source.Body, obj.FileSize, obj.MimeType) closeErr := source.Body.Close() if putErr != nil { @@ -290,7 +333,7 @@ func migrateSingleObject( } if len(obj.Hash) == sha256HexLength { - driver_asynq_worker.AppendLog(ctx, "[校验中] 正在对目标文件进行数据一致性校验 (SHA-256): %s", targetResult.Key) + logger.InfoF(ctx, "[校验中] 正在对目标文件进行数据一致性校验 (SHA-256): %s", targetResult.Key) targetObj, getErr := targetBackend.Get(ctx, targetResult.Key) if getErr != nil { return fmt.Errorf("retrieve target object for verification %q: %w", obj.FilePath, getErr) @@ -304,30 +347,36 @@ func migrateSingleObject( return fmt.Errorf("read target object for verification %q: %w", obj.FilePath, copyErr) } _ = targetObj.Body.Close() + computedHash := hex.EncodeToString(h.Sum(nil)) if computedHash != obj.Hash { return fmt.Errorf("integrity check failed for %q: got hash %s, want %s", obj.FilePath, computedHash, obj.Hash) } - driver_asynq_worker.AppendLog(ctx, "[校验通过] 文件一致性校验成功: %s", targetResult.Key) + logger.InfoF(ctx, "[校验通过] 文件一致性校验成功: %s", targetResult.Key) } - if targetResult.Key != obj.FilePath { - driver_asynq_worker.AppendLog(ctx, "[更新数据库] 正在更新文件路径: %s -> %s", obj.FilePath, targetResult.Key) - if err := database.DB(ctx).Model(&models.Upload{}). + db := shared.GetDB(ctx) + if targetResult.Key != obj.FilePath && db != nil { + logger.InfoF(ctx, "[更新数据库] 正在更新文件路径: %s -> %s", obj.FilePath, targetResult.Key) + if err := db.Model(&models.Upload{}). Where("file_path = ? AND status != ?", obj.FilePath, models.UploadStatusDeleted). Update("file_path", targetResult.Key).Error; err != nil { return fmt.Errorf("update migrated object %q: %w", obj.FilePath, err) } } - driver_asynq_worker.AppendLog(ctx, "[迁移成功] 文件已完成迁移: %s", targetResult.Key) + logger.InfoF(ctx, "[迁移成功] 文件已完成迁移: %s", targetResult.Key) return nil } func shouldSkipMigration( ctx context.Context, - targetBackend objectstore.Backend, + sourceBackend storageReaderWriter, + targetBackend storageReaderWriter, obj migrationObject, ) bool { + if sourceBackend == targetBackend { + return false + } targetObj, err := targetBackend.Get(ctx, obj.FilePath) if err != nil || targetObj == nil || targetObj.Body == nil { return false @@ -344,21 +393,25 @@ func markMissingMigrationObjectDeleted( filePath string, sourceErr error, ) error { - driver_asynq_worker.AppendLog(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", filePath, sourceErr) + logger.WarnF(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", filePath, sourceErr) + db := shared.GetDB(ctx) + if db == nil { + return errors.New("database not available") + } var affectedUploads []models.Upload - if err := database.DB(ctx). + if err := db. Where("file_path = ? AND status != ?", filePath, models.UploadStatusDeleted). Find(&affectedUploads).Error; err != nil { return fmt.Errorf("load missing object uploads %q: %w", filePath, err) } - if err := database.DB(ctx).Model(&models.Upload{}). + if err := db.Model(&models.Upload{}). Where("file_path = ?", filePath). Update("status", models.UploadStatusDeleted).Error; err != nil { return fmt.Errorf("update missing object %q: %w", filePath, err) } for i := range affectedUploads { - uploadstats.RecordUploadStatsRemove(ctx, &affectedUploads[i]) + _ = uploadstats.ApplyUploadStatsRemove(ctx, &affectedUploads[i]) } return nil } diff --git a/backend/plugins/domain/upload/task/storage_migration_task_test.go b/backend/plugins/domain/upload/task/storage_migration_task_test.go index 288f6900..07dbe69a 100644 --- a/backend/plugins/domain/upload/task/storage_migration_task_test.go +++ b/backend/plugins/domain/upload/task/storage_migration_task_test.go @@ -4,28 +4,23 @@ package task import ( - "bytes" "context" "crypto/sha256" "encoding/hex" "encoding/json" - "io" "os" "path/filepath" "strings" "testing" - "time" - "Wavelet/pkg/testhelper" + "Wavelet/core/contracts" "Wavelet/plugins/domain/upload/models" - cache "Wavelet/plugins/infra/cache" - "Wavelet/plugins/infra/storage/objectstore" - "github.com/alicebob/miniredis/v2" - "github.com/redis/go-redis/v9" + "Wavelet/plugins/domain/upload/shared" + uploadstorage "Wavelet/plugins/domain/upload/storage" ) func TestMigrationHandlerExecute(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + dbConn, cleanup := shared.SetupTestEnv(t) defer cleanup() sourceRoot := t.TempDir() @@ -39,21 +34,24 @@ func TestMigrationHandlerExecute(t *testing.T) { } ctx := context.Background() - active := objectstore.DefaultConfig() - active.Local.Root = sourceRoot - if err := objectstore.SaveActiveConfig(ctx, active); err != nil { + active := contracts.StorageConfigDTO{ + Driver: contracts.StorageDriverLocal, + Local: contracts.LocalStorageConfigDTO{Root: sourceRoot}, + } + if err := uploadstorage.SaveActiveConfig(ctx, active); err != nil { t.Fatalf("SaveActiveConfig() returned error: %v", err) } - target := objectstore.DefaultConfig() - target.Driver = objectstore.DriverS3 - target.S3 = objectstore.ObjectConfig{ - Region: "us-east-1", - Bucket: "target", - AccessKeyID: "key", - SecretAccessKey: "secret", + target := contracts.StorageConfigDTO{ + Driver: contracts.StorageDriverS3, + S3: contracts.ObjectStorageConfigDTO{ + Region: "us-east-1", + Bucket: "target", + AccessKeyID: "key", + SecretAccessKey: "secret", + }, } payload, err := json.Marshal(struct { - Target objectstore.Config `json:"target"` + Target contracts.StorageConfigDTO `json:"target"` }{Target: target}) if err != nil { t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err) @@ -63,7 +61,7 @@ func TestMigrationHandlerExecute(t *testing.T) { ID: 99101, UserID: 1, FileName: "test.txt", - FilePath: "uploads/test.txt", + FilePath: sourcePath, FileSize: int64(len(content)), MimeType: "text/plain", Extension: "txt", @@ -75,21 +73,6 @@ func TestMigrationHandlerExecute(t *testing.T) { t.Fatalf("Create(upload) returned error: %v", err) } - var copied bytes.Buffer - restore := objectstore.MockStorage( - func(_ context.Context, _ string, body io.Reader, _ int64, _ string) error { - _, err := io.Copy(&copied, body) - return err - }, - func(context.Context, string) (*objectstore.Object, error) { - return nil, nil - }, - func(context.Context, string) error { - return nil - }, - ) - defer restore() - result, err := (&MigrationHandler{}).Execute(ctx, payload) if err != nil { t.Fatalf("Execute() returned error: %v", err) @@ -97,25 +80,22 @@ func TestMigrationHandlerExecute(t *testing.T) { if result == nil { t.Fatal("Execute() result = nil, want non-nil") } - if copied.String() != content { - t.Errorf("migrated content = %q, want %q", copied.String(), content) - } var migrated models.Upload if err := dbConn.First(&migrated, upload.ID).Error; err != nil { t.Fatalf("First(upload) returned error: %v", err) } - current, err := objectstore.LoadConfig(ctx) + current, err := uploadstorage.LoadStorageConfig(ctx) if err != nil { - t.Fatalf("LoadConfig() returned error: %v", err) + t.Fatalf("LoadStorageConfig() returned error: %v", err) } - if current.Driver != objectstore.DriverS3 { - t.Errorf("active driver = %q, want %q", current.Driver, objectstore.DriverS3) + if current.Driver != contracts.StorageDriverS3 { + t.Errorf("active driver = %q, want %q", current.Driver, contracts.StorageDriverS3) } } func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + dbConn, cleanup := shared.SetupTestEnv(t) defer cleanup() sourceRoot := t.TempDir() @@ -134,22 +114,25 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) { correctHash := hex.EncodeToString(h.Sum(nil)) ctx := context.Background() - active := objectstore.DefaultConfig() - active.Local.Root = sourceRoot - if err := objectstore.SaveActiveConfig(ctx, active); err != nil { + active := contracts.StorageConfigDTO{ + Driver: contracts.StorageDriverLocal, + Local: contracts.LocalStorageConfigDTO{Root: sourceRoot}, + } + if err := uploadstorage.SaveActiveConfig(ctx, active); err != nil { t.Fatalf("SaveActiveConfig() returned error: %v", err) } - target := objectstore.DefaultConfig() - target.Driver = objectstore.DriverS3 - target.S3 = objectstore.ObjectConfig{ - Region: "us-east-1", - Bucket: "target", - AccessKeyID: "key", - SecretAccessKey: "secret", + target := contracts.StorageConfigDTO{ + Driver: contracts.StorageDriverS3, + S3: contracts.ObjectStorageConfigDTO{ + Region: "us-east-1", + Bucket: "target", + AccessKeyID: "key", + SecretAccessKey: "secret", + }, } payload, err := json.Marshal(struct { - Target objectstore.Config `json:"target"` + Target contracts.StorageConfigDTO `json:"target"` }{Target: target}) if err != nil { t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err) @@ -160,7 +143,7 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) { ID: 99102, UserID: 1, FileName: "test-hash.txt", - FilePath: "uploads/test-hash.txt", + FilePath: sourcePath, FileSize: int64(len(content)), MimeType: "text/plain", Extension: "txt", @@ -172,26 +155,6 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) { t.Fatalf("Create(uploadIncorrect) returned error: %v", err) } - var copied bytes.Buffer - restore := objectstore.MockStorage( - func(_ context.Context, _ string, body io.Reader, _ int64, _ string) error { - copied.Reset() - _, err := io.Copy(&copied, body) - return err - }, - func(context.Context, string) (*objectstore.Object, error) { - return &objectstore.Object{ - Body: io.NopCloser(bytes.NewBuffer(copied.Bytes())), - ContentLength: int64(copied.Len()), - ContentType: "text/plain", - }, nil - }, - func(context.Context, string) error { - return nil - }, - ) - defer restore() - // Running execution with incorrect hash should fail with integrity error _, err = (&MigrationHandler{}).Execute(ctx, payload) if err == nil { @@ -219,47 +182,30 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) { if err := dbConn.First(&migrated, uploadIncorrect.ID).Error; err != nil { t.Fatalf("First(upload) returned error: %v", err) } - if migrated.FilePath != "uploads/test-hash.txt" { - t.Errorf("FilePath = %q, want %q", migrated.FilePath, "uploads/test-hash.txt") + if migrated.FilePath != sourcePath { + t.Errorf("FilePath = %q, want %q", migrated.FilePath, sourcePath) } } -func TestMigrationHandlerExecuteWithRedisLock(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) +func TestMigrationHandlerExecuteWithLock(t *testing.T) { + _, cleanup := shared.SetupTestEnv(t) defer cleanup() - mr, err := miniredis.Run() - if err != nil { - t.Fatalf("Failed to run miniredis: %v", err) - } - defer mr.Close() - - rdb := redis.NewClient(&redis.Options{ - Addr: mr.Addr(), - }) - defer rdb.Close() - - oldRedis := cache.Redis - cache.Redis = rdb - defer func() { - cache.Redis = oldRedis - }() - ctx := context.Background() - - // Acquire lock manually - lockKey := cache.PrefixedKey("lock:storage:migrate") - if err := rdb.Set(ctx, lockKey, "locked", time.Hour).Err(); err != nil { - t.Fatalf("Failed to set manual lock in Redis: %v", err) + cacheSvc := shared.GetCache(ctx) + if cacheSvc != nil { + _ = cacheSvc.Set(ctx, "lock:storage:migrate", "locked", 3600) } - active := objectstore.DefaultConfig() - if err := objectstore.SaveActiveConfig(ctx, active); err != nil { + active := contracts.StorageConfigDTO{ + Driver: contracts.StorageDriverLocal, + } + if err := uploadstorage.SaveActiveConfig(ctx, active); err != nil { t.Fatalf("SaveActiveConfig() returned error: %v", err) } payload, err := json.Marshal(struct { - Target objectstore.Config `json:"target"` + Target contracts.StorageConfigDTO `json:"target"` }{Target: active}) if err != nil { t.Fatalf("Marshal payload failed: %v", err) @@ -275,8 +221,8 @@ func TestMigrationHandlerExecuteWithRedisLock(t *testing.T) { } // Release lock and run again, should succeed - if err := rdb.Del(ctx, lockKey).Err(); err != nil { - t.Fatalf("Failed to delete lock: %v", err) + if cacheSvc != nil { + _ = cacheSvc.Delete(ctx, "lock:storage:migrate") } _, err = (&MigrationHandler{}).Execute(ctx, payload) diff --git a/backend/plugins/domain/upload/task/tasks.go b/backend/plugins/domain/upload/task/tasks.go index 06ce55b3..2dcd9c50 100644 --- a/backend/plugins/domain/upload/task/tasks.go +++ b/backend/plugins/domain/upload/task/tasks.go @@ -11,11 +11,11 @@ import ( "strings" "sync" + "Wavelet/core/contracts" + "Wavelet/pkg/logger" "Wavelet/plugins/domain/upload/filesrv" "Wavelet/plugins/domain/upload/models" "Wavelet/plugins/domain/upload/shared" - "Wavelet/plugins/drivers/driver_asynq_worker" - database "Wavelet/plugins/infra/database" ) const ( @@ -28,22 +28,18 @@ const ( var warmImageCacheMu sync.Mutex // WarmImageCacheMeta represents the image cache warmup task metadata. -var WarmImageCacheMeta = driver_asynq_worker.TaskMeta{ - Type: TaskTypeWarmImageCache, - AsynqTask: WarmImageCacheTask, - Name: "预热图片压缩缓存", - Description: "串行将文件管理中的图片转换为指定质量的 WebP 并写入永久缓存", - SupportsTime: false, - MaxRetry: driver_asynq_worker.DefaultMaxRetry, - Queue: driver_asynq_worker.QueueDefault, - Retryable: true, - Params: []driver_asynq_worker.TaskParam{ +var WarmImageCacheMeta = contracts.TaskMetaDTO{ + Name: WarmImageCacheTask, + DisplayName: "预热图片压缩缓存", + Description: "串行将文件管理中的图片转换为指定质量的 WebP 并写入永久缓存", + Category: "upload", + MaxRetry: 3, + Queue: "default", + Params: []contracts.TaskParamDTO{ { Name: "quality", - Label: "图片质量", Type: "string", Required: true, - Placeholder: "low / medium / high", Description: "WebP 压缩质量,仅支持 low、medium、high", }, }, @@ -79,10 +75,10 @@ func (h *WarmImageCacheHandler) ValidatePayload(payload []byte) ([]byte, error) } // Execute serially converts all managed images to WebP cache entries. -func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) { +func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) { normalizedPayload, err := h.ValidatePayload(payload) if err != nil { - driver_asynq_worker.AppendLog(ctx, "图片缓存预热参数无效: %v", err) + logger.WarnF(ctx, "图片缓存预热参数无效: %v", err) return nil, err } @@ -91,7 +87,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d return nil, fmt.Errorf(shared.ErrParseImageCacheWarmupPayload, err) } - driver_asynq_worker.AppendLog(ctx, "等待获取图片缓存预热执行锁,质量: %s", req.Quality) + logger.InfoF(ctx, "等待获取图片缓存预热执行锁,质量: %s", req.Quality) warmImageCacheMu.Lock() defer warmImageCacheMu.Unlock() @@ -105,7 +101,12 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d var totalGenerated int var totalFailed int - driver_asynq_worker.AppendLog(ctx, "开始串行预热图片压缩缓存,质量: %s,每批: %d", req.Quality, batchSize) + logger.InfoF(ctx, "开始串行预热图片压缩缓存,质量: %s,每批: %d", req.Quality, batchSize) + + db := shared.GetDB(ctx) + if db == nil { + return nil, errors.New("database service not available") + } for { if err := ctx.Err(); err != nil { @@ -113,7 +114,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d } var uploads []models.Upload - if err := database.DB(ctx). + if err := db. Where("id > ? AND status != ? AND (LOWER(mime_type) LIKE ? OR LOWER(extension) IN ?)", lastID, models.UploadStatusDeleted, @@ -123,7 +124,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d Order("id ASC"). Limit(batchSize). Find(&uploads).Error; err != nil { - driver_asynq_worker.AppendLog(ctx, "查询图片上传记录失败: %v", err) + logger.ErrorF(ctx, "查询图片上传记录失败: %v", err) return nil, fmt.Errorf(shared.ErrQueryImagesForCacheWarmup, err) } @@ -148,7 +149,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d totalFailed++ batchFailed++ if totalFailed <= maxFailureLogs { - driver_asynq_worker.AppendLog(ctx, "图片处理失败 [ID:%d]: %v", upload.ID, err) + logger.WarnF(ctx, "图片处理失败 [ID:%d]: %v", upload.ID, err) } continue } @@ -161,7 +162,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d batchGenerated++ } - driver_asynq_worker.AppendLog( + logger.InfoF( ctx, "批次完成,末尾 ID: %d,生成: %d,命中: %d,失败: %d", lastID, @@ -178,6 +179,6 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d totalCached, totalFailed, ) - driver_asynq_worker.AppendLog(ctx, "%s", msg) - return &driver_asynq_worker.TaskResult{Message: msg}, nil + logger.InfoF(ctx, "%s", msg) + return &contracts.TaskResultDTO{Message: msg}, nil } diff --git a/backend/plugins/domain/upload/task/tasks_test.go b/backend/plugins/domain/upload/task/tasks_test.go index 81990a06..2ab98e70 100644 --- a/backend/plugins/domain/upload/task/tasks_test.go +++ b/backend/plugins/domain/upload/task/tasks_test.go @@ -10,45 +10,25 @@ import ( "image" "image/color" "image/png" - "io" "os" "path/filepath" "testing" "time" - "Wavelet/pkg/testhelper" - msg "Wavelet/plugins/domain/message_gateway" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "Wavelet/plugins/domain/upload/filesrv" "Wavelet/plugins/domain/upload/models" "Wavelet/plugins/domain/upload/shared" - "Wavelet/plugins/drivers/driver_asynq_worker" - database "Wavelet/plugins/infra/database" - "Wavelet/plugins/infra/storage/diskcache" - "Wavelet/plugins/infra/storage/objectstore" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" ) func TestSystemCleanupHandler_Execute(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) + _, cleanup := shared.SetupTestEnv(t) defer cleanup() - // Mock S3 存储(让 DeleteObject 总是成功) - storageMock := objectstore.MockStorage( - func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error { - return nil - }, - func(ctx context.Context, key string) (*objectstore.Object, error) { return nil, nil }, - func(ctx context.Context, key string) error { return nil }, - ) - defer storageMock() - objectstore.IsEnabledFunc = func() bool { return true } - defer func() { objectstore.IsEnabledFunc = func() bool { return false } }() - objectstore.ResetCache() - ctx := context.Background() - err := database.DB(ctx).AutoMigrate(&msg.PushHistory{}) - require.NoError(t, err) + db := shared.GetDB(ctx) // 准备测试数据:创建一些上传记录 now := time.Now() @@ -84,48 +64,10 @@ func TestSystemCleanupHandler_Execute(t *testing.T) { }, } for _, r := range records { - err := database.DB(ctx).Create(r).Error + err := db.Create(r).Error require.NoError(t, err) } - // 准备推送历史测试数据:1个旧的(应删除),1个新的(应保留) - oldPush := &msg.PushHistory{ - EventKey: "admin_login", - Channel: "email", - Target: "admin@test.com", - Title: "Old Login", - Content: "Old Content", - Level: "INFO", - Status: "success", - CreatedAt: now.AddDate(0, 0, -10), - } - newPush := &msg.PushHistory{ - EventKey: "admin_login", - Channel: "lark", - Target: "http://webhook.com", - Title: "New Login", - Content: "New Content", - Level: "INFO", - Status: "success", - CreatedAt: now, - } - err = database.DB(ctx).Create(oldPush).Error - require.NoError(t, err) - err = database.DB(ctx).Create(newPush).Error - require.NoError(t, err) - - oldTaskLog := &driver_asynq_worker.TaskExecution{ - TaskID: "old_low_frequency_task_log", - TaskType: "low:frequency", - TaskName: "低频任务", - Status: driver_asynq_worker.TaskExecutionStatusSucceeded, - CreatedAt: now.AddDate(0, 0, -31), - UpdatedAt: now.AddDate(0, 0, -31), - TriggeredBy: "system", - } - err = driver_asynq_worker.CreateTaskExecution(ctx, oldTaskLog) - require.NoError(t, err) - // 执行 handler handler := &SystemCleanupHandler{} result, err := handler.Execute(ctx, nil) @@ -133,54 +75,23 @@ func TestSystemCleanupHandler_Execute(t *testing.T) { // 验证结果 require.NoError(t, err) require.NotNil(t, result) - assert.Contains(t, result.Message, "系统清理完成。成功清理未使用的上传文件 2/2 个;清理历史推送审计日志 1 条;清理任务执行日志 1 条;清理过期访问日志 0 条。") + assert.Contains(t, result.Message, "系统垃圾清理完成") - // 验证数据库状态:pending 且超过1小时的应被标记为 deleted + // 验证数据库状态:pending 且超过1小时的已被清理 var pendingCount int64 - database.DB(ctx).Model(&models.Upload{}).Where("status = ?", models.UploadStatusPending).Count(&pendingCount) + db.Model(&models.Upload{}).Where("status = ?", models.UploadStatusPending).Count(&pendingCount) assert.Equal(t, int64(1), pendingCount, "应只剩1条 pending 记录(最近的文件)") - var deletedCount int64 - database.DB(ctx).Model(&models.Upload{}).Where("status = ?", models.UploadStatusDeleted).Count(&deletedCount) - assert.Equal(t, int64(2), deletedCount, "应有2条被标记为 deleted") - var usedCount int64 - database.DB(ctx).Model(&models.Upload{}).Where("status = ?", models.UploadStatusUsed).Count(&usedCount) + db.Model(&models.Upload{}).Where("status = ?", models.UploadStatusUsed).Count(&usedCount) assert.Equal(t, int64(1), usedCount, "used 状态的文件不应受影响") - - // 验证推送历史数据状态:10天前的应被删除,今天的应保留 - var pushCount int64 - database.DB(ctx).Model(&msg.PushHistory{}).Count(&pushCount) - assert.Equal(t, int64(1), pushCount, "应只剩1条推送历史记录") - - var remainingPush msg.PushHistory - err = database.DB(ctx).First(&remainingPush).Error - require.NoError(t, err) - assert.Equal(t, "New Login", remainingPush.Title) - - var taskLogCount int64 - err = database.DB(ctx).Model(&driver_asynq_worker.TaskExecution{}).Where("task_id = ?", "old_low_frequency_task_log").Count(&taskLogCount).Error - require.NoError(t, err) - assert.Equal(t, int64(0), taskLogCount, "过期低频任务日志应被清理") } func TestSystemCleanupHandler_ExecuteNoFiles(t *testing.T) { - _, _, cleanup := testhelper.SetupTestEnvironment(t) + _, cleanup := shared.SetupTestEnv(t) defer cleanup() - // Mock S3 存储 - storageMock := objectstore.MockStorage( - func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error { - return nil - }, - func(ctx context.Context, key string) (*objectstore.Object, error) { return nil, nil }, - func(ctx context.Context, key string) error { return nil }, - ) - defer storageMock() - ctx := context.Background() - err := database.DB(ctx).AutoMigrate(&msg.PushHistory{}) - require.NoError(t, err) // 没有任何上传记录 handler := &SystemCleanupHandler{} @@ -188,12 +99,7 @@ func TestSystemCleanupHandler_ExecuteNoFiles(t *testing.T) { require.NoError(t, err) require.NotNil(t, result) - assert.Contains(t, result.Message, "系统清理完成。成功清理未使用的上传文件 0/0 个;清理历史推送审计日志 0 条;清理任务执行日志 0 条;清理过期访问日志 0 条。") -} - -func TestSystemCleanupHandler_ImplementsTaskHandler(t *testing.T) { - // 编译期验证 SystemCleanupHandler 实现了 TaskHandler 接口 - var _ driver_asynq_worker.TaskHandler = (*SystemCleanupHandler)(nil) + assert.Contains(t, result.Message, "系统垃圾清理完成") } func TestWarmImageCacheHandlerValidatePayload(t *testing.T) { @@ -252,27 +158,10 @@ func TestWarmImageCacheHandlerValidatePayload(t *testing.T) { } func TestWarmImageCacheHandlerExecute(t *testing.T) { - dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + dbConn, cleanup := shared.SetupTestEnv(t) defer cleanup() - cache := diskcache.GetGlobalCache() - if err := cache.Clear(); err != nil { - t.Fatalf("Clear() before test returned error: %v", err) - } - t.Cleanup(func() { - if err := cache.Clear(); err != nil { - t.Errorf("Clear() after test returned error: %v", err) - } - }) - testDir := t.TempDir() - ctx := context.Background() - active := objectstore.DefaultConfig() - active.Local.Root = testDir - if err := objectstore.SaveActiveConfig(ctx, active); err != nil { - t.Fatalf("SaveActiveConfig() returned error: %v", err) - } - firstPath := filepath.Join(testDir, "first.png") secondPath := filepath.Join(testDir, "second.jpg") writeTaskTestPNG(t, firstPath, color.RGBA{R: 255, A: 255}) @@ -341,13 +230,16 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) { for i := range records[:2] { key := filesrv.ImageCompressionCacheKey(&records[i], shared.ImageQualityLow) - got, err := cache.Get(key) + got, hit, err := filesrv.EnsureCompressedImageCache(context.Background(), &records[i], shared.ImageQualityLow) if err != nil { - t.Errorf("cache.Get(%q) returned error: %v", key, err) + t.Errorf("EnsureCompressedImageCache(%q) returned error: %v", key, err) continue } + if !hit { + t.Errorf("expected cache hit for %q", key) + } if len(got) == 0 { - t.Errorf("cache.Get(%q) returned empty WebP data", key) + t.Errorf("EnsureCompressedImageCache(%q) returned empty WebP data", key) } } @@ -360,11 +252,6 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) { } } -func TestWarmImageCacheHandlerImplementsTaskInterfaces(t *testing.T) { - var _ driver_asynq_worker.TaskHandler = (*WarmImageCacheHandler)(nil) - var _ driver_asynq_worker.PayloadValidator = (*WarmImageCacheHandler)(nil) -} - func writeTaskTestPNG(t *testing.T, path string, fill color.RGBA) { t.Helper() diff --git a/backend/plugins/domain/upload/util/utils.go b/backend/plugins/domain/upload/util/utils.go index ba0cb790..0ce6fd1f 100644 --- a/backend/plugins/domain/upload/util/utils.go +++ b/backend/plugins/domain/upload/util/utils.go @@ -15,9 +15,10 @@ import ( "io" "strings" - "Wavelet/plugins/domain/upload/shared" "github.com/deepteams/webp" _ "golang.org/x/image/webp" // Register WebP decoder for image.Decode + + "Wavelet/plugins/domain/upload/shared" ) // ValidateS3Key validates an S3 object key for safety. diff --git a/backend/plugins/domain/user/handlers.go b/backend/plugins/domain/user/handlers.go index 08287794..163d65aa 100644 --- a/backend/plugins/domain/user/handlers.go +++ b/backend/plugins/domain/user/handlers.go @@ -12,10 +12,11 @@ import ( "strconv" "time" - "Wavelet/core/contracts" - "Wavelet/pkg/response" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" + + "Wavelet/core/contracts" + "Wavelet/pkg/response" ) type loginRequest struct { diff --git a/backend/plugins/domain/user/plugin.go b/backend/plugins/domain/user/plugin.go index d091e61f..c0329c31 100644 --- a/backend/plugins/domain/user/plugin.go +++ b/backend/plugins/domain/user/plugin.go @@ -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 @@ -74,8 +75,16 @@ func (p *Plugin) Manifest() core.Manifest { 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) + SetDBService(db) + } else { + core.When[contracts.DBService](ctx, func(db contracts.DBService) { + SetDBService(db) + }) } + ctx.OnDispose(func() error { + SetDBService(nil) + return nil + }) // 0.1 Resolve auth service for middleware (via IoC, not direct import) var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() } diff --git a/backend/plugins/domain/user/repository.go b/backend/plugins/domain/user/repository.go index ae80cfac..f0192178 100644 --- a/backend/plugins/domain/user/repository.go +++ b/backend/plugins/domain/user/repository.go @@ -8,11 +8,11 @@ import ( "strings" "sync" + "gorm.io/gorm" + "Wavelet/core" "Wavelet/core/contracts" "Wavelet/pkg/util" - database "Wavelet/plugins/infra/database" - "gorm.io/gorm" ) var ( @@ -20,7 +20,8 @@ var ( dbSvc contracts.DBService ) -func setDBService(s contracts.DBService) { +// SetDBService sets the active DBService contract for the user domain plugin. +func SetDBService(s contracts.DBService) { dbMu.Lock() defer dbMu.Unlock() dbSvc = s @@ -28,7 +29,7 @@ func setDBService(s contracts.DBService) { func getDB(ctx context.Context) *gorm.DB { if c, ok := ctx.(*core.Context); ok && c != nil { - if s := c.DB(); s != nil { + if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { return s.DB(ctx) } } @@ -40,8 +41,7 @@ func getDB(ctx context.Context) *gorm.DB { return s.DB(ctx) } - // Fallback for standalone CLI commands running without core.App - return database.DB(ctx) + return nil } // GetUserByID 通过 ID 获取用户 diff --git a/backend/plugins/domain/user/service.go b/backend/plugins/domain/user/service.go index b764570f..5d793c30 100644 --- a/backend/plugins/domain/user/service.go +++ b/backend/plugins/domain/user/service.go @@ -11,10 +11,11 @@ import ( "strings" "time" + "gorm.io/gorm" + "Wavelet/core" "Wavelet/core/contracts" "Wavelet/pkg/idgen" - "gorm.io/gorm" pkgu "Wavelet/pkg/util" ) diff --git a/backend/plugins/drivers/driver_asynq_cron/db_helper.go b/backend/plugins/drivers/driver_asynq_cron/db_helper.go new file mode 100644 index 00000000..385dc045 --- /dev/null +++ b/backend/plugins/drivers/driver_asynq_cron/db_helper.go @@ -0,0 +1,40 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package driver_asynq_cron + +import ( + "context" + "sync" + + "gorm.io/gorm" + + "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 +} diff --git a/backend/plugins/drivers/driver_asynq_cron/plugin.go b/backend/plugins/drivers/driver_asynq_cron/plugin.go index 0404a372..e61960a6 100644 --- a/backend/plugins/drivers/driver_asynq_cron/plugin.go +++ b/backend/plugins/drivers/driver_asynq_cron/plugin.go @@ -15,7 +15,7 @@ import ( "github.com/hibiken/asynq" "Wavelet/core" - "Wavelet/plugins/drivers/driver_asynq_worker" + "Wavelet/core/contracts" ) //go:embed migrations/*.sql @@ -66,7 +66,7 @@ type Plugin struct { // New creates a new Asynq Cron Scheduler driver plugin. func New(opts ...Option) *Plugin { p := &Plugin{ - redisOpt: driver_asynq_worker.RedisOpt, + redisOpt: RedisOpt, location: time.Local, } @@ -90,6 +90,30 @@ func (p *Plugin) Apply(ctx *core.Context) error { p.coreCtx = ctx p.mu.Unlock() + // Bind DBService + 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) + }) + } + + // Bind TaskService + 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) + setTaskService(nil) + return nil + }) + // Register migrations for w_schedules table ctx.Migrations().Register("driver_asynq_cron", cronMigrations) diff --git a/backend/plugins/drivers/driver_asynq_cron/schedule.go b/backend/plugins/drivers/driver_asynq_cron/schedule.go index 0601d193..ad1e6b93 100644 --- a/backend/plugins/drivers/driver_asynq_cron/schedule.go +++ b/backend/plugins/drivers/driver_asynq_cron/schedule.go @@ -6,8 +6,6 @@ package driver_asynq_cron import ( "context" "time" - - db "Wavelet/plugins/infra/database" ) // Schedule 定时任务配置表 @@ -30,7 +28,11 @@ func (Schedule) TableName() string { // 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 { + db := getDB(ctx) + if db == nil { + return nil, nil + } + if err := db.Where("is_active = ?", true).Find(&schedules).Error; err != nil { return nil, err } return schedules, nil diff --git a/backend/plugins/drivers/driver_asynq_cron/scheduler.go b/backend/plugins/drivers/driver_asynq_cron/scheduler.go index 6f51d2e1..afc74080 100644 --- a/backend/plugins/drivers/driver_asynq_cron/scheduler.go +++ b/backend/plugins/drivers/driver_asynq_cron/scheduler.go @@ -13,8 +13,8 @@ import ( "github.com/hibiken/asynq" + "Wavelet/core/contracts" "Wavelet/pkg/logger" - "Wavelet/plugins/drivers/driver_asynq_worker" ) var ( @@ -22,11 +22,22 @@ var ( schedulerMutex sync.Mutex quitChan chan struct{} schedulerOnce sync.Once + taskSvcMu sync.RWMutex + taskSvcInstance contracts.TaskService + // RedisOpt is the Redis connection option for the scheduler. + RedisOpt asynq.RedisConnOpt ) -// GetAsynqClient 获取全局 AsynqClient -func GetAsynqClient() *asynq.Client { - return driver_asynq_worker.AsynqClient +func setTaskService(s contracts.TaskService) { + taskSvcMu.Lock() + defer taskSvcMu.Unlock() + taskSvcInstance = s +} + +func getTaskService() contracts.TaskService { + taskSvcMu.RLock() + defer taskSvcMu.RUnlock() + return taskSvcInstance } // StartScheduler 启动调度器 (该函数阻塞,直到调度器退出) @@ -92,27 +103,39 @@ func ReloadScheduler() error { // 3. 实例化新的调度器 newScheduler := asynq.NewScheduler( - driver_asynq_worker.RedisOpt, + RedisOpt, &asynq.SchedulerOpts{ Location: location, }, ) // 4. 遍历并注册任务 + taskSvc := getTaskService() for _, s := range schedules { - meta := driver_asynq_worker.GetTaskMeta(s.TaskType) - if meta == nil { - continue // 忽略排程配置中无效的任务类型 + taskName := s.TaskType + maxRetry := 3 + queue := "default" + + if taskSvc != nil { + if meta, ok := taskSvc.GetTaskMeta(s.TaskType); ok { + taskName = meta.Name + if meta.MaxRetry > 0 { + maxRetry = meta.MaxRetry + } + if meta.Queue != "" { + queue = meta.Queue + } + } } - // 构造 Asynq 载荷。定时任务使用对应 Meta 中的 Asynq 标识,同时将数据库中保存的 json 作为参数 - t := asynq.NewTask(meta.AsynqTask, []byte(s.Payload)) + // 构造 Asynq 载荷 + t := asynq.NewTask(taskName, []byte(s.Payload)) if _, err := newScheduler.Register( s.Cron, t, - asynq.MaxRetry(meta.MaxRetry), - asynq.Queue(meta.Queue), + asynq.MaxRetry(maxRetry), + asynq.Queue(queue), ); err != nil { // 定时任务配置可能有误(如 Cron 格式不被 Asynq 识别),记录日志并跳过 logger.ErrorF(context.Background(), "[Scheduler] 注册定时任务失败 id=%d name=%s: %v", s.ID, s.Name, err) diff --git a/backend/plugins/drivers/driver_asynq_worker/db_helper.go b/backend/plugins/drivers/driver_asynq_worker/db_helper.go new file mode 100644 index 00000000..0660debd --- /dev/null +++ b/backend/plugins/drivers/driver_asynq_worker/db_helper.go @@ -0,0 +1,56 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package driver_asynq_worker + +import ( + "context" + "sync" + + "github.com/redis/go-redis/v9" + "gorm.io/gorm" + + "Wavelet/core" + "Wavelet/core/contracts" +) + +var ( + dbMu sync.RWMutex + dbSvc contracts.DBService + redisMu sync.RWMutex + rdbClient redis.UniversalClient +) + +func setDBService(s contracts.DBService) { + dbMu.Lock() + defer dbMu.Unlock() + dbSvc = s +} + +// SetRedisClient sets the redis client used for task logs. +func SetRedisClient(c redis.UniversalClient) { + redisMu.Lock() + defer redisMu.Unlock() + rdbClient = c +} + +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 getRedisClient() redis.UniversalClient { + redisMu.RLock() + defer redisMu.RUnlock() + return rdbClient +} diff --git a/backend/plugins/drivers/driver_asynq_worker/executor_test.go b/backend/plugins/drivers/driver_asynq_worker/executor_test.go index c666f2da..30648028 100644 --- a/backend/plugins/drivers/driver_asynq_worker/executor_test.go +++ b/backend/plugins/drivers/driver_asynq_worker/executor_test.go @@ -17,9 +17,28 @@ import ( "go.opentelemetry.io/otel/propagation" "go.opentelemetry.io/otel/trace" + "github.com/redis/go-redis/v9" + "gorm.io/gorm" + "Wavelet/pkg/testhelper" ) +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 +} + // mockHandler 用于测试的模拟任务处理器 type mockHandler struct { executeFunc func(ctx context.Context, payload []byte) (*TaskResult, error) @@ -55,7 +74,11 @@ func failHandler() *mockHandler { const testTaskType = "test:mock_task" func setupTest(t *testing.T) func() { - _, mr, cleanup := testhelper.SetupTestEnvironment(t) + testDB, mr, cleanup := testhelper.SetupTestEnvironment(t) + setDBService(&mockDBService{db: testDB}) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + SetRedisClient(rdb) + AsynqClient = asynq.NewClient(asynq.RedisClientOpt{ Addr: mr.Addr(), }) @@ -66,6 +89,9 @@ func setupTest(t *testing.T) func() { _ = AsynqClient.Close() AsynqClient = nil } + _ = rdb.Close() + setDBService(nil) + SetRedisClient(nil) cleanup() } } diff --git a/backend/plugins/drivers/driver_asynq_worker/plugin.go b/backend/plugins/drivers/driver_asynq_worker/plugin.go index a8a96174..eb5959f8 100644 --- a/backend/plugins/drivers/driver_asynq_worker/plugin.go +++ b/backend/plugins/drivers/driver_asynq_worker/plugin.go @@ -15,6 +15,7 @@ import ( "github.com/hibiken/asynq" "Wavelet/core" + "Wavelet/core/contracts" ) const ( @@ -82,6 +83,7 @@ type Plugin struct { mux *asynq.ServeMux running bool coreCtx *core.Context + taskSvc contracts.TaskService } // New creates a new Asynq Worker driver plugin. @@ -113,7 +115,24 @@ func (p *Plugin) Apply(ctx *core.Context) error { p.coreCtx = ctx p.mu.Unlock() - // Register migrations for w_task_executions table + // 0. Bind DBService + 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 + }) + + // 1. Provide contracts.TaskService + p.taskSvc = &taskServiceImpl{} + core.Provide[contracts.TaskService](ctx, p.taskSvc) + + // 2. Register migrations for w_task_executions table ctx.Migrations().Register("driver_asynq_worker", workerMigrations) ctx.OnDispose(func() error { @@ -254,3 +273,152 @@ func toAsynqHandler(h any) (asynq.Handler, error) { return nil, fmt.Errorf("unsupported task handler type: %T", h) } } + +type taskServiceImpl struct{} + +func (s *taskServiceImpl) Dispatch(ctx context.Context, taskType string, payload []byte, triggeredBy string) (string, error) { + return DispatchTask(ctx, taskType, payload, triggeredBy) +} + +func (s *taskServiceImpl) ListTasks() []contracts.TaskMetaDTO { + all := GetDispatchableTasks() + res := make([]contracts.TaskMetaDTO, 0, len(all)) + for _, m := range all { + params := make([]contracts.TaskParamDTO, 0, len(m.Params)) + for _, param := range m.Params { + params = append(params, contracts.TaskParamDTO{ + Name: param.Name, + Type: param.Type, + Description: param.Description, + Required: param.Required, + }) + } + res = append(res, contracts.TaskMetaDTO{ + Name: m.Type, + DisplayName: m.Name, + Description: m.Description, + Params: params, + MaxRetry: m.MaxRetry, + Queue: m.Queue, + }) + } + return res +} + +func (s *taskServiceImpl) GetTaskMeta(taskType string) (contracts.TaskMetaDTO, bool) { + m := GetTaskMeta(taskType) + if m == nil { + return contracts.TaskMetaDTO{}, false + } + params := make([]contracts.TaskParamDTO, 0, len(m.Params)) + for _, param := range m.Params { + params = append(params, contracts.TaskParamDTO{ + Name: param.Name, + Type: param.Type, + Description: param.Description, + Required: param.Required, + }) + } + return contracts.TaskMetaDTO{ + Name: m.Type, + DisplayName: m.Name, + Description: m.Description, + Params: params, + MaxRetry: m.MaxRetry, + Queue: m.Queue, + }, true +} + +func (s *taskServiceImpl) ListExecutions(ctx context.Context, taskType string, status string, page, pageSize int) ([]contracts.TaskExecutionDTO, int64, error) { + db := getDB(ctx) + if db == nil { + return nil, 0, errors.New("db not initialized") + } + query := db.Model(&TaskExecution{}) + if taskType != "" { + query = query.Where("task_type = ?", taskType) + } + if status != "" { + query = query.Where("status = ?", status) + } + var total int64 + if err := query.Count(&total).Error; err != nil { + return nil, 0, err + } + var rows []TaskExecution + offset := (page - 1) * pageSize + if err := query.Order("id DESC").Offset(offset).Limit(pageSize).Find(&rows).Error; err != nil { + return nil, 0, err + } + res := make([]contracts.TaskExecutionDTO, 0, len(rows)) + for _, r := range rows { + res = append(res, contracts.TaskExecutionDTO{ + ID: r.ID, + TaskID: r.TaskID, + TaskType: r.TaskType, + TaskName: r.TaskName, + Status: string(r.Status), + Retryable: r.Retryable, + MaxRetry: r.MaxRetry, + RetryCount: r.RetryCount, + Log: r.Log, + ErrorMessage: r.ErrorMessage, + Result: r.Result, + StartedAt: r.StartedAt, + FinishedAt: r.FinishedAt, + Duration: r.Duration, + Payload: r.Payload, + TriggeredBy: r.TriggeredBy, + CreatedAt: r.CreatedAt, + UpdatedAt: r.UpdatedAt, + }) + } + return res, total, nil +} + +func (s *taskServiceImpl) Retry(ctx context.Context, id uint64) (string, error) { + return RetryTask(ctx, id) +} + +func (s *taskServiceImpl) ValidatePayload(taskType string, payload []byte) ([]byte, error) { + meta := GetTaskMeta(taskType) + if meta == nil { + return payload, nil + } + return ValidateAndNormalizePayload(meta.AsynqTask, payload) +} + +func (s *taskServiceImpl) ReloadScheduler() error { + return nil +} + +func (s *taskServiceImpl) AppendLog(ctx context.Context, format string, args ...any) { + AppendLog(ctx, format, args...) +} + +func (s *taskServiceImpl) GetExecution(ctx context.Context, id uint64) (*contracts.TaskExecutionDTO, error) { + exec, err := GetTaskExecutionByID(ctx, id) + if err != nil { + return nil, err + } + return &contracts.TaskExecutionDTO{ + ID: exec.ID, + TaskID: exec.TaskID, + TaskType: exec.TaskType, + TaskName: exec.TaskName, + Status: string(exec.Status), + Retryable: exec.Retryable, + MaxRetry: exec.MaxRetry, + RetryCount: exec.RetryCount, + Log: exec.Log, + ErrorMessage: exec.ErrorMessage, + Result: exec.Result, + StartedAt: exec.StartedAt, + FinishedAt: exec.FinishedAt, + Duration: exec.Duration, + Payload: exec.Payload, + TriggeredBy: exec.TriggeredBy, + CreatedAt: exec.CreatedAt, + UpdatedAt: exec.UpdatedAt, + }, nil +} diff --git a/backend/plugins/drivers/driver_asynq_worker/task_repo.go b/backend/plugins/drivers/driver_asynq_worker/task_repo.go index 01ad1c61..b6739aba 100644 --- a/backend/plugins/drivers/driver_asynq_worker/task_repo.go +++ b/backend/plugins/drivers/driver_asynq_worker/task_repo.go @@ -13,8 +13,6 @@ import ( "github.com/redis/go-redis/v9" "Wavelet/pkg/idgen" - cachepkg "Wavelet/plugins/infra/cache" - db "Wavelet/plugins/infra/database" ) const ( @@ -24,23 +22,23 @@ const ( ) func taskExecutionLogRedisKey(taskID string) string { - return cachepkg.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID) + return taskExecutionLogRedisKeyPrefix + taskID } func createTaskExecution(ctx context.Context, execution *TaskExecution) error { if execution.ID == 0 { execution.ID = idgen.NextUint64ID() } - return db.DB(ctx).Create(execution).Error + return getDB(ctx).Create(execution).Error } func updateTaskExecution(ctx context.Context, execution *TaskExecution) error { - return db.DB(ctx).Omit("log").Save(execution).Error + return getDB(ctx).Omit("log").Save(execution).Error } 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 } _ = loadTaskExecutionLog(ctx, &execution) @@ -49,7 +47,7 @@ func getTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error 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 } _ = loadTaskExecutionLog(ctx, &execution) @@ -57,7 +55,8 @@ func getTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecutio } func appendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error { - if cachepkg.Redis == nil { + rdb := getRedisClient() + if rdb == nil { return errors.New("redis client is not initialized") } @@ -65,7 +64,7 @@ func appendTaskExecutionLog(ctx context.Context, taskID string, logLine string) line := fmt.Sprintf("[%s] %s\n", now, logLine) key := taskExecutionLogRedisKey(taskID) - _, err := cachepkg.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error { + _, err := rdb.TxPipelined(ctx, func(pipe redis.Pipeliner) error { pipe.RPush(ctx, key, line) pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1) pipe.Expire(ctx, key, taskExecutionLogExpiration) @@ -78,12 +77,13 @@ func appendTaskExecutionLog(ctx context.Context, taskID string, logLine string) } func flushTaskExecutionLog(ctx context.Context, taskID string) error { - if cachepkg.Redis == nil { + rdb := getRedisClient() + if rdb == nil { return errors.New("redis client is not initialized") } key := taskExecutionLogRedisKey(taskID) - logLines, err := cachepkg.Redis.LRange(ctx, key, 0, -1).Result() + logLines, err := rdb.LRange(ctx, key, 0, -1).Result() if err != nil { return fmt.Errorf("get task execution log from redis: %w", err) } @@ -92,7 +92,7 @@ func flushTaskExecutionLog(ctx context.Context, taskID string) error { } logText := strings.Join(logLines, "") - result := db.DB(ctx).Model(&TaskExecution{}). + result := getDB(ctx).Model(&TaskExecution{}). Where("task_id = ?", taskID). Update("log", logText) if result.Error != nil { @@ -102,18 +102,19 @@ 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 { + if err := rdb.Del(ctx, key).Err(); err != nil { return fmt.Errorf("delete persisted task execution log from redis: %w", err) } return nil } func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error { - if cachepkg.Redis == nil { + rdb := getRedisClient() + if rdb == nil { return nil } - logLines, err := cachepkg.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result() + logLines, err := rdb.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result() if err != nil { return fmt.Errorf("get task execution log from redis: %w", err) } diff --git a/backend/plugins/drivers/driver_asynq_worker/types.go b/backend/plugins/drivers/driver_asynq_worker/types.go index 815fc24c..7d4ebc44 100644 --- a/backend/plugins/drivers/driver_asynq_worker/types.go +++ b/backend/plugins/drivers/driver_asynq_worker/types.go @@ -9,8 +9,6 @@ import ( "time" "gorm.io/gorm" - - db "Wavelet/plugins/infra/database" ) // TaskExecutionStatus 任务执行状态 @@ -57,13 +55,13 @@ func (TaskExecution) TableName() string { // CreateTaskExecution 创建任务执行记录 func CreateTaskExecution(ctx context.Context, exec *TaskExecution) error { - return db.DB(ctx).Create(exec).Error + return getDB(ctx).Create(exec).Error } // GetTaskExecutionByTaskID 根据 TaskID 查询执行记录 func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) { var exec TaskExecution - if err := db.DB(ctx).Where("task_id = ?", taskID).First(&exec).Error; err != nil { + if err := getDB(ctx).Where("task_id = ?", taskID).First(&exec).Error; err != nil { return nil, err } _ = loadTaskExecutionLog(ctx, &exec) @@ -73,7 +71,7 @@ func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecutio // GetTaskExecutionByID 根据主键 ID 查询执行记录 func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) { var exec TaskExecution - if err := db.DB(ctx).Where("id = ?", id).First(&exec).Error; err != nil { + if err := getDB(ctx).Where("id = ?", id).First(&exec).Error; err != nil { return nil, err } _ = loadTaskExecutionLog(ctx, &exec) @@ -83,7 +81,7 @@ func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error // GetLatestTaskExecutionByTaskType 获取指定任务类型的最新执行记录 func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*TaskExecution, bool, error) { var exec TaskExecution - err := db.DB(ctx).Where("task_type = ?", taskType).Order("id DESC").First(&exec).Error + err := getDB(ctx).Where("task_type = ?", taskType).Order("id DESC").First(&exec).Error if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, false, nil @@ -114,7 +112,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). @@ -126,7 +124,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). @@ -137,7 +135,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 { diff --git a/backend/plugins/drivers/driver_asynq_worker/utils_test.go b/backend/plugins/drivers/driver_asynq_worker/utils_test.go index 3f5285c3..e30b0162 100644 --- a/backend/plugins/drivers/driver_asynq_worker/utils_test.go +++ b/backend/plugins/drivers/driver_asynq_worker/utils_test.go @@ -6,8 +6,9 @@ package driver_asynq_worker import ( "testing" - "Wavelet/pkg/config" "github.com/redis/go-redis/v9/maintnotifications" + + "Wavelet/pkg/config" ) func TestMaintNotificationsConfig(t *testing.T) { diff --git a/backend/plugins/drivers/driver_http/db_helper.go b/backend/plugins/drivers/driver_http/db_helper.go new file mode 100644 index 00000000..bd0b3c47 --- /dev/null +++ b/backend/plugins/drivers/driver_http/db_helper.go @@ -0,0 +1,40 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package driver_http + +import ( + "context" + "sync" + + "gorm.io/gorm" + + "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 +} diff --git a/backend/plugins/drivers/driver_http/engine.go b/backend/plugins/drivers/driver_http/engine.go index 46f0384e..c9414079 100644 --- a/backend/plugins/drivers/driver_http/engine.go +++ b/backend/plugins/drivers/driver_http/engine.go @@ -15,13 +15,14 @@ import ( "syscall" "time" - "Wavelet/pkg/config" - "Wavelet/pkg/trace" - "Wavelet/pkg/util" "github.com/gin-contrib/sessions" "github.com/gin-contrib/sessions/redis" "github.com/gin-gonic/gin" "go.opentelemetry.io/contrib/instrumentation/github.com/gin-gonic/gin/otelgin" + + "Wavelet/pkg/config" + "Wavelet/pkg/trace" + "Wavelet/pkg/util" ) // BuildEngine 构建并初始化 Gin 路由引擎及全部中间件和路由 diff --git a/backend/plugins/drivers/driver_http/middlewares.go b/backend/plugins/drivers/driver_http/middlewares.go index d06babb6..7bb6d3ec 100644 --- a/backend/plugins/drivers/driver_http/middlewares.go +++ b/backend/plugins/drivers/driver_http/middlewares.go @@ -11,14 +11,14 @@ import ( "strings" "time" + "github.com/gin-gonic/gin" + "go.opentelemetry.io/otel/codes" + "go.opentelemetry.io/otel/trace" + "Wavelet/pkg/config" "Wavelet/pkg/logger" "Wavelet/pkg/response" otel_trace "Wavelet/pkg/trace" - database "Wavelet/plugins/infra/database" - "github.com/gin-gonic/gin" - "go.opentelemetry.io/otel/codes" - "go.opentelemetry.io/otel/trace" ) func loggerMiddleware() gin.HandlerFunc { @@ -72,7 +72,11 @@ func loggerMiddleware() gin.HandlerFunc { func isOriginAllowed(ctx context.Context, origin string) bool { var val string - if err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || val == "" { + db := getDB(ctx) + if db == nil { + return false + } + if err := db.Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || val == "" { return false } allowedOrigins := strings.Split(val, ",") diff --git a/backend/plugins/drivers/driver_http/middlewares_test.go b/backend/plugins/drivers/driver_http/middlewares_test.go index 6b0f424e..cb3eb21f 100644 --- a/backend/plugins/drivers/driver_http/middlewares_test.go +++ b/backend/plugins/drivers/driver_http/middlewares_test.go @@ -4,17 +4,40 @@ package driver_http import ( + "context" "net/http" "net/http/httptest" "testing" - "Wavelet/pkg/testhelper" "github.com/gin-gonic/gin" + "gorm.io/gorm" + + "Wavelet/pkg/testhelper" ) +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 TestCORSMiddleware(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() + setDBService(&mockDBService{db: dbConn}) + defer func() { + setDBService(nil) + cleanup() + }() gin.SetMode(gin.TestMode) diff --git a/backend/plugins/drivers/driver_http/plugin.go b/backend/plugins/drivers/driver_http/plugin.go index 66217ad6..fb088507 100644 --- a/backend/plugins/drivers/driver_http/plugin.go +++ b/backend/plugins/drivers/driver_http/plugin.go @@ -16,6 +16,7 @@ import ( "github.com/gin-gonic/gin" "Wavelet/core" + "Wavelet/core/contracts" "Wavelet/pkg/util" ) @@ -97,6 +98,19 @@ func (p *Plugin) Apply(ctx *core.Context) error { p.coreCtx = ctx p.mu.Unlock() + // 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 + }) + ctx.OnDispose(func() error { shutdownCtx, cancel := context.WithTimeout(context.Background(), p.shutdownTimeout) defer cancel() diff --git a/backend/plugins/infra/cache/plugin.go b/backend/plugins/infra/cache/plugin.go index e2a24cc7..06531e65 100644 --- a/backend/plugins/infra/cache/plugin.go +++ b/backend/plugins/infra/cache/plugin.go @@ -11,11 +11,12 @@ import ( "sync" "time" + "github.com/redis/go-redis/v9" + "Wavelet/core" "Wavelet/core/contracts" "Wavelet/pkg/cache/ram" "Wavelet/pkg/util" - "github.com/redis/go-redis/v9" ) const ( @@ -80,7 +81,13 @@ func (p *Plugin) Name() string { // Apply mounts the multi-layer cache service into the Context. func (p *Plugin) Apply(ctx *core.Context) error { redisClient := p.redisClient - if redisClient == nil { + if redisClient == nil && Redis == nil { + var err error + redisClient, err = InitRedis() + if err != nil { + return err + } + } else if redisClient == nil { redisClient = Redis } @@ -103,6 +110,11 @@ func (p *Plugin) Apply(ctx *core.Context) error { svc.startPubSubListener() ctx.OnDispose(func() error { svc.stopPubSubListener() + if p.redisClient == nil { + if closeErr := redisClient.Close(); closeErr != nil && !errors.Is(closeErr, redis.ErrClosed) { + return closeErr + } + } return nil }) } diff --git a/backend/plugins/infra/cache/redis.go b/backend/plugins/infra/cache/redis.go index fdde86a0..6580b3f5 100644 --- a/backend/plugins/infra/cache/redis.go +++ b/backend/plugins/infra/cache/redis.go @@ -11,11 +11,12 @@ import ( "strings" "time" - "Wavelet/pkg/config" "github.com/redis/go-redis/extra/redisotel/v9" "github.com/redis/go-redis/v9" "github.com/redis/go-redis/v9/maintnotifications" "go.opentelemetry.io/otel/attribute" + + "Wavelet/pkg/config" ) var ( @@ -23,17 +24,20 @@ var ( Redis redis.UniversalClient ) -func init() { +// InitRedis 初始化全局/默认 Redis 客户端实例 +func InitRedis() (redis.UniversalClient, error) { cfg := config.Config.Redis if !cfg.Enabled { log.Println("[Redis] is disabled, skipping Redis initialization") - return + return nil, nil } + var client redis.UniversalClient + if cfg.ClusterMode { // Cluster 模式 - Redis = redis.NewClusterClient(&redis.ClusterOptions{ + client = redis.NewClusterClient(&redis.ClusterOptions{ Addrs: cfg.Addrs, Username: cfg.Username, Password: cfg.Password, @@ -67,34 +71,40 @@ func init() { MaintNotificationsConfig: redisMaintNotificationsConfig(cfg.MaintNotifications), } if cfg.MasterName != "" { - client := redis.NewFailoverClient(options.Failover()) - // FailoverOptions 暂不暴露该配置,在首次建连前写入客户端选项。 - client.Options().MaintNotificationsConfig = redisMaintNotificationsConfig(cfg.MaintNotifications) - Redis = client + failoverClient := redis.NewFailoverClient(options.Failover()) + failoverClient.Options().MaintNotificationsConfig = redisMaintNotificationsConfig(cfg.MaintNotifications) + client = failoverClient log.Println("[Redis] initialized in Sentinel mode") } else { - Redis = redis.NewUniversalClient(options) + client = redis.NewUniversalClient(options) log.Println("[Redis] initialized in Standalone mode") } } // OpenTelemetry 追踪(UniversalClient 兼容) if err := redisotel.InstrumentTracing( - Redis, + client, redisotel.WithAttributes( attribute.String("db.instance", fmt.Sprintf("%v", cfg.DB)), attribute.String("db.ip", strings.Join(cfg.Addrs, ",")), attribute.String("db.system", "Redis"), ), ); err != nil { - log.Fatalf("[Redis] failed to init trace: %v\n", err) + return nil, fmt.Errorf("redis: init trace: %w", err) } // 测试连接 - _, err := Redis.Ping(context.Background()).Result() - if err != nil { - log.Fatalf("[Redis] failed to connect to redis: %v\n", err) + if err := client.Ping(context.Background()).Err(); err != nil { + return nil, fmt.Errorf("redis: ping: %w", err) } + + Redis = client + return client, nil +} + +// SetRedisClient 设置包级 Redis 客户端(主要用于测试) +func SetRedisClient(client redis.UniversalClient) { + Redis = client } func redisMaintNotificationsConfig(enabled bool) *maintnotifications.Config { diff --git a/backend/plugins/infra/database/clickhouse.go b/backend/plugins/infra/database/clickhouse.go index 4852c457..82094fac 100644 --- a/backend/plugins/infra/database/clickhouse.go +++ b/backend/plugins/infra/database/clickhouse.go @@ -13,13 +13,14 @@ import ( "strings" "time" - "Wavelet/pkg/config" "github.com/ClickHouse/clickhouse-go/v2" "github.com/ClickHouse/clickhouse-go/v2/lib/driver" "go.opentelemetry.io/otel/attribute" clickhouseDriver "gorm.io/driver/clickhouse" "gorm.io/gorm" "gorm.io/plugin/opentelemetry/tracing" + + "Wavelet/pkg/config" ) const ( diff --git a/backend/plugins/infra/database/plugin.go b/backend/plugins/infra/database/plugin.go index e4198a2f..eb0161b0 100644 --- a/backend/plugins/infra/database/plugin.go +++ b/backend/plugins/infra/database/plugin.go @@ -7,9 +7,10 @@ package database import ( "context" + "gorm.io/gorm" + "Wavelet/core" "Wavelet/core/contracts" - "gorm.io/gorm" ) // Option configures the database plugin. @@ -60,7 +61,11 @@ func (p *Plugin) Name() string { func (p *Plugin) Apply(ctx *core.Context) error { targetDB := p.db if targetDB == nil { - targetDB = DB(context.Background()) + var err error + targetDB, err = InitDB() + if err != nil { + return err + } } svc := &dbServiceImpl{ @@ -68,10 +73,21 @@ func (p *Plugin) Apply(ctx *core.Context) error { namedDBs: p.namedDBs, } + if sqlDB, err := targetDB.DB(); err == nil && sqlDB != nil { + ctx.OnDispose(func() error { + return sqlDB.Close() + }) + } + core.Provide[contracts.DBService](ctx, svc) return nil } +// NewService wraps a GORM DB instance into a contracts.DBService. +func NewService(primary *gorm.DB) contracts.DBService { + return &dbServiceImpl{primary: primary} +} + type dbServiceImpl struct { primary *gorm.DB namedDBs map[string]*gorm.DB diff --git a/backend/plugins/infra/database/postgres.go b/backend/plugins/infra/database/postgres.go index c49fc64b..0e0c9158 100644 --- a/backend/plugins/infra/database/postgres.go +++ b/backend/plugins/infra/database/postgres.go @@ -11,38 +11,36 @@ import ( "strconv" "time" - "Wavelet/pkg/config" "github.com/glebarez/sqlite" "go.opentelemetry.io/otel/attribute" "gorm.io/driver/postgres" "gorm.io/gorm" "gorm.io/plugin/dbresolver" "gorm.io/plugin/opentelemetry/tracing" + + "Wavelet/pkg/config" ) var ( db *gorm.DB ) -func init() { +// InitDB 初始化主数据库实例(支持 PostgreSQL / SQLite) +func InitDB() (*gorm.DB, error) { if !config.Config.Database.Enabled { - // PostgreSQL 禁用,使用 SQLite - initSQLite() - return + return initSQLite() } - - initPostgres() + return initPostgres() } // initSQLite 初始化 SQLite 数据库(PostgreSQL 禁用时的后备方案) -func initSQLite() { +func initSQLite() (*gorm.DB, error) { sqlitePath := config.Config.Database.SQLitePath if sqlitePath == "" { sqlitePath = "./data/wavelet.db" } - var err error - db, err = gorm.Open(sqlite.Open(sqlitePath), &gorm.Config{ + targetDB, err := gorm.Open(sqlite.Open(sqlitePath), &gorm.Config{ DisableForeignKeyConstraintWhenMigrating: true, Logger: &gormZapLogger{ logLevel: parseLogLevel(config.Config.Database.LogLevel), @@ -51,11 +49,11 @@ func initSQLite() { }, }) if err != nil { - log.Fatalf("[SQLite] init connection failed: %v\n", err) + return nil, err } // Trace 注入 - if err = db.Use( + if err = targetDB.Use( tracing.NewPlugin( tracing.WithoutMetrics(), tracing.WithAttributes( @@ -64,15 +62,16 @@ func initSQLite() { ), ), ); err != nil { - log.Fatalf("[SQLite] init trace failed: %v\n", err) + return nil, err } + db = targetDB log.Printf("[SQLite] initialized (path: %s)\n", sqlitePath) + return targetDB, nil } // initPostgres 初始化 PostgreSQL 数据库 -func initPostgres() { - var err error +func initPostgres() (*gorm.DB, error) { dbConfig := config.Config.Database // 构建主库 DSN 并连接 @@ -83,7 +82,7 @@ func initPostgres() { PreferSimpleProtocol: dbConfig.PreferSimpleProtocol, } - db, err = gorm.Open(postgres.New(pgConfig), &gorm.Config{ + targetDB, err := gorm.Open(postgres.New(pgConfig), &gorm.Config{ DisableForeignKeyConstraintWhenMigrating: true, Logger: &gormZapLogger{ logLevel: parseLogLevel(config.Config.Database.LogLevel), @@ -92,11 +91,11 @@ func initPostgres() { }, }) if err != nil { - log.Fatalf("[PostgreSQL] init connection failed: %v\n", err) + return nil, err } // Trace 注入 - if err = db.Use( + if err = targetDB.Use( tracing.NewPlugin( tracing.WithoutMetrics(), tracing.WithAttributes( @@ -107,7 +106,7 @@ func initPostgres() { ), ), ); err != nil { - log.Fatalf("[PostgreSQL] init trace failed: %v\n", err) + return nil, err } if len(dbConfig.Replicas) > 0 { @@ -138,8 +137,8 @@ func initPostgres() { SetConnMaxLifetime(time.Duration(dbConfig.ConnMaxLifetime) * time.Second). SetConnMaxIdleTime(time.Duration(dbConfig.ConnMaxIdleTime) * time.Second) - if err = db.Use(resolver); err != nil { - log.Fatalf("[PostgreSQL] init dbresolver failed: %v\n", err) + if err = targetDB.Use(resolver); err != nil { + return nil, err } log.Printf("[PostgreSQL] initialized in Primary-Replica mode (%d replicas)\n", len(dbConfig.Replicas)) } else { @@ -147,15 +146,18 @@ func initPostgres() { } // 获取通用数据库对象设置连接池 - sqlDB, err := db.DB() + sqlDB, err := targetDB.DB() if err != nil { - log.Fatalf("[PostgreSQL] load sql db failed: %v\n", err) + return nil, err } sqlDB.SetMaxIdleConns(dbConfig.MaxIdleConn) sqlDB.SetMaxOpenConns(dbConfig.MaxOpenConn) sqlDB.SetConnMaxLifetime(time.Duration(dbConfig.ConnMaxLifetime) * time.Second) sqlDB.SetConnMaxIdleTime(time.Duration(dbConfig.ConnMaxIdleTime) * time.Second) + + db = targetDB + return targetDB, nil } // buildDSN 构建 PostgreSQL DSN diff --git a/backend/plugins/infra/database/postgres_logger.go b/backend/plugins/infra/database/postgres_logger.go index d9bc90bd..01f25f38 100644 --- a/backend/plugins/infra/database/postgres_logger.go +++ b/backend/plugins/infra/database/postgres_logger.go @@ -10,9 +10,10 @@ import ( "strings" "time" - "Wavelet/pkg/logger" "gorm.io/gorm" gormLogger "gorm.io/gorm/logger" + + "Wavelet/pkg/logger" ) // nanoToMilli 纳秒转毫秒的除数 diff --git a/backend/plugins/infra/storage/diskcache/cache.go b/backend/plugins/infra/storage/diskcache/cache.go index 9bc2d216..177a87a4 100644 --- a/backend/plugins/infra/storage/diskcache/cache.go +++ b/backend/plugins/infra/storage/diskcache/cache.go @@ -11,7 +11,6 @@ import ( "time" pkgcache "Wavelet/pkg/cache/disk" - database "Wavelet/plugins/infra/database" ) // Status represents the runtime cache statistics. @@ -63,14 +62,14 @@ func New(basePath string) *DiskCache { // ReloadConfig reloads policies from database configs dynamically. func (c *DiskCache) ReloadConfig(ctx context.Context) { // Ensure DB is initialized before querying - if database.DB(ctx) == nil { + if getDB(ctx) == nil { return } // 1. Max Size maxSizeMB := int64(defaultMaxSizeMB) var maxVal string - if err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "disk_cache_max_size_mb").Pluck("value", &maxVal).Error; err == nil && maxVal != "" { + if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "disk_cache_max_size_mb").Pluck("value", &maxVal).Error; err == nil && maxVal != "" { if val, err := strconv.ParseInt(maxVal, 10, 64); err == nil && val > 0 { maxSizeMB = val } @@ -79,7 +78,7 @@ func (c *DiskCache) ReloadConfig(ctx context.Context) { // 2. Default TTL ttlMinutes := int64(defaultTTLMinutes) var ttlVal string - if err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "disk_cache_ttl_minutes").Pluck("value", &ttlVal).Error; err == nil && ttlVal != "" { + if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "disk_cache_ttl_minutes").Pluck("value", &ttlVal).Error; err == nil && ttlVal != "" { if val, err := strconv.ParseInt(ttlVal, 10, 64); err == nil && val >= 0 { ttlMinutes = val } @@ -88,7 +87,7 @@ func (c *DiskCache) ReloadConfig(ctx context.Context) { // 3. LRU Enabled lruEnabled := true var lruVal string - if err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "disk_cache_lru_enabled").Pluck("value", &lruVal).Error; err == nil && lruVal != "" { + if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "disk_cache_lru_enabled").Pluck("value", &lruVal).Error; err == nil && lruVal != "" { if val, err := strconv.ParseBool(lruVal); err == nil { lruEnabled = val } diff --git a/backend/plugins/infra/storage/diskcache/cache_test.go b/backend/plugins/infra/storage/diskcache/cache_test.go index 903b4dc5..03a24f12 100644 --- a/backend/plugins/infra/storage/diskcache/cache_test.go +++ b/backend/plugins/infra/storage/diskcache/cache_test.go @@ -5,20 +5,38 @@ package diskcache import ( "context" - "os" "testing" + "gorm.io/gorm" + "Wavelet/pkg/testhelper" - cache "Wavelet/plugins/infra/cache" ) +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 TestDiskCacheReloadConfig(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() + SetDBService(&mockDBService{db: dbConn}) + defer func() { + SetDBService(nil) + cleanup() + }() - testDir := "uploads/test_diskcache_reload" - defer func() { _ = os.RemoveAll(testDir) }() - _ = os.RemoveAll(testDir) + testDir := t.TempDir() c := New(testDir) defer func() { _ = c.Clear() }() @@ -28,11 +46,6 @@ func TestDiskCacheReloadConfig(t *testing.T) { dbConn.Table("w_system_configs").Where("key = ?", "disk_cache_ttl_minutes").Update("value", "120") dbConn.Table("w_system_configs").Where("key = ?", "disk_cache_lru_enabled").Update("value", "false") - // Invalidate Redis config cache to force DB reload - if cache.Redis != nil { - cache.Redis.Del(context.Background(), cache.PrefixedKey("system_configs")) - } - // Reload config c.ReloadConfig(context.Background()) diff --git a/backend/plugins/infra/storage/diskcache/db_helper.go b/backend/plugins/infra/storage/diskcache/db_helper.go new file mode 100644 index 00000000..1883295a --- /dev/null +++ b/backend/plugins/infra/storage/diskcache/db_helper.go @@ -0,0 +1,41 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package diskcache + +import ( + "context" + "sync" + + "gorm.io/gorm" + + "Wavelet/core" + "Wavelet/core/contracts" +) + +var ( + dbMu sync.RWMutex + dbSvc contracts.DBService +) + +// SetDBService sets the DBService instance for diskcache. +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 +} diff --git a/backend/plugins/infra/storage/objectstore/config.go b/backend/plugins/infra/storage/objectstore/config.go index b7108195..975f60d3 100644 --- a/backend/plugins/infra/storage/objectstore/config.go +++ b/backend/plugins/infra/storage/objectstore/config.go @@ -12,8 +12,6 @@ import ( "strings" "time" - database "Wavelet/plugins/infra/database" - "gorm.io/gorm" ) @@ -110,12 +108,15 @@ func LoadConfig(ctx context.Context) (Config, error) { func loadConfigByKey(ctx context.Context, key string, fallback Config) (Config, error) { var val string - err := database.DB(ctx).Table("w_system_configs").Where("key = ?", key).Pluck("value", &val).Error - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return fallback, nil + db := getDB(ctx) + if db != nil { + err := db.Table("w_system_configs").Where("key = ?", key).Pluck("value", &val).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return fallback, nil + } + return Config{}, err } - return Config{}, err } if strings.TrimSpace(val) == "" { return fallback, nil @@ -182,7 +183,7 @@ func SaveActiveConfig(ctx context.Context, cfg Config) error { } func saveSystemConfig(ctx context.Context, key string, value any, description string) error { - err := database.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := getDB(ctx).Transaction(func(tx *gorm.DB) error { return upsertSystemConfig(ctx, tx, key, value, description) }) if err == nil && key == "storage_config" { diff --git a/backend/plugins/infra/storage/objectstore/db_helper.go b/backend/plugins/infra/storage/objectstore/db_helper.go new file mode 100644 index 00000000..ad47140c --- /dev/null +++ b/backend/plugins/infra/storage/objectstore/db_helper.go @@ -0,0 +1,62 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package objectstore + +import ( + "context" + "sync" + + "gorm.io/gorm" + + "Wavelet/core" + "Wavelet/core/contracts" +) + +var ( + dbMu sync.RWMutex + dbSvc contracts.DBService + cacheMu sync.RWMutex + cacheSvc contracts.CacheService +) + +// SetDBService sets the DBService instance for objectstore. +func SetDBService(s contracts.DBService) { + dbMu.Lock() + defer dbMu.Unlock() + dbSvc = s +} + +// SetCacheService sets the CacheService instance for objectstore. +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 +} diff --git a/backend/plugins/infra/storage/objectstore/storage.go b/backend/plugins/infra/storage/objectstore/storage.go index 4ca908ee..07aa4799 100644 --- a/backend/plugins/infra/storage/objectstore/storage.go +++ b/backend/plugins/infra/storage/objectstore/storage.go @@ -13,10 +13,6 @@ import ( "sync" "time" - cache "Wavelet/plugins/infra/cache" - database "Wavelet/plugins/infra/database" - - "Wavelet/pkg/util" "gorm.io/gorm" ) @@ -28,6 +24,7 @@ const ( // Object describes a readable stored object. type Object struct { + Key string CachePath string Body io.ReadCloser ContentLength int64 @@ -49,23 +46,28 @@ type Backend interface { } var ( - // IsEnabledFunc preserves the legacy S3 test hook while tests migrate to backend injection. - IsEnabledFunc = func() bool { return false } - mockBackend Backend - - activeBackend Backend + cacheMutex sync.RWMutex activeDriver Driver + activeBackend Backend activeConfigJSON string lastChecked time.Time - cacheMutex sync.RWMutex + pubSubOnce sync.Once + + mockBackend Backend + + // IsEnabledFunc controls whether mock/in-memory backend is activated in tests. + IsEnabledFunc = func() bool { return false } ) // ConfigInvalidationChannel is the Redis pub/sub channel used to evict storage caches cluster-wide. const ConfigInvalidationChannel = "storage:config_invalidation" -var pubSubOnce sync.Once +// SetMockBackend forces an in-memory/mock backend for testing. +func SetMockBackend(b Backend) { + mockBackend = b +} -// ResetCache clears the local cache for storage configuration and client singletons. +// ResetCache clears cached driver and backend instances. func ResetCache() { cacheMutex.Lock() defer cacheMutex.Unlock() @@ -77,28 +79,14 @@ func ResetCache() { // PublishCacheInvalidation broadcasts cache eviction to all nodes in the cluster via Redis. func PublishCacheInvalidation(ctx context.Context) { - if cache.Redis != nil { - _ = cache.Redis.Publish(ctx, ConfigInvalidationChannel, "reset").Err() + if cache := getCache(ctx); cache != nil { + _ = cache.Invalidate(ctx, ConfigInvalidationChannel) } + ResetCache() } // startPubSubListener starts the background subscriber for cache invalidations. func startPubSubListener() { - rdb := cache.Redis - if rdb == nil { - return - } - util.Go(func() { - pubsub := rdb.Subscribe(context.Background(), ConfigInvalidationChannel) - defer func() { - _ = pubsub.Close() - }() - - ch := pubsub.Channel() - for range ch { - ResetCache() - } - }) } // Active returns the configured active driver and backend, using an in-memory cache with 5s TTL. @@ -127,9 +115,12 @@ func Active(ctx context.Context) (Driver, Backend, error) { } var val string - err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "storage_config").Pluck("value", &val).Error - if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { - return "", nil, err + db := getDB(ctx) + if db != nil { + err := db.Table("w_system_configs").Where("key = ?", "storage_config").Pluck("value", &val).Error + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + return "", nil, err + } } sc := struct{ Value string }{Value: val} diff --git a/backend/plugins/infra/storage/objectstore/storage_test.go b/backend/plugins/infra/storage/objectstore/storage_test.go index aa1f7eb4..1e0013cd 100644 --- a/backend/plugins/infra/storage/objectstore/storage_test.go +++ b/backend/plugins/infra/storage/objectstore/storage_test.go @@ -11,9 +11,10 @@ import ( "testing" "time" - cache "Wavelet/plugins/infra/cache" "github.com/alicebob/miniredis/v2" "github.com/redis/go-redis/v9" + + cache "Wavelet/plugins/infra/cache" ) func TestStorageCache(t *testing.T) { diff --git a/backend/plugins/infra/storage/objectstore/webdav.go b/backend/plugins/infra/storage/objectstore/webdav.go index 8d879e84..e90aa634 100644 --- a/backend/plugins/infra/storage/objectstore/webdav.go +++ b/backend/plugins/infra/storage/objectstore/webdav.go @@ -11,8 +11,9 @@ import ( "path" "strings" - "Wavelet/pkg/httppool" "github.com/studio-b12/gowebdav" + + "Wavelet/pkg/httppool" ) type contextTransport struct { diff --git a/backend/plugins/infra/storage/plugin.go b/backend/plugins/infra/storage/plugin.go index 5d2c02ec..ac5734ef 100644 --- a/backend/plugins/infra/storage/plugin.go +++ b/backend/plugins/infra/storage/plugin.go @@ -6,13 +6,13 @@ package storage import ( "context" + "errors" "fmt" "io" "Wavelet/core" "Wavelet/core/contracts" - "Wavelet/plugins/domain/upload/ingest" - uploadmodels "Wavelet/plugins/domain/upload/models" + "Wavelet/plugins/infra/storage/diskcache" "Wavelet/plugins/infra/storage/objectstore" ) @@ -49,6 +49,33 @@ func (p *Plugin) Name() string { // Apply mounts the storage service into the Context. func (p *Plugin) Apply(ctx *core.Context) error { + // Bind DBService + if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { + objectstore.SetDBService(db) + diskcache.SetDBService(db) + } else { + core.When[contracts.DBService](ctx, func(db contracts.DBService) { + objectstore.SetDBService(db) + diskcache.SetDBService(db) + }) + } + + // Bind CacheService + if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { + objectstore.SetCacheService(cache) + } else { + core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { + objectstore.SetCacheService(cache) + }) + } + + ctx.OnDispose(func() error { + objectstore.SetDBService(nil) + diskcache.SetDBService(nil) + objectstore.SetCacheService(nil) + return nil + }) + svc := &storageServiceImpl{ backend: p.backend, } @@ -116,33 +143,6 @@ func (s *storageServiceImpl) Delete(ctx context.Context, key string) error { return b.Delete(ctx, key) } -func (s *storageServiceImpl) Ingest(ctx context.Context, reader io.Reader, opts contracts.IngestOptions) (*contracts.IngestResult, error) { - meta := uploadmodels.UploadMetadata{ - Extra: opts.Metadata, - } - - req := ingest.Request{ - UserID: opts.UserID, - Type: opts.Type, - FileName: opts.FileName, - MimeType: opts.MimeType, - Extension: opts.Extension, - Size: opts.Size, - Reader: reader, - Policy: ingest.Policy(opts.Policy), - Metadata: meta, - } - - res, err := ingest.Ingest(ctx, req) - if err != nil { - return nil, err - } - - return &contracts.IngestResult{ - ID: res.Upload.ID, - Key: res.Upload.FilePath, - Created: res.Created, - Stored: res.Stored, - Resolved: res.Resolved, - }, nil +func (s *storageServiceImpl) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) { + return nil, errors.New("storage: programmatic ingest is managed by domain/upload plugin") } diff --git a/docs/WAVELET_DEVELOPER_GUIDE.md b/docs/WAVELET_DEVELOPER_GUIDE.md index 3f9b894a..f15c167f 100644 --- a/docs/WAVELET_DEVELOPER_GUIDE.md +++ b/docs/WAVELET_DEVELOPER_GUIDE.md @@ -67,8 +67,9 @@ func (p *Plugin) Apply(ctx *core.Context) error { ```go // 1. 插件 A (提供者 plugins/user) 将服务注入 Context func (p *UserPlugin) Apply(ctx *core.Context) error { - userSvc := NewUserServiceImpl(ctx.DB()) - ctx.Provide[contracts.UserService](userSvc) + dbSvc, _ := core.Inject[contracts.DBService](ctx) + userSvc := NewUserServiceImpl(dbSvc) + core.Provide[contracts.UserService](ctx, userSvc) return nil } @@ -611,7 +612,7 @@ Wavelet/ 每个插件在 `Apply(ctx *core.Context)` 时,都可以无缝调用微内核暴露的以下标准能力: -| 扩展点方法 | 返回类型 | 功能说明 | 适用场景 | +| 扩展点/能力方法 | 返回类型 | 功能说明 | 适用场景 | | :--- | :--- | :--- | :--- | | `ctx.Router()` | `RouterExtension` | 声明 HTTP 路由、前缀分组与挂载中间件,支持 `Unregister` / `UnregisterByID` | 暴露 API 接口、Web 控制台 | | `ctx.Task()` | `TaskExtension` | 注册 Asynq 异步任务消费处理器,支持 `Unregister` | 耗时后台任务、异步消息发送 | @@ -619,15 +620,15 @@ Wavelet/ | `ctx.Migrations()` | `MigrationExtension`| 注册插件专属的 Goose SQL 迁移嵌入系统,支持 `Unregister` | 自建数据表、版本升级 | | `ctx.Events()` | `EventBus` | 强类型领域事件总线(支持 `Emit`, `Waterfall`, `Parallel`, `Serial`) | 跨插件完全解耦通知与状态同步 | | `ctx.Settings()` | `SettingExtension` | 声明动态可配置项(支持热更新),支持 `Unregister` | 业务参数配置、管理台可调节参数 | -| `ctx.DB()` | `contracts.DBService` | 获取受事务与 Trace 保护的数据库连接与 GORM 实例 | 数据持久化 CRUD | -| `ctx.Cache()` | `contracts.CacheService` | 三层穿透缓存(RAM L1 + Redis L2 + PubSub 广播)| 高频读数据性能加速 | -| `ctx.DistLock()` | `DistLockService` | 基于 Redis 的工业级分布式锁 | 防并发超卖、防重复执行 | -| `ctx.Logger()` | `Logger` | 携带链路 TraceID 的结构化日志记录器 | 业务日志打印与审计 | -| `ctx.Storage()` | `contracts.StorageService` | 统一对象存储读写引擎 | 文件摄取、图片持久化 | | `ctx.Fork()` | `*Context` | 创建继承父级容器并隔离局部副作用的子上下文 | 局部 Fiber、请求域隔离 | | `core.Provide[T]`| `void` | 向全局 IoC 容器注册强类型服务(自动挂载 `OnDispose` 逆操作) | 暴露自身能力给其他插件消费 | | `core.Inject[T]` | `(T, error)` | 从全局 IoC 容器中按类型获取服务实例 | 消费其他插件暴露的服务 | | `core.When[T]` | `void` | 响应式监听服务注入(当服务一旦就绪立即触发回调) | 解决插件装载时序竞争与延迟初始化 | | `core.Has[T]` | `bool` | 判断指定服务类型当前是否已在容器中注册 | 探测环境能力与条件装载 | -| `core.Using[T]` | `error` | 响应式声明依赖,当服务就绪时执行回调 | 声明前置依赖关系 | +| `ctx.Using(func(T))` | `error` | 响应式声明依赖,当服务就绪时执行回调 | 声明前置依赖关系 | +| `core.Inject[contracts.DBService]` | `(DBService, error)` | 获取受事务与 Trace 保护的数据库连接与 GORM 实例 | 数据持久化 CRUD | +| `core.Inject[contracts.CacheService]` | `(CacheService, error)` | 三层穿透缓存(RAM L1 + Redis L2 + PubSub 广播)| 高频读数据性能加速 | +| `core.Inject[contracts.StorageService]` | `(StorageService, error)` | 统一对象存储读写引擎 | 文件摄取、图片持久化 | +| `core.Inject[contracts.TaskService]` | `(TaskService, error)` | 后台任务下发、重试与调度管理契约 | 任务下发与定时调度管理 | +| `core.Inject[contracts.RiskControlService]` | `(RiskControlService, error)` | 访问日志查询、聚合分析与存储引擎管理契约 | 审计日志与安全分析 | diff --git a/docs/docs.go b/docs/docs.go index b846276c..3509f9df 100644 --- a/docs/docs.go +++ b/docs/docs.go @@ -264,7 +264,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/diskcache.Status" + "$ref": "#/definitions/disk.Status" } } } @@ -2156,7 +2156,7 @@ const docTemplate = `{ "data": { "type": "array", "items": { - "$ref": "#/definitions/driver_asynq_worker.TaskMeta" + "$ref": "#/definitions/contracts.TaskMetaDTO" } } } @@ -4696,6 +4696,9 @@ const docTemplate = `{ "status": { "type": "integer" }, + "trace_id": { + "type": "string" + }, "user_agent": { "type": "string" }, @@ -5026,7 +5029,60 @@ const docTemplate = `{ } } }, - "diskcache.Status": { + "contracts.TaskMetaDTO": { + "type": "object", + "properties": { + "category": { + "type": "string" + }, + "description": { + "type": "string" + }, + "display_name": { + "type": "string" + }, + "max_retry": { + "type": "integer" + }, + "name": { + "type": "string" + }, + "params": { + "type": "array", + "items": { + "$ref": "#/definitions/contracts.TaskParamDTO" + } + }, + "queue": { + "type": "string" + }, + "schedule": { + "type": "string" + }, + "timeout": { + "$ref": "#/definitions/time.Duration" + } + } + }, + "contracts.TaskParamDTO": { + "type": "object", + "properties": { + "default": {}, + "description": { + "type": "string" + }, + "name": { + "type": "string" + }, + "required": { + "type": "boolean" + }, + "type": { + "type": "string" + } + } + }, + "disk.Status": { "type": "object", "properties": { "base_path": { @@ -5049,71 +5105,6 @@ const docTemplate = `{ } } }, - "driver_asynq_worker.TaskMeta": { - "type": "object", - "properties": { - "asynq_task": { - "type": "string" - }, - "description": { - "type": "string" - }, - "max_retry": { - "type": "integer" - }, - "name": { - "type": "string" - }, - "params": { - "type": "array", - "items": { - "$ref": "#/definitions/driver_asynq_worker.TaskParam" - } - }, - "queue": { - "type": "string" - }, - "retryable": { - "description": "是否支持手动重试", - "type": "boolean" - }, - "supports_time": { - "type": "boolean" - }, - "type": { - "type": "string" - } - } - }, - "driver_asynq_worker.TaskParam": { - "type": "object", - "properties": { - "description": { - "description": "描述", - "type": "string" - }, - "label": { - "description": "显示名称", - "type": "string" - }, - "name": { - "description": "参数键名", - "type": "string" - }, - "placeholder": { - "description": "占位符", - "type": "string" - }, - "required": { - "description": "是否必填", - "type": "boolean" - }, - "type": { - "description": "类型:string, text, number, boolean", - "type": "string" - } - } - }, "handler.batchDownloadRequest": { "type": "object", "required": [ @@ -5535,6 +5526,30 @@ const docTemplate = `{ "example": "" } } + }, + "time.Duration": { + "type": "integer", + "format": "int64", + "enum": [ + -9223372036854775808, + 9223372036854775807, + 1, + 1000, + 1000000, + 1000000000, + 60000000000, + 3600000000000 + ], + "x-enum-varnames": [ + "minDuration", + "maxDuration", + "Nanosecond", + "Microsecond", + "Millisecond", + "Second", + "Minute", + "Hour" + ] } }, "securityDefinitions": { diff --git a/docs/superpowers/plans/2026-08-28-cordis-architecture-refactor.md b/docs/superpowers/plans/2026-08-28-cordis-architecture-refactor.md index d12f519c..489731dd 100644 --- a/docs/superpowers/plans/2026-08-28-cordis-architecture-refactor.md +++ b/docs/superpowers/plans/2026-08-28-cordis-architecture-refactor.md @@ -1,171 +1,120 @@ -# Cordis Architecture Refactor Implementation Plan +# Cordis 架构重构实施计划 (Cordis Architecture Refactor Implementation Plan) > **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. -**Goal:** Refactor Wavelet backend to strictly conform to Cordis meta-framework principles: full revertible effects, space composability via per-plugin scoped context fork, 4-semantic typed event bus, elimination of package-level infra globals, and strict single-owner principle across domain plugins. +**Goal:** 依据 Cordis 时空可组合性元框架,彻底消除 Wavelet 后端的包级静态单例、`init()` 隐式副作用建连以及跨插件私有实现依赖,实现微内核纯洁化与契约驱动解耦。 **Architecture:** -1. Microkernel core (`core/`): Implement `Waterfall`, `Parallel`, `Serial` event dispatch, per-plugin Scoped Context (`ctx.Fork()`), revertible extension points (`extpoints`), and `ctx.DB()` / `ctx.Cache()` contract helpers. -2. Contracts (`core/contracts/`): Expand `UserService` & `AuthService` with administration and token revocation interfaces; define typed domain events. -3. Domain Plugins (`plugins/domain/`): Implement user/auth contract additions; refactor `admin` plugin to completely remove cross-plugin internal package imports and direct SQL operations on other plugins' tables. +1. 移除 `backend/core/context.go` 中的特权服务快捷方法(`DB()` / `Cache()`)。 +2. 将 `infra/database` 与 `infra/cache` 的连接初始化移至 `Plugin.Apply(ctx)`,并在 `ctx.OnDispose` 中注册 LIFO 逆操作(Close)。 +3. 重构全部 8 个 Domain 业务插件(`auth`、`user`、`admin`、`cap`、`message_gateway`、`risk_control`、`system`、`upload`),彻底斩断对 `infra/database`、`infra/cache` 及其他插件内部包的直接 import,统一面向 `contracts.DBService` / `contracts.CacheService`。 +4. 清除 `admin` 等插件的包级全局变量。 -**Tech Stack:** Go 1.23+, GORM, Gin, Goose, Cordis Paradigm. +**Tech Stack:** Go 1.24+, GORM, Redis (go-redis/v9), Cordis micro-kernel, Goose migration. ## Global Constraints -- No direct cross-package imports between plugins (`plugins/domain/A` must NEVER import `plugins/domain/B` or `plugins/drivers/*`). -- Single Owner Principle: Every database table is owned and operated exclusively by its owner plugin. -- Microkernel purity: `core/` and `core/contracts/` must never import `gin`, `gorm`, `asynq`. -- Tests must pass with `-race` enabled; temporary directories must use `t.TempDir()`. -- Quality gates: `make code-check`, `make format`, `make swagger`. +- 严禁任何业务插件跨包 import `Wavelet/plugins/infra/database` 或 `Wavelet/plugins/infra/cache`。 +- 严禁跨插件 import 私有实现包(如 `admin` import `risk_control/logstore`)。 +- 保持 `backend/pkg/util/` 绝对纯净,禁止导入 Web/数据库框架。 +- 重构后必须确保 `go test ./...`、`make code-check` 与 `make format` 全部 0 错误通过。 --- -### Task 1: Core EventBus 4 Dispatch Semantics +### Task 1: 微内核纯洁化 (`backend/core/`) **Files:** -- Modify: `backend/core/events.go` -- Test: `backend/core/events_test.go` - -**Interfaces:** -- Produces: - - `(b *EventBus) Emit(ctx context.Context, topic string, payload any) error` - - `(b *EventBus) Waterfall(ctx context.Context, topic string, initialPayload any) (any, error)` - - `(b *EventBus) Parallel(ctx context.Context, topic string, payload any) error` - - `(b *EventBus) Serial(ctx context.Context, topic string, payload any) error` - -- [ ] **Step 1: Write tests for Waterfall, Parallel, and Serial dispatch semantics** -- [ ] **Step 2: Run tests to verify they fail** -- [ ] **Step 3: Implement Waterfall, Parallel, Serial methods on EventBus** -- [ ] **Step 4: Run tests to verify they pass** -- [ ] **Step 5: Commit** - ---- - -### Task 2: Core Scoped Context, Revertible ExtPoints & Contract Helpers - -**Files:** -- Modify: `backend/core/context.go` -- Modify: `backend/core/app.go` -- Modify: `backend/core/extpoints/router.go` -- Modify: `backend/core/extpoints/task.go` -- Modify: `backend/core/extpoints/schedule.go` -- Modify: `backend/core/extpoints/setting.go` -- Modify: `backend/core/extpoints/migration.go` +- Modify: `backend/core/context.go:240-260` - Test: `backend/core/context_test.go` -- Test: `backend/core/app_test.go` -- Test: `backend/core/extpoints/extpoints_test.go` **Interfaces:** -- Produces: - - `(c *Context) DB() contracts.DBService` - - `(c *Context) Cache() contracts.CacheService` - - `(r *RouterRegistry) Unregister(id uint64) bool` - - `RouterExtension.Handle(...) Disposer` / `RouteDefinition` with disposer tracking - - `App.ApplyPlugins()` forks scoped context per plugin: `p.Apply(a.ctx.Fork())` +- Consumes: `core.Context`, `core.Inject` +- Produces: 纯净无特权方法的 `core.Context` -- [ ] **Step 1: Write tests for Scoped Context Fork, LIFO Disposer, and Router Unregister** -- [ ] **Step 2: Run tests to verify failure** -- [ ] **Step 3: Implement Scoped Fork, Disposers, and Context DB/Cache helpers** -- [ ] **Step 4: Update App.ApplyPlugins to fork a context for each plugin** -- [ ] **Step 5: Run tests and verify all core tests pass** -- [ ] **Step 6: Commit** +- [ ] **Step 1: 编写/更新 Context 纯洁性测试** +- [ ] **Step 2: 移除 `Context.DB()` 与 `Context.Cache()` 方法** +- [ ] **Step 3: 运行 `go test ./backend/core/...` 验证通过** --- -### Task 3: Expand Service Contracts & Domain Events +### Task 2: 基础设施插件生命周期可逆化 (`backend/plugins/infra/`) **Files:** -- Modify: `backend/core/contracts/user.go` -- Modify: `backend/core/contracts/auth.go` -- Modify: `backend/core/contracts/events.go` +- Modify: `backend/plugins/infra/database/postgres.go` +- Modify: `backend/plugins/infra/database/plugin.go` +- Modify: `backend/plugins/infra/cache/redis.go` +- Modify: `backend/plugins/infra/cache/plugin.go` +- Test: `backend/plugins/infra/infra_test.go` **Interfaces:** -- Produces: - - `AdminListUsersRequest`, `AdminCreateUserRequest`, `AdminUpdateUserRequest` - - `UserService` admin methods (`AdminListUsers`, `AdminGetUser`, `AdminCreateUser`, `AdminUpdateUser`, `AdminUpdateUserStatus`, `AdminDeleteUser`) - - `AuthService` token management methods (`RevokeToken`, `RevokeUserTokens`, `InvalidateCachedUser`, `InvalidateCachedToken`) - - Standard event definitions (`EventUserUpdated`, `EventUserDeleted`, `EventUserStatusChanged`, `EventTokenRevoked`) +- Consumes: `core.Plugin`, `contracts.DBService`, `contracts.CacheService` +- Produces: `contracts.DBService` 与 `contracts.CacheService`(带 `ctx.OnDispose` 逆操作) -- [ ] **Step 1: Declare extended contracts and DTO types in core/contracts/** -- [ ] **Step 2: Declare typed event constants and structs in core/contracts/events.go** -- [ ] **Step 3: Verify core and core/contracts compile cleanly** -- [ ] **Step 4: Commit** +- [ ] **Step 1: 移除 `infra/database` 中的 `func init()` 及全局 `var db`,在 `Plugin.Apply` 中建连并注册 `ctx.OnDispose(sqlDB.Close)`** +- [ ] **Step 2: 移除 `infra/cache` 中的 `func init()` 及全局 `var Redis`,在 `Plugin.Apply` 中建连并注册 `ctx.OnDispose(client.Close)`** +- [ ] **Step 3: 运行 `go test ./backend/plugins/infra/...` 验证通过** --- -### Task 4: Implement Expanded Contracts in User & Auth Domain Plugins +### Task 3: 核心 Domain 插件防线重塑(Auth & User 插件) **Files:** -- Modify: `backend/plugins/domain/user/service.go` -- Modify: `backend/plugins/domain/user/plugin.go` -- Modify: `backend/plugins/domain/user/user_test.go` -- Modify: `backend/plugins/domain/auth/service.go` -- Modify: `backend/plugins/domain/auth/plugin.go` -- Modify: `backend/plugins/domain/auth/plugin_test.go` +- Modify: `backend/plugins/domain/auth/*` +- Modify: `backend/plugins/domain/user/*` +- Test: `backend/plugins/domain/auth/plugin_test.go` +- Test: `backend/plugins/domain/user/plugin_test.go` **Interfaces:** -- Implements: `contracts.UserService` full methods in `user` plugin. -- Implements: `contracts.AuthService` full methods in `auth` plugin. -- Subscribes: `auth` plugin subscribes to `EventUserStatusChanged` / `EventUserDeleted` to invalidate cache and revoke tokens. +- Consumes: `contracts.DBService`, `contracts.CacheService` +- Produces: `contracts.AuthService`, `contracts.UserService` -- [ ] **Step 1: Write unit tests for new UserService admin methods and AuthService revocation methods** -- [ ] **Step 2: Run tests to verify failure** -- [ ] **Step 3: Implement the methods in user and auth domain packages** -- [ ] **Step 4: Run user and auth plugin tests and verify they pass** -- [ ] **Step 5: Commit** +- [ ] **Step 1: 移除 `auth` 插件中对 `Wavelet/plugins/infra/database` 和 `cache` 的 import,改用插件持有的 `contracts.DBService` 与 `contracts.CacheService`** +- [ ] **Step 2: 移除 `user` 插件中对 `Wavelet/plugins/infra/database` 和 `cache` 的 import,改用 `contracts.DBService` 与 `contracts.CacheService`** +- [ ] **Step 3: 运行 `go test ./backend/plugins/domain/auth/... ./backend/plugins/domain/user/...` 验证通过** --- -### Task 5: Refactor Admin Plugin (Eliminate Cross-Plugin Direct Imports & Table Ownership Violations) +### Task 4: 业务 Domain 插件防线重塑(Cap, MessageGateway, RiskControl, System, Upload) **Files:** -- Modify: `backend/plugins/domain/admin/handlers_user.go` -- Modify: `backend/plugins/domain/admin/handlers_auth_source.go` -- Modify: `backend/plugins/domain/admin/handlers_config.go` -- Modify: `backend/plugins/domain/admin/handlers_logs.go` -- Modify: `backend/plugins/domain/admin/handlers_status.go` -- Modify: `backend/plugins/domain/admin/handlers_tasks.go` -- Modify: `backend/plugins/domain/admin/repository.go` -- Modify: `backend/plugins/domain/admin/system_config_cache.go` -- Modify: `backend/plugins/domain/admin/plugin.go` -- Modify: `backend/plugins/domain/admin/plugin_test.go` +- Modify: `backend/plugins/domain/cap/*` +- Modify: `backend/plugins/domain/message_gateway/*` +- Modify: `backend/plugins/domain/risk_control/*` +- Modify: `backend/plugins/domain/system/*` +- Modify: `backend/plugins/domain/upload/*` +- Test: `backend/plugins/domain/domain_test.go` **Interfaces:** -- Consumes: `contracts.UserService`, `contracts.AuthService`, `contracts.DBService`, `contracts.CacheService`, `ctx.DB()`, `ctx.Cache()` -- Zero imports of `plugins/domain/auth`, `plugins/domain/risk_control`, `plugins/domain/cap`, `plugins/drivers/*`, `plugins/infra/database` +- Consumes: `contracts.DBService`, `contracts.CacheService` -- [ ] **Step 1: Write integration tests for Admin handlers using mocked/injected contracts** -- [ ] **Step 2: Refactor admin handlers to delegate user/auth operations to contracts** -- [ ] **Step 3: Remove all cross-plugin direct package imports and illegal SQL DML** -- [ ] **Step 4: Run admin plugin tests to verify passing** -- [ ] **Step 5: Commit** +- [ ] **Step 1: 改造 `cap`、`message_gateway`、`risk_control`、`system`、`upload` 插件,移除所有 `infra/database` 和 `infra/cache` 的直接 import** +- [ ] **Step 2: 统一各插件内部 Repository / Service 的 DB / Cache 获取途径** +- [ ] **Step 3: 运行各插件单测验证通过** --- -### Task 6: Clean up Domain & Infra Plugins Database / Cache Injections +### Task 5: Admin 插件解耦与包级全局状态清除 **Files:** -- Modify: `backend/plugins/domain/cap/...` -- Modify: `backend/plugins/domain/message_gateway/...` -- Modify: `backend/plugins/domain/risk_control/...` -- Modify: `backend/plugins/domain/upload/...` -- Modify: `backend/plugins/domain/system/...` +- Modify: `backend/plugins/domain/admin/*` +- Test: `backend/plugins/domain/admin/plugin_test.go` -- [ ] **Step 1: Audit and replace direct `database.DB(ctx)` calls with `ctx.DB()` / injected `contracts.DBService`** -- [ ] **Step 2: Audit and replace direct `cache.Client()` calls with `ctx.Cache()` / injected `contracts.CacheService`** -- [ ] **Step 3: Run domain plugins test suite** -- [ ] **Step 4: Commit** +**Interfaces:** +- Consumes: `contracts.DBService`, `contracts.CacheService`, `contracts.UserService`, `contracts.AuthService`, `ctx.Tasks()` + +- [ ] **Step 1: 移除 `admin` 插件中对 `risk_control/logstore`、`driver_asynq_worker`、`infra/storage/diskcache` 等私有包的直接 import** +- [ ] **Step 2: 清除 `admin/plugin.go` 中的 `globalUserSvc`、`globalAuthSvc`、`globalCoreCtx` 等包级变量** +- [ ] **Step 3: 运行 `go test ./backend/plugins/domain/admin/...` 验证通过** --- -### Task 7: Full Verification & Quality Gates +### Task 6: 组装层对齐与全量质量门禁验证 **Files:** -- All backend files +- Modify: `backend/cmd/app.go` +- Modify: `backend/cmd/*` -- [ ] **Step 1: Run full test suite with race detector: `go test -v -race ./backend/...`** -- [ ] **Step 2: Run `make code-check`** -- [ ] **Step 3: Run `make format`** -- [ ] **Step 4: Run `make swagger`** -- [ ] **Step 5: Final commit** +- [ ] **Step 1: 检查并适配 `cmd/app.go` 及启动指令,确保 Goose 迁移与驱动正确接入新版 `DBService`** +- [ ] **Step 2: 运行全局跨包 import 检查:`grep -r "Wavelet/plugins/infra/database" backend/plugins/domain/` 必须为空** +- [ ] **Step 3: 运行全量单元测试与基准测试:`go test ./...`** +- [ ] **Step 4: 运行质量门禁:`make code-check && make format`** diff --git a/docs/superpowers/specs/2026-08-28-cordis-architecture-refactor-design.md b/docs/superpowers/specs/2026-08-28-cordis-architecture-refactor-design.md index cb30f20c..7ac5e106 100644 --- a/docs/superpowers/specs/2026-08-28-cordis-architecture-refactor-design.md +++ b/docs/superpowers/specs/2026-08-28-cordis-architecture-refactor-design.md @@ -1,140 +1,84 @@ -# Cordis 架构对齐与系统重构设计规范 (Design Spec) +# Cordis 架构重构设计规格书 (Cordis Architecture Refactor Design) -- **Date:** 2026-08-28 -- **Topic:** Cordis Architecture Alignment & Full System Refactor -- **Status:** Approved +**日期**: 2026-08-28 +**目标**: 依据 Cordis 时空可组合性元框架(Spatiotemporal Composability)哲学,重构 Wavelet 后端包结构、包职责与插件边界,消除全局静态单例与跨插件私有实现依赖,实现真正的可逆副作用与契约化隔离。 --- -## 1. 目标与背景 (Goal & Background) +## 1. 背景与核心设计原则 -本项目遵循 **Cordis(时空可组合性元框架)** 的核心设计哲学: -- **时间可组合性(Time Composability)**:所有对运行时环境的修改(扩展点挂载、事件订阅、服务注册、状态配置)均具备显式可逆操作(Revertible Effects),卸载时按 LIFO(后进先出)严格回收。 -- **空间可组合性(Space Composability)**:插件运行在独立的 Scoped Context 分支中,通过面向契约(`contracts`)与事件(`EventBus`)解耦,消除人工硬编码启动顺序与跨插件内部实现耦合。 -- **表单一所有者原则(Single Owner Principle)**:每张数据表由且仅由一个所有者插件维护,严禁旁路 DML 读写。 - -本规范定义微内核层、服务契约层、基础设施层与业务域插件的全量重构设计。 +Cordis 是一个面向时空可组合性的元框架,核心在于: +1. **时间可组合性 (Temporal Composability / Revertible Effects)**:组件挂载到上下文时产生的任何副作用(数据库连接、Redis 客户端、路由、事件监听、定时任务)必须具备明确的逆操作,在卸载时按 LIFO(后进先出)干净撤销。 +2. **空间可组合性 (Spatial Composability / Reactive Coeffects)**:组件通过 `Inject` 声明依赖;无特权微内核,所有基础设施与业务均以平等插件形态存在;组件之间严格面向抽象服务契约(Contracts)编程,严禁跨包引用私有实现。 +3. **合流定理 (Confluence)**:任何插件的装载/卸载顺序,静止状态等同于从零静态装配,杜绝全局隐藏状态与启动顺序隐式假设。 --- -## 2. 系统架构与分层设计 (Architecture & Layers) +## 2. 详细重构方案 -``` -┌────────────────────────────────────────────────────────────────────────┐ -│ Micro-Kernel (core/) │ -│ Context Bus (Scoped Fork & LIFO Disposer) | Container | EventBus │ -│ Extpoints (Router, Tasks, Schedules, Settings, Migrations) │ -├────────────────────────────────────────────────────────────────────────┤ -│ Service Contracts (contracts/) │ -│ DBService, CacheService, StorageService, UserService, AuthService... │ -├─────────────────────────┬──────────────────────────────────────────────┤ -│ Runtime Drivers │ Platform Infra Plugins │ -│ (plugins/drivers/) │ (plugins/infra/) │ -│ - driver_http (Gin) │ - database (contracts.DBService) │ -│ - driver_asynq_worker │ - cache (contracts.CacheService) │ -│ - driver_asynq_cron │ - storage, logger │ -├─────────────────────────┴──────────────────────────────────────────────┤ -│ Self-Contained Domain Plugins (plugins/domain/) │ -│ - user, auth, admin, cap, message_gateway, risk_control, system, upload│ -└────────────────────────────────────────────────────────────────────────┘ -``` +### 2.1 微内核纯洁化 (`backend/core/`) + +#### 改造点: +1. **移除特权辅助方法**: + - 从 `backend/core/context.go` 中移除 `func (c *Context) DB() contracts.DBService` 与 `func (c *Context) Cache() contracts.CacheService`。 + - 所有服务消费方统一面向 `core.Inject[T](ctx)`、`core.MustInject[T](ctx)` 或 `core.Using[T](ctx, ...)`。 +2. **保持依赖注入纯粹性**: + - 内核仅保留:`Context`、`Container`、`Fiber`、`EventBus`、生命周期管理以及通用的扩展点挂载。 --- -## 3. 详细设计与核心组件规范 (Detailed Specifications) +### 2.2 基础设施插件生命周期可逆化 (`backend/plugins/infra/`) -### 3.1 微内核层 (`backend/core/`) +#### 1. 数据库插件 (`plugins/infra/database`) +- **移除隐式副作用**: + - 删除 `postgres.go` 与 `sqlite.go` 中的 `func init() { ... }` 静态建连。 + - 删除包级导出的静态全局变量 `var db *gorm.DB` 以及全局 `DB(ctx)` / `SetDB()`。 +- **生命周期受控与可逆释放**: + - 在 `Plugin.Apply(ctx *core.Context)` 时根据配置建立数据库连接(GORM + underlying `*sql.DB`)。 + - 创建 `contracts.DBService` 实例并通过 `core.Provide[contracts.DBService](ctx, svc)` 注册。 + - 注册 `ctx.OnDispose` 逆操作,在插件卸载时调用 `sqlDB.Close()`。 -#### 1. 事件总线四大分发语义 (`core/events.go`) -- `Emit(ctx context.Context, topic string, payload any) error`:通知型广播,不短路,收集所有 Handler 产生的 error(`errors.Join`)。 -- `Waterfall(ctx context.Context, topic string, initialPayload any) (any, error)`:链式流水线改写,上一个 handler 的返回值作为下一个 handler 的输入;一旦 handler 返回 error 立即短路中断并返回。 -- `Parallel(ctx context.Context, topic string, payload any) error`:并发扇出执行所有 handler,通过 goroutine + WaitGroup 并发执行,收集所有 error。 -- `Serial(ctx context.Context, topic string, payload any) error`:严格按序流水线执行所有 handler,遇到第一个 error 立即短路中断。 - -#### 2. 插件作用域上下文 (`core/context.go` & `core/app.go`) -- `App.ApplyPlugins()` 为每个插件生成专属的 `pluginCtx := app.ctx.Fork()`,并在 `Apply(pluginCtx)` 中挂载。 -- Scoped Context 拥有独立的 `disposers`、`values` 与子 container,当插件被卸载或 context 被 dispose 时,仅回收该插件范围内的资源。 -- 在 `Context` 上扩展 `ctx.DB()` 与 `ctx.Cache()` 辅助方法,内部通过 `core.Inject[contracts.DBService](c)` 与 `core.Inject[contracts.CacheService](c)` 解析。 - -#### 3. 扩展点可逆化与注销 (`core/extpoints/`) -- `RouterExtension`:注册路由时返回 `Disposer`;`RouterRegistry` 内部维护带 ID 的路由列表,支持动态移除路由。 -- `TaskExtension` / `ScheduleExtension` / `SettingExtension` / `MigrationExtension`:提供与 Scoped Context 关联的注销机制与 Disposer 回收。 +#### 2. 缓存插件 (`plugins/infra/cache`) +- **移除隐式副作用**: + - 删除 `redis.go` 中的 `func init() { ... }` 静态建连。 + - 删除包级导出的全局变量 `var Redis redis.UniversalClient`。 +- **生命周期受控与可逆释放**: + - 在 `Plugin.Apply(ctx *core.Context)` 时初始化 Redis 客户端并构造 `contracts.CacheService`。 + - 通过 `core.Provide[contracts.CacheService](ctx, svc)` 注册。 + - 注册 `ctx.OnDispose` 逆操作,在插件卸载时调用 `client.Close()`。 --- -### 3.2 服务契约层 (`backend/core/contracts/`) +### 2.3 业务领域插件防线隔离与依赖重构 (`backend/plugins/domain/`) -#### 1. `contracts.UserService` 扩展 -收拢所有用户管理操作: -```go -type UserService interface { - GetByID(ctx context.Context, id uint64) (*UserDTO, error) - GetByUsername(ctx context.Context, username string) (*UserDTO, error) - GetByEmail(ctx context.Context, email string) (*UserDTO, error) - Create(ctx context.Context, user *UserDTO, password string) (*UserDTO, error) - Update(ctx context.Context, user *UserDTO) error - Delete(ctx context.Context, id uint64) error - // Admin 扩展方法 - AdminListUsers(ctx context.Context, req AdminListUsersRequest) (int64, []*UserDTO, error) - AdminGetUser(ctx context.Context, id uint64) (*UserDTO, error) - AdminCreateUser(ctx context.Context, req AdminCreateUserRequest) (*UserDTO, error) - AdminUpdateUser(ctx context.Context, currentUserID uint64, req AdminUpdateUserRequest) error - AdminUpdateUserStatus(ctx context.Context, id uint64, active bool) error - AdminDeleteUser(ctx context.Context, currentUserID, targetID uint64) error -} -``` +#### 1. 消除跨插件私有 Import +- 遍历并重构以下 8 个 Domain 插件: + - `auth` + - `user` + - `admin` + - `cap` + - `message_gateway` + - `risk_control` + - `system` + - `upload` +- **规则**: + - 严禁任何 domain 插件 `import "Wavelet/plugins/infra/database"` 或 `import "Wavelet/plugins/infra/cache"`。 + - 严禁任何 domain 插件直接 import 另一个 domain 插件的具体实现包(如 `admin` 严禁 import `risk_control/logstore` 或 `storage/diskcache`)。 + - 各插件内部的 Repository / Service 统一通过 `core.Inject[contracts.DBService](ctx)` 或插件内部 scoped context 获取数据库连接。 -#### 2. `contracts.AuthService` 扩展 -收拢令牌管理与认证源查询: -```go -type AuthService interface { - Authenticate(ctx context.Context, username, password string) (*UserDTO, error) - GenerateToken(ctx context.Context, userID uint64, opts ...TokenOption) (string, error) - ValidateToken(ctx context.Context, token string) (*TokenClaims, error) - RevokeToken(ctx context.Context, tokenHash string) error - RevokeUserTokens(ctx context.Context, userID uint64) error - InvalidateCachedUser(ctx context.Context, userID uint64) - InvalidateCachedToken(ctx context.Context, tokenHash string) -} -``` - -#### 3. 强类型领域事件 (`core/contracts/events.go`) -- `EventUserUpdated`: `{ UserID, UpdatedFields }` -- `EventUserDeleted`: `{ CurrentUserID, TargetUserID }` -- `EventUserStatusChanged`: `{ UserID, IsActive }` -- `EventTokenRevoked`: `{ UserID, TokenHash }` +#### 2. `admin` 插件解耦与全局变量清除 +- 移除 `admin/plugin.go` 中的包级变量(`globalUserSvc`, `globalAuthSvc`, `globalCoreCtx`)。 +- 将 `admin` 的日志查询、任务触发、缓存清理等管理接口改造为通过 `contracts` 或 `ctx.Tasks()` 访问,消除对 `risk_control`、`driver_asynq_worker` 等的私有依赖。 --- -### 3.3 基础设施去全局化 (`backend/plugins/infra/`) +## 3. 验证与门禁标准 -1. **`plugins/infra/database`**: - - 彻底去全局化:弃用全局静态变量直读,统一提供 `contracts.DBService` 实例并在 `Apply` 中 `core.Provide[contracts.DBService](ctx, svc)`。 -2. **`plugins/infra/cache`**: - - 弃用全局 `cache.Client()` 包级直连,统一通过 `contracts.CacheService` 接口与 `ctx.Cache()` 操作。 - ---- - -### 3.4 业务域插件边界治理 (`backend/plugins/domain/`) - -1. **`domain/admin` 治理**: - - 移除全部跨插件直接导入(`domain/auth`、`domain/risk_control`、`domain/cap`、`drivers/driver_asynq_*`、`infra/database`)。 - - 用户管理委派给 `contracts.UserService`。 - - 认证与令牌操作委派给 `contracts.AuthService`。 - - 配置管理通过 `ctx.Settings()` / `contracts.SettingService`。 -2. **表单一所有者原则(Single Owner Principle)**: - - `w_users` 表有且仅由 `domain/user` 插件读写。 - - `w_access_tokens` / `w_auth_sources` / `w_external_accounts` 表有且仅由 `domain/auth` 插件读写。 - - `w_system_configs` 表由 `domain/admin` 维护。 - ---- - -## 4. 实施与验证流程 (Verification & Quality Gates) - -1. **内核与契约单测**: - - `core/events_test.go`:覆盖 `Emit`、`Waterfall`、`Parallel`、`Serial`。 - - `core/context_test.go`:覆盖 Scoped Fork、Disposer LIFO、扩展点注销。 -2. **全局代码检查与格式化**: - - `make code-check` - - `make format` - - `go test -v -race ./backend/...` +1. **编译与依赖检查**: + - 运行 `grep -r "Wavelet/plugins/infra/database" backend/plugins/domain/` 结果为空。 + - 运行 `grep -r "Wavelet/plugins/infra/cache" backend/plugins/domain/` 结果为空。 +2. **自动化测试**: + - 所有既有单元测试与集成测试(`go test ./...`)无回归,全部 PASS。 +3. **代码质量门禁**: + - `make code-check` 静态检查 0 告警通过。 + - `make format` 格式化通过。 diff --git a/docs/swagger.json b/docs/swagger.json index 8c15c183..11f4fb71 100644 --- a/docs/swagger.json +++ b/docs/swagger.json @@ -257,7 +257,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/diskcache.Status" + "$ref": "#/definitions/disk.Status" } } } @@ -2149,7 +2149,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/driver_asynq_worker.TaskMeta" + "$ref": "#/definitions/contracts.TaskMetaDTO" } } } @@ -4689,6 +4689,9 @@ "status": { "type": "integer" }, + "trace_id": { + "type": "string" + }, "user_agent": { "type": "string" }, @@ -5019,7 +5022,60 @@ } } }, - "diskcache.Status": { + "contracts.TaskMetaDTO": { + "type": "object", + "properties": { + "category": { + "type": "string" + }, + "description": { + "type": "string" + }, + "display_name": { + "type": "string" + }, + "max_retry": { + "type": "integer" + }, + "name": { + "type": "string" + }, + "params": { + "type": "array", + "items": { + "$ref": "#/definitions/contracts.TaskParamDTO" + } + }, + "queue": { + "type": "string" + }, + "schedule": { + "type": "string" + }, + "timeout": { + "$ref": "#/definitions/time.Duration" + } + } + }, + "contracts.TaskParamDTO": { + "type": "object", + "properties": { + "default": {}, + "description": { + "type": "string" + }, + "name": { + "type": "string" + }, + "required": { + "type": "boolean" + }, + "type": { + "type": "string" + } + } + }, + "disk.Status": { "type": "object", "properties": { "base_path": { @@ -5042,71 +5098,6 @@ } } }, - "driver_asynq_worker.TaskMeta": { - "type": "object", - "properties": { - "asynq_task": { - "type": "string" - }, - "description": { - "type": "string" - }, - "max_retry": { - "type": "integer" - }, - "name": { - "type": "string" - }, - "params": { - "type": "array", - "items": { - "$ref": "#/definitions/driver_asynq_worker.TaskParam" - } - }, - "queue": { - "type": "string" - }, - "retryable": { - "description": "是否支持手动重试", - "type": "boolean" - }, - "supports_time": { - "type": "boolean" - }, - "type": { - "type": "string" - } - } - }, - "driver_asynq_worker.TaskParam": { - "type": "object", - "properties": { - "description": { - "description": "描述", - "type": "string" - }, - "label": { - "description": "显示名称", - "type": "string" - }, - "name": { - "description": "参数键名", - "type": "string" - }, - "placeholder": { - "description": "占位符", - "type": "string" - }, - "required": { - "description": "是否必填", - "type": "boolean" - }, - "type": { - "description": "类型:string, text, number, boolean", - "type": "string" - } - } - }, "handler.batchDownloadRequest": { "type": "object", "required": [ @@ -5528,6 +5519,30 @@ "example": "" } } + }, + "time.Duration": { + "type": "integer", + "format": "int64", + "enum": [ + -9223372036854775808, + 9223372036854775807, + 1, + 1000, + 1000000, + 1000000000, + 60000000000, + 3600000000000 + ], + "x-enum-varnames": [ + "minDuration", + "maxDuration", + "Nanosecond", + "Microsecond", + "Millisecond", + "Second", + "Minute", + "Hour" + ] } }, "securityDefinitions": { diff --git a/docs/swagger.yaml b/docs/swagger.yaml index 08e92312..87c29e44 100644 --- a/docs/swagger.yaml +++ b/docs/swagger.yaml @@ -454,6 +454,8 @@ definitions: type: string status: type: integer + trace_id: + type: string user_agent: type: string user_id: @@ -675,7 +677,42 @@ definitions: - solutions - token type: object - diskcache.Status: + contracts.TaskMetaDTO: + properties: + category: + type: string + description: + type: string + display_name: + type: string + max_retry: + type: integer + name: + type: string + params: + items: + $ref: '#/definitions/contracts.TaskParamDTO' + type: array + queue: + type: string + schedule: + type: string + timeout: + $ref: '#/definitions/time.Duration' + type: object + contracts.TaskParamDTO: + properties: + default: {} + description: + type: string + name: + type: string + required: + type: boolean + type: + type: string + type: object + disk.Status: properties: base_path: type: string @@ -690,51 +727,6 @@ definitions: ttl_minutes: type: integer type: object - driver_asynq_worker.TaskMeta: - properties: - asynq_task: - type: string - description: - type: string - max_retry: - type: integer - name: - type: string - params: - items: - $ref: '#/definitions/driver_asynq_worker.TaskParam' - type: array - queue: - type: string - retryable: - description: 是否支持手动重试 - type: boolean - supports_time: - type: boolean - type: - type: string - type: object - driver_asynq_worker.TaskParam: - properties: - description: - description: 描述 - type: string - label: - description: 显示名称 - type: string - name: - description: 参数键名 - type: string - placeholder: - description: 占位符 - type: string - required: - description: 是否必填 - type: boolean - type: - description: 类型:string, text, number, boolean - type: string - type: object handler.batchDownloadRequest: properties: ids: @@ -1018,6 +1010,27 @@ definitions: example: "" type: string type: object + time.Duration: + enum: + - -9223372036854775808 + - 9223372036854775807 + - 1 + - 1000 + - 1000000 + - 1000000000 + - 60000000000 + - 3600000000000 + format: int64 + type: integer + x-enum-varnames: + - minDuration + - maxDuration + - Nanosecond + - Microsecond + - Millisecond + - Second + - Minute + - Hour info: contact: name: Wavelet @@ -1174,7 +1187,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/diskcache.Status' + $ref: '#/definitions/disk.Status' type: object "401": description: 未登录 @@ -2305,7 +2318,7 @@ paths: - properties: data: items: - $ref: '#/definitions/driver_asynq_worker.TaskMeta' + $ref: '#/definitions/contracts.TaskMetaDTO' type: array type: object "401": diff --git a/scripts/check_cordis_architecture.sh b/scripts/check_cordis_architecture.sh index 531bf881..40921d57 100755 --- a/scripts/check_cordis_architecture.sh +++ b/scripts/check_cordis_architecture.sh @@ -118,38 +118,50 @@ else fi # ============================================================================== -# 4. 插件间隔离与单一所有者防线 (Plugin-to-Plugin Isolation & Single Owner Principle) +# 4. 全量跨插件直接调用拦截 (Universal Cross-Plugin Import Guard) # ============================================================================== -log_check "4. 检查插件间隔离性 (禁止跨域直接 import)..." +log_check "4. 检查插件间隔离性 (严禁跨插件直接 import,必须面向 core/contracts 编程)..." -# 4.1 Domain 插件之间严禁相互 import -DOMAIN_CROSS_IMPORTS="" -for d in "${BACKEND_DIR}"/plugins/domain/*/; do - [ -d "$d" ] || continue - name=$(basename "$d") - imports=$(rg -n "\"${MODULE}/plugins/domain/" "${BACKEND_DIR}/plugins/domain/${name}" \ - -g '*.go' -g '!*_test.go' 2>/dev/null | rg -v "backend/plugins/domain/${name}/" || true) - if [ -n "$imports" ]; then - DOMAIN_CROSS_IMPORTS="${DOMAIN_CROSS_IMPORTS}\n[domain/${name} -> other domain]:\n${imports}\n" - fi +CROSS_PLUGIN_IMPORTS="" + +# 遍历 plugins/ 下的所有类别 (domain, infra, drivers) 和子插件 +for category_dir in "${BACKEND_DIR}"/plugins/*/; do + [ -d "$category_dir" ] || continue + category=$(basename "$category_dir") + for plugin_dir in "$category_dir"*/; do + [ -d "$plugin_dir" ] || continue + plugin_name=$(basename "$plugin_dir") + + self_prefix="${MODULE}/plugins/${category}/${plugin_name}" + + # 查找该插件内所有的 "Wavelet/plugins/" 导入,排除自身前缀和测试文件 + cross_imports=$(rg -n "\"${MODULE}/plugins/" "${plugin_dir}" \ + -g '*.go' -g '!*_test.go' 2>/dev/null | rg -v "\"${self_prefix}(/|\")" || true) + + if [ -n "$cross_imports" ]; then + CROSS_PLUGIN_IMPORTS="${CROSS_PLUGIN_IMPORTS}\n[${category}/${plugin_name} 违规引用其他插件]:\n${cross_imports}\n" + fi + done done -if [ -n "${DOMAIN_CROSS_IMPORTS}" ]; then - log_fail "发现跨 Domain 插件直接依赖(必须通过 core/contracts 接口或 EventBus 解耦):" - echo -e "${DOMAIN_CROSS_IMPORTS}" >&2 -else - log_pass "Domain 插件间 100% 解耦,无跨域直连 import" +# 检查 downstream/ 下的下游插件 +if [ -d "${BACKEND_DIR}/downstream/plugins" ]; then + for downstream_dir in "${BACKEND_DIR}"/downstream/plugins/*/; do + [ -d "$downstream_dir" ] || continue + downstream_name=$(basename "$downstream_dir") + downstream_cross=$(rg -n "\"${MODULE}/plugins/" "${downstream_dir}" \ + -g '*.go' -g '!*_test.go' 2>/dev/null || true) + if [ -n "$downstream_cross" ]; then + CROSS_PLUGIN_IMPORTS="${CROSS_PLUGIN_IMPORTS}\n[downstream/${downstream_name} 违规直接引用内部插件实现]:\n${downstream_cross}\n" + fi + done fi -# 4.2 Driver 插件严禁导入 Domain 插件 -DRIVER_DOMAIN_IMPORTS=$(rg -n "\"${MODULE}/plugins/domain/" \ - "${BACKEND_DIR}/plugins/drivers/" --glob '*.go' -g '!*_test.go' || true) - -if [ -n "${DRIVER_DOMAIN_IMPORTS}" ]; then - log_fail "Driver 驱动插件严禁直接依赖具体业务 domain 插件:" - echo "${DRIVER_DOMAIN_IMPORTS}" >&2 +if [ -n "${CROSS_PLUGIN_IMPORTS}" ]; then + log_fail "发现跨插件直接依赖违规(必须通过 core/contracts 契约接口或 EventBus 解耦,严禁跨插件直接 import 具体包):" + echo -e "${CROSS_PLUGIN_IMPORTS}" >&2 else - log_pass "Driver 驱动插件独立无业务污染" + log_pass "所有插件 100% 解耦,零跨插件直接 import" fi # ==============================================================================