refactor(api): unify error handling and split upload/oauth god modules

Replace c.JSON(200, response.Err) and middleware gin.H bypasses with
response.Abort* helpers so errors flow through Gin Error chain and
ErrorHandlerMiddleware for OTel trace correlation.

Split oauth/sources.go into domain-focused files and decompose upload
into handler/filesrv/stats/task/cache/storage/util subpackages with a
root facade preserving existing import paths.
This commit is contained in:
ryan
2026-06-18 10:51:25 +08:00
parent f826807cf1
commit 9af84c8ed6
62 changed files with 2040 additions and 1696 deletions
@@ -1,7 +1,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
// Package cache provides in-process upload access-control caches.
package cache
import (
"context"
@@ -10,32 +11,19 @@ import (
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
)
const accessCacheTTL = 5 * time.Second
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
type migrationAccessState struct {
readOnly bool
target storage.Config
hasTarget bool
targetErr error
loadErr error
}
var (
accessCacheOnce sync.Once
migrationAccessMu sync.RWMutex
migrationAccessCached migrationAccessState
migrationAccessValid bool
migrationAccessCheckedAt time.Time
fileAccessWhitelistMu sync.RWMutex
fileAccessWhitelistMu sync.RWMutex
fileAccessWhitelistTypes map[string]struct{}
fileAccessWhitelistValid bool
fileAccessWhitelistCheckedAt time.Time
@@ -43,9 +31,7 @@ var (
// ResetAccessCaches clears in-process upload access caches.
func ResetAccessCaches() {
migrationAccessMu.Lock()
migrationAccessValid = false
migrationAccessMu.Unlock()
uploadstorage.ResetMigrationAccessCache()
fileAccessWhitelistMu.Lock()
fileAccessWhitelistValid = false
@@ -85,62 +71,18 @@ func startAccessCacheInvalidationListener() {
}()
}
func loadMigrationAccessState(ctx context.Context) migrationAccessState {
ensureAccessCacheListener()
migrationAccessMu.RLock()
if migrationAccessValid && time.Since(migrationAccessCheckedAt) < accessCacheTTL {
state := migrationAccessCached
migrationAccessMu.RUnlock()
return state
}
migrationAccessMu.RUnlock()
migrationAccessMu.Lock()
defer migrationAccessMu.Unlock()
if migrationAccessValid && time.Since(migrationAccessCheckedAt) < accessCacheTTL {
return migrationAccessCached
}
migrationAccessCached = buildMigrationAccessState(ctx)
migrationAccessValid = true
migrationAccessCheckedAt = time.Now()
return migrationAccessCached
}
func buildMigrationAccessState(ctx context.Context) migrationAccessState {
execution, ok, err := latestStorageMigrationExecution(ctx)
if err != nil {
return migrationAccessState{loadErr: err, readOnly: true}
}
if !ok {
return migrationAccessState{}
}
state := migrationAccessState{
readOnly: execution.Status != model.TaskExecutionStatusSucceeded,
}
if execution.Status == model.TaskExecutionStatusSucceeded {
return state
}
target, err := parseMigrationTargetConfig(ctx, []byte(execution.Payload))
if err != nil {
state.targetErr = err
return state
}
state.target = target
state.hasTarget = true
return state
// IsFilePublic reports whether uploadType is in the public access whitelist.
func IsFilePublic(ctx context.Context, uploadType string) bool {
whitelist := loadFileAccessWhitelist(ctx)
_, ok := whitelist[strings.ToLower(uploadType)]
return ok
}
func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} {
ensureAccessCacheListener()
fileAccessWhitelistMu.RLock()
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < accessCacheTTL {
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
types := fileAccessWhitelistTypes
fileAccessWhitelistMu.RUnlock()
return types
@@ -150,7 +92,7 @@ func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} {
fileAccessWhitelistMu.Lock()
defer fileAccessWhitelistMu.Unlock()
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < accessCacheTTL {
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
return fileAccessWhitelistTypes
}
@@ -172,7 +114,7 @@ func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} {
func parseFileAccessWhitelist(ctx context.Context) []string {
var sc model.SystemConfig
if err := sc.GetByKey(ctx, model.ConfigKeyFileAccessWhitelist); err != nil || sc.Value == "" {
return []string{defaultPublicUploadType}
return []string{shared.DefaultPublicUploadType}
}
var whitelist []string
@@ -182,7 +124,7 @@ func parseFileAccessWhitelist(ctx context.Context) []string {
whitelist = parseCommaSeparatedWhitelist(sc.Value)
if len(whitelist) == 0 {
return []string{defaultPublicUploadType}
return []string{shared.DefaultPublicUploadType}
}
return whitelist
}
@@ -1,13 +1,15 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package cache
import (
"context"
"testing"
"time"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
@@ -19,14 +21,14 @@ func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
ResetAccessCaches()
ctx := context.Background()
first := loadMigrationAccessState(ctx)
second := loadMigrationAccessState(ctx)
first := uploadstorage.LoadMigrationAccessState(ctx)
second := uploadstorage.LoadMigrationAccessState(ctx)
if first.readOnly != second.readOnly {
t.Fatalf("readOnly mismatch: first=%v second=%v", first.readOnly, second.readOnly)
if first.ReadOnly != second.ReadOnly {
t.Fatalf("readOnly mismatch: first=%v second=%v", first.ReadOnly, second.ReadOnly)
}
if first.hasTarget != second.hasTarget {
t.Fatalf("hasTarget mismatch: first=%v second=%v", first.hasTarget, second.hasTarget)
if first.HasTarget != second.HasTarget {
t.Fatalf("hasTarget mismatch: first=%v second=%v", first.HasTarget, second.HasTarget)
}
}
@@ -36,13 +38,13 @@ func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
ResetAccessCaches()
ctx := context.Background()
if !isFilePublic(ctx, "avatar") {
if !IsFilePublic(ctx, "avatar") {
t.Fatal("expected avatar to be public by default")
}
if isFilePublic(ctx, "attachment") {
if IsFilePublic(ctx, "attachment") {
t.Fatal("expected attachment to be private by default")
}
if !isFilePublic(ctx, "AVATAR") {
if !IsFilePublic(ctx, "AVATAR") {
t.Fatal("expected whitelist lookup to be case-insensitive")
}
}
@@ -53,7 +55,7 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
ResetAccessCaches()
ctx := context.Background()
if !isFilePublic(ctx, "avatar") {
if !IsFilePublic(ctx, "avatar") {
t.Fatal("expected seeded avatar whitelist before reset")
}
@@ -68,12 +70,13 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
if err := db.HSetJSON(ctx, model.SystemConfigRedisHashKey, model.ConfigKeyFileAccessWhitelist, &sc); err != nil {
t.Fatalf("refresh whitelist redis cache: %v", err)
}
model.ResetSystemConfigRAMCacheForTest()
ResetAccessCaches()
if !isFilePublic(ctx, "attachment") {
if !IsFilePublic(ctx, "attachment") {
t.Fatal("expected attachment to be public after whitelist refresh")
}
if isFilePublic(ctx, "avatar") {
if IsFilePublic(ctx, "avatar") {
t.Fatal("expected avatar to be private after whitelist refresh")
}
}
@@ -87,11 +90,11 @@ func TestAccessCacheTTLExpires(t *testing.T) {
_ = loadFileAccessWhitelist(ctx)
fileAccessWhitelistMu.Lock()
fileAccessWhitelistCheckedAt = time.Now().Add(-accessCacheTTL - time.Second)
fileAccessWhitelistCheckedAt = time.Now().Add(-time.Duration(shared.AccessCacheTTL)*time.Second - time.Second)
fileAccessWhitelistMu.Unlock()
// Should still work after TTL by reloading from config.
if !isFilePublic(ctx, "avatar") {
if !IsFilePublic(ctx, "avatar") {
t.Fatal("expected whitelist reload after TTL expiration")
}
}
-20
View File
@@ -1,20 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
import "github.com/Rain-kl/Wavelet/internal/storage"
const (
maxUploadSize = 32 * 1024 * 1024 // 32MB
detectContentBytes = 512 // http.DetectContentType 需要的最小字节数
uploadDirPerm = 0755 // 上传目录权限
uploadFilePerm = 0644 // 上传文件权限
imageQualityLow = "low"
imageQualityMedium = "medium"
imageQualityHigh = "high"
imageQualityOrigin = "origin"
storageDriverLocal = string(storage.DriverLocal)
defaultPublicUploadType = "avatar"
fileStatsTrendDays = 7
)
+29 -32
View File
@@ -5,37 +5,34 @@
// Package upload 提供文件上传与下载功能
package upload
import "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
// 文件管理常量
const (
ErrNoFileSelected = "请选择要上传的文件"
ErrUnsupportedFormat = "只支持 JPG、PNG、WEBP 格式的图片"
ErrProcessFileFailed = "处理文件失败"
ErrSaveFileFailed = "保存文件失败"
ErrOpenFileFailed = "打开文件失败"
ErrSaveUploadRecordFailed = "保存上传记录失败"
ErrGenericFileTooLarge = "文件大小不能超过 32MB"
ErrFileContentExtensionMismatch = "文件内容与扩展名不匹配,可能包含安全风险"
ErrFileValidationFailed = "文件校验失败"
ErrInvalidMetadataJSON = "元数据 JSON 格式不合法"
ErrInvalidFileID = "无效的文件 ID"
ErrQueryUploadRecordFailed = "查询文件记录失败"
ErrInvalidBatchDownloadRequest = "参数绑定失败,请传入有效的文件 ID 数组"
ErrInvalidIDValueFormat = "无效的 ID 值: %s"
ErrRetrieveUploadRecordsFailed = "检索文件记录失败"
ErrNoValidFilesForArchive = "没有找到任何有效的文件记录进行打包"
ErrInvalidParams = "参数错误"
ErrQueryFileCountFailed = "查询文件数量失败"
ErrQueryFileListFailed = "查询文件列表失败"
ErrDeleteFileFailed = "删除文件失败"
ErrStorageReadOnly = "存储迁移维护中,当前仅允许读取文件"
ErrS3KeyRequired = "s3 key must not be empty"
ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d"
ErrS3KeyStartsWithSlash = "s3 key must not start with /"
ErrS3KeyContainsNullBytes = "s3 key must not contain null bytes"
ErrQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w"
errImageCacheWarmupPayloadRequired = "图片缓存预热参数不能为空"
errInvalidImageCacheWarmupPayload = "图片缓存预热参数格式无效: %w"
errInvalidImageCacheWarmupQuality = "图片质量仅支持 low、medium、high"
errParseImageCacheWarmupPayload = "解析图片缓存预热参数失败: %w"
errQueryImagesForCacheWarmup = "查询待预热图片失败: %w"
)
ErrNoFileSelected = shared.ErrNoFileSelected
ErrUnsupportedFormat = shared.ErrUnsupportedFormat
ErrProcessFileFailed = shared.ErrProcessFileFailed
ErrSaveFileFailed = shared.ErrSaveFileFailed
ErrOpenFileFailed = shared.ErrOpenFileFailed
ErrSaveUploadRecordFailed = shared.ErrSaveUploadRecordFailed
ErrGenericFileTooLarge = shared.ErrGenericFileTooLarge
ErrFileContentExtensionMismatch = shared.ErrFileContentExtensionMismatch
ErrFileValidationFailed = shared.ErrFileValidationFailed
ErrInvalidMetadataJSON = shared.ErrInvalidMetadataJSON
ErrInvalidFileID = shared.ErrInvalidFileID
ErrQueryUploadRecordFailed = shared.ErrQueryUploadRecordFailed
ErrInvalidBatchDownloadRequest = shared.ErrInvalidBatchDownloadRequest
ErrInvalidIDValueFormat = shared.ErrInvalidIDValueFormat
ErrRetrieveUploadRecordsFailed = shared.ErrRetrieveUploadRecordsFailed
ErrNoValidFilesForArchive = shared.ErrNoValidFilesForArchive
ErrInvalidParams = shared.ErrInvalidParams
ErrQueryFileCountFailed = shared.ErrQueryFileCountFailed
ErrQueryFileListFailed = shared.ErrQueryFileListFailed
ErrDeleteFileFailed = shared.ErrDeleteFileFailed
ErrStorageReadOnly = shared.ErrStorageReadOnly
ErrS3KeyRequired = shared.ErrS3KeyRequired
ErrS3KeyTooLongFormat = shared.ErrS3KeyTooLongFormat
ErrS3KeyStartsWithSlash = shared.ErrS3KeyStartsWithSlash
ErrS3KeyContainsNullBytes = shared.ErrS3KeyContainsNullBytes
ErrQueryUnusedUploadsFailed = shared.ErrQueryUnusedUploadsFailed
)
+86
View File
@@ -0,0 +1,86 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
import (
"github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/handler"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/task"
)
// HTTP handlers
var (
UploadFile = handler.UploadFile
DownloadFile = handler.DownloadFile
BatchDownloadFiles = handler.BatchDownloadFiles
ListFiles = handler.ListFiles
DeleteFile = handler.DeleteFile
GetDistinctUploadTypes = handler.GetDistinctUploadTypes
ListMyFiles = handler.ListMyFiles
DeleteMyFile = handler.DeleteMyFile
UpdateMyFile = handler.UpdateMyFile
GetFileStats = handler.GetFileStats
ServeFileByID = filesrv.ServeFileByID
)
// Cache management
var (
ResetAccessCaches = cache.ResetAccessCaches
PublishAccessCacheInvalidation = cache.PublishAccessCacheInvalidation
)
// Stats
var (
ApplyUploadStatsAdd = uploadstats.ApplyUploadStatsAdd
ApplyUploadStatsRemove = uploadstats.ApplyUploadStatsRemove
RebuildUploadStats = uploadstats.RebuildUploadStats
)
// Utilities
var (
CompressImageToWebP = util.CompressImageToWebP
ValidateS3Key = util.ValidateS3Key
)
// Task identifiers and metadata
const (
StorageMigrationTask = uploadtask.StorageMigrationTask
SystemCleanupTask = uploadtask.SystemCleanupTask
WarmImageCacheTask = uploadtask.WarmImageCacheTask
)
var (
// StorageMigrationMeta describes the storage migration async task.
StorageMigrationMeta = uploadtask.StorageMigrationMeta
// SystemCleanupMeta describes the orphaned upload cleanup task.
SystemCleanupMeta = uploadtask.SystemCleanupMeta
// WarmImageCacheMeta describes the image compression cache warmup task.
WarmImageCacheMeta = uploadtask.WarmImageCacheMeta
)
// MigrationHandler executes storage migration tasks.
type MigrationHandler = uploadtask.MigrationHandler
// SystemCleanupHandler removes orphaned upload files.
type SystemCleanupHandler = uploadtask.SystemCleanupHandler
// WarmImageCacheHandler pre-warms compressed image caches.
type WarmImageCacheHandler = uploadtask.WarmImageCacheHandler
// WarmImageCachePayload is the payload for image cache warmup tasks.
type WarmImageCachePayload = uploadtask.WarmImageCachePayload
// Ensure task handler types implement required interfaces.
var (
_ task.TaskHandler = (*MigrationHandler)(nil)
_ task.TaskHandler = (*SystemCleanupHandler)(nil)
_ interface {
task.TaskHandler
ValidatePayload([]byte) ([]byte, error)
} = (*WarmImageCacheHandler)(nil)
)
@@ -2,7 +2,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
// Package filesrv serves uploaded files with access control and image compression.
package filesrv
import (
"bytes"
@@ -15,11 +16,16 @@ import (
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
apputil "github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"golang.org/x/sync/singleflight"
@@ -34,6 +40,15 @@ type compressedImageCacheResult struct {
err error
}
type fileTypeCategory string
const (
fileTypeImage fileTypeCategory = "image"
fileTypeVideo fileTypeCategory = "video"
fileTypeAudio fileTypeCategory = "audio"
fileTypeOther fileTypeCategory = "other"
)
// ServeFileByID 根据 ID 获取并提供已上传的文件
// @Summary 获取已上传文件
// @Description 根据文件 ID 获取并提供已上传的临时或正式文件,若配置了缓存则优先走本地缓存,否则从 S3 等后端存储读取并流式返回
@@ -48,7 +63,7 @@ type compressedImageCacheResult struct {
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /f/{id} [get]
func ServeFileByID(c *gin.Context) {
upload, err := getUploadRecordByID(c)
upload, err := GetUploadRecordByID(c)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.AbortWithStatus(http.StatusNotFound)
@@ -62,18 +77,16 @@ func ServeFileByID(c *gin.Context) {
return
}
// 校验业务白名单与访问权限
if err := checkFileAccessPermission(c, upload); err != nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
if err := CheckFileAccessPermission(c, upload); err != nil {
response.AbortUnauthorized(c, common.UnAuthorized)
return
}
ServeUpload(c, upload)
}
// getUploadRecordByID 从请求路径参数中解析文件 ID 并从数据库中检索处于 Pending 或 Used 状态的上传记录。
// 同时会自动设置通用的安全响应头。
func getUploadRecordByID(c *gin.Context) (*model.Upload, error) {
// GetUploadRecordByID 从请求路径参数中解析文件 ID 并从数据库中检索处于 Pending 或 Used 状态的上传记录。
func GetUploadRecordByID(c *gin.Context) (*model.Upload, error) {
c.Header("X-Content-Type-Options", "nosniff")
c.Header("Content-Security-Policy", "sandbox")
@@ -93,22 +106,11 @@ func getUploadRecordByID(c *gin.Context) (*model.Upload, error) {
return &upload, nil
}
// fileTypeCategory 定义文件的大类,便于未来扩展不同的处理方式
type fileTypeCategory string
const (
fileTypeImage fileTypeCategory = "image"
fileTypeVideo fileTypeCategory = "video"
fileTypeAudio fileTypeCategory = "audio"
fileTypeOther fileTypeCategory = "other"
)
// getFileTypeCategory 判断并返回文件的大类
func getFileTypeCategory(upload *model.Upload) fileTypeCategory {
mime := strings.ToLower(upload.MimeType)
ext := strings.ToLower(upload.Extension)
if strings.HasPrefix(mime, "image/") || isImageExtension(ext) {
if strings.HasPrefix(mime, "image/") || util.IsImageExtension(ext) {
return fileTypeImage
}
if strings.HasPrefix(mime, "video/") {
@@ -120,32 +122,27 @@ func getFileTypeCategory(upload *model.Upload) fileTypeCategory {
return fileTypeOther
}
// ServeUpload 将已存在的文件内容读取并流式响应给客户端,支持本地和 S3/CDN 驱动,并可选支持 WebP 图片压缩与本地缓存。
// ServeUpload 将已存在的文件内容读取并流式响应给客户端。
func ServeUpload(c *gin.Context, upload *model.Upload) {
// 设置通用的缓存控制响应头
setCacheHeaders(c, upload)
category := getFileTypeCategory(upload)
quality := normalizeImageQuality(c.Query("quality"))
quality := util.NormalizeImageQuality(c.Query("quality"))
switch category {
case fileTypeImage:
// 如果是图片且不是原图质量,则提供压缩优化后的图片预览
if quality != imageQualityOrigin {
if quality != shared.ImageQualityOrigin {
serveCompressedImage(c, upload, quality)
return
}
// 请求原图质量时,退化到默认提供原文件
fallthrough
default:
// 默认提供原文件,并执行协商缓存校验
serveOriginalWithConditionalCheck(c, upload)
}
}
func setCacheHeaders(c *gin.Context, upload *model.Upload) {
if isFilePublic(c.Request.Context(), upload.Type) {
if cache.IsFilePublic(c.Request.Context(), upload.Type) {
c.Header("Cache-Control", "public, max-age=31536000")
} else {
c.Header("Cache-Control", "private, no-cache")
@@ -173,7 +170,7 @@ func serveCompressedImage(c *gin.Context, upload *model.Upload, quality string)
return
}
webpBytes, _, err := ensureCompressedImageCache(c.Request.Context(), upload, quality)
webpBytes, _, err := EnsureCompressedImageCache(c.Request.Context(), upload, quality)
if err != nil {
if len(webpBytes) > 0 {
logger.WarnF(c.Request.Context(), "failed to cache compressed image: %v", err)
@@ -188,14 +185,15 @@ func serveCompressedImage(c *gin.Context, upload *model.Upload, quality string)
c.Data(http.StatusOK, "image/webp", webpBytes)
}
func ensureCompressedImageCache(
// EnsureCompressedImageCache returns cached or freshly generated WebP bytes for an upload.
func EnsureCompressedImageCache(
ctx context.Context,
upload *model.Upload,
quality string,
) ([]byte, bool, error) {
cache := diskcache.GetGlobalCache()
cacheKey := imageCompressionCacheKey(upload, quality)
webpBytes, err := cache.Get(cacheKey)
cacheStore := diskcache.GetGlobalCache()
cacheKey := ImageCompressionCacheKey(upload, quality)
webpBytes, err := cacheStore.Get(cacheKey)
if err == nil {
return webpBytes, true, nil
}
@@ -220,9 +218,9 @@ func generateCompressedImageCache(
quality string,
cacheKey string,
) (compressedImageCacheResult, error) {
cache := diskcache.GetGlobalCache()
cacheStore := diskcache.GetGlobalCache()
webpBytes, err := cache.Get(cacheKey)
webpBytes, err := cacheStore.Get(cacheKey)
if err == nil {
return compressedImageCacheResult{bytes: webpBytes, cached: true}, nil
}
@@ -235,12 +233,12 @@ func generateCompressedImageCache(
return compressedImageCacheResult{}, fmt.Errorf("read original image: %w", err)
}
webpBytes, err = CompressImageToWebP(bytes.NewReader(origBytes), quality)
webpBytes, err = util.CompressImageToWebP(bytes.NewReader(origBytes), quality)
if err != nil {
return compressedImageCacheResult{}, fmt.Errorf("compress image to WebP: %w", err)
}
if err := cache.Set(cacheKey, webpBytes, diskcache.NoExpiration); err != nil {
if err := cacheStore.Set(cacheKey, webpBytes, diskcache.NoExpiration); err != nil {
return compressedImageCacheResult{
bytes: webpBytes,
err: fmt.Errorf("write compressed image cache: %w", err),
@@ -250,7 +248,8 @@ func generateCompressedImageCache(
return compressedImageCacheResult{bytes: webpBytes}, nil
}
func imageCompressionCacheKey(upload *model.Upload, quality string) string {
// ImageCompressionCacheKey returns the disk cache key for a compressed upload image.
func ImageCompressionCacheKey(upload *model.Upload, quality string) string {
return fmt.Sprintf(
"upload_webp_v1_%d_%d_%d_%s_%s",
upload.ID,
@@ -261,18 +260,8 @@ func imageCompressionCacheKey(upload *model.Upload, quality string) string {
)
}
func normalizeImageQuality(quality string) string {
switch strings.ToLower(quality) {
case imageQualityLow, imageQualityMedium, imageQualityHigh:
return strings.ToLower(quality)
default:
return imageQualityOrigin
}
}
// serveOriginal 原始文件的流式响应逻辑
func serveOriginal(c *gin.Context, upload *model.Upload) {
obj, err := openStoredObject(c.Request.Context(), upload)
obj, err := uploadstorage.OpenStoredObject(c.Request.Context(), upload)
if err != nil {
c.AbortWithStatus(http.StatusNotFound)
return
@@ -281,9 +270,8 @@ func serveOriginal(c *gin.Context, upload *model.Upload) {
c.DataFromReader(http.StatusOK, obj.ContentLength, obj.ContentType, obj.Body, nil)
}
// getOriginalFileBytes 获取原始文件所有字节
func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, error) {
obj, err := openStoredObject(ctx, upload)
obj, err := uploadstorage.OpenStoredObject(ctx, upload)
if err != nil {
return nil, err
}
@@ -291,17 +279,10 @@ func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, er
return io.ReadAll(obj.Body)
}
// isFilePublic 校验文件类型是否在公开访问白名单中
func isFilePublic(ctx context.Context, uploadType string) bool {
whitelist := loadFileAccessWhitelist(ctx)
_, ok := whitelist[strings.ToLower(uploadType)]
return ok
}
func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
var currUser *model.User
var err error
if u, ok := util.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil {
if u, ok := apputil.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil {
currUser = u
} else {
currUser, err = oauth.GetUserFromRequest(c)
@@ -318,21 +299,18 @@ func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
return nil
}
// checkFileAccessPermission 校验文件是否可以被当前请求访问
func checkFileAccessPermission(c *gin.Context, upload *model.Upload) error {
// 1. 私有文件校验(优先级高于当前白名单逻辑)
// CheckFileAccessPermission 校验文件是否可以被当前请求访问
func CheckFileAccessPermission(c *gin.Context, upload *model.Upload) error {
if upload.AccessMode == 0 {
return checkPrivateFileOwner(c, upload.UserID)
}
// 2. 如果类型为公开的则再进行校验白名单
if !isFilePublic(c.Request.Context(), upload.Type) {
// 必须进行鉴权
if _, ok := util.GetFromContext[*model.User](c, oauth.UserObjKey); !ok {
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
if _, ok := apputil.GetFromContext[*model.User](c, oauth.UserObjKey); !ok {
if _, err := oauth.GetUserFromRequest(c); err != nil {
return err
}
}
}
return nil
}
}
@@ -2,7 +2,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package filesrv
import (
"bytes"
@@ -15,7 +15,11 @@ import (
"os"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
@@ -27,6 +31,7 @@ import (
func TestServeFileByIDAccessControl(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
cache.ResetAccessCaches()
// Ensure uploads dir is cleaned up
defer func() { _ = os.RemoveAll("uploads") }()
@@ -91,6 +96,7 @@ func TestServeFileByIDAccessControl(t *testing.T) {
// Set up router
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(response.ErrorHandlerMiddleware())
store := cookie.NewStore([]byte("secret"))
r.Use(sessions.Sessions("test_session", store))
r.GET("/f/:id", ServeFileByID)
@@ -151,58 +157,6 @@ func TestServeFileByIDAccessControl(t *testing.T) {
})
}
func TestGetDistinctUploadTypes(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
// Seed some uploads with new custom types
user := model.User{ID: 2222, Username: "test_user_2"}
dbConn.Create(&user)
customUpload := model.Upload{
ID: 9001,
UserID: user.ID,
FileName: "custom.txt",
FilePath: "uploads/custom.txt",
FileSize: 10,
MimeType: "text/plain",
Extension: "txt",
StorageDriver: "local",
Type: "custom_type_xyz",
Status: model.UploadStatusUsed,
}
dbConn.Create(&customUpload)
gin.SetMode(gin.TestMode)
r := gin.New()
r.GET("/api/v1/admin/uploads/types", GetDistinctUploadTypes)
req, _ := http.NewRequest("GET", "/api/v1/admin/uploads/types", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200, got %d", w.Code)
}
var resp struct {
ErrorMsg string `json:"error_msg"`
Data []string `json:"data"`
}
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to parse JSON: %v", err)
}
if resp.ErrorMsg != "" {
t.Fatalf("unexpected error: %s", resp.ErrorMsg)
}
// Verify that only custom_type_xyz is present
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)
}
}
func TestImageCompression(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
@@ -295,7 +249,7 @@ func TestImageCompression(t *testing.T) {
t.Errorf("expected Content-Type image/webp, got %s", w.Header().Get("Content-Type"))
}
cacheKey := imageCompressionCacheKey(&uploadRecord, imageQualityMedium)
cacheKey := ImageCompressionCacheKey(&uploadRecord, shared.ImageQualityMedium)
cachedBytes, err := cache.Get(cacheKey)
if err != nil {
t.Fatalf("disk cache Get(%q) returned error: %v", cacheKey, err)
@@ -376,19 +330,19 @@ func TestNormalizeImageQuality(t *testing.T) {
quality string
want string
}{
{name: imageQualityLow, quality: imageQualityLow, want: imageQualityLow},
{name: imageQualityMedium, quality: imageQualityMedium, want: imageQualityMedium},
{name: imageQualityHigh, quality: imageQualityHigh, want: imageQualityHigh},
{name: shared.ImageQualityLow, quality: shared.ImageQualityLow, want: shared.ImageQualityLow},
{name: shared.ImageQualityMedium, quality: shared.ImageQualityMedium, want: shared.ImageQualityMedium},
{name: shared.ImageQualityHigh, quality: shared.ImageQualityHigh, want: shared.ImageQualityHigh},
{name: "origin", quality: "origin", want: "origin"},
{name: "uppercase", quality: "LOW", want: imageQualityLow},
{name: "uppercase", quality: "LOW", want: shared.ImageQualityLow},
{name: "empty", quality: "", want: "origin"},
{name: "invalid", quality: "maximum", want: "origin"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := normalizeImageQuality(tt.quality); got != tt.want {
t.Errorf("normalizeImageQuality(%q) = %q, want %q", tt.quality, got, tt.want)
if got := util.NormalizeImageQuality(tt.quality); got != tt.want {
t.Errorf("NormalizeImageQuality(%q) = %q, want %q", tt.quality, got, tt.want)
}
})
}
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package handler
import (
"errors"
@@ -11,10 +11,13 @@ import (
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
apputil "github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
@@ -56,7 +59,7 @@ func ListFiles(c *gin.Context) {
var req listFilesRequest
if err := c.ShouldBindQuery(&req); err != nil {
c.JSON(http.StatusOK, response.Err(ErrInvalidParams))
response.AbortBadRequest(c, shared.ErrInvalidParams)
return
}
if req.Page <= 0 {
@@ -84,14 +87,14 @@ func ListFiles(c *gin.Context) {
var total int64
if err := query.Count(&total).Error; err != nil {
c.JSON(http.StatusOK, response.Err(ErrQueryFileCountFailed))
response.AbortBadRequest(c, shared.ErrQueryFileCountFailed)
return
}
var items []model.Upload
offset := (req.Page - 1) * req.PageSize
if err := query.Order("created_at DESC").Offset(offset).Limit(req.PageSize).Find(&items).Error; err != nil {
c.JSON(http.StatusOK, response.Err(ErrQueryFileListFailed))
response.AbortBadRequest(c, shared.ErrQueryFileListFailed)
return
}
@@ -116,14 +119,14 @@ func ListFiles(c *gin.Context) {
// @Router /api/v1/admin/uploads/{id} [delete]
func DeleteFile(c *gin.Context) {
ctx := c.Request.Context()
if StorageReadOnly(ctx) {
c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly))
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
return
}
uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusOK, response.Err(ErrInvalidFileID))
response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
@@ -133,14 +136,14 @@ func DeleteFile(c *gin.Context) {
c.AbortWithStatus(http.StatusNotFound)
return
}
c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed))
response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
return
}
if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil {
c.JSON(http.StatusOK, response.Err(ErrDeleteFileFailed))
response.AbortBadRequest(c, shared.ErrDeleteFileFailed)
return
}
recordUploadStatsRemove(ctx, &upload)
uploadstats.RecordUploadStatsRemove(ctx, &upload)
c.JSON(http.StatusOK, response.OKNil())
}
@@ -161,7 +164,7 @@ func GetDistinctUploadTypes(c *gin.Context) {
Where("type IS NOT NULL AND type != ''").
Distinct().
Pluck("type", &dbTypes).Error; err != nil {
c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
response.AbortInternal(c, err.Error())
return
}
sort.Strings(dbTypes)
@@ -198,12 +201,12 @@ type listMyFilesResponse struct {
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/upload/my [get]
func ListMyFiles(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
var req listMyFilesRequest
if err := c.ShouldBindQuery(&req); err != nil {
c.JSON(http.StatusOK, response.Err(ErrInvalidParams))
response.AbortBadRequest(c, shared.ErrInvalidParams)
return
}
if req.Page <= 0 {
@@ -228,14 +231,14 @@ func ListMyFiles(c *gin.Context) {
var total int64
if err := query.Count(&total).Error; err != nil {
c.JSON(http.StatusOK, response.Err(ErrQueryFileCountFailed))
response.AbortBadRequest(c, shared.ErrQueryFileCountFailed)
return
}
var items []model.Upload
offset := (req.Page - 1) * req.PageSize
if err := query.Order("created_at DESC").Offset(offset).Limit(req.PageSize).Find(&items).Error; err != nil {
c.JSON(http.StatusOK, response.Err(ErrQueryFileListFailed))
response.AbortBadRequest(c, shared.ErrQueryFileListFailed)
return
}
@@ -259,16 +262,16 @@ func ListMyFiles(c *gin.Context) {
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [delete]
func DeleteMyFile(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
if StorageReadOnly(ctx) {
c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly))
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
return
}
uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusOK, response.Err(ErrInvalidFileID))
response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
@@ -278,7 +281,7 @@ func DeleteMyFile(c *gin.Context) {
c.AbortWithStatus(http.StatusNotFound)
return
}
c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed))
response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
return
}
@@ -288,10 +291,10 @@ func DeleteMyFile(c *gin.Context) {
}
if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil {
c.JSON(http.StatusOK, response.Err(ErrDeleteFileFailed))
response.AbortBadRequest(c, shared.ErrDeleteFileFailed)
return
}
recordUploadStatsRemove(ctx, &upload)
uploadstats.RecordUploadStatsRemove(ctx, &upload)
c.JSON(http.StatusOK, response.OKNil())
}
@@ -314,22 +317,22 @@ type updateMyFileRequest struct {
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [put]
func UpdateMyFile(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
if StorageReadOnly(ctx) {
c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly))
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
return
}
uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusOK, response.Err(ErrInvalidFileID))
response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
var req updateMyFileRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusOK, response.Err(ErrInvalidParams))
response.AbortBadRequest(c, shared.ErrInvalidParams)
return
}
@@ -339,7 +342,7 @@ func UpdateMyFile(c *gin.Context) {
c.AbortWithStatus(http.StatusNotFound)
return
}
c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed))
response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
return
}
@@ -358,10 +361,10 @@ func UpdateMyFile(c *gin.Context) {
if len(updates) > 0 {
if err := db.DB(ctx).Model(&upload).Updates(updates).Error; err != nil {
c.JSON(http.StatusOK, response.Err("更新文件记录失败"))
response.AbortBadRequest(c, "更新文件记录失败")
return
}
}
c.JSON(http.StatusOK, response.OK(upload))
}
}
@@ -0,0 +1,65 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
)
func TestGetDistinctUploadTypes(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
user := model.User{ID: 2222, Username: "test_user_2"}
dbConn.Create(&user)
customUpload := model.Upload{
ID: 9001,
UserID: user.ID,
FileName: "custom.txt",
FilePath: "uploads/custom.txt",
FileSize: 10,
MimeType: "text/plain",
Extension: "txt",
StorageDriver: "local",
Type: "custom_type_xyz",
Status: model.UploadStatusUsed,
}
dbConn.Create(&customUpload)
gin.SetMode(gin.TestMode)
r := gin.New()
r.GET("/api/v1/admin/uploads/types", GetDistinctUploadTypes)
req, _ := http.NewRequest("GET", "/api/v1/admin/uploads/types", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200, got %d", w.Code)
}
var resp struct {
ErrorMsg string `json:"error_msg"`
Data []string `json:"data"`
}
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to parse JSON: %v", err)
}
if resp.ErrorMsg != "" {
t.Fatalf("unexpected error: %s", resp.ErrorMsg)
}
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)
}
}
@@ -2,9 +2,11 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
// Package handler provides upload HTTP API handlers.
package handler
import ("archive/zip"
import (
"archive/zip"
"bytes"
"context"
"crypto/sha256"
@@ -22,16 +24,21 @@ import ("archive/zip"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/util"
apputil "github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
type batchDownloadRequest struct {
@@ -59,59 +66,53 @@ func UploadFile(c *gin.Context) {
c.Header("X-Content-Type-Options", "nosniff")
c.Header("Content-Security-Policy", "sandbox")
// 限制请求体大小以防止 DoS
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, maxUploadSize)
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize)
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
header, err := c.FormFile("file")
if err != nil {
c.JSON(http.StatusOK, response.Err(ErrNoFileSelected))
response.AbortBadRequest(c, shared.ErrNoFileSelected)
return
}
file, err := header.Open()
if err != nil {
c.JSON(http.StatusOK, response.Err(ErrOpenFileFailed))
response.AbortBadRequest(c, shared.ErrOpenFileFailed)
return
}
defer func() { _ = file.Close() }()
// 校验大小
if header.Size > maxUploadSize {
c.JSON(http.StatusOK, response.Err(ErrGenericFileTooLarge))
if header.Size > shared.MaxUploadSize {
response.AbortBadRequest(c, shared.ErrGenericFileTooLarge)
return
}
// 2. 提取文件基本元数据
origName := header.Filename
ext := strings.ToLower(strings.TrimPrefix(filepath.Ext(origName), "."))
if ext == "" {
ext = "bin"
}
// 3. 校验文件后缀是否在允许的系统配置列表中
if errMsg := validateUploadExtension(ctx, ext); errMsg != "" {
c.JSON(http.StatusOK, response.Err(errMsg))
response.AbortBadRequest(c, errMsg)
return
}
// 4. 读取文件并计算 Hash
hashWriter := sha256.New()
var buf bytes.Buffer
size, err := io.Copy(&buf, io.TeeReader(file, hashWriter))
if err != nil {
c.JSON(http.StatusOK, response.Err(ErrProcessFileFailed))
response.AbortBadRequest(c, shared.ErrProcessFileFailed)
return
}
fileHash := hex.EncodeToString(hashWriter.Sum(nil))
mimeType := detectMimeType(&buf, header, size)
// 校验真实 MIME Type 是否与常见图片扩展名匹配,防止 Polyglot / HTML 注入攻击
if isImageExtension(ext) && !strings.HasPrefix(mimeType, "image/") {
c.JSON(http.StatusOK, response.Err(ErrFileContentExtensionMismatch))
if util.IsImageExtension(ext) && !strings.HasPrefix(mimeType, "image/") {
response.AbortBadRequest(c, shared.ErrFileContentExtensionMismatch)
return
}
@@ -119,38 +120,34 @@ func UploadFile(c *gin.Context) {
accessMode, errMsg := resolveUploadAccessMode(c, uploadType)
if errMsg != "" {
c.JSON(http.StatusOK, response.Err(errMsg))
response.AbortBadRequest(c, errMsg)
return
}
// 6. 秒传匹配校验:校验数据库中是否存在相同 Hash 且大小一致的可用文件
handled, lookupErr := tryInstantUpload(ctx, c, currUser, fileHash, size, mimeType, ext, origName, accessMode)
if handled {
return
}
if lookupErr != nil && !errors.Is(lookupErr, gorm.ErrRecordNotFound) {
c.JSON(http.StatusOK, response.Err(ErrFileValidationFailed))
response.AbortBadRequest(c, shared.ErrFileValidationFailed)
return
}
// 7. 解析可选元数据字段
meta, errMsg := parseUploadMetadata(c, mimeType)
if errMsg != "" {
c.JSON(http.StatusOK, response.Err(errMsg))
response.AbortBadRequest(c, errMsg)
return
}
id := idgen.NextUint64ID()
subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext)
// 8. 写入当前活动存储驱动。
storageDriver, subPath, errMsg := storeUploadFile(ctx, subPath, size, mimeType, &buf, &meta)
if errMsg != "" {
c.JSON(http.StatusOK, response.Err(errMsg))
response.AbortBadRequest(c, errMsg)
return
}
// 9. 保存文件记录至数据库
newUpload := model.Upload{
ID: id,
UserID: currUser.ID,
@@ -168,7 +165,7 @@ func UploadFile(c *gin.Context) {
}
if err := saveUploadRecord(ctx, &newUpload, storageDriver, subPath); err != "" {
c.JSON(http.StatusOK, response.Err(err))
response.AbortBadRequest(c, err)
return
}
@@ -189,31 +186,30 @@ func UploadFile(c *gin.Context) {
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/uploads/download/{id} [get]
func DownloadFile(c *gin.Context) {
upload, err := getUploadRecordByID(c)
upload, err := filesrv.GetUploadRecordByID(c)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.AbortWithStatus(http.StatusNotFound)
return
}
if _, ok := err.(*strconv.NumError); ok {
c.JSON(http.StatusOK, response.Err(ErrInvalidFileID))
response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
c.JSON(http.StatusOK, response.Err(ErrQueryUploadRecordFailed))
response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
return
}
// 校验文件访问权限
if err := checkFileAccessPermission(c, upload); err != nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
if err := filesrv.CheckFileAccessPermission(c, upload); err != nil {
response.AbortUnauthorized(c, common.UnAuthorized)
return
}
fileName := upload.FileName
quality := normalizeImageQuality(c.Query("quality"))
isImage := strings.HasPrefix(strings.ToLower(upload.MimeType), "image/") || isImageExtension(strings.ToLower(upload.Extension))
quality := util.NormalizeImageQuality(c.Query("quality"))
isImage := strings.HasPrefix(strings.ToLower(upload.MimeType), "image/") || util.IsImageExtension(strings.ToLower(upload.Extension))
if quality != imageQualityOrigin && isImage {
if quality != shared.ImageQualityOrigin && isImage {
ext := filepath.Ext(fileName)
if ext != "" {
fileName = strings.TrimSuffix(fileName, ext) + ".webp"
@@ -222,9 +218,8 @@ func DownloadFile(c *gin.Context) {
}
}
// 设置下载 Attachment 响应头 (支持 UTF-8 中文文件名转义)
c.Header("Content-Disposition", fmt.Sprintf("attachment; filename*=UTF-8''%s", url.PathEscape(fileName)))
ServeUpload(c, upload)
filesrv.ServeUpload(c, upload)
}
// BatchDownloadFiles 批量打包 ZIP 下载接口
@@ -233,7 +228,7 @@ func DownloadFile(c *gin.Context) {
// @Tags admin
// @Accept json
// @Produce octet-stream
// @Param request body upload.batchDownloadRequest true "包含文件 ID 数组 of string 的请求体"
// @Param request body handler.batchDownloadRequest true "包含文件 ID 数组 of string 的请求体"
// @Security SessionCookie
// @Success 200 {file} file "成功下载打包后的 ZIP"
// @Failure 400 {object} response.Any "参数错误"
@@ -244,52 +239,45 @@ func BatchDownloadFiles(c *gin.Context) {
var req batchDownloadRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusOK, response.Err(ErrInvalidBatchDownloadRequest))
response.AbortBadRequest(c, shared.ErrInvalidBatchDownloadRequest)
return
}
// 转换 ID 列表
var ids []uint64
for _, idStr := range req.IDs {
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
c.JSON(http.StatusOK, response.Err(fmt.Sprintf(ErrInvalidIDValueFormat, idStr)))
response.AbortBadRequest(c, fmt.Sprintf(shared.ErrInvalidIDValueFormat, idStr))
return
}
ids = append(ids, id)
}
// 查库获取所有匹配且正常的文件记录
var uploads []model.Upload
if err := db.DB(ctx).Where("id IN ? AND status IN (?, ?)", ids, model.UploadStatusPending, model.UploadStatusUsed).Find(&uploads).Error; err != nil {
c.JSON(http.StatusOK, response.Err(ErrRetrieveUploadRecordsFailed))
response.AbortBadRequest(c, shared.ErrRetrieveUploadRecordsFailed)
return
}
if len(uploads) == 0 {
c.JSON(http.StatusOK, response.Err(ErrNoValidFilesForArchive))
response.AbortBadRequest(c, shared.ErrNoValidFilesForArchive)
return
}
// 设置 ZIP 格式流的响应头
c.Header("Content-Type", "application/zip")
c.Header("Content-Disposition", "attachment; filename=\"batch_download.zip\"")
// 开启实时 ZIP 压缩器并直接输出给 Response Writer
zipWriter := zip.NewWriter(c.Writer)
defer func() { _ = zipWriter.Close() }()
// 用于解决 ZIP 内部文件名称发生碰撞冲突的问题
usedNames := make(map[string]int)
for _, upload := range uploads {
// 校验文件访问权限
if err := checkFileAccessPermission(c, &upload); err != nil {
if err := filesrv.CheckFileAccessPermission(c, &upload); err != nil {
logger.WarnF(ctx, "Batch download: skip file %d due to permission denied: %v", upload.ID, err)
continue
}
// 校验防冲突重命名逻辑
fileName := upload.FileName
if count, exists := usedNames[fileName]; exists {
usedNames[fileName] = count + 1
@@ -300,23 +288,19 @@ func BatchDownloadFiles(c *gin.Context) {
usedNames[fileName] = 1
}
// 在 ZIP 包内建新条目
zipFileEntry, err := zipWriter.Create(fileName)
if err != nil {
logger.ErrorF(ctx, "ZIP 添加条目失败 [%s]: %v", fileName, err)
continue
}
// 打开底层文件数据源
var rc io.ReadCloser
obj, err := openStoredObject(ctx, &upload)
obj, err := uploadstorage.OpenStoredObject(ctx, &upload)
if err != nil {
logger.ErrorF(ctx, "打包时读取文件失败: %v", err)
continue
}
rc = obj.Body
rc := obj.Body
// 流式拷贝到 ZIP entry
_, err = io.Copy(zipFileEntry, rc)
_ = rc.Close()
if err != nil {
@@ -328,7 +312,7 @@ func BatchDownloadFiles(c *gin.Context) {
func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) {
accessModeStr := c.PostForm("access_mode")
if accessModeStr == "" {
if uploadType == defaultPublicUploadType {
if uploadType == shared.DefaultPublicUploadType {
return 1, ""
}
return 0, ""
@@ -341,7 +325,6 @@ func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) {
return accessMode, ""
}
// validateUploadExtension 校验文件后缀是否在系统允许的上传扩展名列表中
func validateUploadExtension(ctx context.Context, ext string) string {
var sc model.SystemConfig
if err := sc.GetByKey(ctx, model.ConfigKeyUploadAllowedExtensions); err == nil && sc.Value != "" {
@@ -354,21 +337,20 @@ func validateUploadExtension(ctx context.Context, ext string) string {
}
}
if !allowed {
return ErrUnsupportedFormat
return shared.ErrUnsupportedFormat
}
}
return ""
}
// tryInstantUpload 尝试秒传:若数据库已存在相同 Hash 且大小一致的可用文件,直接生成新记录
func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User, fileHash string, size int64, mimeType, ext, origName string, accessMode int) (bool, error) {
var existing model.Upload
err := db.DB(ctx).Where("hash = ? AND file_size = ? AND status IN (?, ?)", fileHash, size, model.UploadStatusPending, model.UploadStatusUsed).First(&existing).Error
if err != nil {
return false, err
}
if StorageReadOnly(ctx) {
c.JSON(http.StatusConflict, response.Err(ErrStorageReadOnly))
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
return true, nil
}
@@ -390,52 +372,40 @@ func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User,
}
if err := db.DB(ctx).Create(&newUpload).Error; err != nil {
c.JSON(http.StatusOK, response.Err(ErrSaveUploadRecordFailed))
response.AbortBadRequest(c, shared.ErrSaveUploadRecordFailed)
return true, err
}
recordUploadStatsAdd(ctx, &newUpload)
uploadstats.RecordUploadStatsAdd(ctx, &newUpload)
logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", id, existing.FilePath)
c.JSON(http.StatusOK, response.OK(newUpload))
return true, nil
}
// storeUploadFile 将文件写入当前活动存储驱动。
func storeUploadFile(ctx context.Context, subPath string, size int64, mimeType string, buf *bytes.Buffer, meta *model.UploadMetadata) (string, string, string) {
if StorageReadOnly(ctx) {
return "", "", ErrStorageReadOnly
if uploadstorage.ReadOnly(ctx) {
return "", "", shared.ErrStorageReadOnly
}
driver, backend, err := storage.Active(ctx)
if err != nil {
logger.ErrorF(ctx, "初始化活动存储失败: %v", err)
return "", "", ErrSaveFileFailed
return "", "", shared.ErrSaveFileFailed
}
result, err := backend.Put(ctx, subPath, bytes.NewReader(buf.Bytes()), size, mimeType)
if err != nil {
logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err)
return "", "", ErrSaveFileFailed
return "", "", shared.ErrSaveFileFailed
}
meta.Bucket = result.Bucket
return string(driver), result.Key, ""
}
// isImageExtension 判断文件扩展名是否属于常见图片格式
func isImageExtension(ext string) bool {
for _, imgExt := range []string{"jpg", "jpeg", "png", "webp", "gif"} {
if ext == imgExt {
return true
}
}
return false
}
// parseUploadMetadata 解析上传元数据字段
func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata, string) {
var meta model.UploadMetadata
metadataStr := c.DefaultPostForm("metadata", "")
if metadataStr != "" {
if err := json.Unmarshal([]byte(metadataStr), &meta); err != nil {
return meta, ErrInvalidMetadataJSON
return meta, shared.ErrInvalidMetadataJSON
}
}
meta.OriginalMime = mimeType
@@ -444,16 +414,14 @@ func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata,
return meta, ""
}
// detectMimeType 检测文件的 MIME 类型,优先使用 Content-Type 头部信息
func detectMimeType(buf *bytes.Buffer, header *multipart.FileHeader, size int64) string {
mimeType := http.DetectContentType(buf.Bytes()[:min(detectContentBytes, int(size))])
mimeType := http.DetectContentType(buf.Bytes()[:min(shared.DetectContentBytes, int(size))])
if mimeType == "application/octet-stream" && header.Header.Get("Content-Type") != "" {
mimeType = header.Header.Get("Content-Type")
}
return mimeType
}
// saveUploadRecord 保存上传记录到数据库,失败时清理本地垃圾文件
func saveUploadRecord(ctx context.Context, upload *model.Upload, storageDriver, filePath string) string {
if err := db.DB(ctx).Create(upload).Error; err != nil {
backend, backendErr := storage.ForDriver(ctx, storage.Driver(storageDriver))
@@ -462,8 +430,8 @@ func saveUploadRecord(ctx context.Context, upload *model.Upload, storageDriver,
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
}
}
return ErrSaveUploadRecordFailed
return shared.ErrSaveUploadRecordFailed
}
recordUploadStatsAdd(ctx, upload)
uploadstats.RecordUploadStatsAdd(ctx, upload)
return ""
}
}
@@ -2,7 +2,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package handler
import (
"archive/zip"
@@ -20,6 +20,9 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/common/response"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
@@ -36,6 +39,7 @@ type testResponse struct {
func setupTestRouter(authUser *model.User) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(response.ErrorHandlerMiddleware())
authMiddleware := func(c *gin.Context) {
if authUser != nil {
@@ -211,13 +215,13 @@ func TestUploadFile(t *testing.T) {
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String())
if w.Code != http.StatusBadRequest {
t.Fatalf("expected status 400, got %d. Body: %s", w.Code, w.Body.String())
}
var resp testResponse
_ = json.Unmarshal(w.Body.Bytes(), &resp)
if resp.ErrorMsg == "" || !strings.Contains(resp.ErrorMsg, ErrUnsupportedFormat) {
if resp.ErrorMsg == "" || !strings.Contains(resp.ErrorMsg, shared.ErrUnsupportedFormat) {
t.Errorf("expected unsupported format error, got: %v", resp)
}
})
@@ -294,6 +298,7 @@ func TestUploadFile(t *testing.T) {
sc.Value = "jpg,png,webp,txt"
dbConn.Save(&sc)
_ = db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, sc.Key, &sc)
model.ResetSystemConfigRAMCacheForTest()
contentType, body := createMultipartRequest(t, "file", "doc.txt", []byte("hello world generic document file"), map[string]string{
"type": "document",
@@ -806,7 +811,7 @@ func TestGetFileStats(t *testing.T) {
t.Fatalf("failed to create upload: %v", err)
}
}
if err := RebuildUploadStats(context.Background()); err != nil {
if err := uploadstats.RebuildUploadStats(context.Background()); err != nil {
t.Fatalf("failed to rebuild upload stats: %v", err)
}
@@ -1,27 +1,17 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package handler
import (
"net/http"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
const (
catImage = "图片"
catVideo = "视频"
catAudio = "音频"
catDocument = "文档"
catArchive = "压缩包"
catOther = "其他"
)
type trendItem struct {
@@ -60,15 +50,15 @@ func GetFileStats(c *gin.Context) {
var stats []model.UploadStat
if err := db.DB(ctx).Find(&stats).Error; err != nil {
c.JSON(http.StatusOK, response.Err(err.Error()))
response.AbortBadRequest(c, err.Error())
return
}
now := time.Now()
trendDates := make([]string, 0, fileStatsTrendDays)
trendCountMap := make(map[string]int64, fileStatsTrendDays)
trendSizeMap := make(map[string]int64, fileStatsTrendDays)
for i := fileStatsTrendDays - 1; i >= 0; i-- {
trendDates := make([]string, 0, shared.FileStatsTrendDays)
trendCountMap := make(map[string]int64, shared.FileStatsTrendDays)
trendSizeMap := make(map[string]int64, shared.FileStatsTrendDays)
for i := shared.FileStatsTrendDays - 1; i >= 0; i-- {
date := now.AddDate(0, 0, -i).Format("2006-01-02")
trendDates = append(trendDates, date)
trendCountMap[date] = 0
@@ -82,7 +72,7 @@ func GetFileStats(c *gin.Context) {
categories []distributionItem
)
categoriesList := []string{catImage, catVideo, catAudio, catDocument, catArchive, catOther}
categoriesList := []string{"图片", "视频", "音频", "文档", "压缩包", "其他"}
categoryMap := make(map[string]distributionItem, len(categoriesList))
for _, cat := range categoriesList {
categoryMap[cat] = distributionItem{Name: cat}
@@ -134,44 +124,4 @@ func GetFileStats(c *gin.Context) {
Categories: categories,
Types: types,
}))
}
func getFileCategory(mimeType, ext string) string {
mimeType = strings.ToLower(mimeType)
ext = strings.ToLower(ext)
if strings.HasPrefix(mimeType, "image/") || isImageExtension(ext) {
return catImage
}
if strings.HasPrefix(mimeType, "video/") {
return catVideo
}
if strings.HasPrefix(mimeType, "audio/") {
return catAudio
}
if isArchiveExtension(ext) || strings.Contains(mimeType, "zip") || strings.Contains(mimeType, "tar") || strings.Contains(mimeType, "gzip") {
return catArchive
}
if isDocumentExtension(ext) || strings.HasPrefix(mimeType, "text/") || mimeType == "application/pdf" {
return catDocument
}
return catOther
}
func isArchiveExtension(ext string) bool {
for _, e := range []string{"zip", "rar", "7z", "tar", "gz", "tgz", "bz2", "xz"} {
if ext == e {
return true
}
}
return false
}
func isDocumentExtension(ext string) bool {
for _, e := range []string{"pdf", "doc", "docx", "xls", "xlsx", "ppt", "pptx", "txt", "md", "csv", "json", "yaml", "yml", "xml"} {
if ext == e {
return true
}
}
return false
}
+23
View File
@@ -0,0 +1,23 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package shared
import "github.com/Rain-kl/Wavelet/internal/storage"
// Upload size, path, media quality, and cache constants shared across subpackages.
const (
MaxUploadSize = 32 * 1024 * 1024 // 32MB
DetectContentBytes = 512 // http.DetectContentType 需要的最小字节数
UploadDirPerm = 0755 // 上传目录权限
UploadFilePerm = 0644 // 上传文件权限
ImageQualityLow = "low"
ImageQualityMedium = "medium"
ImageQualityHigh = "high"
ImageQualityOrigin = "origin"
StorageDriverLocal = string(storage.DriverLocal)
DefaultPublicUploadType = "avatar"
FileStatsTrendDays = 7
MaxS3KeyLength = 1024
AccessCacheTTL = 5 // seconds; multiplied by time.Second at use site
)
+41
View File
@@ -0,0 +1,41 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package shared holds upload error and configuration constants shared across subpackages.
package shared
// 文件管理常量
const (
ErrNoFileSelected = "请选择要上传的文件"
ErrUnsupportedFormat = "只支持 JPG、PNG、WEBP 格式的图片"
ErrProcessFileFailed = "处理文件失败"
ErrSaveFileFailed = "保存文件失败"
ErrOpenFileFailed = "打开文件失败"
ErrSaveUploadRecordFailed = "保存上传记录失败"
ErrGenericFileTooLarge = "文件大小不能超过 32MB"
ErrFileContentExtensionMismatch = "文件内容与扩展名不匹配,可能包含安全风险"
ErrFileValidationFailed = "文件校验失败"
ErrInvalidMetadataJSON = "元数据 JSON 格式不合法"
ErrInvalidFileID = "无效的文件 ID"
ErrQueryUploadRecordFailed = "查询文件记录失败"
ErrInvalidBatchDownloadRequest = "参数绑定失败,请传入有效的文件 ID 数组"
ErrInvalidIDValueFormat = "无效的 ID 值: %s"
ErrRetrieveUploadRecordsFailed = "检索文件记录失败"
ErrNoValidFilesForArchive = "没有找到任何有效的文件记录进行打包"
ErrInvalidParams = "参数错误"
ErrQueryFileCountFailed = "查询文件数量失败"
ErrQueryFileListFailed = "查询文件列表失败"
ErrDeleteFileFailed = "删除文件失败"
ErrStorageReadOnly = "存储迁移维护中,当前仅允许读取文件"
ErrS3KeyRequired = "s3 key must not be empty"
ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d"
ErrS3KeyStartsWithSlash = "s3 key must not start with /"
ErrS3KeyContainsNullBytes = "s3 key must not contain null bytes"
ErrQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w"
ErrImageCacheWarmupPayloadRequired = "图片缓存预热参数不能为空"
ErrInvalidImageCacheWarmupPayload = "图片缓存预热参数格式无效: %w"
ErrInvalidImageCacheWarmupQuality = "图片质量仅支持 low、medium、high"
ErrParseImageCacheWarmupPayload = "解析图片缓存预热参数失败: %w"
ErrQueryImagesForCacheWarmup = "查询待预热图片失败: %w"
)
+43
View File
@@ -0,0 +1,43 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package stats maintains incremental upload statistics and aggregations.
package stats
import (
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
)
const (
catImage = "图片"
catVideo = "视频"
catAudio = "音频"
catDocument = "文档"
catArchive = "压缩包"
catOther = "其他"
)
// GetFileCategory classifies a file by mime type and extension.
func GetFileCategory(mimeType, ext string) string {
mimeType = strings.ToLower(mimeType)
ext = strings.ToLower(ext)
if strings.HasPrefix(mimeType, "image/") || util.IsImageExtension(ext) {
return catImage
}
if strings.HasPrefix(mimeType, "video/") {
return catVideo
}
if strings.HasPrefix(mimeType, "audio/") {
return catAudio
}
if util.IsArchiveExtension(ext) || strings.Contains(mimeType, "zip") || strings.Contains(mimeType, "tar") || strings.Contains(mimeType, "gzip") {
return catArchive
}
if util.IsDocumentExtension(ext) || strings.HasPrefix(mimeType, "text/") || mimeType == "application/pdf" {
return catDocument
}
return catOther
}
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package stats
import (
"context"
@@ -72,7 +72,7 @@ func applyUploadStatsDeltaTx(tx *gorm.DB, upload *model.Upload, sign int64) erro
}{
{model.UploadStatDimensionTotal, ""},
{model.UploadStatDimensionType, typeKey},
{model.UploadStatDimensionCategory, getFileCategory(upload.MimeType, upload.Extension)},
{model.UploadStatDimensionCategory, GetFileCategory(upload.MimeType, upload.Extension)},
{model.UploadStatDimensionTrend, upload.CreatedAt.Format("2006-01-02")},
}
@@ -111,13 +111,15 @@ func upsertUploadStatDelta(tx *gorm.DB, dimension, key string, countDelta, sizeD
}).Error
}
func recordUploadStatsAdd(ctx context.Context, upload *model.Upload) {
// RecordUploadStatsAdd logs and applies upload stats increment.
func RecordUploadStatsAdd(ctx context.Context, upload *model.Upload) {
if err := ApplyUploadStatsAdd(ctx, upload); err != nil {
logger.WarnF(ctx, "increment upload stats failed: %v", err)
}
}
func recordUploadStatsRemove(ctx context.Context, upload *model.Upload) {
// RecordUploadStatsRemove logs and applies upload stats decrement.
func RecordUploadStatsRemove(ctx context.Context, upload *model.Upload) {
if err := ApplyUploadStatsRemove(ctx, upload); err != nil {
logger.WarnF(ctx, "decrement upload stats failed: %v", err)
}
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package stats
import (
"context"
@@ -0,0 +1,88 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package storage provides upload storage backend operations and migration state.
package storage
import (
"context"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
)
// MigrationAccessState captures cached migration maintenance state.
type MigrationAccessState struct {
ReadOnly bool
Target storage.Config
HasTarget bool
TargetErr error
LoadErr error
}
var (
migrationAccessMu sync.RWMutex
migrationAccessCached MigrationAccessState
migrationAccessValid bool
migrationAccessCheckedAt time.Time
)
// ResetMigrationAccessCache clears the in-process migration access cache.
func ResetMigrationAccessCache() {
migrationAccessMu.Lock()
migrationAccessValid = false
migrationAccessMu.Unlock()
}
// LoadMigrationAccessState returns cached migration maintenance state.
func LoadMigrationAccessState(ctx context.Context) MigrationAccessState {
migrationAccessMu.RLock()
if migrationAccessValid && time.Since(migrationAccessCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
state := migrationAccessCached
migrationAccessMu.RUnlock()
return state
}
migrationAccessMu.RUnlock()
migrationAccessMu.Lock()
defer migrationAccessMu.Unlock()
if migrationAccessValid && time.Since(migrationAccessCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
return migrationAccessCached
}
migrationAccessCached = buildMigrationAccessState(ctx)
migrationAccessValid = true
migrationAccessCheckedAt = time.Now()
return migrationAccessCached
}
func buildMigrationAccessState(ctx context.Context) MigrationAccessState {
execution, ok, err := LatestMigrationExecution(ctx)
if err != nil {
return MigrationAccessState{LoadErr: err, ReadOnly: true}
}
if !ok {
return MigrationAccessState{}
}
state := MigrationAccessState{
ReadOnly: execution.Status != model.TaskExecutionStatusSucceeded,
}
if execution.Status == model.TaskExecutionStatusSucceeded {
return state
}
target, err := ParseMigrationTargetConfig(ctx, []byte(execution.Payload))
if err != nil {
state.TargetErr = err
return state
}
state.Target = target
state.HasTarget = true
return state
}
+93
View File
@@ -0,0 +1,93 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package storage
import (
"context"
"encoding/json"
"errors"
"fmt"
"strings"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"gorm.io/gorm"
)
// StorageMigrationTask is the Asynq task name for storage migration.
const StorageMigrationTask = "storage:migrate"
// LatestMigrationExecution returns the most recent storage migration task execution.
func LatestMigrationExecution(ctx context.Context) (*model.TaskExecution, bool, error) {
var execution model.TaskExecution
err := db.DB(ctx).
Where("task_type = ?", StorageMigrationTask).
Order("id DESC").
First(&execution).Error
if err == nil {
return &execution, true, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, err
}
return nil, false, nil
}
// ParseMigrationTargetConfig parses and validates a storage migration target payload.
func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (storage.Config, error) {
if strings.TrimSpace(string(payload)) == "" {
return storage.Config{}, 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 storage.Config{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
}
if len(raw.Target) == 0 {
return storage.Config{}, errors.New("storage migration target payload is required")
}
var targetBytes []byte
var targetStr string
if err := json.Unmarshal(raw.Target, &targetStr); err == nil {
targetBytes = []byte(targetStr)
} else {
targetBytes = raw.Target
}
var target storage.Config
if err := json.Unmarshal(targetBytes, &target); err != nil {
return storage.Config{}, fmt.Errorf("parse target storage config: %w", err)
}
current, err := storage.LoadConfig(ctx)
if err != nil {
return storage.Config{}, fmt.Errorf("load active storage config: %w", err)
}
target = storage.MergeMaskedSecrets(target, current)
if err := storage.ValidateConfig(target); err != nil {
return storage.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, storage.Config, error) {
target, err := ParseMigrationTargetConfig(ctx, payload)
if err != nil {
return nil, storage.Config{}, err
}
type storageMigrationPayload struct {
Target storage.Config `json:"target"`
}
normalized, err := json.Marshal(storageMigrationPayload{Target: target})
if err != nil {
return nil, storage.Config{}, fmt.Errorf("marshal storage migration payload: %w", err)
}
return normalized, target, nil
}
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package storage
import (
"context"
@@ -12,17 +12,18 @@ import (
"github.com/Rain-kl/Wavelet/pkg/logger"
)
// StorageReadOnly checks if the storage system is in read-only maintenance mode.
func StorageReadOnly(ctx context.Context) bool {
state := loadMigrationAccessState(ctx)
if state.loadErr != nil {
logger.ErrorF(ctx, "读取存储维护状态失败: %v", state.loadErr)
// ReadOnly checks if the storage system is in read-only maintenance mode.
func ReadOnly(ctx context.Context) bool {
state := LoadMigrationAccessState(ctx)
if state.LoadErr != nil {
logger.ErrorF(ctx, "读取存储维护状态失败: %v", state.LoadErr)
return true
}
return state.readOnly
return state.ReadOnly
}
func openStoredObject(ctx context.Context, upload *model.Upload) (*storage.Object, error) {
// OpenStoredObject opens a stored upload object from its configured backend.
func OpenStoredObject(ctx context.Context, upload *model.Upload) (*storage.Object, error) {
driver := storage.Driver(upload.StorageDriver)
if driver == "" {
driver = storage.DriverLocal
@@ -40,7 +41,7 @@ func backendForStoredDriver(ctx context.Context, driver storage.Driver) (storage
return backend, nil
}
target, ok, targetErr := currentMigrationTargetConfig(ctx)
target, ok, targetErr := CurrentMigrationTargetConfig(ctx)
if targetErr != nil {
return nil, targetErr
}
@@ -50,16 +51,17 @@ func backendForStoredDriver(ctx context.Context, driver storage.Driver) (storage
return nil, fmt.Errorf("storage configuration for driver %q is unavailable", driver)
}
func currentMigrationTargetConfig(ctx context.Context) (storage.Config, bool, error) {
state := loadMigrationAccessState(ctx)
if state.loadErr != nil {
return storage.Config{}, false, state.loadErr
// CurrentMigrationTargetConfig returns the pending migration target config when available.
func CurrentMigrationTargetConfig(ctx context.Context) (storage.Config, bool, error) {
state := LoadMigrationAccessState(ctx)
if state.LoadErr != nil {
return storage.Config{}, false, state.LoadErr
}
if state.targetErr != nil {
return storage.Config{}, false, state.targetErr
if state.TargetErr != nil {
return storage.Config{}, false, state.TargetErr
}
if !state.hasTarget {
if !state.HasTarget {
return storage.Config{}, false, nil
}
return state.target, true, nil
}
return state.Target, true, nil
}
@@ -1,8 +1,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package upload implements upload tasks and file cleanup services.
package upload
// Package task provides upload-related async background task handlers.
package task
import (
"context"
@@ -10,6 +10,9 @@ import (
"fmt"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
@@ -18,16 +21,11 @@ import (
"gorm.io/gorm"
)
// 异步任务名称与管理类型定义
const (
// SystemCleanupTask 系统定期垃圾清理任务标识
SystemCleanupTask = "system:cleanup"
// TaskTypeSystemCleanup 系统定期垃圾清理管理类型
TaskTypeSystemCleanup = "system_cleanup"
// 错误描述常量
errStorageReadOnly = "存储迁移维护中,当前仅允许读取文件"
errQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w"
)
// SystemCleanupMeta represents the task metadata.
@@ -47,21 +45,19 @@ type SystemCleanupHandler struct{}
// Execute 执行系统清理(包含文件清理、历史推送日志和任务执行日志清理)
func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
if storageReadOnly(ctx) {
return nil, errors.New(errStorageReadOnly)
if uploadstorage.ReadOnly(ctx) {
return nil, errors.New(shared.ErrStorageReadOnly)
}
const batchSize = 100 // 每批处理100个文件
const batchSize = 100
var lastID uint64
var totalProcessed int
var totalDeleted int
// 计算1小时前的时间
oneHourAgo := time.Now().Add(-1 * time.Hour)
task.AppendLog(ctx, "开始扫描未使用上传文件,阈值: %s", oneHourAgo.Format(time.RFC3339))
for {
// 使用游标分页查询未使用且超过1小时的上传记录
var unusedUploads []model.Upload
if err := db.DB(ctx).
Where("id > ? AND status = ? AND created_at < ?", lastID, model.UploadStatusPending, oneHourAgo).
@@ -69,22 +65,19 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
Limit(batchSize).
Find(&unusedUploads).Error; err != nil {
task.AppendLog(ctx, "查询未使用的上传文件失败: %v", err)
return nil, fmt.Errorf(errQueryUnusedUploadsFailed, err)
return nil, fmt.Errorf(shared.ErrQueryUnusedUploadsFailed, err)
}
// 没有更多数据,退出循环
if len(unusedUploads) == 0 {
break
}
task.AppendLog(ctx, "本批次找到 %d 个需要清理的上传文件", len(unusedUploads))
// 处理每个未使用的上传文件
for _, u := range unusedUploads {
totalProcessed++
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
// 更新上传记录状态
if err := tx.Model(&model.Upload{}).
Where("id = ? AND status = ?", u.ID, model.UploadStatusPending).
Update("status", model.UploadStatusDeleted).Error; err != nil {
@@ -110,13 +103,12 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
continue
}
recordUploadStatsRemove(ctx, &u)
uploadstats.RecordUploadStatsRemove(ctx, &u)
totalDeleted++
lastID = u.ID
}
}
// 2. 清理超过7天的历史推送日志
task.AppendLog(ctx, "开始清理历史推送审计日志,只保留最近7天数据...")
cutoff := time.Now().AddDate(0, 0, -7)
var pushHistoryCount int64
@@ -132,7 +124,6 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
task.AppendLog(ctx, "没有需要清理的历史推送记录 (截止时间: %s)", cutoff.Format("2006-01-02 15:04:05"))
}
// 3. 清理任务执行日志:高频任务保留3天,低频任务保留30天。
task.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...")
taskLogStats, err := model.CleanupTaskExecutionLogs(ctx, time.Now())
if err != nil {
@@ -154,17 +145,4 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
)
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
func storageReadOnly(ctx context.Context) bool {
var execution model.TaskExecution
err := db.DB(ctx).Where("task_type = ?", "storage:migrate").Order("id DESC").First(&execution).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return false
}
logger.ErrorF(ctx, "读取存储维护状态失败: %v", err)
return true
}
return execution.Status != model.TaskExecutionStatusSucceeded
}
}
@@ -1,13 +1,12 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package task
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
@@ -16,17 +15,18 @@ import (
"sync/atomic"
"time"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/task"
"golang.org/x/sync/errgroup"
"gorm.io/gorm"
)
const (
// StorageMigrationTask is the Asynq task name for storage migration.
StorageMigrationTask = "storage:migrate"
StorageMigrationTask = uploadstorage.StorageMigrationTask
// TaskTypeStorageMigration is the task metadata type for storage migration.
TaskTypeStorageMigration = "storage_migration"
@@ -58,13 +58,9 @@ var StorageMigrationMeta = task.TaskMeta{
// MigrationHandler copies stored objects and activates the target backend.
type MigrationHandler struct{}
type storageMigrationPayload struct {
Target storage.Config `json:"target"`
}
// ValidatePayload rejects duplicate active migrations through the task framework.
func (h *MigrationHandler) ValidatePayload(payload []byte) ([]byte, error) {
normalized, _, err := normalizeStorageMigrationPayload(context.Background(), payload)
normalized, _, err := uploadstorage.NormalizeMigrationPayload(context.Background(), payload)
if err != nil {
return payload, err
}
@@ -95,7 +91,6 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
return nil, errors.New("另一个存储迁移任务正在运行中")
}
// 任务结束时清理锁,使用 Background context 避免受任务 context 取消的影响
stopRenewal := make(chan struct{})
//nolint:contextcheck
defer func() {
@@ -105,7 +100,6 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
_ = db.Redis.Del(cleanupCtx, lockKey)
}()
// 启动看门狗续租协程,每 10 分钟将锁的 TTL 自动延长为 1 小时
//nolint:contextcheck,gosec
go func() {
ticker := time.NewTicker(renewalInterval)
@@ -129,7 +123,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
if err != nil {
return nil, fmt.Errorf("load active storage config: %w", err)
}
target, err := parseMigrationTargetConfig(ctx, payload)
target, err := uploadstorage.ParseMigrationTargetConfig(ctx, payload)
if err != nil {
return nil, err
}
@@ -178,62 +172,6 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
return &task.TaskResult{Message: message}, nil
}
func normalizeStorageMigrationPayload(ctx context.Context, payload []byte) ([]byte, storage.Config, error) {
target, err := parseMigrationTargetConfig(ctx, payload)
if err != nil {
return nil, storage.Config{}, err
}
normalized, err := json.Marshal(storageMigrationPayload{Target: target})
if err != nil {
return nil, storage.Config{}, fmt.Errorf("marshal storage migration payload: %w", err)
}
return normalized, target, nil
}
func parseMigrationTargetConfig(ctx context.Context, payload []byte) (storage.Config, error) {
if strings.TrimSpace(string(payload)) == "" {
return storage.Config{}, errors.New("storage migration target payload is required")
}
// Try to parse using raw JSON message to handle both struct and string payload formats
var raw struct {
Target json.RawMessage `json:"target"`
}
if err := json.Unmarshal(payload, &raw); err != nil {
return storage.Config{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
}
if len(raw.Target) == 0 {
return storage.Config{}, errors.New("storage migration target payload is required")
}
var targetBytes []byte
var targetStr string
// Check if Target is a JSON string
if err := json.Unmarshal(raw.Target, &targetStr); err == nil {
// It is a string (e.g. from dynamic form input), parse its content as JSON
targetBytes = []byte(targetStr)
} else {
// It is a JSON object, use directly
targetBytes = raw.Target
}
var target storage.Config
if err := json.Unmarshal(targetBytes, &target); err != nil {
return storage.Config{}, fmt.Errorf("parse target storage config: %w", err)
}
current, err := storage.LoadConfig(ctx)
if err != nil {
return storage.Config{}, fmt.Errorf("load active storage config: %w", err)
}
target = storage.MergeMaskedSecrets(target, current)
if err := storage.ValidateConfig(target); err != nil {
return storage.Config{}, fmt.Errorf("validate target storage config: %w", err)
}
return target, nil
}
func countStorageObjects(ctx context.Context, driver storage.Driver) (int64, error) {
var count int64
err := db.DB(ctx).Model(&model.Upload{}).
@@ -244,28 +182,13 @@ func countStorageObjects(ctx context.Context, driver storage.Driver) (int64, err
}
func hasUnresolvedMigrationTask(ctx context.Context) (bool, error) {
execution, ok, err := latestStorageMigrationExecution(ctx)
execution, ok, err := uploadstorage.LatestMigrationExecution(ctx)
if err != nil || !ok {
return false, err
}
return execution.Status == model.TaskExecutionStatusPending || execution.Status == model.TaskExecutionStatusRunning, nil
}
func latestStorageMigrationExecution(ctx context.Context) (*model.TaskExecution, bool, error) {
var execution model.TaskExecution
err := db.DB(ctx).
Where("task_type = ?", StorageMigrationTask).
Order("id DESC").
First(&execution).Error
if err == nil {
return &execution, true, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, err
}
return nil, false, nil
}
func migrateObjects(
ctx context.Context,
sourceBackend storage.Backend,
@@ -311,7 +234,7 @@ func migrateObjects(
g.SetLimit(migrationConcurrency)
for _, object := range objects {
obj := object // Capture range variable
obj := object
g.Go(func() error {
if err := migrateSingleObject(ctx, sourceBackend, targetBackend, sourceDriver, targetDriver, obj, sha256HexLength); err != nil {
return err
@@ -344,7 +267,6 @@ func migrateSingleObject(
},
sha256HexLength int,
) error {
// Check if the file already exists in target storage and has matching size
if shouldSkipMigration(ctx, targetBackend, obj) {
task.AppendLog(ctx, "[跳过迁移] 目标存储已存在相同文件且校验一致: %s", obj.FilePath)
if err := db.DB(ctx).Model(&model.Upload{}).
@@ -375,7 +297,6 @@ func migrateSingleObject(
return fmt.Errorf("close source object %q: %w", obj.FilePath, closeErr)
}
// Data integrity check (SHA-256 hash verification)
if len(obj.Hash) == sha256HexLength {
task.AppendLog(ctx, "[校验中] 正在对目标文件进行数据一致性校验 (SHA-256): %s", targetResult.Key)
targetObj, getErr := targetBackend.Get(ctx, targetResult.Key)
@@ -450,13 +371,13 @@ func markMissingMigrationObjectDeleted(
if err := db.DB(ctx).Model(&model.Upload{}).
Where("storage_driver = ? AND file_path = ?", sourceDriver, filePath).
Updates(map[string]any{
"status": model.UploadStatusDeleted,
colStorageDriver: targetDriver,
"status": model.UploadStatusDeleted,
colStorageDriver: targetDriver,
}).Error; err != nil {
return fmt.Errorf("update missing object %q: %w", filePath, err)
}
for i := range affectedUploads {
recordUploadStatsRemove(ctx, &affectedUploads[i])
uploadstats.RecordUploadStatsRemove(ctx, &affectedUploads[i])
}
return nil
}
@@ -475,4 +396,4 @@ func isNotFoundError(err error) bool {
}
}
return false
}
}
@@ -1,7 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package task
import (
"bytes"
@@ -52,7 +52,9 @@ func TestMigrationHandlerExecute(t *testing.T) {
AccessKeyID: "key",
SecretAccessKey: "secret",
}
payload, err := json.Marshal(storageMigrationPayload{Target: target})
payload, err := json.Marshal(struct {
Target storage.Config `json:"target"`
}{Target: target})
if err != nil {
t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err)
}
@@ -150,7 +152,9 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
AccessKeyID: "key",
SecretAccessKey: "secret",
}
payload, err := json.Marshal(storageMigrationPayload{Target: target})
payload, err := json.Marshal(struct {
Target storage.Config `json:"target"`
}{Target: target})
if err != nil {
t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err)
}
@@ -259,7 +263,9 @@ func TestMigrationHandlerExecuteWithRedisLock(t *testing.T) {
t.Fatalf("SaveActiveConfig() returned error: %v", err)
}
payload, err := json.Marshal(storageMigrationPayload{Target: active})
payload, err := json.Marshal(struct {
Target storage.Config `json:"target"`
}{Target: active})
if err != nil {
t.Fatalf("Marshal payload failed: %v", err)
}
@@ -2,7 +2,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package task
import (
"context"
@@ -12,12 +12,13 @@ import (
"strings"
"sync"
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
)
// 异步任务名称与管理类型定义
const (
// WarmImageCacheTask 图片压缩缓存预热任务标识
WarmImageCacheTask = "upload:warm_image_cache"
@@ -54,29 +55,25 @@ type WarmImageCachePayload struct {
Quality string `json:"quality"`
}
// SystemCleanupHandler 系统定期垃圾清理异步任务处理器
// WarmImageCacheHandler serially warms compressed image cache entries.
type WarmImageCacheHandler struct{}
// Execute 执行系统清理(包含文件清理和历史消息推送日志清理)
// ValidatePayload validates and normalizes image cache warmup parameters.
func (h *WarmImageCacheHandler) ValidatePayload(payload []byte) ([]byte, error) {
if len(payload) == 0 {
return nil, errors.New(errImageCacheWarmupPayloadRequired)
return nil, errors.New(shared.ErrImageCacheWarmupPayloadRequired)
}
var req WarmImageCachePayload
if err := json.Unmarshal(payload, &req); err != nil {
return nil, fmt.Errorf(errInvalidImageCacheWarmupPayload, err)
return nil, fmt.Errorf(shared.ErrInvalidImageCacheWarmupPayload, err)
}
req.Quality = strings.ToLower(strings.TrimSpace(req.Quality))
if req.Quality != imageQualityLow &&
req.Quality != imageQualityMedium &&
req.Quality != imageQualityHigh {
return nil, errors.New(errInvalidImageCacheWarmupQuality)
if req.Quality != shared.ImageQualityLow &&
req.Quality != shared.ImageQualityMedium &&
req.Quality != shared.ImageQualityHigh {
return nil, errors.New(shared.ErrInvalidImageCacheWarmupQuality)
}
return json.Marshal(req)
@@ -92,7 +89,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
var req WarmImageCachePayload
if err := json.Unmarshal(normalizedPayload, &req); err != nil {
return nil, fmt.Errorf(errParseImageCacheWarmupPayload, err)
return nil, fmt.Errorf(shared.ErrParseImageCacheWarmupPayload, err)
}
task.AppendLog(ctx, "等待获取图片缓存预热执行锁,质量: %s", req.Quality)
@@ -128,7 +125,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
Limit(batchSize).
Find(&uploads).Error; err != nil {
task.AppendLog(ctx, "查询图片上传记录失败: %v", err)
return nil, fmt.Errorf(errQueryImagesForCacheWarmup, err)
return nil, fmt.Errorf(shared.ErrQueryImagesForCacheWarmup, err)
}
if len(uploads) == 0 {
@@ -147,7 +144,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
totalProcessed++
lastID = upload.ID
_, cacheHit, err := ensureCompressedImageCache(ctx, upload, req.Quality)
_, cacheHit, err := filesrv.EnsureCompressedImageCache(ctx, upload, req.Quality)
if err != nil {
totalFailed++
batchFailed++
@@ -184,4 +181,4 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
)
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
}
@@ -2,7 +2,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
package task
import (
"bytes"
@@ -17,6 +17,8 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
@@ -201,7 +203,7 @@ func TestWarmImageCacheHandlerValidatePayload(t *testing.T) {
{
name: "normalizes quality",
payload: []byte(`{"quality":" HIGH "}`),
wantQuality: imageQualityHigh,
wantQuality: shared.ImageQualityHigh,
},
{
name: "empty payload",
@@ -281,7 +283,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: firstPath,
MimeType: "image/png",
Extension: "png",
StorageDriver: storageDriverLocal,
StorageDriver: shared.StorageDriverLocal,
Status: model.UploadStatusUsed,
},
{
@@ -291,7 +293,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: secondPath,
MimeType: "application/octet-stream",
Extension: "jpg",
StorageDriver: storageDriverLocal,
StorageDriver: shared.StorageDriverLocal,
Status: model.UploadStatusPending,
},
{
@@ -301,7 +303,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: filepath.Join(testDir, "notes.txt"),
MimeType: "text/plain",
Extension: "txt",
StorageDriver: storageDriverLocal,
StorageDriver: shared.StorageDriverLocal,
Status: model.UploadStatusUsed,
},
{
@@ -311,7 +313,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: firstPath,
MimeType: "image/png",
Extension: "png",
StorageDriver: storageDriverLocal,
StorageDriver: shared.StorageDriverLocal,
Status: model.UploadStatusDeleted,
},
}
@@ -339,7 +341,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
}
for i := range records[:2] {
key := imageCompressionCacheKey(&records[i], imageQualityLow)
key := filesrv.ImageCompressionCacheKey(&records[i], shared.ImageQualityLow)
got, err := cache.Get(key)
if err != nil {
t.Errorf("cache.Get(%q) returned error: %v", key, err)
+50
View File
@@ -0,0 +1,50 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import (
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
)
// IsImageExtension reports whether ext is a common image format.
func IsImageExtension(ext string) bool {
for _, imgExt := range []string{"jpg", "jpeg", "png", "webp", "gif"} {
if ext == imgExt {
return true
}
}
return false
}
// IsArchiveExtension reports whether ext is a common archive format.
func IsArchiveExtension(ext string) bool {
for _, e := range []string{"zip", "rar", "7z", "tar", "gz", "tgz", "bz2", "xz"} {
if ext == e {
return true
}
}
return false
}
// IsDocumentExtension reports whether ext is a common document format.
func IsDocumentExtension(ext string) bool {
for _, e := range []string{"pdf", "doc", "docx", "xls", "xlsx", "ppt", "pptx", "txt", "md", "csv", "json", "yaml", "yml", "xml"} {
if ext == e {
return true
}
}
return false
}
// NormalizeImageQuality normalizes the requested image quality query parameter.
func NormalizeImageQuality(quality string) string {
switch strings.ToLower(quality) {
case shared.ImageQualityLow, shared.ImageQualityMedium, shared.ImageQualityHigh:
return strings.ToLower(quality)
default:
return shared.ImageQualityOrigin
}
}
@@ -2,7 +2,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
// Package util provides upload media helpers and image utilities.
package util
import (
"bytes"
@@ -15,28 +16,27 @@ import (
"io"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/deepteams/webp"
_ "golang.org/x/image/webp" // Register WebP decoder for image.Decode
)
const maxS3KeyLength = 1024
// ValidateS3Key validates an S3 object key for safety.
func ValidateS3Key(key string) error {
if key == "" {
return errors.New(ErrS3KeyRequired)
return errors.New(shared.ErrS3KeyRequired)
}
if len(key) > maxS3KeyLength {
return fmt.Errorf(ErrS3KeyTooLongFormat, maxS3KeyLength)
if len(key) > shared.MaxS3KeyLength {
return fmt.Errorf(shared.ErrS3KeyTooLongFormat, shared.MaxS3KeyLength)
}
if strings.HasPrefix(key, "/") {
return errors.New(ErrS3KeyStartsWithSlash)
return errors.New(shared.ErrS3KeyStartsWithSlash)
}
if strings.Contains(key, "\x00") {
return errors.New(ErrS3KeyContainsNullBytes)
return errors.New(shared.ErrS3KeyContainsNullBytes)
}
return nil
@@ -45,34 +45,31 @@ func ValidateS3Key(key string) error {
// CompressImageToWebP decodes an image from srcReader and encodes it into WebP format
// using the specified quality (low -> 60, medium -> 75, high -> 85).
func CompressImageToWebP(srcReader io.Reader, quality string) ([]byte, error) {
// Decode the image
img, format, err := image.Decode(srcReader)
if err != nil {
return nil, fmt.Errorf("failed to decode image (format: %s): %w", format, err)
}
// Determine quality
var qualityScore float32
switch strings.ToLower(quality) {
case imageQualityLow:
case shared.ImageQualityLow:
qualityScore = 60
case imageQualityMedium:
case shared.ImageQualityMedium:
qualityScore = 75
case imageQualityHigh, "":
case shared.ImageQualityHigh, "":
qualityScore = 85
default:
qualityScore = 85
}
// Encode to WebP
var buf bytes.Buffer
err = webp.Encode(&buf, img, &webp.EncoderOptions{
Quality: qualityScore,
Method: 4, // Default method
Method: 4,
})
if err != nil {
return nil, fmt.Errorf("failed to encode WebP: %w", err)
}
return buf.Bytes(), nil
}
}