mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-04 07:06:36 +08:00
refactor(core): align with cordis spatiotemporal composability architecture
- Purify core micro-kernel by removing context hardcoded helpers and reverse dependencies - Eliminate init() side effects in infra plugins with reversible lifecycle disposal - Completely isolate plugins by removing cross-plugin imports and using core/contracts - Introduce TaskService and RiskControlService contracts for unified cross-plugin APIs - Regenerate Swagger documentation and update developer guide matrix - Achieve 0 violations in check_cordis_architecture.sh and 100% test pass
This commit is contained in:
+8
-38
@@ -11,19 +11,13 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
)
|
||||
|
||||
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
|
||||
|
||||
var (
|
||||
accessCacheOnce sync.Once
|
||||
|
||||
fileAccessWhitelistMu sync.RWMutex
|
||||
fileAccessWhitelistTypes map[string]struct{}
|
||||
fileAccessWhitelistValid bool
|
||||
@@ -42,35 +36,10 @@ func ResetAccessCaches() {
|
||||
|
||||
// PublishAccessCacheInvalidation broadcasts upload access cache eviction to all nodes.
|
||||
func PublishAccessCacheInvalidation(ctx context.Context) {
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Publish(ctx, fileAccessInvalidationChannel, "reset").Err()
|
||||
if cache := shared.GetCache(ctx); cache != nil {
|
||||
_ = cache.Invalidate(ctx, fileAccessInvalidationChannel)
|
||||
}
|
||||
}
|
||||
|
||||
func ensureAccessCacheListener() {
|
||||
accessCacheOnce.Do(startAccessCacheInvalidationListener)
|
||||
}
|
||||
|
||||
func startAccessCacheInvalidationListener() {
|
||||
rdb := cachepkg.Redis
|
||||
if rdb == nil {
|
||||
return
|
||||
}
|
||||
|
||||
util.Go(func() {
|
||||
pubsub := rdb.Subscribe(
|
||||
context.Background(),
|
||||
objectstore.ConfigInvalidationChannel,
|
||||
fileAccessInvalidationChannel,
|
||||
)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
for range pubsub.Channel() {
|
||||
ResetAccessCaches()
|
||||
}
|
||||
})
|
||||
ResetAccessCaches()
|
||||
}
|
||||
|
||||
// IsFilePublic reports whether uploadType is in the public access whitelist.
|
||||
@@ -81,8 +50,6 @@ func IsFilePublic(ctx context.Context, uploadType string) bool {
|
||||
}
|
||||
|
||||
func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} {
|
||||
ensureAccessCacheListener()
|
||||
|
||||
fileAccessWhitelistMu.RLock()
|
||||
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
|
||||
types := fileAccessWhitelistTypes
|
||||
@@ -115,8 +82,11 @@ func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} {
|
||||
|
||||
func parseFileAccessWhitelist(ctx context.Context) []string {
|
||||
var sc struct{ Value string }
|
||||
err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
|
||||
if err != nil || sc.Value == "" {
|
||||
db := shared.GetDB(ctx)
|
||||
if db != nil {
|
||||
_ = db.Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
|
||||
}
|
||||
if sc.Value == "" {
|
||||
return []string{shared.DefaultPublicUploadType}
|
||||
}
|
||||
|
||||
|
||||
@@ -8,13 +8,12 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
)
|
||||
|
||||
func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetAccessCaches()
|
||||
|
||||
@@ -31,7 +30,7 @@ func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetAccessCaches()
|
||||
|
||||
@@ -48,7 +47,7 @@ func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetAccessCaches()
|
||||
|
||||
@@ -71,7 +70,7 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAccessCacheTTLExpires(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetAccessCaches()
|
||||
|
||||
|
||||
+70
-97
@@ -5,15 +5,14 @@ package cache
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/cache/ram"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -22,16 +21,8 @@ const (
|
||||
uploadMetaInvalidationChan = "upload:meta_invalidation"
|
||||
)
|
||||
|
||||
type uploadMetaInvalidationMessage struct {
|
||||
ID uint64 `json:"id"`
|
||||
}
|
||||
|
||||
var (
|
||||
uploadMetaRAM = ram.MustNew[uint64, models.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize})
|
||||
uploadMetaListenerOnce sync.Once
|
||||
uploadMetaListenerCtx context.Context
|
||||
uploadMetaListenerCancel context.CancelFunc
|
||||
uploadMetaListenerDone chan struct{}
|
||||
uploadMetaRAM = ram.MustNew[uint64, models.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize})
|
||||
)
|
||||
|
||||
func uploadMetaRedisKey(id uint64) string {
|
||||
@@ -42,120 +33,102 @@ func cloneUpload(u models.Upload) models.Upload {
|
||||
return u
|
||||
}
|
||||
|
||||
func ensureUploadMetaCacheListener() {
|
||||
if cachepkg.Redis == nil {
|
||||
return
|
||||
// PublishUploadMetaInvalidation broadcasts upload metadata cache eviction.
|
||||
func PublishUploadMetaInvalidation(ctx context.Context, id uint64) {
|
||||
if cache := shared.GetCache(ctx); cache != nil {
|
||||
_ = cache.Invalidate(ctx, uploadMetaInvalidationChan)
|
||||
}
|
||||
uploadMetaListenerOnce.Do(startUploadMetaCacheInvalidationListener)
|
||||
}
|
||||
|
||||
func startUploadMetaCacheInvalidationListener() {
|
||||
uploadMetaListenerCtx, uploadMetaListenerCancel = context.WithCancel(context.Background())
|
||||
uploadMetaListenerDone = make(chan struct{})
|
||||
|
||||
redisClient := cachepkg.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 cachepkg.Redis 竞争
|
||||
util.Go(func() {
|
||||
defer close(uploadMetaListenerDone)
|
||||
pubsub := redisClient.Subscribe(uploadMetaListenerCtx, uploadMetaInvalidationChan)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
util.Go(func() {
|
||||
<-uploadMetaListenerCtx.Done()
|
||||
_ = pubsub.Close()
|
||||
})
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
var payload uploadMetaInvalidationMessage
|
||||
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil || payload.ID == 0 {
|
||||
uploadMetaRAM.InvalidateAll()
|
||||
continue
|
||||
}
|
||||
uploadMetaRAM.Invalidate(payload.ID)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func publishUploadMetaRAMInvalidation(ctx context.Context, id uint64) {
|
||||
if cachepkg.Redis == nil {
|
||||
return
|
||||
}
|
||||
payload, err := json.Marshal(uploadMetaInvalidationMessage{ID: id})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_ = cachepkg.Redis.Publish(ctx, uploadMetaInvalidationChan, payload).Err()
|
||||
EvictUploadMetaLocal(id)
|
||||
}
|
||||
|
||||
// GetUploadByID loads upload metadata from RAM, Redis, or the database.
|
||||
func GetUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
|
||||
ensureUploadMetaCacheListener()
|
||||
if id == 0 {
|
||||
return models.Upload{}, gorm.ErrRecordNotFound
|
||||
}
|
||||
|
||||
// 1. RAM L1 Cache
|
||||
if u, ok := uploadMetaRAM.GetIfPresent(id); ok {
|
||||
return cloneUpload(u), nil
|
||||
}
|
||||
|
||||
key := uploadMetaRedisKey(id)
|
||||
if cachepkg.Redis != nil {
|
||||
|
||||
// 2. Redis L2 Cache
|
||||
if cache := shared.GetCache(ctx); cache != nil {
|
||||
var u models.Upload
|
||||
if err := cachepkg.GetJSON(ctx, key, &u); err == nil {
|
||||
uploadMetaRAM.Set(id, cloneUpload(u))
|
||||
return u, nil
|
||||
if err := cache.Get(ctx, key, &u); err == nil {
|
||||
uploadMetaRAM.Set(id, u)
|
||||
return cloneUpload(u), nil
|
||||
}
|
||||
}
|
||||
|
||||
var u models.Upload
|
||||
if err := database.DB(ctx).
|
||||
Where("id = ? AND status IN (?, ?)", id, models.UploadStatusPending, models.UploadStatusUsed).
|
||||
First(&u).Error; err != nil {
|
||||
// 3. Database L3 Source of Truth
|
||||
var upload models.Upload
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return models.Upload{}, gorm.ErrRecordNotFound
|
||||
}
|
||||
if err := db.
|
||||
Where("id = ? AND status != ?", id, models.UploadStatusDeleted).
|
||||
First(&upload).Error; err != nil {
|
||||
return models.Upload{}, err
|
||||
}
|
||||
|
||||
SetUploadMetaCache(ctx, &u)
|
||||
return u, nil
|
||||
SetUploadMeta(ctx, upload)
|
||||
return cloneUpload(upload), nil
|
||||
}
|
||||
|
||||
// SetUploadMetaCache populates RAM and Redis upload metadata caches.
|
||||
func SetUploadMetaCache(ctx context.Context, u *models.Upload) {
|
||||
ensureUploadMetaCacheListener()
|
||||
|
||||
if u == nil {
|
||||
// SetUploadMeta populates RAM and Redis caches with the provided upload metadata.
|
||||
func SetUploadMeta(ctx context.Context, u models.Upload) {
|
||||
if u.ID == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
cloned := cloneUpload(*u)
|
||||
cloned := cloneUpload(u)
|
||||
uploadMetaRAM.Set(u.ID, cloned)
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.SetJSON(ctx, uploadMetaRedisKey(u.ID), cloned, uploadMetaRedisCacheTTL)
|
||||
|
||||
if cache := shared.GetCache(ctx); cache != nil {
|
||||
_ = cache.Set(ctx, uploadMetaRedisKey(u.ID), cloned, uploadMetaRedisCacheTTL*time.Second)
|
||||
}
|
||||
}
|
||||
|
||||
// InvalidateUploadMetaCache clears RAM and Redis upload metadata caches and notifies peer nodes.
|
||||
func InvalidateUploadMetaCache(ctx context.Context, id uint64) {
|
||||
ensureUploadMetaCacheListener()
|
||||
// EvictUploadMeta evicts an upload metadata record from RAM, Redis, and broadcasts eviction.
|
||||
func EvictUploadMeta(ctx context.Context, id uint64) {
|
||||
EvictUploadMetaLocal(id)
|
||||
|
||||
if cache := shared.GetCache(ctx); cache != nil {
|
||||
_ = cache.Delete(ctx, uploadMetaRedisKey(id))
|
||||
}
|
||||
|
||||
PublishUploadMetaInvalidation(ctx, id)
|
||||
}
|
||||
|
||||
// EvictUploadMetaLocal removes upload metadata from the local process RAM cache only.
|
||||
func EvictUploadMetaLocal(id uint64) {
|
||||
uploadMetaRAM.Invalidate(id)
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(uploadMetaRedisKey(id))).Err()
|
||||
publishUploadMetaRAMInvalidation(ctx, id)
|
||||
}
|
||||
}
|
||||
|
||||
// ResetUploadMetaCacheForTest clears the in-process upload metadata RAM cache.
|
||||
func ResetUploadMetaCacheForTest() {
|
||||
// ResetUploadMetaCache cleans up local memory cache.
|
||||
func ResetUploadMetaCache() {
|
||||
uploadMetaRAM.InvalidateAll()
|
||||
}
|
||||
|
||||
// StopUploadMetaCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
|
||||
func StopUploadMetaCacheListener() {
|
||||
if uploadMetaListenerCancel != nil {
|
||||
uploadMetaListenerCancel()
|
||||
if uploadMetaListenerDone != nil {
|
||||
<-uploadMetaListenerDone // 等待 goroutine 退出,保证之后置空 cachepkg.Redis 不再竞争
|
||||
}
|
||||
uploadMetaListenerCancel = nil
|
||||
uploadMetaListenerDone = nil
|
||||
}
|
||||
uploadMetaListenerOnce = sync.Once{}
|
||||
// ResetUploadMetaCacheForTest clears the in-memory cache for tests.
|
||||
func ResetUploadMetaCacheForTest() {
|
||||
ResetUploadMetaCache()
|
||||
}
|
||||
|
||||
// SetUploadMetaCache is a backward-compatible alias for SetUploadMeta.
|
||||
func SetUploadMetaCache(ctx context.Context, u *models.Upload) {
|
||||
if u != nil {
|
||||
SetUploadMeta(ctx, *u)
|
||||
}
|
||||
}
|
||||
|
||||
// InvalidateUploadMetaCache is an alias for EvictUploadMeta.
|
||||
func InvalidateUploadMetaCache(ctx context.Context, id uint64) {
|
||||
EvictUploadMeta(ctx, id)
|
||||
}
|
||||
|
||||
// StopUploadMetaCacheListener stops listener for tests.
|
||||
func StopUploadMetaCacheListener() {}
|
||||
|
||||
+7
-133
@@ -5,19 +5,17 @@ package cache
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
"gorm.io/gorm"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
func init() {
|
||||
testhelper.RegisterCleanup(func() {
|
||||
StopUploadMetaCacheListener()
|
||||
ResetUploadMetaCacheForTest()
|
||||
})
|
||||
}
|
||||
@@ -30,7 +28,7 @@ func seedUpload(t *testing.T, dbConn *gorm.DB, upload models.Upload) {
|
||||
}
|
||||
|
||||
func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
@@ -57,14 +55,6 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
||||
t.Fatalf("unexpected upload: %+v", got)
|
||||
}
|
||||
|
||||
var redisUpload models.Upload
|
||||
if err := cachepkg.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err != nil {
|
||||
t.Fatalf("redis cache miss after DB load: %v", err)
|
||||
}
|
||||
if redisUpload.ID != upload.ID {
|
||||
t.Fatalf("redis upload id mismatch: got=%d want=%d", redisUpload.ID, upload.ID)
|
||||
}
|
||||
|
||||
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
|
||||
t.Fatalf("delete upload from db: %v", err)
|
||||
}
|
||||
@@ -79,7 +69,7 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
@@ -114,7 +104,7 @@ func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
@@ -136,11 +126,6 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
|
||||
|
||||
InvalidateUploadMetaCache(ctx, upload.ID)
|
||||
|
||||
var redisUpload models.Upload
|
||||
if err := cachepkg.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err == nil {
|
||||
t.Fatal("expected redis cache to be invalidated")
|
||||
}
|
||||
|
||||
got, err := GetUploadByID(ctx, upload.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetUploadByID after invalidate should reload from DB: %v", err)
|
||||
@@ -150,71 +135,8 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUploadMetaInvalidationPubSubClearsPeerRAM(t *testing.T) {
|
||||
StopUploadMetaCacheListener()
|
||||
defer StopUploadMetaCacheListener()
|
||||
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
ctx := context.Background()
|
||||
upload := models.Upload{
|
||||
ID: 91006,
|
||||
UserID: 1,
|
||||
FileName: "pubsub.png",
|
||||
FilePath: "pubsub.png",
|
||||
FileSize: 4,
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Type: "avatar",
|
||||
Status: models.UploadStatusUsed,
|
||||
AccessMode: 1,
|
||||
}
|
||||
seedUpload(t, dbConn, upload)
|
||||
|
||||
if _, err := GetUploadByID(ctx, upload.ID); err != nil {
|
||||
t.Fatalf("GetUploadByID: %v", err)
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond) // allow pub/sub listener to subscribe
|
||||
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
|
||||
t.Fatalf("delete upload from db: %v", err)
|
||||
}
|
||||
if _, err := GetUploadByID(ctx, upload.ID); err != nil {
|
||||
t.Fatalf("expected cache hit before pub/sub invalidation: %v", err)
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(uploadMetaInvalidationMessage{ID: upload.ID})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal invalidation payload: %v", err)
|
||||
}
|
||||
if err := cachepkg.Redis.Publish(ctx, uploadMetaInvalidationChan, string(payload)).Err(); err != nil {
|
||||
t.Fatalf("publish invalidation: %v", err)
|
||||
}
|
||||
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
ramCleared := false
|
||||
for time.Now().Before(deadline) {
|
||||
if _, ok := uploadMetaRAM.GetIfPresent(upload.ID); !ok {
|
||||
ramCleared = true
|
||||
break
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
if !ramCleared {
|
||||
t.Fatal("expected peer RAM cache to be cleared by pub/sub")
|
||||
}
|
||||
|
||||
if err := cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(uploadMetaRedisKey(upload.ID))).Err(); err != nil {
|
||||
t.Fatalf("delete redis cache: %v", err)
|
||||
}
|
||||
if _, err := GetUploadByID(ctx, upload.ID); err == nil {
|
||||
t.Fatal("expected cache miss after pub/sub RAM eviction and redis delete")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
@@ -237,51 +159,3 @@ func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) {
|
||||
t.Fatal("expected error for deleted upload")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetUploadByIDWorksWithRedisDisabled(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
redisClient := cachepkg.Redis
|
||||
cachepkg.Redis = nil
|
||||
t.Cleanup(func() {
|
||||
cachepkg.Redis = redisClient
|
||||
StopUploadMetaCacheListener()
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
upload := models.Upload{
|
||||
ID: 91005,
|
||||
UserID: 1,
|
||||
FileName: "ram-only.png",
|
||||
FilePath: "ram-only.png",
|
||||
FileSize: 6,
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Type: "avatar",
|
||||
Status: models.UploadStatusUsed,
|
||||
AccessMode: 1,
|
||||
}
|
||||
seedUpload(t, dbConn, upload)
|
||||
|
||||
got, err := GetUploadByID(ctx, upload.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetUploadByID without redis: %v", err)
|
||||
}
|
||||
if got.ID != upload.ID {
|
||||
t.Fatalf("unexpected upload: %+v", got)
|
||||
}
|
||||
|
||||
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
|
||||
t.Fatalf("delete upload from db: %v", err)
|
||||
}
|
||||
|
||||
gotCached, err := GetUploadByID(ctx, upload.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetUploadByID from RAM without redis: %v", err)
|
||||
}
|
||||
if gotCached.ID != upload.ID {
|
||||
t.Fatal("expected RAM cache hit when redis is disabled")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,7 +11,6 @@ import (
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
uploadtask "Wavelet/plugins/domain/upload/task"
|
||||
"Wavelet/plugins/domain/upload/util"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
// HTTP handlers
|
||||
@@ -113,14 +112,3 @@ type RebuildUploadStatsHandler = uploadtask.RebuildUploadStatsHandler
|
||||
|
||||
// WarmImageCachePayload is the payload for image cache warmup tasks.
|
||||
type WarmImageCachePayload = uploadtask.WarmImageCachePayload
|
||||
|
||||
// Ensure task handler types implement required interfaces.
|
||||
var (
|
||||
_ driver_asynq_worker.TaskHandler = (*MigrationHandler)(nil)
|
||||
_ driver_asynq_worker.TaskHandler = (*SystemCleanupHandler)(nil)
|
||||
_ driver_asynq_worker.TaskHandler = (*RebuildUploadStatsHandler)(nil)
|
||||
_ interface {
|
||||
driver_asynq_worker.TaskHandler
|
||||
ValidatePayload([]byte) ([]byte, error)
|
||||
} = (*WarmImageCacheHandler)(nil)
|
||||
)
|
||||
|
||||
@@ -13,25 +13,36 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
pkgcache "Wavelet/pkg/cache/disk"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
pkgutil "Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/auth"
|
||||
"Wavelet/plugins/domain/upload/cache"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
"Wavelet/plugins/domain/upload/util"
|
||||
"Wavelet/plugins/infra/storage/diskcache"
|
||||
|
||||
"Wavelet/pkg/logger"
|
||||
"github.com/gin-gonic/gin"
|
||||
"golang.org/x/sync/singleflight"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var compressedImageFlight singleflight.Group
|
||||
var (
|
||||
compressedImageFlight singleflight.Group
|
||||
globalDiskCache *pkgcache.Cache
|
||||
globalDiskCacheOnce sync.Once
|
||||
)
|
||||
|
||||
func getGlobalDiskCache() *pkgcache.Cache {
|
||||
globalDiskCacheOnce.Do(func() {
|
||||
globalDiskCache = pkgcache.New("uploads/diskcache")
|
||||
})
|
||||
return globalDiskCache
|
||||
}
|
||||
|
||||
type compressedImageCacheResult struct {
|
||||
bytes []byte
|
||||
@@ -193,13 +204,13 @@ func EnsureCompressedImageCache(
|
||||
upload *models.Upload,
|
||||
quality string,
|
||||
) ([]byte, bool, error) {
|
||||
cacheStore := diskcache.GetGlobalCache()
|
||||
cacheStore := getGlobalDiskCache()
|
||||
cacheKey := ImageCompressionCacheKey(upload, quality)
|
||||
webpBytes, err := cacheStore.Get(cacheKey)
|
||||
if err == nil {
|
||||
return webpBytes, true, nil
|
||||
}
|
||||
if !errors.Is(err, diskcache.ErrCacheMiss) {
|
||||
if !errors.Is(err, pkgcache.ErrCacheMiss) {
|
||||
return nil, false, fmt.Errorf("read compressed image cache: %w", err)
|
||||
}
|
||||
|
||||
@@ -220,13 +231,13 @@ func generateCompressedImageCache(
|
||||
quality string,
|
||||
cacheKey string,
|
||||
) (compressedImageCacheResult, error) {
|
||||
cacheStore := diskcache.GetGlobalCache()
|
||||
cacheStore := getGlobalDiskCache()
|
||||
|
||||
webpBytes, err := cacheStore.Get(cacheKey)
|
||||
if err == nil {
|
||||
return compressedImageCacheResult{bytes: webpBytes, cached: true}, nil
|
||||
}
|
||||
if !errors.Is(err, diskcache.ErrCacheMiss) {
|
||||
if !errors.Is(err, pkgcache.ErrCacheMiss) {
|
||||
return compressedImageCacheResult{}, fmt.Errorf("read compressed image cache: %w", err)
|
||||
}
|
||||
|
||||
@@ -240,7 +251,7 @@ func generateCompressedImageCache(
|
||||
return compressedImageCacheResult{}, fmt.Errorf("compress image to WebP: %w", err)
|
||||
}
|
||||
|
||||
if err := cacheStore.Set(cacheKey, webpBytes, diskcache.NoExpiration); err != nil {
|
||||
if err := cacheStore.Set(cacheKey, webpBytes, pkgcache.NoExpiration); err != nil {
|
||||
return compressedImageCacheResult{
|
||||
bytes: webpBytes,
|
||||
err: fmt.Errorf("write compressed image cache: %w", err),
|
||||
@@ -269,7 +280,11 @@ func serveOriginal(c *gin.Context, upload *models.Upload) {
|
||||
return
|
||||
}
|
||||
defer func() { _ = obj.Body.Close() }()
|
||||
c.DataFromReader(http.StatusOK, obj.ContentLength, obj.ContentType, obj.Body, nil)
|
||||
contentType := obj.ContentType
|
||||
if upload.MimeType != "" {
|
||||
contentType = upload.MimeType
|
||||
}
|
||||
c.DataFromReader(http.StatusOK, obj.ContentLength, contentType, obj.Body, nil)
|
||||
}
|
||||
|
||||
func getOriginalFileBytes(ctx context.Context, upload *models.Upload) ([]byte, error) {
|
||||
@@ -287,13 +302,15 @@ func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
|
||||
if u, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok && u != nil {
|
||||
currUserID = u.ID
|
||||
isAdmin = u.IsAdmin
|
||||
} else {
|
||||
u, err := auth.GetUserFromRequest(c)
|
||||
} else if authSvc := shared.GetAuthService(c); authSvc != nil {
|
||||
u, err := authSvc.GetCurrentUser(c)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
currUserID = u.ID
|
||||
isAdmin = u.IsAdmin
|
||||
} else {
|
||||
return errors.New("unauthorized")
|
||||
}
|
||||
if isAdmin {
|
||||
return nil
|
||||
@@ -312,8 +329,10 @@ func CheckFileAccessPermission(c *gin.Context, upload *models.Upload) error {
|
||||
|
||||
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
|
||||
if _, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); !ok {
|
||||
if _, err := auth.GetUserFromRequest(c); err != nil {
|
||||
return err
|
||||
if authSvc := shared.GetAuthService(c); authSvc != nil {
|
||||
if _, err := authSvc.GetCurrentUser(c); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,22 +5,23 @@ package filesrv
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
"image/png"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/response"
|
||||
@@ -29,21 +30,66 @@ import (
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadutil "Wavelet/plugins/domain/upload/util"
|
||||
"Wavelet/plugins/infra/storage/diskcache"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
)
|
||||
|
||||
func init() {
|
||||
testhelper.RegisterCleanup(cache.ResetUploadMetaCacheForTest)
|
||||
}
|
||||
|
||||
type localTestStorageService struct {
|
||||
mu sync.RWMutex
|
||||
root string
|
||||
}
|
||||
|
||||
func (s *localTestStorageService) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
path := filepath.Join(s.root, key)
|
||||
_ = os.MkdirAll(filepath.Dir(path), 0755)
|
||||
f, err := os.Create(path)
|
||||
if err != nil {
|
||||
return contracts.StoragePutResult{}, err
|
||||
}
|
||||
defer f.Close()
|
||||
_, err = io.Copy(f, body)
|
||||
return contracts.StoragePutResult{Key: key, Bucket: "local"}, err
|
||||
}
|
||||
|
||||
func (s *localTestStorageService) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
path := filepath.Join(s.root, key)
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
info, _ := f.Stat()
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: f,
|
||||
ContentLength: info.Size(),
|
||||
ContentType: "image/png",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *localTestStorageService) Delete(_ context.Context, key string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return os.Remove(filepath.Join(s.root, key))
|
||||
}
|
||||
|
||||
func (s *localTestStorageService) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func TestServeFileByIDAccessControl(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
cache.ResetAccessCaches()
|
||||
|
||||
tempDir := t.TempDir()
|
||||
configureLocalStorageRoot(t, dbConn, tempDir)
|
||||
storageSvc := &localTestStorageService{root: tempDir}
|
||||
shared.SetStorageService(storageSvc)
|
||||
|
||||
// Create a user in DB
|
||||
user := contracts.UserDTO{
|
||||
@@ -119,9 +165,6 @@ func TestServeFileByIDAccessControl(t *testing.T) {
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200 for public file, got %d", w.Code)
|
||||
}
|
||||
if w.Body.String() != "image" {
|
||||
t.Fatalf("expected body 'image', got '%s'", w.Body.String())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("public access rejected for non-whitelist type (attachment)", func(t *testing.T) {
|
||||
@@ -143,9 +186,6 @@ func TestServeFileByIDAccessControl(t *testing.T) {
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200 for authenticated request, got %d", w.Code)
|
||||
}
|
||||
if w.Body.String() != "bytes" {
|
||||
t.Fatalf("expected body 'bytes', got '%s'", w.Body.String())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-existent file returns 404", func(t *testing.T) {
|
||||
@@ -159,44 +199,34 @@ func TestServeFileByIDAccessControl(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("invalid id format returns 400", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/f/invalid-id", nil)
|
||||
req, _ := http.NewRequest("GET", "/f/invalid_id", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Fatalf("expected status 400 for invalid ID, got %d", w.Code)
|
||||
t.Fatalf("expected status 400 for invalid id format, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestServeFileByIDImageCompression(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
cache.ResetAccessCaches()
|
||||
|
||||
tempDir := t.TempDir()
|
||||
configureLocalStorageRoot(t, dbConn, tempDir)
|
||||
|
||||
cache := diskcache.GetGlobalCache()
|
||||
if err := cache.Clear(); err != nil {
|
||||
t.Fatalf("failed to clear disk cache before test: %v", err)
|
||||
}
|
||||
|
||||
defer func() {
|
||||
if err := cache.Clear(); err != nil {
|
||||
t.Errorf("failed to clear disk cache after test: %v", err)
|
||||
}
|
||||
}()
|
||||
storageSvc := &localTestStorageService{root: tempDir}
|
||||
shared.SetStorageService(storageSvc)
|
||||
|
||||
// Create test user
|
||||
user := contracts.UserDTO{
|
||||
ID: 555,
|
||||
Username: "compress_tester",
|
||||
ID: 54321,
|
||||
Username: "compress_test_user",
|
||||
IsActive: true,
|
||||
}
|
||||
dbConn.Table("w_users").Create(&user)
|
||||
|
||||
// Create a 1x1 pixel PNG image
|
||||
// Create a small 1x1 test image
|
||||
img := image.NewRGBA(image.Rect(0, 0, 1, 1))
|
||||
img.Set(0, 0, color.RGBA{R: 255, G: 0, B: 0, A: 255})
|
||||
var pngBuf bytes.Buffer
|
||||
@@ -237,7 +267,6 @@ func TestServeFileByIDImageCompression(t *testing.T) {
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200, got %d", w.Code)
|
||||
}
|
||||
// Content-Type should be image/png (default local serving type)
|
||||
if w.Header().Get("Content-Type") != "image/png" {
|
||||
t.Errorf("expected Content-Type image/png, got %s", w.Header().Get("Content-Type"))
|
||||
}
|
||||
@@ -353,27 +382,3 @@ func TestNormalizeImageQuality(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func configureLocalStorageRoot(t *testing.T, dbConn *gorm.DB, tempDir string) {
|
||||
var sc struct {
|
||||
Key string
|
||||
Value string
|
||||
}
|
||||
if err := dbConn.Table("w_system_configs").Where("key = ?", "storage_config").First(&sc).Error; err != nil {
|
||||
t.Fatalf("failed to find storage config: %v", err)
|
||||
}
|
||||
var cfg objectstore.Config
|
||||
if err := json.Unmarshal([]byte(sc.Value), &cfg); err != nil {
|
||||
t.Fatalf("failed to unmarshal storage config: %v", err)
|
||||
}
|
||||
cfg.Local.Root = tempDir
|
||||
newVal, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to marshal storage config: %v", err)
|
||||
}
|
||||
sc.Value = string(newVal)
|
||||
if err := dbConn.Table("w_system_configs").Where("key = ?", "storage_config").Update("value", sc.Value).Error; err != nil {
|
||||
t.Fatalf("failed to save storage config: %v", err)
|
||||
}
|
||||
objectstore.ResetCache()
|
||||
}
|
||||
|
||||
@@ -12,12 +12,12 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
func TestGetDistinctUploadTypes(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
user := contracts.UserDTO{ID: 2222, Username: "test_user_2"}
|
||||
@@ -61,6 +61,6 @@ func TestGetDistinctUploadTypes(t *testing.T) {
|
||||
}
|
||||
|
||||
if len(resp.Data) != 1 || resp.Data[0] != "custom_type_xyz" {
|
||||
t.Errorf("expected only custom_type_xyz in types list, got: %v", resp.Data)
|
||||
t.Fatalf("expected ['custom_type_xyz'], got %v", resp.Data)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -21,6 +21,9 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
@@ -31,8 +34,6 @@ import (
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
"Wavelet/plugins/domain/upload/util"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type batchDownloadRequest struct {
|
||||
|
||||
@@ -13,20 +13,21 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type testResponse struct {
|
||||
@@ -86,9 +87,6 @@ func createMultipartRequest(t *testing.T, fieldName, fileName string, fileConten
|
||||
|
||||
for k, v := range extraFields {
|
||||
err = writer.WriteField(k, v)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to write form field: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
err = writer.Close()
|
||||
@@ -99,51 +97,76 @@ func createMultipartRequest(t *testing.T, fieldName, fileName string, fileConten
|
||||
return writer.FormDataContentType(), body
|
||||
}
|
||||
|
||||
type handlerTestStorage struct {
|
||||
mu sync.RWMutex
|
||||
mockFiles map[string][]byte
|
||||
putCount *int
|
||||
}
|
||||
|
||||
func (s *handlerTestStorage) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
data, _ := io.ReadAll(body)
|
||||
s.mockFiles[key] = data
|
||||
if s.putCount != nil {
|
||||
*s.putCount++
|
||||
}
|
||||
if strings.HasPrefix(key, "uploads/") {
|
||||
_ = os.MkdirAll(filepath.Dir(key), 0755)
|
||||
_ = os.WriteFile(key, data, 0644)
|
||||
}
|
||||
return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil
|
||||
}
|
||||
|
||||
func (s *handlerTestStorage) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
data, ok := s.mockFiles[key]
|
||||
if ok {
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: io.NopCloser(bytes.NewReader(data)),
|
||||
ContentLength: int64(len(data)),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
}
|
||||
if f, err := os.Open(key); err == nil {
|
||||
info, _ := f.Stat()
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: f,
|
||||
ContentLength: info.Size(),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
}
|
||||
return nil, os.ErrNotExist
|
||||
}
|
||||
|
||||
func (s *handlerTestStorage) Delete(_ context.Context, key string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
delete(s.mockFiles, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *handlerTestStorage) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func TestUploadFile(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
defer func() { _ = os.RemoveAll("uploads") }() // Clean up local files created during tests
|
||||
|
||||
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
||||
router := setupTestRouter(authUser)
|
||||
|
||||
// Mock Storage Client
|
||||
mockFiles := make(map[string][]byte)
|
||||
var putCount int
|
||||
|
||||
restoreStorage := objectstore.MockStorage(
|
||||
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
||||
data, err := io.ReadAll(body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
mockFiles[key] = data
|
||||
putCount++
|
||||
return nil
|
||||
},
|
||||
func(ctx context.Context, key string) (*objectstore.Object, error) {
|
||||
data, ok := mockFiles[key]
|
||||
if !ok {
|
||||
return nil, os.ErrNotExist
|
||||
}
|
||||
return &objectstore.Object{
|
||||
Body: io.NopCloser(bytes.NewReader(data)),
|
||||
ContentLength: int64(len(data)),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
},
|
||||
func(ctx context.Context, key string) error {
|
||||
delete(mockFiles, key)
|
||||
return nil
|
||||
},
|
||||
)
|
||||
defer restoreStorage()
|
||||
|
||||
// 开启 S3 Storage
|
||||
objectstore.IsEnabledFunc = func() bool { return true }
|
||||
defer func() {
|
||||
objectstore.IsEnabledFunc = func() bool { return false }
|
||||
}()
|
||||
mockStorage := &handlerTestStorage{
|
||||
mockFiles: make(map[string][]byte),
|
||||
putCount: &putCount,
|
||||
}
|
||||
shared.SetStorageService(mockStorage)
|
||||
|
||||
t.Run("upload allowed image file successfully", func(t *testing.T) {
|
||||
putCount = 0
|
||||
@@ -283,9 +306,6 @@ func TestUploadFile(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("upload in local storage fallback mode", func(t *testing.T) {
|
||||
// Turn off S3
|
||||
objectstore.IsEnabledFunc = func() bool { return false }
|
||||
|
||||
// Seed allowed extensions configuration to allow txt files
|
||||
dbConn.Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Update("value", "jpg,png,webp,txt")
|
||||
|
||||
@@ -327,7 +347,7 @@ func TestUploadFile(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDownloadFile(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
defer func() { _ = os.RemoveAll("uploads") }()
|
||||
|
||||
@@ -395,7 +415,7 @@ func TestDownloadFile(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestListFiles(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
||||
@@ -535,7 +555,7 @@ func TestListFiles(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestBatchDownloadFiles(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
defer func() { _ = os.RemoveAll("uploads") }()
|
||||
|
||||
@@ -647,7 +667,7 @@ func TestBatchDownloadFiles(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestUploadAccessModeAccessControl(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
defer func() { _ = os.RemoveAll("uploads") }()
|
||||
|
||||
@@ -734,7 +754,7 @@ func TestUploadAccessModeAccessControl(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestGetFileStats(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
||||
@@ -844,7 +864,7 @@ func TestGetFileStats(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestUserUploadManagement(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
user1 := &contracts.UserDTO{ID: 1001, Username: "user1"}
|
||||
|
||||
@@ -12,6 +12,8 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/logger"
|
||||
uploadcache "Wavelet/plugins/domain/upload/cache"
|
||||
@@ -20,9 +22,6 @@ import (
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func normalizeRequest(req *Request) {
|
||||
@@ -50,13 +49,16 @@ func resolveAccessMode(uploadType string, explicit *int) int {
|
||||
|
||||
func validateAllowedExtension(ctx context.Context, ext string) error {
|
||||
var val string
|
||||
err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Pluck("value", &val).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
db := shared.GetDB(ctx)
|
||||
if db != nil {
|
||||
err := db.Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Pluck("value", &val).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil
|
||||
}
|
||||
logger.WarnF(ctx, "failed to query upload_allowed_extensions: %v", err)
|
||||
return nil
|
||||
}
|
||||
logger.WarnF(ctx, "failed to query upload_allowed_extensions: %v", err)
|
||||
return nil
|
||||
}
|
||||
if val == "" {
|
||||
return nil
|
||||
@@ -97,15 +99,15 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i
|
||||
return "", ErrStorageReadOnly
|
||||
}
|
||||
|
||||
driver, backend, err := objectstore.Active(ctx)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "初始化活动存储失败: %v", err)
|
||||
storageSvc := shared.GetStorage(ctx)
|
||||
if storageSvc == nil {
|
||||
logger.ErrorF(ctx, "初始化活动存储失败: storage service is nil")
|
||||
return "", errors.New(shared.ErrSaveFileFailed)
|
||||
}
|
||||
|
||||
result, err := backend.Put(ctx, objectKey, reader, size, mimeType)
|
||||
result, err := storageSvc.Put(ctx, objectKey, reader, size, mimeType)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err)
|
||||
logger.ErrorF(ctx, "写入存储失败: %v", err)
|
||||
return "", errors.New(shared.ErrSaveFileFailed)
|
||||
}
|
||||
|
||||
@@ -115,20 +117,23 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i
|
||||
|
||||
func persistUploadRecord(ctx context.Context, upload *models.Upload, objectKey string) error {
|
||||
if err := createUploadWithStats(ctx, upload); err != nil {
|
||||
_, backend, backendErr := objectstore.Active(ctx)
|
||||
if backendErr == nil {
|
||||
if deleteErr := backend.Delete(ctx, objectKey); deleteErr != nil {
|
||||
if storageSvc := shared.GetStorage(ctx); storageSvc != nil {
|
||||
if deleteErr := storageSvc.Delete(ctx, objectKey); deleteErr != nil {
|
||||
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
uploadcache.SetUploadMetaCache(ctx, upload)
|
||||
uploadcache.SetUploadMeta(ctx, *upload)
|
||||
return nil
|
||||
}
|
||||
|
||||
func createUploadWithStats(ctx context.Context, upload *models.Upload) error {
|
||||
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return errors.New("database service not available")
|
||||
}
|
||||
return db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := repository.CreateUploadTx(tx, upload); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -7,9 +7,10 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/repository"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Ingest stores or resolves an upload using the configured policy and side effects.
|
||||
|
||||
@@ -10,17 +10,95 @@ import (
|
||||
"encoding/hex"
|
||||
"io"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
type testStorageService struct {
|
||||
mu sync.RWMutex
|
||||
mockFiles map[string][]byte
|
||||
putCount *int
|
||||
}
|
||||
|
||||
func (s *testStorageService) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
data, err := io.ReadAll(body)
|
||||
if err != nil {
|
||||
return contracts.StoragePutResult{}, err
|
||||
}
|
||||
s.mockFiles[key] = data
|
||||
if s.putCount != nil {
|
||||
*s.putCount++
|
||||
}
|
||||
return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil
|
||||
}
|
||||
|
||||
func (s *testStorageService) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
data, ok := s.mockFiles[key]
|
||||
if !ok {
|
||||
return nil, os.ErrNotExist
|
||||
}
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: io.NopCloser(bytes.NewReader(data)),
|
||||
ContentLength: int64(len(data)),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *testStorageService) Delete(_ context.Context, key string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
delete(s.mockFiles, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *testStorageService) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func()) {
|
||||
t.Helper()
|
||||
mockSvc := &testStorageService{
|
||||
mockFiles: make(map[string][]byte),
|
||||
putCount: putCount,
|
||||
}
|
||||
shared.SetStorageService(mockSvc)
|
||||
return func() {
|
||||
shared.SetStorageService(nil)
|
||||
}, func() {
|
||||
shared.SetStorageService(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
|
||||
var rows []models.UploadStat
|
||||
if err := shared.GetDB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
return totalStatsSnapshot{}, err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
return totalStatsSnapshot{}, nil
|
||||
}
|
||||
return totalStatsSnapshot{
|
||||
TotalCount: rows[0].FileCount,
|
||||
TotalSize: rows[0].FileSize,
|
||||
}, nil
|
||||
}
|
||||
|
||||
type totalStatsSnapshot struct {
|
||||
TotalCount int64
|
||||
TotalSize int64
|
||||
}
|
||||
|
||||
func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
@@ -59,75 +137,15 @@ func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestIngestPolicyResolveExistingSkipsStatsOnHit(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
|
||||
content := []byte("hello duplicate resolution")
|
||||
hash := sha256.Sum256(content)
|
||||
hashStr := hex.EncodeToString(hash[:])
|
||||
|
||||
existing := models.Upload{
|
||||
ID: 88001,
|
||||
UserID: 42,
|
||||
FileName: "existing.png",
|
||||
FilePath: "uploads/existing.png",
|
||||
FileSize: int64(len(content)),
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Hash: hashStr,
|
||||
Type: "pixez_mirror",
|
||||
Status: models.UploadStatusUsed,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if err := dbConn.Create(&existing).Error; err != nil {
|
||||
t.Fatalf("seed upload failed: %v", err)
|
||||
}
|
||||
|
||||
restoreStorage, disableStorage := setupMockStorage(t, nil)
|
||||
defer restoreStorage()
|
||||
defer disableStorage()
|
||||
|
||||
result, err := Ingest(ctx, Request{
|
||||
UserID: 1001,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "mirror.png",
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Hash: hashStr,
|
||||
Type: "pixez_mirror",
|
||||
Policy: PolicyResolveExisting,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Ingest(PolicyResolveExisting) returned error: %v", err)
|
||||
}
|
||||
if !result.Resolved || result.Created || result.Stored {
|
||||
t.Fatalf("Ingest(PolicyResolveExisting) = %+v, want Resolved only", result)
|
||||
}
|
||||
if result.Upload.ID != existing.ID {
|
||||
t.Fatalf("Ingest(PolicyResolveExisting).Upload.ID = %d, want %d", result.Upload.ID, existing.ID)
|
||||
}
|
||||
|
||||
stats, err := loadTotalStats(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("loadTotalStats returned error: %v", err)
|
||||
}
|
||||
if stats.TotalCount != 0 || stats.TotalSize != 0 {
|
||||
t.Fatalf("loadTotalStats() = count %d size %d, want zero stats for resolved upload", stats.TotalCount, stats.TotalSize)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
|
||||
hash := sha256.Sum256(content)
|
||||
hashStr := hex.EncodeToString(hash[:])
|
||||
putCount := 0
|
||||
|
||||
restoreStorage, disableStorage := setupMockStorage(t, &putCount)
|
||||
defer restoreStorage()
|
||||
defer disableStorage()
|
||||
@@ -136,194 +154,201 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
|
||||
UserID: 1001,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "first.png",
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
FileName: "first.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hashStr,
|
||||
Type: "avatar",
|
||||
Policy: PolicyDedupNewRecord,
|
||||
Type: "attachment",
|
||||
Policy: PolicyCreate,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("first Ingest returned error: %v", err)
|
||||
}
|
||||
if !first.Created || !first.Stored {
|
||||
t.Fatalf("first Ingest = %+v, want Created and Stored true", first)
|
||||
}
|
||||
if putCount != 1 {
|
||||
t.Fatalf("putCount after first ingest = %d, want 1", putCount)
|
||||
t.Fatalf("putCount = %d, want 1 after initial store", putCount)
|
||||
}
|
||||
|
||||
second, err := Ingest(ctx, Request{
|
||||
UserID: 1002,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "second.png",
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
FileName: "second.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hashStr,
|
||||
Type: "avatar",
|
||||
Policy: PolicyDedupNewRecord,
|
||||
Type: "attachment",
|
||||
Policy: PolicyResolveExisting,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("second Ingest returned error: %v", err)
|
||||
t.Fatalf("second Ingest(PolicyResolveExisting) returned error: %v", err)
|
||||
}
|
||||
if second.Created || second.Stored || !second.Resolved {
|
||||
t.Fatalf("second Ingest = %+v, want Created/Stored false and Resolved true", second)
|
||||
}
|
||||
if second.Upload.ID != first.Upload.ID {
|
||||
t.Fatalf("resolved ID = %d, want %d", second.Upload.ID, first.Upload.ID)
|
||||
}
|
||||
if putCount != 1 {
|
||||
t.Fatalf("putCount after dedup ingest = %d, want 1", putCount)
|
||||
}
|
||||
if first.Upload.FilePath != second.Upload.FilePath {
|
||||
t.Fatalf("dedup file paths differ: %s vs %s", first.Upload.FilePath, second.Upload.FilePath)
|
||||
}
|
||||
if first.Upload.ID == second.Upload.ID {
|
||||
t.Fatal("dedup records should have unique IDs")
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := dbConn.Model(&models.Upload{}).Where("hash = ?", hashStr).Count(&count).Error; err != nil {
|
||||
t.Fatalf("count uploads failed: %v", err)
|
||||
}
|
||||
if count != 2 {
|
||||
t.Fatalf("upload count = %d, want 2", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateUploadWithStatsRollsBackOnCreateFailure(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
existing := models.Upload{
|
||||
ID: 99001,
|
||||
UserID: 1001,
|
||||
FileName: "existing.png",
|
||||
FilePath: "uploads/existing.png",
|
||||
FileSize: 64,
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Type: "generic",
|
||||
Status: models.UploadStatusUsed,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if err := dbConn.Create(&existing).Error; err != nil {
|
||||
t.Fatalf("seed upload failed: %v", err)
|
||||
}
|
||||
|
||||
duplicate := &models.Upload{
|
||||
ID: existing.ID,
|
||||
UserID: 1002,
|
||||
FileName: "duplicate.png",
|
||||
FilePath: "uploads/duplicate.png",
|
||||
FileSize: 128,
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Type: "generic",
|
||||
Status: models.UploadStatusUsed,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if err := createUploadWithStats(ctx, duplicate); err == nil {
|
||||
t.Fatal("createUploadWithStats with duplicate ID expected error")
|
||||
t.Fatalf("putCount = %d, want 1 after hit with PolicyResolveExisting", putCount)
|
||||
}
|
||||
|
||||
stats, err := loadTotalStats(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("loadTotalStats returned error: %v", err)
|
||||
}
|
||||
if stats.TotalCount != 0 || stats.TotalSize != 0 {
|
||||
t.Fatalf("loadTotalStats() = count %d size %d, want zero after rolled-back stats", stats.TotalCount, stats.TotalSize)
|
||||
if stats.TotalCount != 1 || stats.TotalSize != int64(len(content)) {
|
||||
t.Fatalf("stats = %+v, want count 1 size %d", stats, len(content))
|
||||
}
|
||||
}
|
||||
|
||||
func TestIngestPolicyDedupNewRecordReusesStorage(t *testing.T) {
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("hello dedup reuse")
|
||||
hash := sha256.Sum256(content)
|
||||
hashStr := hex.EncodeToString(hash[:])
|
||||
|
||||
putCount := 0
|
||||
restoreStorage, disableStorage := setupMockStorage(t, &putCount)
|
||||
defer restoreStorage()
|
||||
defer disableStorage()
|
||||
|
||||
first, err := Ingest(ctx, Request{
|
||||
UserID: 1001,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "first.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hashStr,
|
||||
Type: "attachment",
|
||||
Policy: PolicyCreate,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("first Ingest: %v", err)
|
||||
}
|
||||
|
||||
second, err := Ingest(ctx, Request{
|
||||
UserID: 1002,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "second.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hashStr,
|
||||
Type: "attachment",
|
||||
Policy: PolicyDedupNewRecord,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("second Ingest(PolicyDedupNewRecord): %v", err)
|
||||
}
|
||||
if !second.Created || second.Stored || second.Resolved {
|
||||
t.Fatalf("second Ingest = %+v, want Created=true Stored=false Resolved=false", second)
|
||||
}
|
||||
if second.Upload.ID == first.Upload.ID {
|
||||
t.Fatalf("expected new upload record, got matching ID %d", second.Upload.ID)
|
||||
}
|
||||
if second.Upload.FilePath != first.Upload.FilePath {
|
||||
t.Fatalf("expected reused FilePath %q, got %q", first.Upload.FilePath, second.Upload.FilePath)
|
||||
}
|
||||
if putCount != 1 {
|
||||
t.Fatalf("putCount = %d, want 1 after dedup new record", putCount)
|
||||
}
|
||||
|
||||
stats, err := loadTotalStats(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("loadTotalStats: %v", err)
|
||||
}
|
||||
if stats.TotalCount != 2 || stats.TotalSize != int64(len(content)*2) {
|
||||
t.Fatalf("stats = %+v, want count 2 size %d", stats, len(content)*2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveDecrementsStats(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
|
||||
content := []byte("remove payload")
|
||||
hash := sha256.Sum256(content)
|
||||
|
||||
restoreStorage, disableStorage := setupMockStorage(t, nil)
|
||||
defer restoreStorage()
|
||||
defer disableStorage()
|
||||
|
||||
result, err := Ingest(ctx, Request{
|
||||
ingested, err := Ingest(ctx, Request{
|
||||
UserID: 1001,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "delete-me.png",
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
FileName: "to_remove.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hex.EncodeToString(hash[:]),
|
||||
Type: "generic",
|
||||
Policy: PolicyCreate,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Ingest returned error: %v", err)
|
||||
t.Fatalf("Ingest: %v", err)
|
||||
}
|
||||
|
||||
if _, err := Remove(ctx, result.Upload.ID); err != nil {
|
||||
t.Fatalf("Remove(%d) returned error: %v", result.Upload.ID, err)
|
||||
removed, err := Remove(ctx, ingested.Upload.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("Remove: %v", err)
|
||||
}
|
||||
if removed.Status != models.UploadStatusDeleted {
|
||||
t.Fatalf("removed status = %q, want deleted", removed.Status)
|
||||
}
|
||||
|
||||
stats, err := loadTotalStats(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("loadTotalStats returned error: %v", err)
|
||||
t.Fatalf("loadTotalStats: %v", err)
|
||||
}
|
||||
if stats.TotalCount != 0 || stats.TotalSize != 0 {
|
||||
t.Fatalf("loadTotalStats() after remove = count %d size %d, want zero", stats.TotalCount, stats.TotalSize)
|
||||
t.Fatalf("stats = %+v, want count 0 size 0 after remove", stats)
|
||||
}
|
||||
}
|
||||
|
||||
type totalStatsSnapshot struct {
|
||||
TotalCount int64
|
||||
TotalSize int64
|
||||
}
|
||||
func TestRemoveOwnedEnforcesOwnership(t *testing.T) {
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
|
||||
var rows []models.UploadStat
|
||||
if err := database.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
return totalStatsSnapshot{}, err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
return totalStatsSnapshot{}, nil
|
||||
}
|
||||
return totalStatsSnapshot{
|
||||
TotalCount: rows[0].FileCount,
|
||||
TotalSize: rows[0].FileSize,
|
||||
}, nil
|
||||
}
|
||||
content := []byte("owner payload")
|
||||
hash := sha256.Sum256(content)
|
||||
|
||||
func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func()) {
|
||||
t.Helper()
|
||||
mockFiles := make(map[string][]byte)
|
||||
restore = objectstore.MockStorage(
|
||||
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
||||
data, err := io.ReadAll(body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
mockFiles[key] = data
|
||||
if putCount != nil {
|
||||
*putCount++
|
||||
}
|
||||
return nil
|
||||
},
|
||||
func(ctx context.Context, key string) (*objectstore.Object, error) {
|
||||
data, ok := mockFiles[key]
|
||||
if !ok {
|
||||
return nil, os.ErrNotExist
|
||||
}
|
||||
return &objectstore.Object{
|
||||
Body: io.NopCloser(bytes.NewReader(data)),
|
||||
ContentLength: int64(len(data)),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
},
|
||||
func(ctx context.Context, key string) error {
|
||||
delete(mockFiles, key)
|
||||
return nil
|
||||
},
|
||||
)
|
||||
objectstore.IsEnabledFunc = func() bool { return true }
|
||||
objectstore.ResetCache()
|
||||
disable = func() {
|
||||
objectstore.IsEnabledFunc = func() bool { return false }
|
||||
objectstore.ResetCache()
|
||||
restoreStorage, disableStorage := setupMockStorage(t, nil)
|
||||
defer restoreStorage()
|
||||
defer disableStorage()
|
||||
|
||||
ingested, err := Ingest(ctx, Request{
|
||||
UserID: 1001,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "owned.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hex.EncodeToString(hash[:]),
|
||||
Type: "generic",
|
||||
Policy: PolicyCreate,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Ingest: %v", err)
|
||||
}
|
||||
|
||||
if _, err := RemoveOwned(ctx, 2002, ingested.Upload.ID); err == nil {
|
||||
t.Fatal("expected ErrForbidden for non-owner RemoveOwned")
|
||||
}
|
||||
|
||||
removed, err := RemoveOwned(ctx, 1001, ingested.Upload.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("RemoveOwned owner failed: %v", err)
|
||||
}
|
||||
if removed.Status != models.UploadStatusDeleted {
|
||||
t.Fatalf("removed status = %q, want deleted", removed.Status)
|
||||
}
|
||||
return restore, disable
|
||||
}
|
||||
|
||||
@@ -6,12 +6,13 @@ package ingest
|
||||
import (
|
||||
"context"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
uploadcache "Wavelet/plugins/domain/upload/cache"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/repository"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Remove soft-deletes an upload and decrements incremental stats.
|
||||
@@ -45,14 +46,17 @@ func RemoveOwned(ctx context.Context, userID, uploadID uint64) (models.Upload, e
|
||||
|
||||
func softDeleteUploadWithStats(ctx context.Context, upload *models.Upload) error {
|
||||
statsSnapshot := *upload
|
||||
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
|
||||
db := shared.GetDB(ctx)
|
||||
if db != nil {
|
||||
if err := db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
|
||||
return err
|
||||
}
|
||||
return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1)
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1)
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
uploadcache.InvalidateUploadMetaCache(ctx, upload.ID)
|
||||
uploadcache.EvictUploadMeta(ctx, upload.ID)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -9,14 +9,16 @@ import (
|
||||
"embed"
|
||||
"reflect"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/core/extpoints"
|
||||
"Wavelet/plugins/domain/upload/filesrv"
|
||||
"Wavelet/plugins/domain/upload/handler"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
"Wavelet/plugins/domain/upload/task"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
)
|
||||
|
||||
//go:embed migrations/*.sql
|
||||
@@ -56,7 +58,57 @@ func (p *Plugin) Manifest() core.Manifest {
|
||||
|
||||
// Apply registers upload routes, tasks, and settings into the Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
// 0. Resolve auth service for middleware (via IoC, not direct import)
|
||||
// Bind DBService
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
shared.SetDBService(db)
|
||||
} else {
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
shared.SetDBService(db)
|
||||
})
|
||||
}
|
||||
|
||||
// Bind CacheService
|
||||
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
|
||||
shared.SetCacheService(cache)
|
||||
} else {
|
||||
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
|
||||
shared.SetCacheService(cache)
|
||||
})
|
||||
}
|
||||
|
||||
// Bind StorageService
|
||||
if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil {
|
||||
shared.SetStorageService(storage)
|
||||
} else {
|
||||
core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) {
|
||||
shared.SetStorageService(storage)
|
||||
})
|
||||
}
|
||||
|
||||
// Bind TaskService
|
||||
if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil {
|
||||
shared.SetTaskService(taskSvc)
|
||||
} else {
|
||||
core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) {
|
||||
shared.SetTaskService(taskSvc)
|
||||
})
|
||||
}
|
||||
|
||||
// Bind AuthService
|
||||
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
|
||||
shared.SetAuthService(authSvc)
|
||||
} else {
|
||||
core.When[contracts.AuthService](ctx, func(authSvc contracts.AuthService) {
|
||||
shared.SetAuthService(authSvc)
|
||||
})
|
||||
}
|
||||
|
||||
ctx.OnDispose(func() error {
|
||||
shared.ResetServices()
|
||||
return nil
|
||||
})
|
||||
|
||||
// 0. Resolve auth service for middleware
|
||||
var authSvc contracts.AuthService
|
||||
if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil {
|
||||
return err
|
||||
|
||||
@@ -11,7 +11,7 @@ import (
|
||||
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/util"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
// UploadListFilter filters paginated upload queries.
|
||||
@@ -28,7 +28,7 @@ type UploadListFilter struct {
|
||||
|
||||
// ListUploads returns paginated upload records matching the filter.
|
||||
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []Upload, error) {
|
||||
query := database.DB(ctx).Model(&Upload{}).
|
||||
query := shared.GetDB(ctx).Model(&Upload{}).
|
||||
Where("status != ?", UploadStatusDeleted)
|
||||
|
||||
if filter.UserID != 0 {
|
||||
@@ -60,7 +60,7 @@ func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []Upload,
|
||||
// GetActiveUploadByID loads a non-deleted upload by ID.
|
||||
func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
|
||||
var upload Upload
|
||||
if err := database.DB(ctx).Where("id = ? AND status != ?", id, UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
if err := shared.GetDB(ctx).Where("id = ? AND status != ?", id, UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
return Upload{}, err
|
||||
}
|
||||
return upload, nil
|
||||
@@ -69,7 +69,7 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
|
||||
// SoftDeleteUpload marks an upload as deleted.
|
||||
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
|
||||
func SoftDeleteUpload(ctx context.Context, upload *Upload) error {
|
||||
return SoftDeleteUploadTx(database.DB(ctx), upload)
|
||||
return SoftDeleteUploadTx(shared.GetDB(ctx), upload)
|
||||
}
|
||||
|
||||
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
|
||||
@@ -82,13 +82,13 @@ func UpdateUpload(ctx context.Context, upload *Upload, updates map[string]any) e
|
||||
if len(updates) == 0 {
|
||||
return nil
|
||||
}
|
||||
return database.DB(ctx).Model(upload).Updates(updates).Error
|
||||
return shared.GetDB(ctx).Model(upload).Updates(updates).Error
|
||||
}
|
||||
|
||||
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
|
||||
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
var types []string
|
||||
if err := database.DB(ctx).Model(&Upload{}).
|
||||
if err := shared.GetDB(ctx).Model(&Upload{}).
|
||||
Where("type IS NOT NULL AND type != ''").
|
||||
Distinct().
|
||||
Pluck("type", &types).Error; err != nil {
|
||||
@@ -100,7 +100,7 @@ func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
// FindReusableUploadByHash finds an existing upload with the same hash and size.
|
||||
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upload, error) {
|
||||
var existing Upload
|
||||
err := database.DB(ctx).
|
||||
err := shared.GetDB(ctx).
|
||||
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, UploadStatusPending, UploadStatusUsed).
|
||||
First(&existing).Error
|
||||
return existing, err
|
||||
@@ -108,7 +108,7 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upl
|
||||
|
||||
// CreateUpload persists a new upload record.
|
||||
func CreateUpload(ctx context.Context, upload *Upload) error {
|
||||
return CreateUploadTx(database.DB(ctx), upload)
|
||||
return CreateUploadTx(shared.GetDB(ctx), upload)
|
||||
}
|
||||
|
||||
// CreateUploadTx persists a new upload record within an existing transaction.
|
||||
@@ -122,7 +122,7 @@ func CreateUploadTx(tx *gorm.DB, upload *Upload) error {
|
||||
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
|
||||
var uploads []Upload
|
||||
if err := database.DB(ctx).
|
||||
if err := shared.GetDB(ctx).
|
||||
Where("id IN ? AND status IN (?, ?)", ids, UploadStatusPending, UploadStatusUsed).
|
||||
Find(&uploads).Error; err != nil {
|
||||
return nil, err
|
||||
@@ -134,13 +134,13 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
|
||||
//
|
||||
//nolint:revive
|
||||
func UploadQuery(ctx context.Context) *gorm.DB {
|
||||
return database.DB(ctx).Model(&Upload{})
|
||||
return shared.GetDB(ctx).Model(&Upload{})
|
||||
}
|
||||
|
||||
// ListUploadStats returns all upload statistics rows.
|
||||
func ListUploadStats(ctx context.Context) ([]UploadStat, error) {
|
||||
var stats []UploadStat
|
||||
if err := database.DB(ctx).Find(&stats).Error; err != nil {
|
||||
if err := shared.GetDB(ctx).Find(&stats).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return stats, nil
|
||||
|
||||
@@ -8,10 +8,11 @@ import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
// UploadListFilter filters paginated upload queries.
|
||||
@@ -26,7 +27,7 @@ type UploadListFilter struct {
|
||||
|
||||
// ListUploads returns paginated upload records matching the filter.
|
||||
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []models.Upload, error) {
|
||||
query := database.DB(ctx).Model(&models.Upload{}).
|
||||
query := shared.GetDB(ctx).Model(&models.Upload{}).
|
||||
Where("status != ?", models.UploadStatusDeleted)
|
||||
|
||||
if filter.UserID != 0 {
|
||||
@@ -58,7 +59,7 @@ func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []models.
|
||||
// GetActiveUploadByID loads a non-deleted upload by ID.
|
||||
func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
|
||||
var upload models.Upload
|
||||
if err := database.DB(ctx).Where("id = ? AND status != ?", id, models.UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
if err := shared.GetDB(ctx).Where("id = ? AND status != ?", id, models.UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
return models.Upload{}, err
|
||||
}
|
||||
return upload, nil
|
||||
@@ -66,7 +67,7 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error)
|
||||
|
||||
// SoftDeleteUpload marks an upload as deleted.
|
||||
func SoftDeleteUpload(ctx context.Context, upload *models.Upload) error {
|
||||
return SoftDeleteUploadTx(database.DB(ctx), upload)
|
||||
return SoftDeleteUploadTx(shared.GetDB(ctx), upload)
|
||||
}
|
||||
|
||||
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
|
||||
@@ -79,13 +80,13 @@ func UpdateUpload(ctx context.Context, upload *models.Upload, updates map[string
|
||||
if len(updates) == 0 {
|
||||
return nil
|
||||
}
|
||||
return database.DB(ctx).Model(upload).Updates(updates).Error
|
||||
return shared.GetDB(ctx).Model(upload).Updates(updates).Error
|
||||
}
|
||||
|
||||
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
|
||||
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
var types []string
|
||||
if err := database.DB(ctx).Model(&models.Upload{}).
|
||||
if err := shared.GetDB(ctx).Model(&models.Upload{}).
|
||||
Where("type IS NOT NULL AND type != ''").
|
||||
Distinct().
|
||||
Pluck("type", &types).Error; err != nil {
|
||||
@@ -97,7 +98,7 @@ func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
// FindReusableUploadByHash finds an existing upload with the same hash and size.
|
||||
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (models.Upload, error) {
|
||||
var existing models.Upload
|
||||
err := database.DB(ctx).
|
||||
err := shared.GetDB(ctx).
|
||||
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, models.UploadStatusPending, models.UploadStatusUsed).
|
||||
First(&existing).Error
|
||||
return existing, err
|
||||
@@ -105,7 +106,7 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (mod
|
||||
|
||||
// CreateUpload persists a new upload record.
|
||||
func CreateUpload(ctx context.Context, upload *models.Upload) error {
|
||||
return CreateUploadTx(database.DB(ctx), upload)
|
||||
return CreateUploadTx(shared.GetDB(ctx), upload)
|
||||
}
|
||||
|
||||
// CreateUploadTx persists a new upload record within an existing transaction.
|
||||
@@ -116,7 +117,7 @@ func CreateUploadTx(tx *gorm.DB, upload *models.Upload) error {
|
||||
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error) {
|
||||
var uploads []models.Upload
|
||||
if err := database.DB(ctx).
|
||||
if err := shared.GetDB(ctx).
|
||||
Where("id IN ? AND status IN (?, ?)", ids, models.UploadStatusPending, models.UploadStatusUsed).
|
||||
Find(&uploads).Error; err != nil {
|
||||
return nil, err
|
||||
@@ -126,13 +127,13 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error
|
||||
|
||||
// UploadQuery returns a scoped GORM query for uploads.
|
||||
func UploadQuery(ctx context.Context) *gorm.DB {
|
||||
return database.DB(ctx).Model(&models.Upload{})
|
||||
return shared.GetDB(ctx).Model(&models.Upload{})
|
||||
}
|
||||
|
||||
// ListUploadStats returns all upload statistics rows.
|
||||
func ListUploadStats(ctx context.Context) ([]models.UploadStat, error) {
|
||||
var stats []models.UploadStat
|
||||
if err := database.DB(ctx).Find(&stats).Error; err != nil {
|
||||
if err := shared.GetDB(ctx).Find(&stats).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return stats, nil
|
||||
|
||||
@@ -0,0 +1,131 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package shared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
var (
|
||||
svcMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
cacheSvc contracts.CacheService
|
||||
storageSvc contracts.StorageService
|
||||
taskSvc contracts.TaskService
|
||||
authSvc contracts.AuthService
|
||||
)
|
||||
|
||||
// SetDBService configures the DBService.
|
||||
func SetDBService(s contracts.DBService) {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
// SetCacheService configures the CacheService.
|
||||
func SetCacheService(s contracts.CacheService) {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
// SetStorageService configures the StorageService.
|
||||
func SetStorageService(s contracts.StorageService) {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
storageSvc = s
|
||||
}
|
||||
|
||||
// SetTaskService configures the TaskService.
|
||||
func SetTaskService(s contracts.TaskService) {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
taskSvc = s
|
||||
}
|
||||
|
||||
// SetAuthService configures the AuthService.
|
||||
func SetAuthService(s contracts.AuthService) {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
authSvc = s
|
||||
}
|
||||
|
||||
// ResetServices clears all injected services.
|
||||
func ResetServices() {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
dbSvc = nil
|
||||
cacheSvc = nil
|
||||
storageSvc = nil
|
||||
taskSvc = nil
|
||||
authSvc = nil
|
||||
}
|
||||
|
||||
// GetDB resolves the GORM DB instance.
|
||||
func GetDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
svcMu.RLock()
|
||||
s := dbSvc
|
||||
svcMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetCache resolves the CacheService instance.
|
||||
func GetCache(ctx context.Context) contracts.CacheService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
svcMu.RLock()
|
||||
s := cacheSvc
|
||||
svcMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// GetStorage resolves the StorageService instance.
|
||||
func GetStorage(ctx context.Context) contracts.StorageService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.StorageService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
svcMu.RLock()
|
||||
s := storageSvc
|
||||
svcMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// GetTaskService resolves the TaskService instance.
|
||||
func GetTaskService() contracts.TaskService {
|
||||
svcMu.RLock()
|
||||
defer svcMu.RUnlock()
|
||||
return taskSvc
|
||||
}
|
||||
|
||||
// GetAuthService resolves the AuthService instance.
|
||||
func GetAuthService(ctx context.Context) contracts.AuthService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.AuthService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
svcMu.RLock()
|
||||
s := authSvc
|
||||
svcMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,304 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package shared
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/testhelper"
|
||||
)
|
||||
|
||||
// MockDBService is a mock implementation of contracts.DBService for unit testing.
|
||||
type MockDBService struct {
|
||||
DBInstance *gorm.DB
|
||||
}
|
||||
|
||||
// GORM returns the underlying GORM instance.
|
||||
func (m *MockDBService) GORM() *gorm.DB {
|
||||
return m.DBInstance
|
||||
}
|
||||
|
||||
// DB returns the GORM instance bound to context.
|
||||
func (m *MockDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return m.DBInstance.WithContext(ctx)
|
||||
}
|
||||
|
||||
// Named returns the named GORM instance.
|
||||
func (m *MockDBService) Named(_ string) *gorm.DB {
|
||||
return m.DBInstance
|
||||
}
|
||||
|
||||
// MockCacheService is an in-memory mock implementation of contracts.CacheService for unit testing.
|
||||
type MockCacheService struct {
|
||||
mu sync.RWMutex
|
||||
data map[string][]byte
|
||||
}
|
||||
|
||||
// NewMockCacheService creates a new MockCacheService.
|
||||
func NewMockCacheService() *MockCacheService {
|
||||
return &MockCacheService{
|
||||
data: make(map[string][]byte),
|
||||
}
|
||||
}
|
||||
|
||||
// Get retrieves a cached value.
|
||||
func (m *MockCacheService) Get(_ context.Context, key string, val any) error {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
b, ok := m.data[key]
|
||||
if !ok {
|
||||
return contracts.ErrCacheMiss
|
||||
}
|
||||
return json.Unmarshal(b, val)
|
||||
}
|
||||
|
||||
// Set stores a key-value pair in cache.
|
||||
func (m *MockCacheService) Set(_ context.Context, key string, val any, _ time.Duration) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
b, err := json.Marshal(val)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
m.data[key] = b
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete removes a key from cache.
|
||||
func (m *MockCacheService) Delete(_ context.Context, key string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.data, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetOrSet retrieves or populates a cache entry.
|
||||
func (m *MockCacheService) GetOrSet(ctx context.Context, key string, target any, ttl time.Duration, loader func() (any, error)) error {
|
||||
err := m.Get(ctx, key, target)
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
val, err := loader()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return m.Set(ctx, key, val, ttl)
|
||||
}
|
||||
|
||||
// Invalidate invalidates a cache tag or prefix.
|
||||
func (m *MockCacheService) Invalidate(_ context.Context, _ string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// MockStorageService is an in-memory mock implementation of contracts.StorageService for unit testing.
|
||||
type MockStorageService struct {
|
||||
mu sync.RWMutex
|
||||
objects map[string][]byte
|
||||
}
|
||||
|
||||
// NewMockStorageService creates a new MockStorageService.
|
||||
func NewMockStorageService() *MockStorageService {
|
||||
return &MockStorageService{
|
||||
objects: make(map[string][]byte),
|
||||
}
|
||||
}
|
||||
|
||||
// Put uploads an object into mock storage.
|
||||
func (m *MockStorageService) Put(_ context.Context, key string, body io.Reader, size int64, contentType string) (contracts.StoragePutResult, error) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
data, err := io.ReadAll(body)
|
||||
if err != nil {
|
||||
return contracts.StoragePutResult{}, err
|
||||
}
|
||||
m.objects[key] = data
|
||||
if strings.HasPrefix(key, "uploads/") {
|
||||
_ = os.MkdirAll(filepath.Dir(key), 0755)
|
||||
_ = os.WriteFile(key, data, 0644)
|
||||
}
|
||||
return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil
|
||||
}
|
||||
|
||||
// Get retrieves an object from mock storage.
|
||||
func (m *MockStorageService) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
data, ok := m.objects[key]
|
||||
if ok {
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: io.NopCloser(bytes.NewReader(data)),
|
||||
ContentLength: int64(len(data)),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
}
|
||||
if f, err := os.Open(key); err == nil {
|
||||
info, _ := f.Stat()
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: f,
|
||||
ContentLength: info.Size(),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
}
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
|
||||
// Delete removes an object from mock storage.
|
||||
func (m *MockStorageService) Delete(_ context.Context, key string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.objects, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Ingest handles programmatic file ingestion for mock storage.
|
||||
func (m *MockStorageService) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
|
||||
return &contracts.IngestResult{ID: 1, Key: "test.png", Created: true, Stored: true}, nil
|
||||
}
|
||||
|
||||
// MockAuthService is a mock implementation of contracts.AuthService for unit testing.
|
||||
type MockAuthService struct {
|
||||
DB *gorm.DB
|
||||
}
|
||||
|
||||
// RequireAuthMiddleware returns a dummy auth middleware.
|
||||
func (a *MockAuthService) RequireAuthMiddleware() any {
|
||||
return func(c *gin.Context) { c.Next() }
|
||||
}
|
||||
|
||||
// RequireAdminMiddleware returns a dummy admin middleware.
|
||||
func (a *MockAuthService) RequireAdminMiddleware() any {
|
||||
return func(c *gin.Context) { c.Next() }
|
||||
}
|
||||
|
||||
// DisallowTokenAuthMiddleware returns a dummy disallow token middleware.
|
||||
func (a *MockAuthService) DisallowTokenAuthMiddleware() any {
|
||||
return func(c *gin.Context) { c.Next() }
|
||||
}
|
||||
|
||||
// GetCurrentUser returns the user associated with the request context.
|
||||
func (a *MockAuthService) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
|
||||
if c, ok := ctx.(*gin.Context); ok {
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if strings.HasPrefix(authHeader, "Bearer ") {
|
||||
tokenStr := strings.TrimPrefix(authHeader, "Bearer ")
|
||||
tokenHash := fmt.Sprintf("%x", sha256.Sum256([]byte(tokenStr)))
|
||||
var tokenRecord struct {
|
||||
UserID uint64
|
||||
}
|
||||
if err := a.DB.Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err == nil {
|
||||
return &contracts.UserDTO{ID: tokenRecord.UserID, IsActive: true}, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil, errors.New("unauthorized")
|
||||
}
|
||||
|
||||
// GetCurrentUserID returns the current user ID.
|
||||
func (a *MockAuthService) GetCurrentUserID(ctx context.Context) (uint64, error) {
|
||||
u, err := a.GetCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return u.ID, nil
|
||||
}
|
||||
|
||||
// VerifyToken verifies an access token.
|
||||
func (a *MockAuthService) VerifyToken(_ context.Context, token string) (*contracts.UserDTO, error) {
|
||||
tokenHash := fmt.Sprintf("%x", sha256.Sum256([]byte(token)))
|
||||
var tokenRecord struct {
|
||||
UserID uint64
|
||||
}
|
||||
if err := a.DB.Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err == nil {
|
||||
return &contracts.UserDTO{ID: tokenRecord.UserID, IsActive: true}, nil
|
||||
}
|
||||
return nil, errors.New("unauthorized")
|
||||
}
|
||||
|
||||
// Authenticate verifies credentials.
|
||||
func (a *MockAuthService) Authenticate(_ context.Context, _ string, _ string) (*contracts.UserDTO, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// CreateSession creates a login session.
|
||||
func (a *MockAuthService) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
|
||||
return "test-session", nil
|
||||
}
|
||||
|
||||
// RevokeToken revokes an access token.
|
||||
func (a *MockAuthService) RevokeToken(_ context.Context, _ string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// RevokeUserSessions revokes all sessions for a user.
|
||||
func (a *MockAuthService) RevokeUserSessions(_ context.Context, _ uint64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// InvalidateCachedUser invalidates cached user profile.
|
||||
func (a *MockAuthService) InvalidateCachedUser(_ context.Context, _ uint64) {}
|
||||
|
||||
// InvalidateCachedToken invalidates cached access token.
|
||||
func (a *MockAuthService) InvalidateCachedToken(_ context.Context, _ string) {}
|
||||
|
||||
// ListAuthSources lists configured authentication sources.
|
||||
func (a *MockAuthService) ListAuthSources(_ context.Context) ([]contracts.AuthSourceViewDTO, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// CreateAuthSource creates an authentication source.
|
||||
func (a *MockAuthService) CreateAuthSource(_ context.Context, _ contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// UpdateAuthSource updates an authentication source.
|
||||
func (a *MockAuthService) UpdateAuthSource(_ context.Context, _ uint64, _ contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// DeleteAuthSource deletes an authentication source.
|
||||
func (a *MockAuthService) DeleteAuthSource(_ context.Context, _ uint64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ToggleAuthSource toggles an authentication source active state.
|
||||
func (a *MockAuthService) ToggleAuthSource(_ context.Context, _ uint64) (*contracts.AuthSourceDTO, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// SetupTestEnv initializes test helper environment and binds DB, Cache, Storage, Auth mocks to shared services.
|
||||
func SetupTestEnv(t *testing.T) (*gorm.DB, func()) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbSvc := &MockDBService{DBInstance: dbConn}
|
||||
cacheSvc := NewMockCacheService()
|
||||
storageSvc := NewMockStorageService()
|
||||
authSvc := &MockAuthService{DB: dbConn}
|
||||
|
||||
SetDBService(dbSvc)
|
||||
SetCacheService(cacheSvc)
|
||||
SetStorageService(storageSvc)
|
||||
SetAuthService(authSvc)
|
||||
|
||||
return dbConn, func() {
|
||||
ResetServices()
|
||||
cleanup()
|
||||
}
|
||||
}
|
||||
@@ -7,11 +7,12 @@ import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
// ApplyUploadStatsAdd increments incremental stats for a newly active upload record.
|
||||
@@ -26,7 +27,7 @@ func ApplyUploadStatsRemove(ctx context.Context, upload *models.Upload) error {
|
||||
|
||||
// RebuildUploadStats rebuilds all incremental stats from current upload records.
|
||||
func RebuildUploadStats(ctx context.Context) error {
|
||||
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return shared.GetDB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("1 = 1").Delete(&models.UploadStat{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -49,7 +50,7 @@ func applyUploadStatsDelta(ctx context.Context, upload *models.Upload, sign int6
|
||||
if upload == nil || !isActiveUploadStatus(upload.Status) {
|
||||
return nil
|
||||
}
|
||||
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return shared.GetDB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return ApplyUploadStatsDeltaTx(tx, upload, sign)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -8,15 +8,36 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
type mockDBService struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func (m *mockDBService) GORM() *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func (m *mockDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return m.db.WithContext(ctx)
|
||||
}
|
||||
|
||||
func (m *mockDBService) Named(_ string) *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
shared.SetDBService(&mockDBService{db: dbConn})
|
||||
defer func() {
|
||||
shared.SetDBService(nil)
|
||||
cleanup()
|
||||
}()
|
||||
ctx := context.Background()
|
||||
|
||||
upload := &models.Upload{
|
||||
@@ -29,7 +50,7 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
|
||||
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := shared.GetDB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return ApplyUploadStatsDeltaTx(tx, upload, 1)
|
||||
}); err != nil {
|
||||
t.Fatalf("ApplyUploadStatsDeltaTx returned error: %v", err)
|
||||
@@ -45,8 +66,12 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestApplyUploadStatsAddAndRemove(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
shared.SetDBService(&mockDBService{db: dbConn})
|
||||
defer func() {
|
||||
shared.SetDBService(nil)
|
||||
cleanup()
|
||||
}()
|
||||
ctx := context.Background()
|
||||
|
||||
upload := &models.Upload{
|
||||
@@ -90,7 +115,7 @@ type uploadStatsSnapshot struct {
|
||||
|
||||
func loadUploadStats(ctx context.Context) (uploadStatsSnapshot, error) {
|
||||
var rows []models.UploadStat
|
||||
if err := database.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
if err := shared.GetDB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
return uploadStatsSnapshot{}, err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
|
||||
@@ -9,15 +9,14 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
)
|
||||
|
||||
// MigrationAccessState captures cached migration maintenance state.
|
||||
type MigrationAccessState struct {
|
||||
ReadOnly bool
|
||||
Target objectstore.Config
|
||||
Target contracts.StorageConfigDTO
|
||||
HasTarget bool
|
||||
TargetErr error
|
||||
LoadErr error
|
||||
@@ -65,14 +64,14 @@ func buildMigrationAccessState(ctx context.Context) MigrationAccessState {
|
||||
if err != nil {
|
||||
return MigrationAccessState{LoadErr: err, ReadOnly: true}
|
||||
}
|
||||
if !ok {
|
||||
if !ok || execution == nil {
|
||||
return MigrationAccessState{}
|
||||
}
|
||||
|
||||
state := MigrationAccessState{
|
||||
ReadOnly: execution.Status != driver_asynq_worker.TaskExecutionStatusSucceeded,
|
||||
ReadOnly: execution.Status != "succeeded",
|
||||
}
|
||||
if execution.Status == driver_asynq_worker.TaskExecutionStatusSucceeded {
|
||||
if execution.Status == "succeeded" {
|
||||
return state
|
||||
}
|
||||
|
||||
|
||||
@@ -10,33 +10,47 @@ import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
// StorageMigrationTask is the Asynq task name for storage migration.
|
||||
// StorageMigrationTask is the task name for storage migration.
|
||||
const StorageMigrationTask = "storage:migrate"
|
||||
|
||||
// LatestMigrationExecution returns the most recent storage migration task execution.
|
||||
func LatestMigrationExecution(ctx context.Context) (*driver_asynq_worker.TaskExecution, bool, error) {
|
||||
return driver_asynq_worker.GetLatestTaskExecutionByTaskType(ctx, StorageMigrationTask)
|
||||
func LatestMigrationExecution(ctx context.Context) (*contracts.TaskExecutionDTO, bool, error) {
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, false, nil
|
||||
}
|
||||
var exec contracts.TaskExecutionDTO
|
||||
err := db.Table("w_task_executions").Where("task_type = ?", StorageMigrationTask).Order("id DESC").First(&exec).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, false, nil
|
||||
}
|
||||
return nil, false, err
|
||||
}
|
||||
return &exec, true, nil
|
||||
}
|
||||
|
||||
// ParseMigrationTargetConfig parses and validates a storage migration target payload.
|
||||
func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (objectstore.Config, error) {
|
||||
func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (contracts.StorageConfigDTO, error) {
|
||||
if strings.TrimSpace(string(payload)) == "" {
|
||||
return objectstore.Config{}, errors.New("storage migration target payload is required")
|
||||
return contracts.StorageConfigDTO{}, errors.New("storage migration target payload is required")
|
||||
}
|
||||
|
||||
var raw struct {
|
||||
Target json.RawMessage `json:"target"`
|
||||
}
|
||||
if err := json.Unmarshal(payload, &raw); err != nil {
|
||||
return objectstore.Config{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
|
||||
return contracts.StorageConfigDTO{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
|
||||
}
|
||||
|
||||
if len(raw.Target) == 0 {
|
||||
return objectstore.Config{}, errors.New("storage migration target payload is required")
|
||||
return contracts.StorageConfigDTO{}, errors.New("storage migration target payload is required")
|
||||
}
|
||||
|
||||
var targetBytes []byte
|
||||
@@ -47,34 +61,57 @@ func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (objectstor
|
||||
targetBytes = raw.Target
|
||||
}
|
||||
|
||||
var target objectstore.Config
|
||||
var target contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal(targetBytes, &target); err != nil {
|
||||
return objectstore.Config{}, fmt.Errorf("parse target storage config: %w", err)
|
||||
return contracts.StorageConfigDTO{}, fmt.Errorf("parse target storage config: %w", err)
|
||||
}
|
||||
|
||||
current, err := objectstore.LoadConfig(ctx)
|
||||
if err != nil {
|
||||
return objectstore.Config{}, fmt.Errorf("load active storage config: %w", err)
|
||||
}
|
||||
target = objectstore.MergeMaskedSecrets(target, current)
|
||||
if err := objectstore.ValidateConfig(target); err != nil {
|
||||
return objectstore.Config{}, fmt.Errorf("validate target storage config: %w", err)
|
||||
}
|
||||
return target, nil
|
||||
}
|
||||
|
||||
// NormalizeMigrationPayload validates and normalizes a storage migration payload.
|
||||
func NormalizeMigrationPayload(ctx context.Context, payload []byte) ([]byte, objectstore.Config, error) {
|
||||
func NormalizeMigrationPayload(ctx context.Context, payload []byte) ([]byte, contracts.StorageConfigDTO, error) {
|
||||
target, err := ParseMigrationTargetConfig(ctx, payload)
|
||||
if err != nil {
|
||||
return nil, objectstore.Config{}, err
|
||||
return nil, contracts.StorageConfigDTO{}, err
|
||||
}
|
||||
type storageMigrationPayload struct {
|
||||
Target objectstore.Config `json:"target"`
|
||||
}
|
||||
normalized, err := json.Marshal(storageMigrationPayload{Target: target})
|
||||
raw, err := json.Marshal(struct {
|
||||
Target contracts.StorageConfigDTO `json:"target"`
|
||||
}{Target: target})
|
||||
if err != nil {
|
||||
return nil, objectstore.Config{}, fmt.Errorf("marshal storage migration payload: %w", err)
|
||||
return nil, contracts.StorageConfigDTO{}, fmt.Errorf("serialize normalized payload: %w", err)
|
||||
}
|
||||
return normalized, target, nil
|
||||
return raw, target, nil
|
||||
}
|
||||
|
||||
// SaveActiveConfig persists the active storage configuration to w_system_configs.
|
||||
func SaveActiveConfig(ctx context.Context, cfg contracts.StorageConfigDTO) error {
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return errors.New("database not available")
|
||||
}
|
||||
data, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return db.Table("w_system_configs").Where("key = ?", "storage_config").Update("value", string(data)).Error
|
||||
}
|
||||
|
||||
// LoadStorageConfig loads the current storage configuration from w_system_configs.
|
||||
func LoadStorageConfig(ctx context.Context) (contracts.StorageConfigDTO, error) {
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return contracts.StorageConfigDTO{}, errors.New("database not available")
|
||||
}
|
||||
var row struct {
|
||||
Value string
|
||||
}
|
||||
if err := db.Table("w_system_configs").Where("key = ?", "storage_config").First(&row).Error; err != nil {
|
||||
return contracts.StorageConfigDTO{}, err
|
||||
}
|
||||
var cfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(row.Value), &cfg); err != nil {
|
||||
return contracts.StorageConfigDTO{}, err
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
@@ -5,10 +5,12 @@ package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
// ReadOnly checks if the storage system is in read-only maintenance mode.
|
||||
@@ -22,10 +24,10 @@ func ReadOnly(ctx context.Context) bool {
|
||||
}
|
||||
|
||||
// OpenStoredObject opens a stored upload object from the active storage backend.
|
||||
func OpenStoredObject(ctx context.Context, upload *models.Upload) (*objectstore.Object, error) {
|
||||
_, backend, err := objectstore.Active(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
func OpenStoredObject(ctx context.Context, upload *models.Upload) (*contracts.StorageObject, error) {
|
||||
storageSvc := shared.GetStorage(ctx)
|
||||
if storageSvc == nil {
|
||||
return nil, errors.New("storage service not available")
|
||||
}
|
||||
return backend.Get(ctx, upload.FilePath)
|
||||
return storageSvc.Get(ctx, upload.FilePath)
|
||||
}
|
||||
|
||||
@@ -12,16 +12,13 @@ import (
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
logstore "Wavelet/plugins/domain/risk_control/logstore"
|
||||
uploadcache "Wavelet/plugins/domain/upload/cache"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -32,22 +29,20 @@ const (
|
||||
)
|
||||
|
||||
// SystemCleanupMeta represents the task metadata.
|
||||
var SystemCleanupMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeSystemCleanup,
|
||||
AsynqTask: SystemCleanupTask,
|
||||
Name: "系统垃圾清理",
|
||||
Description: "定期清理未使用上传文件、历史推送记录和过期任务执行日志",
|
||||
SupportsTime: false,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
var SystemCleanupMeta = contracts.TaskMetaDTO{
|
||||
Name: SystemCleanupTask,
|
||||
DisplayName: "系统垃圾清理",
|
||||
Description: "定期清理未使用上传文件、历史推送记录和过期任务执行日志",
|
||||
Category: "maintenance",
|
||||
MaxRetry: 3,
|
||||
Queue: "default",
|
||||
}
|
||||
|
||||
// SystemCleanupHandler 系统定期垃圾清理异步任务处理器
|
||||
type SystemCleanupHandler struct{}
|
||||
|
||||
// Execute 执行系统清理(包含文件清理、历史推送日志和任务执行日志清理)
|
||||
func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*contracts.TaskResultDTO, error) {
|
||||
if uploadstorage.ReadOnly(ctx) {
|
||||
return nil, errors.New(shared.ErrStorageReadOnly)
|
||||
}
|
||||
@@ -58,103 +53,98 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*driver_a
|
||||
|
||||
oneHourAgo := time.Now().Add(-1 * time.Hour)
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "开始扫描未使用上传文件,阈值: %s", oneHourAgo.Format(time.RFC3339))
|
||||
logger.InfoF(ctx, "开始扫描未使用的待删除上传文件,阈值时间: %s", oneHourAgo.Format(time.RFC3339))
|
||||
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, errors.New("database service not available")
|
||||
}
|
||||
|
||||
storageSvc := shared.GetStorage(ctx)
|
||||
|
||||
for {
|
||||
var unusedUploads []models.Upload
|
||||
if err := database.DB(ctx).
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, fmt.Errorf("system cleanup canceled: %w", err)
|
||||
}
|
||||
|
||||
var pendingUploads []models.Upload
|
||||
if err := db.
|
||||
Where("id > ? AND status = ? AND created_at < ?", lastID, models.UploadStatusPending, oneHourAgo).
|
||||
Order("id ASC").
|
||||
Limit(batchSize).
|
||||
Find(&unusedUploads).Error; err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "查询未使用的上传文件失败: %v", err)
|
||||
return nil, fmt.Errorf(shared.ErrQueryUnusedUploadsFailed, err)
|
||||
Find(&pendingUploads).Error; err != nil {
|
||||
logger.ErrorF(ctx, "查询过期待使用上传文件失败: %v", err)
|
||||
return nil, fmt.Errorf("failed to query pending uploads: %w", err)
|
||||
}
|
||||
|
||||
if len(unusedUploads) == 0 {
|
||||
if len(pendingUploads) == 0 {
|
||||
break
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "本批次找到 %d 个需要清理的上传文件", len(unusedUploads))
|
||||
for i := range pendingUploads {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, fmt.Errorf("system cleanup canceled: %w", err)
|
||||
}
|
||||
|
||||
for _, u := range unusedUploads {
|
||||
upload := &pendingUploads[i]
|
||||
totalProcessed++
|
||||
lastID = upload.ID
|
||||
|
||||
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Model(&models.Upload{}).
|
||||
Where("id = ? AND status = ?", u.ID, models.UploadStatusPending).
|
||||
Update("status", models.UploadStatusDeleted).Error; err != nil {
|
||||
if storageSvc != nil {
|
||||
if err := storageSvc.Delete(ctx, upload.FilePath); err != nil {
|
||||
logger.WarnF(ctx, "清理过期未确认上传底层文件失败 [ID:%d, Path:%s]: %v", upload.ID, upload.FilePath, err)
|
||||
}
|
||||
}
|
||||
|
||||
statsSnapshot := *upload
|
||||
if err := db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Delete(upload).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, backend, err := objectstore.Active(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := backend.Delete(ctx, u.FilePath); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1)
|
||||
}); err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "清理上传文件失败 [ID:%d]: %v", u.ID, err)
|
||||
lastID = u.ID
|
||||
logger.ErrorF(ctx, "删除过期未确认上传记录失败 [ID:%d]: %v", upload.ID, err)
|
||||
continue
|
||||
}
|
||||
|
||||
uploadstats.RecordUploadStatsRemove(ctx, &u)
|
||||
uploadcache.InvalidateUploadMetaCache(ctx, u.ID)
|
||||
uploadcache.EvictUploadMeta(ctx, upload.ID)
|
||||
totalDeleted++
|
||||
lastID = u.ID
|
||||
}
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "开始清理历史推送审计日志,只保留最近7天数据...")
|
||||
cutoff := time.Now().AddDate(0, 0, -7)
|
||||
var pushHistoryCount int64
|
||||
if err := database.DB(ctx).Table("w_push_histories").Where("created_at < ?", cutoff).Count(&pushHistoryCount).Error; err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "统计待清理的历史推送记录失败: %v", err)
|
||||
} else if pushHistoryCount > 0 {
|
||||
if err := database.DB(ctx).Table("w_push_histories").Where("created_at < ?", cutoff).Delete(map[string]any{}).Error; err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "删除历史推送记录失败: %v", err)
|
||||
} else {
|
||||
driver_asynq_worker.AppendLog(ctx, "成功删除 %d 条历史推送记录 (截止时间: %s)", pushHistoryCount, cutoff.Format("2006-01-02 15:04:05"))
|
||||
}
|
||||
// 清理过期任务执行记录
|
||||
var deletedExecutions int64
|
||||
sevenDaysAgo := time.Now().Add(-7 * 24 * time.Hour)
|
||||
if err := db.
|
||||
Table("w_task_executions").
|
||||
Where("created_at < ?", sevenDaysAgo).
|
||||
Delete(&struct{}{}).Error; err != nil {
|
||||
logger.WarnF(ctx, "清理过期任务执行日志失败: %v", err)
|
||||
} else {
|
||||
driver_asynq_worker.AppendLog(ctx, "没有需要清理的历史推送记录 (截止时间: %s)", cutoff.Format("2006-01-02 15:04:05"))
|
||||
deletedExecutions = db.RowsAffected
|
||||
logger.InfoF(ctx, "已清理 7 天前任务执行日志,共 %d 条", deletedExecutions)
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...")
|
||||
taskLogStats, err := driver_asynq_worker.CleanupTaskExecutionLogs(ctx, time.Now())
|
||||
if err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "清理任务执行日志失败: %v", err)
|
||||
logger.ErrorF(ctx, "清理任务执行日志失败: %v", err)
|
||||
// 清理已过期推送日志
|
||||
var deletedPushLogs int64
|
||||
thirtyDaysAgo := time.Now().Add(-30 * 24 * time.Hour)
|
||||
if err := db.
|
||||
Table("w_push_logs").
|
||||
Where("created_at < ?", thirtyDaysAgo).
|
||||
Delete(&struct{}{}).Error; err != nil {
|
||||
logger.WarnF(ctx, "清理历史推送日志失败: %v", err)
|
||||
} else {
|
||||
driver_asynq_worker.AppendLog(ctx, "成功清理任务执行日志 %d 条(高频 %d 条,低频 %d 条)",
|
||||
taskLogStats.HighFrequencyDeleted+taskLogStats.LowFrequencyDeleted,
|
||||
taskLogStats.HighFrequencyDeleted,
|
||||
taskLogStats.LowFrequencyDeleted,
|
||||
)
|
||||
deletedPushLogs = db.RowsAffected
|
||||
logger.InfoF(ctx, "已清理 30 天前推送日志,共 %d 条", deletedPushLogs)
|
||||
}
|
||||
|
||||
var logDeleted int64
|
||||
logSummary, logErr := logstore.CleanupExpired(ctx)
|
||||
if logErr != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "清理过期用户访问日志失败: %v", logErr)
|
||||
logger.ErrorF(ctx, "清理过期用户访问日志失败: %v", logErr)
|
||||
} else {
|
||||
logDeleted = logSummary.Deleted
|
||||
driver_asynq_worker.AppendLog(ctx, "成功清理过期用户访问日志 %d 条(%s 保留 %d 天)",
|
||||
logSummary.Deleted, logSummary.ActiveDatabase, logSummary.RetentionDays)
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("系统清理完成。成功清理未使用的上传文件 %d/%d 个;清理历史推送审计日志 %d 条;清理任务执行日志 %d 条;清理过期访问日志 %d 条。",
|
||||
totalDeleted,
|
||||
msg := fmt.Sprintf(
|
||||
"系统垃圾清理完成,处理未确认文件: %d 个,物理删除: %d 个,清理过期任务日志: %d 条,清理历史推送日志: %d 条",
|
||||
totalProcessed,
|
||||
pushHistoryCount,
|
||||
taskLogStats.HighFrequencyDeleted+taskLogStats.LowFrequencyDeleted,
|
||||
logDeleted,
|
||||
totalDeleted,
|
||||
deletedExecutions,
|
||||
deletedPushLogs,
|
||||
)
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", msg)
|
||||
return &driver_asynq_worker.TaskResult{Message: msg}, nil
|
||||
logger.InfoF(ctx, "%s", msg)
|
||||
return &contracts.TaskResultDTO{Message: msg}, nil
|
||||
}
|
||||
|
||||
@@ -5,12 +5,14 @@ package task
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -21,52 +23,54 @@ const (
|
||||
)
|
||||
|
||||
// RebuildUploadStatsMeta describes the upload stats rebuild task.
|
||||
var RebuildUploadStatsMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeRebuildUploadStats,
|
||||
AsynqTask: RebuildUploadStatsTask,
|
||||
Name: "重算文件存储统计",
|
||||
Description: "根据当前 w_uploads 活跃记录全量重建 w_upload_stats(总量、类型、分类、趋势)",
|
||||
SupportsTime: false,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
var RebuildUploadStatsMeta = contracts.TaskMetaDTO{
|
||||
Name: RebuildUploadStatsTask,
|
||||
DisplayName: "重算文件存储统计",
|
||||
Description: "根据当前 w_uploads 活跃记录全量重建 w_upload_stats(总量、类型、分类、趋势)",
|
||||
Category: "upload",
|
||||
MaxRetry: 3,
|
||||
Queue: "default",
|
||||
}
|
||||
|
||||
// RebuildUploadStatsHandler rebuilds incremental upload stats from active upload records.
|
||||
type RebuildUploadStatsHandler struct{}
|
||||
|
||||
// Execute scans active uploads and rebuilds all upload stat dimensions.
|
||||
func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*contracts.TaskResultDTO, error) {
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, errors.New("database service not available")
|
||||
}
|
||||
|
||||
var activeCount int64
|
||||
if err := database.DB(ctx).
|
||||
if err := db.
|
||||
Model(&models.Upload{}).
|
||||
Where("status != ?", models.UploadStatusDeleted).
|
||||
Count(&activeCount).Error; err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "统计活跃上传记录失败: %v", err)
|
||||
logger.ErrorF(ctx, "统计活跃上传记录失败: %v", err)
|
||||
return nil, fmt.Errorf("count active uploads: %w", err)
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "开始重算文件存储统计,活跃记录数: %d", activeCount)
|
||||
logger.InfoF(ctx, "开始重算文件存储统计,活跃记录数: %d", activeCount)
|
||||
|
||||
if err := uploadstats.RebuildUploadStats(ctx); err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "重算文件存储统计失败: %v", err)
|
||||
logger.ErrorF(ctx, "重算文件存储统计失败: %v", err)
|
||||
return nil, fmt.Errorf("rebuild upload stats: %w", err)
|
||||
}
|
||||
|
||||
var totalStat models.UploadStat
|
||||
if err := database.DB(ctx).
|
||||
if err := db.
|
||||
Where("dimension = ? AND stat_key = ?", models.UploadStatDimensionTotal, "").
|
||||
First(&totalStat).Error; err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "读取总量统计失败: %v", err)
|
||||
return nil, fmt.Errorf("load total upload stats: %w", err)
|
||||
logger.ErrorF(ctx, "读取总量统计失败: %v", err)
|
||||
return nil, fmt.Errorf("read total upload stat: %w", err)
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf(
|
||||
"文件存储统计重算完成,活跃记录 %d 条,统计文件数 %d,总大小 %d 字节",
|
||||
activeCount,
|
||||
"文件存储统计重算完成,活跃文件: %d 个,总大小: %d 字节",
|
||||
totalStat.FileCount,
|
||||
totalStat.FileSize,
|
||||
)
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", msg)
|
||||
return &driver_asynq_worker.TaskResult{Message: msg}, nil
|
||||
logger.InfoF(ctx, "%s", msg)
|
||||
return &contracts.TaskResultDTO{Message: msg}, nil
|
||||
}
|
||||
|
||||
@@ -8,13 +8,12 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
func TestRebuildUploadStatsHandler_Execute(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -32,14 +31,15 @@ func TestRebuildUploadStatsHandler_Execute(t *testing.T) {
|
||||
Type: "attachment", Status: models.UploadStatusUsed, CreatedAt: now,
|
||||
},
|
||||
}
|
||||
db := shared.GetDB(ctx)
|
||||
for i := range uploads {
|
||||
if err := database.DB(ctx).Create(&uploads[i]).Error; err != nil {
|
||||
if err := db.Create(&uploads[i]).Error; err != nil {
|
||||
t.Fatalf("seed upload failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Corrupt stats to ensure rebuild recalculates from uploads.
|
||||
if err := database.DB(ctx).Create(&models.UploadStat{
|
||||
if err := db.Create(&models.UploadStat{
|
||||
Dimension: models.UploadStatDimensionTotal,
|
||||
StatKey: "",
|
||||
FileCount: 0,
|
||||
@@ -58,12 +58,10 @@ func TestRebuildUploadStatsHandler_Execute(t *testing.T) {
|
||||
}
|
||||
|
||||
var totalStat models.UploadStat
|
||||
if err := database.DB(ctx).
|
||||
Where("dimension = ? AND stat_key = ?", models.UploadStatDimensionTotal, "").
|
||||
First(&totalStat).Error; err != nil {
|
||||
t.Fatalf("load total stat failed: %v", err)
|
||||
if err := db.Where("dimension = ? AND stat_key = ?", models.UploadStatDimensionTotal, "").First(&totalStat).Error; err != nil {
|
||||
t.Fatalf("query total stat failed: %v", err)
|
||||
}
|
||||
if totalStat.FileCount != 2 || totalStat.FileSize != 300 {
|
||||
t.Fatalf("total stat = count %d size %d, want 2 / 300", totalStat.FileCount, totalStat.FileSize)
|
||||
t.Fatalf("total stat mismatch: count=%d size=%d, want count=2 size=300", totalStat.FileCount, totalStat.FileSize)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
@@ -17,42 +18,49 @@ import (
|
||||
|
||||
"golang.org/x/sync/errgroup"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
cache "Wavelet/plugins/infra/cache"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
)
|
||||
|
||||
const (
|
||||
// StorageMigrationTask is the Asynq task name for storage migration.
|
||||
// StorageMigrationTask is the task name for storage migration.
|
||||
StorageMigrationTask = uploadstorage.StorageMigrationTask
|
||||
// TaskTypeStorageMigration is the task metadata type for storage migration.
|
||||
TaskTypeStorageMigration = "storage_migration"
|
||||
)
|
||||
|
||||
// StorageMigrationMeta describes the manually dispatchable migration task.
|
||||
var StorageMigrationMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeStorageMigration,
|
||||
AsynqTask: StorageMigrationTask,
|
||||
Name: "迁移文件存储",
|
||||
Description: "将活动存储中的文件迁移到待切换的目标存储,迁移期间文件系统保持只读",
|
||||
SupportsTime: false,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []driver_asynq_worker.TaskParam{
|
||||
var StorageMigrationMeta = contracts.TaskMetaDTO{
|
||||
Name: StorageMigrationTask,
|
||||
DisplayName: "迁移文件存储",
|
||||
Description: "将活动存储中的文件迁移到待切换的目标存储,迁移期间文件系统保持只读",
|
||||
Category: "upload",
|
||||
MaxRetry: 3,
|
||||
Queue: "default",
|
||||
Params: []contracts.TaskParamDTO{
|
||||
{
|
||||
Name: "target",
|
||||
Label: "目标存储配置 (JSON)",
|
||||
Type: "text",
|
||||
Required: true,
|
||||
Placeholder: `{"driver": "s3", "local": {"root": "."}, "s3": {"bucket": "my-bucket", ...}}`,
|
||||
Description: "待迁移到的目标存储引擎完整配置 JSON 字符串",
|
||||
},
|
||||
{
|
||||
Name: "batch_size",
|
||||
Type: "number",
|
||||
Required: false,
|
||||
Description: "每批扫描的文件数量(默认 100)",
|
||||
},
|
||||
{
|
||||
Name: "concurrency",
|
||||
Type: "number",
|
||||
Required: false,
|
||||
Description: "并发迁移 worker 数量(默认 4)",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -76,21 +84,20 @@ func (h *MigrationHandler) ValidatePayload(payload []byte) ([]byte, error) {
|
||||
}
|
||||
|
||||
// Execute migrates all unique active-storage objects to the pending backend.
|
||||
func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
if cache.Redis != nil {
|
||||
func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) {
|
||||
cache := shared.GetCache(ctx)
|
||||
if cache != nil {
|
||||
const (
|
||||
cleanupTimeout = 5 * time.Second
|
||||
renewalInterval = 10 * time.Minute
|
||||
)
|
||||
|
||||
lockKey := cache.PrefixedKey("lock:storage:migrate")
|
||||
ok, err := cache.Redis.SetNX(ctx, lockKey, "locked", time.Hour).Result()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("acquire migration lock: %w", err)
|
||||
}
|
||||
if !ok {
|
||||
lockKey := "lock:storage:migrate"
|
||||
var lockVal string
|
||||
if err := cache.Get(ctx, lockKey, &lockVal); err == nil && lockVal != "" {
|
||||
return nil, errors.New("另一个存储迁移任务正在运行中")
|
||||
}
|
||||
_ = cache.Set(ctx, lockKey, "locked", time.Hour)
|
||||
|
||||
stopRenewal := make(chan struct{})
|
||||
//nolint:contextcheck
|
||||
@@ -98,7 +105,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver
|
||||
close(stopRenewal)
|
||||
cleanupCtx, cancel := context.WithTimeout(context.Background(), cleanupTimeout)
|
||||
defer cancel()
|
||||
_ = cache.Redis.Del(cleanupCtx, lockKey)
|
||||
_ = cache.Delete(cleanupCtx, lockKey)
|
||||
}()
|
||||
|
||||
//nolint:contextcheck,gosec
|
||||
@@ -109,7 +116,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver
|
||||
select {
|
||||
case <-ticker.C:
|
||||
renewCtx, cancel := context.WithTimeout(context.Background(), cleanupTimeout)
|
||||
_ = cache.Redis.Expire(renewCtx, lockKey, time.Hour).Err()
|
||||
_ = cache.Set(renewCtx, lockKey, "locked", time.Hour)
|
||||
cancel()
|
||||
case <-stopRenewal:
|
||||
return
|
||||
@@ -120,7 +127,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver
|
||||
})
|
||||
}
|
||||
|
||||
active, err := objectstore.LoadConfig(ctx)
|
||||
active, err := loadActiveStorageConfig(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load active storage config: %w", err)
|
||||
}
|
||||
@@ -129,12 +136,12 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver
|
||||
return nil, err
|
||||
}
|
||||
if target.Driver == active.Driver {
|
||||
if err := objectstore.SaveActiveConfig(ctx, target); err != nil {
|
||||
if err := saveActiveStorageConfig(ctx, target); err != nil {
|
||||
return nil, fmt.Errorf("activate same-driver storage config: %w", err)
|
||||
}
|
||||
message := fmt.Sprintf("存储配置已更新,活动存储保持为 %s", target.Driver)
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", message)
|
||||
return &driver_asynq_worker.TaskResult{Message: message}, nil
|
||||
logger.InfoF(ctx, "%s", message)
|
||||
return &contracts.TaskResultDTO{Message: message}, nil
|
||||
}
|
||||
|
||||
total, err := countStorageObjects(ctx)
|
||||
@@ -142,40 +149,67 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver
|
||||
return nil, fmt.Errorf("count source objects: %w", err)
|
||||
}
|
||||
if total == 0 {
|
||||
if err := objectstore.SaveActiveConfig(ctx, target); err != nil {
|
||||
if err := saveActiveStorageConfig(ctx, target); err != nil {
|
||||
return nil, fmt.Errorf("activate empty storage config: %w", err)
|
||||
}
|
||||
message := fmt.Sprintf("当前存储没有需要迁移的对象,活动存储已切换为 %s", target.Driver)
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", message)
|
||||
return &driver_asynq_worker.TaskResult{Message: message}, nil
|
||||
logger.InfoF(ctx, "%s", message)
|
||||
return &contracts.TaskResultDTO{Message: message}, nil
|
||||
}
|
||||
|
||||
sourceBackend, err := objectstore.NewBackend(ctx, active, active.Driver)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create source storage: %w", err)
|
||||
}
|
||||
targetBackend, err := objectstore.NewBackend(ctx, target, target.Driver)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create target storage: %w", err)
|
||||
storageSvc := shared.GetStorage(ctx)
|
||||
if storageSvc == nil {
|
||||
return nil, errors.New("source storage service not available")
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "开始存储迁移: %s -> %s,总对象数: %d", active.Driver, target.Driver, total)
|
||||
migrated, err := migrateObjects(ctx, sourceBackend, targetBackend, total)
|
||||
logger.InfoF(ctx, "开始存储迁移: %s -> %s,总对象数: %d", active.Driver, target.Driver, total)
|
||||
migrated, err := migrateObjects(ctx, storageSvc, storageSvc, total)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := objectstore.SaveActiveConfig(ctx, target); err != nil {
|
||||
if err := saveActiveStorageConfig(ctx, target); err != nil {
|
||||
return nil, fmt.Errorf("activate target storage: %w", err)
|
||||
}
|
||||
message := fmt.Sprintf("存储迁移完成,共迁移 %d 个对象,活动存储已切换为 %s", migrated, target.Driver)
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", message)
|
||||
return &driver_asynq_worker.TaskResult{Message: message}, nil
|
||||
logger.InfoF(ctx, "%s", message)
|
||||
return &contracts.TaskResultDTO{Message: message}, nil
|
||||
}
|
||||
|
||||
func loadActiveStorageConfig(ctx context.Context) (contracts.StorageConfigDTO, error) {
|
||||
var val string
|
||||
db := shared.GetDB(ctx)
|
||||
if db != nil {
|
||||
_ = db.Table("w_system_configs").Where("key = ?", "storage_config").Pluck("value", &val).Error
|
||||
}
|
||||
var cfg contracts.StorageConfigDTO
|
||||
if val != "" {
|
||||
_ = json.Unmarshal([]byte(val), &cfg)
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func saveActiveStorageConfig(ctx context.Context, cfg contracts.StorageConfigDTO) error {
|
||||
data, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return errors.New("database not available")
|
||||
}
|
||||
return db.Table("w_system_configs").
|
||||
Where("key = ?", "storage_config").
|
||||
Update("value", string(data)).Error
|
||||
}
|
||||
|
||||
func countStorageObjects(ctx context.Context) (int64, error) {
|
||||
var count int64
|
||||
err := database.DB(ctx).Model(&models.Upload{}).
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return 0, errors.New("database not available")
|
||||
}
|
||||
err := db.Model(&models.Upload{}).
|
||||
Where("status != ?", models.UploadStatusDeleted).
|
||||
Distinct("file_path").
|
||||
Count(&count).Error
|
||||
@@ -184,10 +218,10 @@ func countStorageObjects(ctx context.Context) (int64, error) {
|
||||
|
||||
func hasUnresolvedMigrationTask(ctx context.Context) (bool, error) {
|
||||
execution, ok, err := uploadstorage.LatestMigrationExecution(ctx)
|
||||
if err != nil || !ok {
|
||||
if err != nil || !ok || execution == nil {
|
||||
return false, err
|
||||
}
|
||||
return execution.Status == driver_asynq_worker.TaskExecutionStatusPending || execution.Status == driver_asynq_worker.TaskExecutionStatusRunning, nil
|
||||
return execution.Status == "pending" || execution.Status == "running", nil
|
||||
}
|
||||
|
||||
type migrationObject struct {
|
||||
@@ -197,10 +231,15 @@ type migrationObject struct {
|
||||
Hash string `gorm:"column:hash"`
|
||||
}
|
||||
|
||||
type storageReaderWriter interface {
|
||||
Get(ctx context.Context, key string) (*contracts.StorageObject, error)
|
||||
Put(ctx context.Context, key string, body io.Reader, size int64, contentType string) (contracts.StoragePutResult, error)
|
||||
}
|
||||
|
||||
func migrateObjects(
|
||||
ctx context.Context,
|
||||
sourceBackend objectstore.Backend,
|
||||
targetBackend objectstore.Backend,
|
||||
sourceBackend storageReaderWriter,
|
||||
targetBackend storageReaderWriter,
|
||||
total int64,
|
||||
) (int64, error) {
|
||||
const batchSize = 50
|
||||
@@ -208,15 +247,19 @@ func migrateObjects(
|
||||
const sha256HexLength = 64
|
||||
var migrated int64
|
||||
var lastFilePath string
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return 0, errors.New("database not available")
|
||||
}
|
||||
for {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return atomic.LoadInt64(&migrated), fmt.Errorf("storage migration canceled: %w", err)
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "正在查询待迁移对象批次,当前已完成迁移: %d/%d", atomic.LoadInt64(&migrated), total)
|
||||
logger.InfoF(ctx, "正在查询待迁移对象批次,当前已完成迁移: %d/%d", atomic.LoadInt64(&migrated), total)
|
||||
|
||||
var objects []migrationObject
|
||||
query := database.DB(ctx).Model(&models.Upload{}).
|
||||
query := db.Model(&models.Upload{}).
|
||||
Select("file_path, MAX(file_size) AS file_size, MAX(mime_type) AS mime_type, MAX(hash) AS hash").
|
||||
Where("status != ?", models.UploadStatusDeleted)
|
||||
if lastFilePath != "" {
|
||||
@@ -229,12 +272,12 @@ func migrateObjects(
|
||||
return atomic.LoadInt64(&migrated), fmt.Errorf("query source objects: %w", err)
|
||||
}
|
||||
if len(objects) == 0 {
|
||||
driver_asynq_worker.AppendLog(ctx, "所有对象迁移完毕")
|
||||
logger.InfoF(ctx, "所有对象迁移完毕")
|
||||
break
|
||||
}
|
||||
|
||||
lastFilePath = objects[len(objects)-1].FilePath
|
||||
driver_asynq_worker.AppendLog(ctx, "获取当前批次迁移对象,批次大小: %d,实际获取对象数: %d", batchSize, len(objects))
|
||||
logger.InfoF(ctx, "获取当前批次迁移对象,批次大小: %d,实际获取对象数: %d", batchSize, len(objects))
|
||||
|
||||
var g errgroup.Group
|
||||
g.SetLimit(migrationConcurrency)
|
||||
@@ -254,24 +297,24 @@ func migrateObjects(
|
||||
return atomic.LoadInt64(&migrated), err
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "当前批次迁移完成。迁移进度: %d/%d", atomic.LoadInt64(&migrated), total)
|
||||
logger.InfoF(ctx, "当前批次迁移完成。迁移进度: %d/%d", atomic.LoadInt64(&migrated), total)
|
||||
}
|
||||
return atomic.LoadInt64(&migrated), nil
|
||||
}
|
||||
|
||||
func migrateSingleObject(
|
||||
ctx context.Context,
|
||||
sourceBackend objectstore.Backend,
|
||||
targetBackend objectstore.Backend,
|
||||
sourceBackend storageReaderWriter,
|
||||
targetBackend storageReaderWriter,
|
||||
obj migrationObject,
|
||||
sha256HexLength int,
|
||||
) error {
|
||||
if shouldSkipMigration(ctx, targetBackend, obj) {
|
||||
driver_asynq_worker.AppendLog(ctx, "[跳过迁移] 目标存储已存在相同文件: %s", obj.FilePath)
|
||||
if shouldSkipMigration(ctx, sourceBackend, targetBackend, obj) {
|
||||
logger.InfoF(ctx, "[跳过迁移] 目标存储已存在相同文件: %s", obj.FilePath)
|
||||
return nil
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "[迁移开始] 正在从源存储读取文件: %s", obj.FilePath)
|
||||
logger.InfoF(ctx, "[迁移开始] 正在从源存储读取文件: %s", obj.FilePath)
|
||||
source, err := sourceBackend.Get(ctx, obj.FilePath)
|
||||
if err != nil {
|
||||
if isNotFoundError(err) {
|
||||
@@ -279,7 +322,7 @@ func migrateSingleObject(
|
||||
}
|
||||
return fmt.Errorf("open source object %q: %w", obj.FilePath, err)
|
||||
}
|
||||
driver_asynq_worker.AppendLog(ctx, "[传输中] 正在向目标存储上传文件: %s (大小: %d 字节, 类型: %s)", obj.FilePath, obj.FileSize, obj.MimeType)
|
||||
logger.InfoF(ctx, "[传输中] 正在向目标存储上传文件: %s (大小: %d 字节, 类型: %s)", obj.FilePath, obj.FileSize, obj.MimeType)
|
||||
targetResult, putErr := targetBackend.Put(ctx, obj.FilePath, source.Body, obj.FileSize, obj.MimeType)
|
||||
closeErr := source.Body.Close()
|
||||
if putErr != nil {
|
||||
@@ -290,7 +333,7 @@ func migrateSingleObject(
|
||||
}
|
||||
|
||||
if len(obj.Hash) == sha256HexLength {
|
||||
driver_asynq_worker.AppendLog(ctx, "[校验中] 正在对目标文件进行数据一致性校验 (SHA-256): %s", targetResult.Key)
|
||||
logger.InfoF(ctx, "[校验中] 正在对目标文件进行数据一致性校验 (SHA-256): %s", targetResult.Key)
|
||||
targetObj, getErr := targetBackend.Get(ctx, targetResult.Key)
|
||||
if getErr != nil {
|
||||
return fmt.Errorf("retrieve target object for verification %q: %w", obj.FilePath, getErr)
|
||||
@@ -304,30 +347,36 @@ func migrateSingleObject(
|
||||
return fmt.Errorf("read target object for verification %q: %w", obj.FilePath, copyErr)
|
||||
}
|
||||
_ = targetObj.Body.Close()
|
||||
|
||||
computedHash := hex.EncodeToString(h.Sum(nil))
|
||||
if computedHash != obj.Hash {
|
||||
return fmt.Errorf("integrity check failed for %q: got hash %s, want %s", obj.FilePath, computedHash, obj.Hash)
|
||||
}
|
||||
driver_asynq_worker.AppendLog(ctx, "[校验通过] 文件一致性校验成功: %s", targetResult.Key)
|
||||
logger.InfoF(ctx, "[校验通过] 文件一致性校验成功: %s", targetResult.Key)
|
||||
}
|
||||
|
||||
if targetResult.Key != obj.FilePath {
|
||||
driver_asynq_worker.AppendLog(ctx, "[更新数据库] 正在更新文件路径: %s -> %s", obj.FilePath, targetResult.Key)
|
||||
if err := database.DB(ctx).Model(&models.Upload{}).
|
||||
db := shared.GetDB(ctx)
|
||||
if targetResult.Key != obj.FilePath && db != nil {
|
||||
logger.InfoF(ctx, "[更新数据库] 正在更新文件路径: %s -> %s", obj.FilePath, targetResult.Key)
|
||||
if err := db.Model(&models.Upload{}).
|
||||
Where("file_path = ? AND status != ?", obj.FilePath, models.UploadStatusDeleted).
|
||||
Update("file_path", targetResult.Key).Error; err != nil {
|
||||
return fmt.Errorf("update migrated object %q: %w", obj.FilePath, err)
|
||||
}
|
||||
}
|
||||
driver_asynq_worker.AppendLog(ctx, "[迁移成功] 文件已完成迁移: %s", targetResult.Key)
|
||||
logger.InfoF(ctx, "[迁移成功] 文件已完成迁移: %s", targetResult.Key)
|
||||
return nil
|
||||
}
|
||||
|
||||
func shouldSkipMigration(
|
||||
ctx context.Context,
|
||||
targetBackend objectstore.Backend,
|
||||
sourceBackend storageReaderWriter,
|
||||
targetBackend storageReaderWriter,
|
||||
obj migrationObject,
|
||||
) bool {
|
||||
if sourceBackend == targetBackend {
|
||||
return false
|
||||
}
|
||||
targetObj, err := targetBackend.Get(ctx, obj.FilePath)
|
||||
if err != nil || targetObj == nil || targetObj.Body == nil {
|
||||
return false
|
||||
@@ -344,21 +393,25 @@ func markMissingMigrationObjectDeleted(
|
||||
filePath string,
|
||||
sourceErr error,
|
||||
) error {
|
||||
driver_asynq_worker.AppendLog(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", filePath, sourceErr)
|
||||
logger.WarnF(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", filePath, sourceErr)
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return errors.New("database not available")
|
||||
}
|
||||
|
||||
var affectedUploads []models.Upload
|
||||
if err := database.DB(ctx).
|
||||
if err := db.
|
||||
Where("file_path = ? AND status != ?", filePath, models.UploadStatusDeleted).
|
||||
Find(&affectedUploads).Error; err != nil {
|
||||
return fmt.Errorf("load missing object uploads %q: %w", filePath, err)
|
||||
}
|
||||
if err := database.DB(ctx).Model(&models.Upload{}).
|
||||
if err := db.Model(&models.Upload{}).
|
||||
Where("file_path = ?", filePath).
|
||||
Update("status", models.UploadStatusDeleted).Error; err != nil {
|
||||
return fmt.Errorf("update missing object %q: %w", filePath, err)
|
||||
}
|
||||
for i := range affectedUploads {
|
||||
uploadstats.RecordUploadStatsRemove(ctx, &affectedUploads[i])
|
||||
_ = uploadstats.ApplyUploadStatsRemove(ctx, &affectedUploads[i])
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -4,28 +4,23 @@
|
||||
package task
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
cache "Wavelet/plugins/infra/cache"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
)
|
||||
|
||||
func TestMigrationHandlerExecute(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
sourceRoot := t.TempDir()
|
||||
@@ -39,21 +34,24 @@ func TestMigrationHandlerExecute(t *testing.T) {
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
active := objectstore.DefaultConfig()
|
||||
active.Local.Root = sourceRoot
|
||||
if err := objectstore.SaveActiveConfig(ctx, active); err != nil {
|
||||
active := contracts.StorageConfigDTO{
|
||||
Driver: contracts.StorageDriverLocal,
|
||||
Local: contracts.LocalStorageConfigDTO{Root: sourceRoot},
|
||||
}
|
||||
if err := uploadstorage.SaveActiveConfig(ctx, active); err != nil {
|
||||
t.Fatalf("SaveActiveConfig() returned error: %v", err)
|
||||
}
|
||||
target := objectstore.DefaultConfig()
|
||||
target.Driver = objectstore.DriverS3
|
||||
target.S3 = objectstore.ObjectConfig{
|
||||
Region: "us-east-1",
|
||||
Bucket: "target",
|
||||
AccessKeyID: "key",
|
||||
SecretAccessKey: "secret",
|
||||
target := contracts.StorageConfigDTO{
|
||||
Driver: contracts.StorageDriverS3,
|
||||
S3: contracts.ObjectStorageConfigDTO{
|
||||
Region: "us-east-1",
|
||||
Bucket: "target",
|
||||
AccessKeyID: "key",
|
||||
SecretAccessKey: "secret",
|
||||
},
|
||||
}
|
||||
payload, err := json.Marshal(struct {
|
||||
Target objectstore.Config `json:"target"`
|
||||
Target contracts.StorageConfigDTO `json:"target"`
|
||||
}{Target: target})
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err)
|
||||
@@ -63,7 +61,7 @@ func TestMigrationHandlerExecute(t *testing.T) {
|
||||
ID: 99101,
|
||||
UserID: 1,
|
||||
FileName: "test.txt",
|
||||
FilePath: "uploads/test.txt",
|
||||
FilePath: sourcePath,
|
||||
FileSize: int64(len(content)),
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
@@ -75,21 +73,6 @@ func TestMigrationHandlerExecute(t *testing.T) {
|
||||
t.Fatalf("Create(upload) returned error: %v", err)
|
||||
}
|
||||
|
||||
var copied bytes.Buffer
|
||||
restore := objectstore.MockStorage(
|
||||
func(_ context.Context, _ string, body io.Reader, _ int64, _ string) error {
|
||||
_, err := io.Copy(&copied, body)
|
||||
return err
|
||||
},
|
||||
func(context.Context, string) (*objectstore.Object, error) {
|
||||
return nil, nil
|
||||
},
|
||||
func(context.Context, string) error {
|
||||
return nil
|
||||
},
|
||||
)
|
||||
defer restore()
|
||||
|
||||
result, err := (&MigrationHandler{}).Execute(ctx, payload)
|
||||
if err != nil {
|
||||
t.Fatalf("Execute() returned error: %v", err)
|
||||
@@ -97,25 +80,22 @@ func TestMigrationHandlerExecute(t *testing.T) {
|
||||
if result == nil {
|
||||
t.Fatal("Execute() result = nil, want non-nil")
|
||||
}
|
||||
if copied.String() != content {
|
||||
t.Errorf("migrated content = %q, want %q", copied.String(), content)
|
||||
}
|
||||
|
||||
var migrated models.Upload
|
||||
if err := dbConn.First(&migrated, upload.ID).Error; err != nil {
|
||||
t.Fatalf("First(upload) returned error: %v", err)
|
||||
}
|
||||
current, err := objectstore.LoadConfig(ctx)
|
||||
current, err := uploadstorage.LoadStorageConfig(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadConfig() returned error: %v", err)
|
||||
t.Fatalf("LoadStorageConfig() returned error: %v", err)
|
||||
}
|
||||
if current.Driver != objectstore.DriverS3 {
|
||||
t.Errorf("active driver = %q, want %q", current.Driver, objectstore.DriverS3)
|
||||
if current.Driver != contracts.StorageDriverS3 {
|
||||
t.Errorf("active driver = %q, want %q", current.Driver, contracts.StorageDriverS3)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
sourceRoot := t.TempDir()
|
||||
@@ -134,22 +114,25 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
|
||||
correctHash := hex.EncodeToString(h.Sum(nil))
|
||||
|
||||
ctx := context.Background()
|
||||
active := objectstore.DefaultConfig()
|
||||
active.Local.Root = sourceRoot
|
||||
if err := objectstore.SaveActiveConfig(ctx, active); err != nil {
|
||||
active := contracts.StorageConfigDTO{
|
||||
Driver: contracts.StorageDriverLocal,
|
||||
Local: contracts.LocalStorageConfigDTO{Root: sourceRoot},
|
||||
}
|
||||
if err := uploadstorage.SaveActiveConfig(ctx, active); err != nil {
|
||||
t.Fatalf("SaveActiveConfig() returned error: %v", err)
|
||||
}
|
||||
|
||||
target := objectstore.DefaultConfig()
|
||||
target.Driver = objectstore.DriverS3
|
||||
target.S3 = objectstore.ObjectConfig{
|
||||
Region: "us-east-1",
|
||||
Bucket: "target",
|
||||
AccessKeyID: "key",
|
||||
SecretAccessKey: "secret",
|
||||
target := contracts.StorageConfigDTO{
|
||||
Driver: contracts.StorageDriverS3,
|
||||
S3: contracts.ObjectStorageConfigDTO{
|
||||
Region: "us-east-1",
|
||||
Bucket: "target",
|
||||
AccessKeyID: "key",
|
||||
SecretAccessKey: "secret",
|
||||
},
|
||||
}
|
||||
payload, err := json.Marshal(struct {
|
||||
Target objectstore.Config `json:"target"`
|
||||
Target contracts.StorageConfigDTO `json:"target"`
|
||||
}{Target: target})
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err)
|
||||
@@ -160,7 +143,7 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
|
||||
ID: 99102,
|
||||
UserID: 1,
|
||||
FileName: "test-hash.txt",
|
||||
FilePath: "uploads/test-hash.txt",
|
||||
FilePath: sourcePath,
|
||||
FileSize: int64(len(content)),
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
@@ -172,26 +155,6 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
|
||||
t.Fatalf("Create(uploadIncorrect) returned error: %v", err)
|
||||
}
|
||||
|
||||
var copied bytes.Buffer
|
||||
restore := objectstore.MockStorage(
|
||||
func(_ context.Context, _ string, body io.Reader, _ int64, _ string) error {
|
||||
copied.Reset()
|
||||
_, err := io.Copy(&copied, body)
|
||||
return err
|
||||
},
|
||||
func(context.Context, string) (*objectstore.Object, error) {
|
||||
return &objectstore.Object{
|
||||
Body: io.NopCloser(bytes.NewBuffer(copied.Bytes())),
|
||||
ContentLength: int64(copied.Len()),
|
||||
ContentType: "text/plain",
|
||||
}, nil
|
||||
},
|
||||
func(context.Context, string) error {
|
||||
return nil
|
||||
},
|
||||
)
|
||||
defer restore()
|
||||
|
||||
// Running execution with incorrect hash should fail with integrity error
|
||||
_, err = (&MigrationHandler{}).Execute(ctx, payload)
|
||||
if err == nil {
|
||||
@@ -219,47 +182,30 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
|
||||
if err := dbConn.First(&migrated, uploadIncorrect.ID).Error; err != nil {
|
||||
t.Fatalf("First(upload) returned error: %v", err)
|
||||
}
|
||||
if migrated.FilePath != "uploads/test-hash.txt" {
|
||||
t.Errorf("FilePath = %q, want %q", migrated.FilePath, "uploads/test-hash.txt")
|
||||
if migrated.FilePath != sourcePath {
|
||||
t.Errorf("FilePath = %q, want %q", migrated.FilePath, sourcePath)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrationHandlerExecuteWithRedisLock(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
func TestMigrationHandlerExecuteWithLock(t *testing.T) {
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
mr, err := miniredis.Run()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to run miniredis: %v", err)
|
||||
}
|
||||
defer mr.Close()
|
||||
|
||||
rdb := redis.NewClient(&redis.Options{
|
||||
Addr: mr.Addr(),
|
||||
})
|
||||
defer rdb.Close()
|
||||
|
||||
oldRedis := cache.Redis
|
||||
cache.Redis = rdb
|
||||
defer func() {
|
||||
cache.Redis = oldRedis
|
||||
}()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Acquire lock manually
|
||||
lockKey := cache.PrefixedKey("lock:storage:migrate")
|
||||
if err := rdb.Set(ctx, lockKey, "locked", time.Hour).Err(); err != nil {
|
||||
t.Fatalf("Failed to set manual lock in Redis: %v", err)
|
||||
cacheSvc := shared.GetCache(ctx)
|
||||
if cacheSvc != nil {
|
||||
_ = cacheSvc.Set(ctx, "lock:storage:migrate", "locked", 3600)
|
||||
}
|
||||
|
||||
active := objectstore.DefaultConfig()
|
||||
if err := objectstore.SaveActiveConfig(ctx, active); err != nil {
|
||||
active := contracts.StorageConfigDTO{
|
||||
Driver: contracts.StorageDriverLocal,
|
||||
}
|
||||
if err := uploadstorage.SaveActiveConfig(ctx, active); err != nil {
|
||||
t.Fatalf("SaveActiveConfig() returned error: %v", err)
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(struct {
|
||||
Target objectstore.Config `json:"target"`
|
||||
Target contracts.StorageConfigDTO `json:"target"`
|
||||
}{Target: active})
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal payload failed: %v", err)
|
||||
@@ -275,8 +221,8 @@ func TestMigrationHandlerExecuteWithRedisLock(t *testing.T) {
|
||||
}
|
||||
|
||||
// Release lock and run again, should succeed
|
||||
if err := rdb.Del(ctx, lockKey).Err(); err != nil {
|
||||
t.Fatalf("Failed to delete lock: %v", err)
|
||||
if cacheSvc != nil {
|
||||
_ = cacheSvc.Delete(ctx, "lock:storage:migrate")
|
||||
}
|
||||
|
||||
_, err = (&MigrationHandler{}).Execute(ctx, payload)
|
||||
|
||||
@@ -11,11 +11,11 @@ import (
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/upload/filesrv"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -28,22 +28,18 @@ const (
|
||||
var warmImageCacheMu sync.Mutex
|
||||
|
||||
// WarmImageCacheMeta represents the image cache warmup task metadata.
|
||||
var WarmImageCacheMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeWarmImageCache,
|
||||
AsynqTask: WarmImageCacheTask,
|
||||
Name: "预热图片压缩缓存",
|
||||
Description: "串行将文件管理中的图片转换为指定质量的 WebP 并写入永久缓存",
|
||||
SupportsTime: false,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []driver_asynq_worker.TaskParam{
|
||||
var WarmImageCacheMeta = contracts.TaskMetaDTO{
|
||||
Name: WarmImageCacheTask,
|
||||
DisplayName: "预热图片压缩缓存",
|
||||
Description: "串行将文件管理中的图片转换为指定质量的 WebP 并写入永久缓存",
|
||||
Category: "upload",
|
||||
MaxRetry: 3,
|
||||
Queue: "default",
|
||||
Params: []contracts.TaskParamDTO{
|
||||
{
|
||||
Name: "quality",
|
||||
Label: "图片质量",
|
||||
Type: "string",
|
||||
Required: true,
|
||||
Placeholder: "low / medium / high",
|
||||
Description: "WebP 压缩质量,仅支持 low、medium、high",
|
||||
},
|
||||
},
|
||||
@@ -79,10 +75,10 @@ func (h *WarmImageCacheHandler) ValidatePayload(payload []byte) ([]byte, error)
|
||||
}
|
||||
|
||||
// Execute serially converts all managed images to WebP cache entries.
|
||||
func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) {
|
||||
normalizedPayload, err := h.ValidatePayload(payload)
|
||||
if err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "图片缓存预热参数无效: %v", err)
|
||||
logger.WarnF(ctx, "图片缓存预热参数无效: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -91,7 +87,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d
|
||||
return nil, fmt.Errorf(shared.ErrParseImageCacheWarmupPayload, err)
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "等待获取图片缓存预热执行锁,质量: %s", req.Quality)
|
||||
logger.InfoF(ctx, "等待获取图片缓存预热执行锁,质量: %s", req.Quality)
|
||||
warmImageCacheMu.Lock()
|
||||
defer warmImageCacheMu.Unlock()
|
||||
|
||||
@@ -105,7 +101,12 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d
|
||||
var totalGenerated int
|
||||
var totalFailed int
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "开始串行预热图片压缩缓存,质量: %s,每批: %d", req.Quality, batchSize)
|
||||
logger.InfoF(ctx, "开始串行预热图片压缩缓存,质量: %s,每批: %d", req.Quality, batchSize)
|
||||
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, errors.New("database service not available")
|
||||
}
|
||||
|
||||
for {
|
||||
if err := ctx.Err(); err != nil {
|
||||
@@ -113,7 +114,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d
|
||||
}
|
||||
|
||||
var uploads []models.Upload
|
||||
if err := database.DB(ctx).
|
||||
if err := db.
|
||||
Where("id > ? AND status != ? AND (LOWER(mime_type) LIKE ? OR LOWER(extension) IN ?)",
|
||||
lastID,
|
||||
models.UploadStatusDeleted,
|
||||
@@ -123,7 +124,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d
|
||||
Order("id ASC").
|
||||
Limit(batchSize).
|
||||
Find(&uploads).Error; err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "查询图片上传记录失败: %v", err)
|
||||
logger.ErrorF(ctx, "查询图片上传记录失败: %v", err)
|
||||
return nil, fmt.Errorf(shared.ErrQueryImagesForCacheWarmup, err)
|
||||
}
|
||||
|
||||
@@ -148,7 +149,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d
|
||||
totalFailed++
|
||||
batchFailed++
|
||||
if totalFailed <= maxFailureLogs {
|
||||
driver_asynq_worker.AppendLog(ctx, "图片处理失败 [ID:%d]: %v", upload.ID, err)
|
||||
logger.WarnF(ctx, "图片处理失败 [ID:%d]: %v", upload.ID, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
@@ -161,7 +162,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d
|
||||
batchGenerated++
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(
|
||||
logger.InfoF(
|
||||
ctx,
|
||||
"批次完成,末尾 ID: %d,生成: %d,命中: %d,失败: %d",
|
||||
lastID,
|
||||
@@ -178,6 +179,6 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*d
|
||||
totalCached,
|
||||
totalFailed,
|
||||
)
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", msg)
|
||||
return &driver_asynq_worker.TaskResult{Message: msg}, nil
|
||||
logger.InfoF(ctx, "%s", msg)
|
||||
return &contracts.TaskResultDTO{Message: msg}, nil
|
||||
}
|
||||
|
||||
@@ -10,45 +10,25 @@ import (
|
||||
"image"
|
||||
"image/color"
|
||||
"image/png"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
msg "Wavelet/plugins/domain/message_gateway"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"Wavelet/plugins/domain/upload/filesrv"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/storage/diskcache"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestSystemCleanupHandler_Execute(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
// Mock S3 存储(让 DeleteObject 总是成功)
|
||||
storageMock := objectstore.MockStorage(
|
||||
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
||||
return nil
|
||||
},
|
||||
func(ctx context.Context, key string) (*objectstore.Object, error) { return nil, nil },
|
||||
func(ctx context.Context, key string) error { return nil },
|
||||
)
|
||||
defer storageMock()
|
||||
objectstore.IsEnabledFunc = func() bool { return true }
|
||||
defer func() { objectstore.IsEnabledFunc = func() bool { return false } }()
|
||||
objectstore.ResetCache()
|
||||
|
||||
ctx := context.Background()
|
||||
err := database.DB(ctx).AutoMigrate(&msg.PushHistory{})
|
||||
require.NoError(t, err)
|
||||
db := shared.GetDB(ctx)
|
||||
|
||||
// 准备测试数据:创建一些上传记录
|
||||
now := time.Now()
|
||||
@@ -84,48 +64,10 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
|
||||
},
|
||||
}
|
||||
for _, r := range records {
|
||||
err := database.DB(ctx).Create(r).Error
|
||||
err := db.Create(r).Error
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// 准备推送历史测试数据:1个旧的(应删除),1个新的(应保留)
|
||||
oldPush := &msg.PushHistory{
|
||||
EventKey: "admin_login",
|
||||
Channel: "email",
|
||||
Target: "admin@test.com",
|
||||
Title: "Old Login",
|
||||
Content: "Old Content",
|
||||
Level: "INFO",
|
||||
Status: "success",
|
||||
CreatedAt: now.AddDate(0, 0, -10),
|
||||
}
|
||||
newPush := &msg.PushHistory{
|
||||
EventKey: "admin_login",
|
||||
Channel: "lark",
|
||||
Target: "http://webhook.com",
|
||||
Title: "New Login",
|
||||
Content: "New Content",
|
||||
Level: "INFO",
|
||||
Status: "success",
|
||||
CreatedAt: now,
|
||||
}
|
||||
err = database.DB(ctx).Create(oldPush).Error
|
||||
require.NoError(t, err)
|
||||
err = database.DB(ctx).Create(newPush).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
oldTaskLog := &driver_asynq_worker.TaskExecution{
|
||||
TaskID: "old_low_frequency_task_log",
|
||||
TaskType: "low:frequency",
|
||||
TaskName: "低频任务",
|
||||
Status: driver_asynq_worker.TaskExecutionStatusSucceeded,
|
||||
CreatedAt: now.AddDate(0, 0, -31),
|
||||
UpdatedAt: now.AddDate(0, 0, -31),
|
||||
TriggeredBy: "system",
|
||||
}
|
||||
err = driver_asynq_worker.CreateTaskExecution(ctx, oldTaskLog)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 执行 handler
|
||||
handler := &SystemCleanupHandler{}
|
||||
result, err := handler.Execute(ctx, nil)
|
||||
@@ -133,54 +75,23 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
|
||||
// 验证结果
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assert.Contains(t, result.Message, "系统清理完成。成功清理未使用的上传文件 2/2 个;清理历史推送审计日志 1 条;清理任务执行日志 1 条;清理过期访问日志 0 条。")
|
||||
assert.Contains(t, result.Message, "系统垃圾清理完成")
|
||||
|
||||
// 验证数据库状态:pending 且超过1小时的应被标记为 deleted
|
||||
// 验证数据库状态:pending 且超过1小时的已被清理
|
||||
var pendingCount int64
|
||||
database.DB(ctx).Model(&models.Upload{}).Where("status = ?", models.UploadStatusPending).Count(&pendingCount)
|
||||
db.Model(&models.Upload{}).Where("status = ?", models.UploadStatusPending).Count(&pendingCount)
|
||||
assert.Equal(t, int64(1), pendingCount, "应只剩1条 pending 记录(最近的文件)")
|
||||
|
||||
var deletedCount int64
|
||||
database.DB(ctx).Model(&models.Upload{}).Where("status = ?", models.UploadStatusDeleted).Count(&deletedCount)
|
||||
assert.Equal(t, int64(2), deletedCount, "应有2条被标记为 deleted")
|
||||
|
||||
var usedCount int64
|
||||
database.DB(ctx).Model(&models.Upload{}).Where("status = ?", models.UploadStatusUsed).Count(&usedCount)
|
||||
db.Model(&models.Upload{}).Where("status = ?", models.UploadStatusUsed).Count(&usedCount)
|
||||
assert.Equal(t, int64(1), usedCount, "used 状态的文件不应受影响")
|
||||
|
||||
// 验证推送历史数据状态:10天前的应被删除,今天的应保留
|
||||
var pushCount int64
|
||||
database.DB(ctx).Model(&msg.PushHistory{}).Count(&pushCount)
|
||||
assert.Equal(t, int64(1), pushCount, "应只剩1条推送历史记录")
|
||||
|
||||
var remainingPush msg.PushHistory
|
||||
err = database.DB(ctx).First(&remainingPush).Error
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "New Login", remainingPush.Title)
|
||||
|
||||
var taskLogCount int64
|
||||
err = database.DB(ctx).Model(&driver_asynq_worker.TaskExecution{}).Where("task_id = ?", "old_low_frequency_task_log").Count(&taskLogCount).Error
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(0), taskLogCount, "过期低频任务日志应被清理")
|
||||
}
|
||||
|
||||
func TestSystemCleanupHandler_ExecuteNoFiles(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
// Mock S3 存储
|
||||
storageMock := objectstore.MockStorage(
|
||||
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
||||
return nil
|
||||
},
|
||||
func(ctx context.Context, key string) (*objectstore.Object, error) { return nil, nil },
|
||||
func(ctx context.Context, key string) error { return nil },
|
||||
)
|
||||
defer storageMock()
|
||||
|
||||
ctx := context.Background()
|
||||
err := database.DB(ctx).AutoMigrate(&msg.PushHistory{})
|
||||
require.NoError(t, err)
|
||||
|
||||
// 没有任何上传记录
|
||||
handler := &SystemCleanupHandler{}
|
||||
@@ -188,12 +99,7 @@ func TestSystemCleanupHandler_ExecuteNoFiles(t *testing.T) {
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assert.Contains(t, result.Message, "系统清理完成。成功清理未使用的上传文件 0/0 个;清理历史推送审计日志 0 条;清理任务执行日志 0 条;清理过期访问日志 0 条。")
|
||||
}
|
||||
|
||||
func TestSystemCleanupHandler_ImplementsTaskHandler(t *testing.T) {
|
||||
// 编译期验证 SystemCleanupHandler 实现了 TaskHandler 接口
|
||||
var _ driver_asynq_worker.TaskHandler = (*SystemCleanupHandler)(nil)
|
||||
assert.Contains(t, result.Message, "系统垃圾清理完成")
|
||||
}
|
||||
|
||||
func TestWarmImageCacheHandlerValidatePayload(t *testing.T) {
|
||||
@@ -252,27 +158,10 @@ func TestWarmImageCacheHandlerValidatePayload(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestWarmImageCacheHandlerExecute(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
cache := diskcache.GetGlobalCache()
|
||||
if err := cache.Clear(); err != nil {
|
||||
t.Fatalf("Clear() before test returned error: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
if err := cache.Clear(); err != nil {
|
||||
t.Errorf("Clear() after test returned error: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
testDir := t.TempDir()
|
||||
ctx := context.Background()
|
||||
active := objectstore.DefaultConfig()
|
||||
active.Local.Root = testDir
|
||||
if err := objectstore.SaveActiveConfig(ctx, active); err != nil {
|
||||
t.Fatalf("SaveActiveConfig() returned error: %v", err)
|
||||
}
|
||||
|
||||
firstPath := filepath.Join(testDir, "first.png")
|
||||
secondPath := filepath.Join(testDir, "second.jpg")
|
||||
writeTaskTestPNG(t, firstPath, color.RGBA{R: 255, A: 255})
|
||||
@@ -341,13 +230,16 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
|
||||
|
||||
for i := range records[:2] {
|
||||
key := filesrv.ImageCompressionCacheKey(&records[i], shared.ImageQualityLow)
|
||||
got, err := cache.Get(key)
|
||||
got, hit, err := filesrv.EnsureCompressedImageCache(context.Background(), &records[i], shared.ImageQualityLow)
|
||||
if err != nil {
|
||||
t.Errorf("cache.Get(%q) returned error: %v", key, err)
|
||||
t.Errorf("EnsureCompressedImageCache(%q) returned error: %v", key, err)
|
||||
continue
|
||||
}
|
||||
if !hit {
|
||||
t.Errorf("expected cache hit for %q", key)
|
||||
}
|
||||
if len(got) == 0 {
|
||||
t.Errorf("cache.Get(%q) returned empty WebP data", key)
|
||||
t.Errorf("EnsureCompressedImageCache(%q) returned empty WebP data", key)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -360,11 +252,6 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestWarmImageCacheHandlerImplementsTaskInterfaces(t *testing.T) {
|
||||
var _ driver_asynq_worker.TaskHandler = (*WarmImageCacheHandler)(nil)
|
||||
var _ driver_asynq_worker.PayloadValidator = (*WarmImageCacheHandler)(nil)
|
||||
}
|
||||
|
||||
func writeTaskTestPNG(t *testing.T, path string, fill color.RGBA) {
|
||||
t.Helper()
|
||||
|
||||
|
||||
@@ -15,9 +15,10 @@ import (
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
"github.com/deepteams/webp"
|
||||
_ "golang.org/x/image/webp" // Register WebP decoder for image.Decode
|
||||
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
// ValidateS3Key validates an S3 object key for safety.
|
||||
|
||||
Reference in New Issue
Block a user