upload+accessKey

This commit is contained in:
ryan
2026-06-07 21:21:06 +08:00
parent 360a26f109
commit 70a13dc107
35 changed files with 4763 additions and 1321 deletions
+37 -10
View File
@@ -18,6 +18,7 @@ package oauth
import (
"net/http"
"time"
"github.com/gin-gonic/gin"
"github.com/linux-do/credit/internal/common"
@@ -44,19 +45,45 @@ func LoginRequired() gin.HandlerFunc {
ctx, span := otel_trace.Start(c.Request.Context(), "LoginRequired")
defer span.End()
// load user
userId := GetUserIDFromContext(c)
if userId <= 0 {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
return
// check token in headers
tokenStr := c.GetHeader("X-Access-Token")
if tokenStr == "" {
authHeader := c.GetHeader("Authorization")
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
tokenStr = authHeader[7:]
}
}
// load user from db to make sure is active
var user model.User
tx := db.DB(ctx).Where("id = ? AND is_active = ?", userId, true).First(&user)
if tx.Error != nil {
c.AbortWithStatusJSON(http.StatusInternalServerError, gin.H{"error_msg": tx.Error.Error(), "data": nil})
return
var authenticated bool
if tokenStr != "" {
tokenHash := model.HashToken(tokenStr)
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err == nil {
if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&user).Error; err == nil {
authenticated = true
// update token last used time
now := time.Now()
db.DB(ctx).Model(&tokenRecord).Update("last_used_at", &now)
}
}
}
if !authenticated {
// load user from session
userId := GetUserIDFromContext(c)
if userId <= 0 {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
return
}
// load user from db to make sure is active
tx := db.DB(ctx).Where("id = ? AND is_active = ?", userId, true).First(&user)
if tx.Error != nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
return
}
}
// log
+2 -2
View File
@@ -18,7 +18,7 @@ package upload
const (
ErrNoFileSelected = "请选择要上传的文件"
ErrInvalidCoverType = "无效的封面类型"
ErrInvalidUploadType = "无效的上传类型"
ErrFileTooLarge = "图片大小不能超过 2MB"
ErrUnsupportedFormat = "只支持 JPG、PNG、WEBP 格式的图片"
ErrInvalidImage = "无效的图片文件"
@@ -28,5 +28,5 @@ const (
ErrOpenFileFailed = "打开文件失败"
ErrInvalidFilePath = "非法文件路径"
ErrSaveUploadRecordFailed = "保存上传记录失败"
ErrQueryHistoryCoverFailed = "查询历史封面失败"
ErrQueryHistoryUploadFailed = "查询历史上传记录失败"
)
+5
View File
@@ -59,6 +59,11 @@ func ServeFileByID(c *gin.Context) {
return
}
if upload.StorageDriver == "local" || (upload.StorageDriver == "" && !storage.IsEnabled()) {
c.File(upload.FilePath)
return
}
// Retrieve file from S3 (via CDN if configured)
obj, err := storage.GetObjectViaCache(c.Request.Context(), upload.FilePath)
if err != nil {
+534
View File
@@ -0,0 +1,534 @@
/*
Copyright 2026 linux.do
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package upload
import (
"archive/zip"
"bytes"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"os"
"path/filepath"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/linux-do/credit/internal/apps/oauth"
"github.com/linux-do/credit/internal/common/response"
"github.com/linux-do/credit/internal/config"
"github.com/linux-do/credit/internal/db"
"github.com/linux-do/credit/internal/db/idgen"
"github.com/linux-do/credit/internal/logger"
"github.com/linux-do/credit/internal/model"
"github.com/linux-do/credit/internal/storage"
"github.com/linux-do/credit/internal/util"
"gorm.io/gorm"
)
const maxUploadSize = 32 * 1024 * 1024 // 32MB
type batchDownloadRequest struct {
IDs []string `json:"ids" binding:"required,min=1"`
}
// UploadFile 通用上传文件接口
// @Summary 上传文件
// @Description 支持各种类型的通用文件上传,支持自动文件类型检测、哈希计算与“秒传”去重
// @Tags upload
// @Accept multipart/form-data
// @Produce json
// @Param file formData file true "要上传的文件"
// @Param type formData string false "业务分类 (例如: avatar, attachment, doc,默认为 generic)"
// @Param metadata formData string false "额外的 JSON 格式元数据"
// @Security SessionCookie
// @Success 200 {object} util.ResponseAny{data=model.Upload} "上传成功"
// @Failure 400 {object} util.ResponseAny "请求参数错误或文件受限"
// @Failure 401 {object} util.ResponseAny "未登录"
// @Failure 500 {object} util.ResponseAny "内部错误"
// @Router /api/v1/upload [post]
func UploadFile(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
header, err := c.FormFile("file")
if err != nil {
response.RespondFailure(c, ErrNoFileSelected)
return
}
file, err := header.Open()
if err != nil {
response.RespondFailure(c, ErrOpenFileFailed)
return
}
defer file.Close()
// 校验大小
if header.Size > maxUploadSize {
response.RespondFailure(c, "文件大小不能超过 32MB")
return
}
// 2. 提取文件基本元数据
origName := header.Filename
ext := strings.ToLower(strings.TrimPrefix(filepath.Ext(origName), "."))
if ext == "" {
ext = "bin"
}
// 3. 校验文件后缀是否在允许的系统配置列表中
var sc model.SystemConfig
if err := sc.GetByKey(ctx, model.ConfigKeyUploadAllowedExtensions); err == nil && sc.Value != "" {
allowedExts := strings.Split(strings.ToLower(sc.Value), ",")
allowed := false
for _, allowedExt := range allowedExts {
if strings.TrimSpace(allowedExt) == ext {
allowed = true
break
}
}
if !allowed {
response.RespondFailure(c, ErrUnsupportedFormat)
return
}
}
// 4. 读取文件并计算 Hash
hashWriter := sha256.New()
var buf bytes.Buffer
size, err := io.Copy(&buf, io.TeeReader(file, hashWriter))
if err != nil {
response.RespondFailure(c, ErrProcessFileFailed)
return
}
fileHash := hex.EncodeToString(hashWriter.Sum(nil))
mimeType := http.DetectContentType(buf.Bytes()[:min(512, int(size))])
if mimeType == "application/octet-stream" && header.Header.Get("Content-Type") != "" {
mimeType = header.Header.Get("Content-Type")
}
// 6. 秒传匹配校验:校验数据库中是否存在相同 Hash 且大小一致的可用文件
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 {
// 命中了相同文件,直接生成新记录指向已有的存储路径(实现秒传)
id := idgen.NextUint64ID()
newUpload := model.Upload{
ID: id,
UserID: currUser.ID,
FileName: origName,
FilePath: existing.FilePath,
FileSize: size,
MimeType: mimeType,
Extension: ext,
Hash: fileHash,
StorageDriver: existing.StorageDriver,
Type: c.DefaultPostForm("type", "generic"),
Status: model.UploadStatusUsed,
Metadata: existing.Metadata,
}
if err := db.DB(ctx).Create(&newUpload).Error; err != nil {
response.RespondFailure(c, ErrSaveUploadRecordFailed)
return
}
logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", id, existing.FilePath)
response.RespondSuccess(c, newUpload)
return
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
response.RespondFailure(c, "文件校验失败")
return
}
// 7. 解析可选元数据字段
metadataStr := c.DefaultPostForm("metadata", "")
var meta model.UploadMetadata
if metadataStr != "" {
if err := json.Unmarshal([]byte(metadataStr), &meta); err != nil {
response.RespondFailure(c, "元数据 JSON 格式不合法")
return
}
}
meta.OriginalMime = mimeType
meta.UserAgent = c.Request.UserAgent()
meta.ClientIP = c.ClientIP()
id := idgen.NextUint64ID()
subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext)
var storageDriver string
// 8. 写入底层存储驱动 (优先 S3 驱动,无配置或未开启则 fallback 至本地文件)
if storage.IsEnabled() {
storageDriver = "s3"
meta.Bucket = config.Config.S3.Bucket
fullKey := storage.BuildKey(subPath)
err = storage.PutObject(ctx, fullKey, bytes.NewReader(buf.Bytes()), size, mimeType)
if err != nil {
logger.ErrorF(ctx, "S3 存储上传失败: %v", err)
response.RespondFailure(c, ErrSaveFileFailed)
return
}
} else {
storageDriver = "local"
localDir := filepath.Join("uploads", time.Now().Format("2006/01/02"))
if err := os.MkdirAll(localDir, 0755); err != nil {
logger.ErrorF(ctx, "创建本地上传目录失败: %v", err)
response.RespondFailure(c, ErrSaveFileFailed)
return
}
localPath := filepath.Join(localDir, fmt.Sprintf("%d.%s", id, ext))
if err := os.WriteFile(localPath, buf.Bytes(), 0644); err != nil {
logger.ErrorF(ctx, "本地磁盘写入文件失败: %v", err)
response.RespondFailure(c, ErrSaveFileFailed)
return
}
// 统一使用相对路径,方便将来环境移植或备份
subPath = localPath
}
// 9. 保存文件记录至数据库
newUpload := model.Upload{
ID: id,
UserID: currUser.ID,
FileName: origName,
FilePath: subPath,
FileSize: size,
MimeType: mimeType,
Extension: ext,
Hash: fileHash,
StorageDriver: storageDriver,
Type: c.DefaultPostForm("type", "generic"),
Status: model.UploadStatusUsed,
Metadata: meta,
}
if err := db.DB(ctx).Create(&newUpload).Error; err != nil {
// 失败时若为本地存储,可以尝试清理已保存的垃圾文件
if storageDriver == "local" {
_ = os.Remove(subPath)
}
response.RespondFailure(c, ErrSaveUploadRecordFailed)
return
}
response.RespondSuccess(c, newUpload)
}
// DownloadFile 通用单文件下载接口
// @Summary 下载单文件
// @Description 根据文件 ID 获取文件,以附件形式 (Attachment) 强制开启客户端浏览器下载
// @Tags upload
// @Produce octet-stream
// @Param id path string true "文件 ID"
// @Security SessionCookie
// @Success 200 {file} file "成功下载文件"
// @Failure 400 {object} util.ResponseAny "参数错误"
// @Failure 404 {object} util.ResponseAny "文件不存在"
// @Failure 500 {object} util.ResponseAny "服务内部错误"
// @Router /api/v1/upload/download/{id} [get]
func DownloadFile(c *gin.Context) {
ctx := c.Request.Context()
idStr := c.Param("id")
uploadID, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.RespondFailure(c, "无效的文件 ID")
return
}
var upload model.Upload
if err := db.DB(ctx).Where("id = ? AND status IN (?, ?)", uploadID, model.UploadStatusPending, model.UploadStatusUsed).First(&upload).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.AbortWithStatus(http.StatusNotFound)
return
}
response.RespondFailure(c, "查询文件记录失败")
return
}
// 设置下载 Attachment 响应头 (支持 UTF-8 中文文件名转义)
c.Header("Content-Disposition", fmt.Sprintf("attachment; filename*=UTF-8''%s", url.PathEscape(upload.FileName)))
c.Header("Content-Type", upload.MimeType)
c.Header("Content-Length", strconv.FormatInt(upload.FileSize, 10))
// 根据存储驱动类型提供流式文件服务
if upload.StorageDriver == "local" || (upload.StorageDriver == "" && !storage.IsEnabled()) {
c.File(upload.FilePath)
return
}
// 从 S3/CDN 加载并返回
obj, err := storage.GetObjectViaCache(ctx, upload.FilePath)
if err != nil {
c.AbortWithStatus(http.StatusNotFound)
return
}
if obj.CachePath != "" {
c.File(obj.CachePath)
return
}
defer obj.Body.Close()
_, _ = io.Copy(c.Writer, obj.Body)
}
// BatchDownloadFiles 批量打包 ZIP 下载接口
// @Summary 批量打包下载
// @Description 传入多个文件 ID,后台实时将其打包压缩为 ZIP 流并输出,自动处理文件名重复冲突
// @Tags upload
// @Accept json
// @Produce octet-stream
// @Param request body upload.batchDownloadRequest true "包含文件 ID 数组的请求体"
// @Security SessionCookie
// @Success 200 {file} file "成功下载打包后的 ZIP"
// @Failure 400 {object} util.ResponseAny "参数错误"
// @Failure 500 {object} util.ResponseAny "打包失败"
// @Router /api/v1/upload/download/batch [post]
func BatchDownloadFiles(c *gin.Context) {
ctx := c.Request.Context()
var req batchDownloadRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.RespondFailure(c, "参数绑定失败,请传入有效的文件 ID 数组")
return
}
// 转换 ID 列表
var ids []uint64
for _, idStr := range req.IDs {
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.RespondFailure(c, fmt.Sprintf("无效的 ID 值: %s", 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 {
response.RespondFailure(c, "检索文件记录失败")
return
}
if len(uploads) == 0 {
response.RespondFailure(c, "没有找到任何有效的文件记录进行打包")
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 zipWriter.Close()
// 用于解决 ZIP 内部文件名称发生碰撞冲突的问题
usedNames := make(map[string]int)
for _, upload := range uploads {
// 校验防冲突重命名逻辑
fileName := upload.FileName
if count, exists := usedNames[fileName]; exists {
usedNames[fileName] = count + 1
ext := filepath.Ext(fileName)
base := strings.TrimSuffix(fileName, ext)
fileName = fmt.Sprintf("%s_%d%s", base, count, ext)
} else {
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
if upload.StorageDriver == "local" || (upload.StorageDriver == "" && !storage.IsEnabled()) {
fileSrc, err := os.Open(upload.FilePath)
if err != nil {
logger.ErrorF(ctx, "打包时读取本地文件失败: %v", err)
continue
}
rc = fileSrc
} else {
obj, err := storage.GetObject(ctx, upload.FilePath)
if err != nil {
logger.ErrorF(ctx, "打包时拉取 S3 文件失败: %v", err)
continue
}
rc = obj.Body
}
// 流式拷贝到 ZIP entry
_, err = io.Copy(zipFileEntry, rc)
_ = rc.Close()
if err != nil {
logger.ErrorF(ctx, "写入 ZIP 流失败: %v", err)
}
}
}
type listMyFilesRequest struct {
Page int `form:"page"`
PageSize int `form:"page_size"`
Keyword string `form:"keyword"`
Type string `form:"type"`
Extension string `form:"extension"`
}
type listMyFilesResponse struct {
Total int64 `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
Items []model.Upload `json:"items"`
}
// ListMyFiles 获取当前用户上传的文件列表
// @Summary 获取我的文件列表
// @Description 分页获取当前登录用户上传的文件,支持文件名关键词、业务类型、扩展名过滤
// @Tags upload
// @Produce json
// @Param page query int false "页码(默认 1)"
// @Param page_size query int false "每页数量(默认 20,最大 100)"
// @Param keyword query string false "文件名关键词(模糊匹配)"
// @Param type query string false "业务分类过滤"
// @Param extension query string false "扩展名过滤"
// @Security SessionCookie
// @Success 200 {object} util.ResponseAny{data=listMyFilesResponse} "查询成功"
// @Failure 401 {object} util.ResponseAny "未登录"
// @Router /api/v1/upload/my [get]
func ListMyFiles(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
var req listMyFilesRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.RespondFailure(c, "参数错误")
return
}
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 || req.PageSize > 100 {
req.PageSize = 20
}
query := db.DB(ctx).Model(&model.Upload{}).
Where("user_id = ? AND status != ?", currUser.ID, model.UploadStatusDeleted)
if req.Keyword != "" {
query = query.Where("file_name ILIKE ?", "%"+req.Keyword+"%")
}
if req.Type != "" {
query = query.Where("type = ?", req.Type)
}
if req.Extension != "" {
query = query.Where("extension = ?", strings.ToLower(req.Extension))
}
var total int64
if err := query.Count(&total).Error; err != nil {
response.RespondFailure(c, "查询文件数量失败")
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 {
response.RespondFailure(c, "查询文件列表失败")
return
}
response.RespondSuccess(c, listMyFilesResponse{
Total: total,
Page: req.Page,
PageSize: req.PageSize,
Items: items,
})
}
// DeleteFile 软删除文件记录
// @Summary 删除文件
// @Description 将文件状态置为 deleted(软删除),不会立即清理底层存储对象
// @Tags upload
// @Produce json
// @Param id path string true "文件 ID"
// @Security SessionCookie
// @Success 200 {object} util.ResponseAny "删除成功"
// @Failure 403 {object} util.ResponseAny "无权操作"
// @Failure 404 {object} util.ResponseAny "文件不存在"
// @Router /api/v1/upload/{id} [delete]
func DeleteFile(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
idStr := c.Param("id")
uploadID, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.RespondFailure(c, "无效的文件 ID")
return
}
var upload model.Upload
if err := db.DB(ctx).Where("id = ? AND status != ?", uploadID, model.UploadStatusDeleted).First(&upload).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.AbortWithStatus(http.StatusNotFound)
return
}
response.RespondFailure(c, "查询文件记录失败")
return
}
// 仅允许文件所有者或管理员删除
if upload.UserID != currUser.ID && !currUser.IsAdmin {
c.AbortWithStatus(http.StatusForbidden)
return
}
if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil {
response.RespondFailure(c, "删除文件失败")
return
}
response.RespondSuccess(c, nil)
}
func min(a, b int) int {
if a < b {
return a
}
return b
}
+516
View File
@@ -0,0 +1,516 @@
/*
Copyright 2026 linux.do
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package upload
import (
"archive/zip"
"bytes"
"context"
"encoding/json"
"io"
"mime/multipart"
"net/http"
"net/http/httptest"
"os"
"strings"
"testing"
"github.com/gin-gonic/gin"
"github.com/linux-do/credit/internal/apps/oauth"
"github.com/linux-do/credit/internal/db"
"github.com/linux-do/credit/internal/model"
"github.com/linux-do/credit/internal/storage"
"github.com/linux-do/credit/internal/testhelper"
"github.com/linux-do/credit/internal/util"
)
type testResponse struct {
Success bool `json:"success"`
Message string `json:"message"`
Data json.RawMessage `json:"data"`
}
func setupTestRouter(authUser *model.User) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
uploadGroup := r.Group("/api/v1/upload")
// Mock authentication middleware
uploadGroup.Use(func(c *gin.Context) {
if authUser != nil {
util.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})
uploadGroup.POST("", UploadFile)
uploadGroup.GET("/download/:id", DownloadFile)
uploadGroup.POST("/download/batch", BatchDownloadFiles)
return r
}
func createMultipartRequest(t *testing.T, fieldName, fileName string, fileContent []byte, extraFields map[string]string) (string, *bytes.Buffer) {
body := &bytes.Buffer{}
writer := multipart.NewWriter(body)
part, err := writer.CreateFormFile(fieldName, fileName)
if err != nil {
t.Fatalf("failed to create form file: %v", err)
}
_, err = part.Write(fileContent)
if err != nil {
t.Fatalf("failed to write file content: %v", err)
}
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()
if err != nil {
t.Fatalf("failed to close multipart writer: %v", err)
}
return writer.FormDataContentType(), body
}
func TestUploadFile(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
defer os.RemoveAll("uploads") // Clean up local files created during tests
authUser := &model.User{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Mock Storage Client
mockFiles := make(map[string][]byte)
var putCount int
restoreStorage := storage.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) (*storage.ObjectInfo, error) {
data, ok := mockFiles[key]
if !ok {
return nil, os.ErrNotExist
}
return &storage.ObjectInfo{
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
storage.IsEnabledFunc = func() bool { return true }
defer func() {
storage.IsEnabledFunc = func() bool { return false }
}()
t.Run("upload allowed image file successfully", func(t *testing.T) {
putCount = 0
imgContent := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01\x08\x06\x00\x00\x00\x1f\x15\xc4\x89") // Valid PNG header
contentType, body := createMultipartRequest(t, "file", "test.png", imgContent, map[string]string{
"type": "avatar",
"metadata": `{"extra":{"source":"test_runner"}}`,
})
req, _ := http.NewRequest("POST", "/api/v1/upload", body)
req.Header.Set("Content-Type", contentType)
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())
}
var resp testResponse
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if !resp.Success {
t.Fatalf("expected success response, got failure: %s", resp.Message)
}
// Verify database record
var uploadRecord model.Upload
if err := json.Unmarshal(resp.Data, &uploadRecord); err != nil {
t.Fatalf("failed to unmarshal upload record: %v", err)
}
var dbRecord model.Upload
if err := dbConn.First(&dbRecord, uploadRecord.ID).Error; err != nil {
t.Fatalf("failed to retrieve database record: %v", err)
}
if dbRecord.FileName != "test.png" || dbRecord.Extension != "png" {
t.Errorf("incorrect filename or extension: %s, %s", dbRecord.FileName, dbRecord.Extension)
}
if dbRecord.MimeType != "image/png" {
t.Errorf("incorrect mime type detected: %s", dbRecord.MimeType)
}
if dbRecord.StorageDriver != "s3" {
t.Errorf("expected storage driver s3, got %s", dbRecord.StorageDriver)
}
if dbRecord.Metadata.Extra["source"] != "test_runner" {
t.Errorf("expected extra meta 'source' to be 'test_runner', got %v", dbRecord.Metadata.Extra)
}
if putCount != 1 {
t.Errorf("expected 1 storage Put operation, got %d", putCount)
}
})
t.Run("upload blocked extension file", func(t *testing.T) {
// System config allowed: jpg,png,webp. Uploading docx should be blocked.
contentType, body := createMultipartRequest(t, "file", "contract.docx", []byte("fake docx content"), nil)
req, _ := http.NewRequest("POST", "/api/v1/upload", body)
req.Header.Set("Content-Type", contentType)
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())
}
var resp testResponse
json.Unmarshal(w.Body.Bytes(), &resp)
if resp.Success || !strings.Contains(resp.Message, ErrUnsupportedFormat) {
t.Errorf("expected unsupported format error, got: %v", resp)
}
})
t.Run("instant upload deduplication (秒传)", func(t *testing.T) {
putCount = 0
imgContent := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
// Upload first time
contentType1, body1 := createMultipartRequest(t, "file", "avatar1.png", imgContent, map[string]string{"type": "avatar"})
req1, _ := http.NewRequest("POST", "/api/v1/upload", body1)
req1.Header.Set("Content-Type", contentType1)
w1 := httptest.NewRecorder()
router.ServeHTTP(w1, req1)
if w1.Code != http.StatusOK {
t.Fatalf("first upload failed: %s", w1.Body.String())
}
if putCount != 1 {
t.Errorf("expected 1 put count on first upload, got %d", putCount)
}
// Upload same file second time (different filename, same content)
contentType2, body2 := createMultipartRequest(t, "file", "avatar2.png", imgContent, map[string]string{"type": "avatar"})
req2, _ := http.NewRequest("POST", "/api/v1/upload", body2)
req2.Header.Set("Content-Type", contentType2)
w2 := httptest.NewRecorder()
router.ServeHTTP(w2, req2)
if w2.Code != http.StatusOK {
t.Fatalf("second upload failed: %s", w2.Body.String())
}
var resp2 testResponse
json.Unmarshal(w2.Body.Bytes(), &resp2)
if !resp2.Success {
t.Fatalf("second upload was unsuccessful: %s", resp2.Message)
}
var uploadRecord2 model.Upload
if err := json.Unmarshal(resp2.Data, &uploadRecord2); err != nil {
t.Fatalf("failed to unmarshal second upload record: %v", err)
}
// Check if it triggered another storage put
if putCount != 1 {
t.Errorf("PutObject was triggered again! Expected deduplication (putCount=1), got putCount=%d", putCount)
}
// Check if database contains both records sharing the same FilePath
var records []model.Upload
dbConn.Where("hash = ?", uploadRecord2.Hash).Find(&records)
if len(records) != 2 {
t.Errorf("expected 2 database records sharing the same hash, got %d", len(records))
}
if records[0].FilePath != records[1].FilePath {
t.Errorf("file paths are different: %s vs %s", records[0].FilePath, records[1].FilePath)
}
if records[0].ID == records[1].ID {
t.Error("database record IDs should be unique")
}
t.Logf("Instant upload success. Record 1: %d, Record 2: %d", records[0].ID, records[1].ID)
})
t.Run("upload in local storage fallback mode", func(t *testing.T) {
// Turn off S3
storage.IsEnabledFunc = func() bool { return false }
// Seed allowed extensions configuration to allow txt files
var sc model.SystemConfig
dbConn.Where("key = ?", model.ConfigKeyUploadAllowedExtensions).First(&sc)
sc.Value = "jpg,png,webp,txt"
dbConn.Save(&sc)
_ = db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, sc.Key, &sc)
contentType, body := createMultipartRequest(t, "file", "doc.txt", []byte("hello world generic document file"), map[string]string{
"type": "document",
})
req, _ := http.NewRequest("POST", "/api/v1/upload", body)
req.Header.Set("Content-Type", contentType)
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())
}
var resp testResponse
json.Unmarshal(w.Body.Bytes(), &resp)
if !resp.Success {
t.Fatalf("local upload failed: %s", resp.Message)
}
var localRecord model.Upload
if err := json.Unmarshal(resp.Data, &localRecord); err != nil {
t.Fatalf("failed to unmarshal local upload record: %v", err)
}
if localRecord.StorageDriver != "local" {
t.Errorf("expected storage driver local, got %s", localRecord.StorageDriver)
}
// Confirm file was actually written to local disk
fileContent, err := os.ReadFile(localRecord.FilePath)
if err != nil {
t.Fatalf("failed to read local file: %v", err)
}
if string(fileContent) != "hello world generic document file" {
t.Errorf("unexpected local file contents: %s", string(fileContent))
}
})
}
func TestDownloadFile(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
defer os.RemoveAll("uploads")
authUser := &model.User{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Seed upload records in DB
localUpload := model.Upload{
ID: 2001,
UserID: 1001,
FileName: "中文文件名.txt",
FilePath: "uploads/test_download.txt",
FileSize: 12,
MimeType: "text/plain",
Extension: "txt",
StorageDriver: "local",
Status: model.UploadStatusUsed,
}
// Create local file
err := os.MkdirAll("uploads", 0755)
if err != nil {
t.Fatalf("failed to create directory: %v", err)
}
err = os.WriteFile(localUpload.FilePath, []byte("hello download"), 0644)
if err != nil {
t.Fatalf("failed to write file: %v", err)
}
dbConn.Create(&localUpload)
t.Run("download file successfully", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/upload/download/2001", nil)
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.Body.String() != "hello download" {
t.Errorf("expected body 'hello download', got '%s'", w.Body.String())
}
// Verify Content-Disposition header (supports UTF-8 escaping)
contentDisp := w.Header().Get("Content-Disposition")
expectedDisp := "attachment; filename*=UTF-8''%E4%B8%AD%E6%96%87%E6%96%87%E4%BB%B6%E5%90%8D.txt"
if contentDisp != expectedDisp {
t.Errorf("expected Content-Disposition header %q, got %q", expectedDisp, contentDisp)
}
if w.Header().Get("Content-Type") != "text/plain" {
t.Errorf("expected Content-Type text/plain, got %s", w.Header().Get("Content-Type"))
}
})
t.Run("download non-existent file", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/upload/download/9999", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusNotFound {
t.Errorf("expected status 404, got %d", w.Code)
}
})
}
func TestBatchDownloadFiles(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
defer os.RemoveAll("uploads")
authUser := &model.User{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Create and write files locally
err := os.MkdirAll("uploads", 0755)
if err != nil {
t.Fatalf("failed to create local dir: %v", err)
}
_ = os.WriteFile("uploads/f1.txt", []byte("file1 content"), 0644)
_ = os.WriteFile("uploads/f2.txt", []byte("file2 content"), 0644)
_ = os.WriteFile("uploads/f3.txt", []byte("duplicate name file content"), 0644)
// Seed upload records. Note f2 and f3 have the same FileName "file_a.txt" to trigger name collision resolution.
uploads := []model.Upload{
{
ID: 3001,
UserID: 1001,
FileName: "file_a.txt",
FilePath: "uploads/f1.txt",
FileSize: 13,
MimeType: "text/plain",
Extension: "txt",
StorageDriver: "local",
Status: model.UploadStatusUsed,
},
{
ID: 3002,
UserID: 1001,
FileName: "file_b.txt",
FilePath: "uploads/f2.txt",
FileSize: 13,
MimeType: "text/plain",
Extension: "txt",
StorageDriver: "local",
Status: model.UploadStatusUsed,
},
{
ID: 3003,
UserID: 1001,
FileName: "file_a.txt", // COLLISION with 3001!
FilePath: "uploads/f3.txt",
FileSize: 28,
MimeType: "text/plain",
Extension: "txt",
StorageDriver: "local",
Status: model.UploadStatusUsed,
},
}
for _, up := range uploads {
dbConn.Create(&up)
}
t.Run("batch download zip successfully and check duplicate renaming", func(t *testing.T) {
reqBody, _ := json.Marshal(batchDownloadRequest{
IDs: []string{"3001", "3002", "3003"},
})
req, _ := http.NewRequest("POST", "/api/v1/upload/download/batch", bytes.NewReader(reqBody))
req.Header.Set("Content-Type", "application/json")
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.Header().Get("Content-Type") != "application/zip" {
t.Errorf("expected Content-Type application/zip, got %s", w.Header().Get("Content-Type"))
}
// Unzip in-memory
zipReader, err := zip.NewReader(bytes.NewReader(w.Body.Bytes()), int64(w.Body.Len()))
if err != nil {
t.Fatalf("failed to read zip buffer: %v", err)
}
if len(zipReader.File) != 3 {
t.Errorf("expected 3 files inside the ZIP, got %d", len(zipReader.File))
}
// Extract files to check their contents and name collision resolutions
extracted := make(map[string]string)
for _, f := range zipReader.File {
rc, err := f.Open()
if err != nil {
t.Fatalf("failed to open zip file entry %s: %v", f.Name, err)
}
content, _ := io.ReadAll(rc)
rc.Close()
extracted[f.Name] = string(content)
}
// Checks
if extracted["file_a.txt"] != "file1 content" {
t.Errorf("file_a.txt content incorrect: %q", extracted["file_a.txt"])
}
if extracted["file_b.txt"] != "file2 content" {
t.Errorf("file_b.txt content incorrect: %q", extracted["file_b.txt"])
}
// The second file_a.txt should be renamed to file_a_1.txt
if extracted["file_a_1.txt"] != "duplicate name file content" {
t.Errorf("file_a_1.txt content incorrect: %q. Extracted files: %v", extracted["file_a_1.txt"], extracted)
}
t.Logf("Successfully unzipped batch. Extracted files: %+v", extracted)
})
}
+219
View File
@@ -0,0 +1,219 @@
/*
Copyright 2025 linux.do
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package user
import (
"strconv"
"strings"
"github.com/gin-gonic/gin"
"github.com/linux-do/credit/internal/apps/oauth"
"github.com/linux-do/credit/internal/common/response"
"github.com/linux-do/credit/internal/db"
"github.com/linux-do/credit/internal/model"
"github.com/linux-do/credit/internal/util"
)
type createTokenRequest struct {
Name string `json:"name"`
}
type tokenResponse struct {
Token string `json:"token"`
Record model.AccessToken `json:"record"`
}
// ListAccessTokens 获取当前用户的 AccessToken 列表
// @Summary 获取当前用户的 AccessToken 列表
// @Description 返回当前登录用户的所有 active access tokens(脱敏后)
// @Tags user
// @Produce json
// @Security SessionCookie
// @Success 200 {object} util.ResponseAny{data=[]model.AccessToken} "令牌列表"
// @Failure 401 {object} util.ResponseAny "未登录"
// @Router /api/v1/user/access-tokens [get]
func ListAccessTokens(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
var tokens []model.AccessToken
if err := db.DB(ctx).Where("user_id = ?", currUser.ID).Order("created_at desc").Find(&tokens).Error; err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, tokens)
}
// CreateAccessToken 创建一个新的 AccessToken
// @Summary 创建一个新的 AccessToken
// @Description 为当前用户新建一个 API 访问令牌,仅在此接口返回一次明文令牌值,请妥善保存。
// @Tags user
// @Accept json
// @Produce json
// @Param request body user.createTokenRequest true "令牌名称"
// @Security SessionCookie
// @Success 200 {object} util.ResponseAny{data=user.tokenResponse} "新建令牌成功"
// @Failure 400 {object} util.ResponseAny "参数错误或超限"
// @Router /api/v1/user/access-tokens [post]
func CreateAccessToken(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
var req createTokenRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.RespondFailure(c, "参数绑定失败")
return
}
req.Name = strings.TrimSpace(req.Name)
if req.Name == "" {
response.RespondFailure(c, "令牌名称不能为空")
return
}
// 检查最大限制(基于 ConfigKeyMaxAPIKeysPerUser 配置,默认值为 5)
maxLimit := 5
if val, err := model.GetIntByKey(ctx, model.ConfigKeyMaxAPIKeysPerUser); err == nil {
maxLimit = val
}
var count int64
if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", currUser.ID).Count(&count).Error; err != nil {
response.RespondFailure(c, err.Error())
return
}
if int(count) >= maxLimit {
response.RespondFailure(c, "已达到访问令牌最大创建数量限制")
return
}
// 生成 Token
tokenStr, err := model.GenerateTokenString()
if err != nil {
response.RespondFailure(c, "生成令牌失败")
return
}
tokenHash := model.HashToken(tokenStr)
maskedToken := model.MaskTokenString(tokenStr)
tokenRecord := model.AccessToken{
UserID: currUser.ID,
Name: req.Name,
TokenHash: tokenHash,
MaskedToken: maskedToken,
}
if err := db.DB(ctx).Create(&tokenRecord).Error; err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, tokenResponse{
Token: tokenStr,
Record: tokenRecord,
})
}
// DeleteAccessToken 删除一个 AccessToken
// @Summary 删除一个 AccessToken
// @Description 撤销并删除一个属于当前用户的 API 访问令牌
// @Tags user
// @Produce json
// @Param id path string true "令牌ID"
// @Security SessionCookie
// @Success 200 {object} util.ResponseAny{data=string} "删除成功"
// @Failure 400 {object} util.ResponseAny "参数错误"
// @Router /api/v1/user/access-tokens/{id} [delete]
func DeleteAccessToken(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.RespondFailure(c, "无效的令牌ID")
return
}
tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).Delete(&model.AccessToken{})
if tx.Error != nil {
response.RespondFailure(c, tx.Error.Error())
return
}
if tx.RowsAffected == 0 {
response.RespondFailure(c, "令牌不存在或无权操作")
return
}
response.RespondSuccess(c, "删除成功")
}
// RotateAccessToken 轮换一个 AccessToken
// @Summary 轮换一个 AccessToken
// @Description 轮换(重新生成)一个属于当前用户的 API 访问令牌的密钥,旧令牌将立即失效
// @Tags user
// @Produce json
// @Param id path string true "令牌ID"
// @Security SessionCookie
// @Success 200 {object} util.ResponseAny{data=user.tokenResponse} "令牌轮换成功"
// @Failure 400 {object} util.ResponseAny "参数错误"
// @Router /api/v1/user/access-tokens/{id}/rotate [post]
func RotateAccessToken(c *gin.Context) {
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.RespondFailure(c, "无效的令牌ID")
return
}
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).First(&tokenRecord).Error; err != nil {
response.RespondFailure(c, "令牌不存在或无权操作")
return
}
// 生成新的 Token
newTokenStr, err := model.GenerateTokenString()
if err != nil {
response.RespondFailure(c, "生成令牌失败")
return
}
newTokenHash := model.HashToken(newTokenStr)
newMaskedToken := model.MaskTokenString(newTokenStr)
tokenRecord.TokenHash = newTokenHash
tokenRecord.MaskedToken = newMaskedToken
tokenRecord.LastUsedAt = nil // 轮换后重置使用时间
if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil {
response.RespondFailure(c, err.Error())
return
}
response.RespondSuccess(c, tokenResponse{
Token: newTokenStr,
Record: tokenRecord,
})
}
-7
View File
@@ -27,7 +27,6 @@ type configModel struct {
Scheduler schedulerConfig `mapstructure:"scheduler"`
Worker workerConfig `mapstructure:"worker"`
ClickHouse clickHouseConfig `mapstructure:"clickhouse"`
LinuxDo linuxDoConfig `mapstructure:"linuxdo"`
OpenAPIRisk openAPIRiskConfig `mapstructure:"openapi_risk"`
Otel otelConfig `mapstructure:"otel"`
S3 s3Config `mapstructure:"s3"`
@@ -42,7 +41,6 @@ type appConfig struct {
APIPrefix string `mapstructure:"api_prefix"`
GracefulShutdownTimeout int `mapstructure:"graceful_shutdown_timeout"`
FrontendURL string `mapstructure:"frontend_url"`
FrontendPayURL string `mapstructure:"frontend_pay_url"`
SessionCookieName string `mapstructure:"session_cookie_name"`
SessionSecret string `mapstructure:"session_secret"`
SessionDomain string `mapstructure:"session_domain"`
@@ -163,11 +161,6 @@ type QueueConfig struct {
Priority int `mapstructure:"priority"`
}
// linuxDoConfig
type linuxDoConfig struct {
ApiKey string `mapstructure:"api_key"`
}
// openAPIRiskConfig OpenAPI 用户风险配置
type openAPIRiskConfig struct {
Enabled bool `mapstructure:"enabled"`
+1
View File
@@ -37,6 +37,7 @@ func Migrate() {
&model.ExternalAccount{},
&model.SystemConfig{},
&model.Upload{},
&model.AccessToken{},
); err != nil {
log.Fatalf("[PostgreSQL] auto migrate failed: %v\n", err)
}
+60
View File
@@ -0,0 +1,60 @@
/*
Copyright 2025 linux.do
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package model
import (
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"fmt"
"time"
)
type AccessToken struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
UserID uint64 `json:"user_id" gorm:"index;not null"`
Name string `json:"name" gorm:"size:128;not null"`
TokenHash string `json:"-" gorm:"size:64;uniqueIndex;not null"`
MaskedToken string `json:"masked_token" gorm:"size:64;not null"`
LastUsedAt *time.Time `json:"last_used_at"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// GenerateTokenString 生成加密安全的随机 Token 值
func GenerateTokenString() (string, error) {
bytes := make([]byte, 24)
if _, err := rand.Read(bytes); err != nil {
return "", err
}
return fmt.Sprintf("at_%s", hex.EncodeToString(bytes)), nil
}
// HashToken 计算 Token 的 SHA-256 哈希值用于数据库存储与查询
func HashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
// MaskTokenString 生成脱敏显示的 Token,仅保留前缀和最后四位
func MaskTokenString(token string) string {
if len(token) <= 8 {
return "at_****"
}
return fmt.Sprintf("%s...%s", token[:7], token[len(token)-4:])
}
+25 -13
View File
@@ -29,20 +29,32 @@ const (
UploadStatusDeleted UploadStatus = "deleted" // 已删除
)
// UploadType 上传类型常量
const (
UploadTypeCover = "cover" // 红包背景封面
UploadTypeHeterotypic = "heterotypic" // 红包异形装饰
)
// UploadMetadata 自定义可扩展的 JSON 字段存储非核心或可选的文件元数据
type UploadMetadata struct {
Width int `json:"width,omitempty"` // 图像/视频宽度 (px)
Height int `json:"height,omitempty"` // 图像/视频高度 (px)
Duration float64 `json:"duration,omitempty"` // 音视频时长 (s)
OriginalMime string `json:"original_mime,omitempty"` // 原始 MIME 类型
UserAgent string `json:"user_agent,omitempty"` // 上传者的 UA
ClientIP string `json:"client_ip,omitempty"` // 上传者 IP
Bucket string `json:"bucket,omitempty"` // 存储桶名称 (适用于 S3 等)
Extra map[string]any `json:"extra,omitempty"` // 其它任意业务自定义元数据
}
// Upload 上传文件记录
type Upload struct {
ID uint64 `json:"id,string" gorm:"primaryKey"`
UserID uint64 `json:"user_id,string" gorm:"index;not null"`
FilePath string `json:"file_path" gorm:"size:500;not null;uniqueIndex"` // 文件路径
FileSize int64 `json:"file_size" gorm:"not null"` // 文件大小(字节)
Type string `json:"type" gorm:"column:type;size:50;not null;index"` // 类型 (cover, heterotypic)
Status UploadStatus `json:"status" gorm:"type:varchar(20);not null"` // 状态
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
ID uint64 `json:"id,string" gorm:"primaryKey"`
UserID uint64 `json:"user_id,string" gorm:"index;not null"`
FileName string `json:"file_name" gorm:"size:255;not null"` // 原始文件名 (例如: image.png)
FilePath string `json:"file_path" gorm:"size:500;not null;index"` // 文件相对路径 / S3 Key
FileSize int64 `json:"file_size" gorm:"not null"` // 文件大小(字节)
MimeType string `json:"mime_type" gorm:"size:100;not null"` // 媒体类型 (MIME, 如 image/png)
Extension string `json:"extension" gorm:"size:50;not null"` // 文件后缀名 (不含点,如 png, pdf)
Hash string `json:"hash" gorm:"size:64;index"` // 文件哈希 (SHA-256/MD5,可用于排重)
StorageDriver string `json:"storage_driver" gorm:"size:50;not null"` // 存储引擎驱动 (如 local, s3, oss)
Type string `json:"type" gorm:"column:type;size:50;not null;index"` // 业务标识类型 (如 avatar, doc, attachment)
Status UploadStatus `json:"status" gorm:"type:varchar(20);not null"` // 状态
Metadata UploadMetadata `json:"metadata" gorm:"serializer:json;type:jsonb"` // 业务扩展元数据
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
+15 -1
View File
@@ -127,13 +127,27 @@ func Serve() {
userRouter.POST("/register", user.Register)
userRouter.GET("/logout", user.Logout)
userRouter.GET("/self", oauth.LoginRequired(), oauth.UserInfo)
// Access Token
tokenRouter := userRouter.Group("/access-tokens")
tokenRouter.Use(oauth.LoginRequired())
{
tokenRouter.GET("", user.ListAccessTokens)
tokenRouter.POST("", user.CreateAccessToken)
tokenRouter.DELETE("/:id", user.DeleteAccessToken)
tokenRouter.POST("/:id/rotate", user.RotateAccessToken)
}
}
// Upload
uploadRouter := apiV1Router.Group("/upload")
uploadRouter.Use(oauth.LoginRequired())
{
// Keep generic uploads if needed
uploadRouter.POST("", upload.UploadFile)
uploadRouter.GET("/my", upload.ListMyFiles)
uploadRouter.DELETE("/:id", upload.DeleteFile)
uploadRouter.GET("/download/:id", upload.DownloadFile)
uploadRouter.POST("/download/batch", upload.BatchDownloadFiles)
}
// Config (public)
+44 -1
View File
@@ -74,17 +74,52 @@ func init() {
log.Printf("[Storage] S3 storage initialized (bucket: %s, prefix: %s, cdn: %s)\n", bucket, keyPrefix, cdnURL)
}
func IsEnabled() bool {
var IsEnabledFunc = func() bool {
return client != nil
}
func IsEnabled() bool {
return IsEnabledFunc()
}
// BuildKey constructs a full S3 object key with the configured prefix.
func BuildKey(path string) string {
return keyPrefix + path
}
var (
// PutObjectFunc enables mocking S3 uploads in tests.
PutObjectFunc = putObjectDefault
// GetObjectFunc enables mocking S3 downloads in tests.
GetObjectFunc = getObjectDefault
// DeleteObjectFunc enables mocking S3 deletion in tests.
DeleteObjectFunc = deleteObjectDefault
)
// MockStorage is a test helper to mock S3 storage operations.
// It returns a function that restores original implementations.
func MockStorage(
mockPut func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error,
mockGet func(ctx context.Context, key string) (*ObjectInfo, error),
mockDelete func(ctx context.Context, key string) error,
) func() {
origPut, origGet, origDelete := PutObjectFunc, GetObjectFunc, DeleteObjectFunc
PutObjectFunc = mockPut
GetObjectFunc = mockGet
DeleteObjectFunc = mockDelete
return func() {
PutObjectFunc = origPut
GetObjectFunc = origGet
DeleteObjectFunc = origDelete
}
}
// PutObject uploads a file to S3.
func PutObject(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
return PutObjectFunc(ctx, key, body, size, contentType)
}
func putObjectDefault(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
ctx, span := otel_trace.Start(ctx, "S3.PutObject", trace.WithSpanKind(trace.SpanKindClient))
defer span.End()
@@ -125,6 +160,10 @@ type ObjectInfo struct {
// GetObject retrieves a file directly from S3.
func GetObject(ctx context.Context, key string) (*ObjectInfo, error) {
return GetObjectFunc(ctx, key)
}
func getObjectDefault(ctx context.Context, key string) (*ObjectInfo, error) {
ctx, span := otel_trace.Start(ctx, "S3.GetObject", trace.WithSpanKind(trace.SpanKindClient))
defer span.End()
@@ -206,6 +245,10 @@ func GetObjectViaProxy(ctx context.Context, key string) (*ObjectInfo, error) {
// DeleteObject deletes a file from S3.
func DeleteObject(ctx context.Context, key string) error {
return DeleteObjectFunc(ctx, key)
}
func deleteObjectDefault(ctx context.Context, key string) error {
ctx, span := otel_trace.Start(ctx, "S3.DeleteObject", trace.WithSpanKind(trace.SpanKindClient))
defer span.End()