refactor(core): align with cordis spatiotemporal composability architecture

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

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