fix(upload): restrict cross-user private file access (UPLOAD-1)

- Add access_mode column to w_uploads table (0 = private, 1 = public) and initialize data in a single migration script
- Enforce strict ownership check for private files during download
- Allow public files to follow whitelisted public-access rules
- Default access_mode to public for avatars and private for generic uploads
- Update frontend service to support optional accessMode parameter
This commit is contained in:
ryan
2026-06-13 10:13:41 +08:00
parent 5412c385dc
commit 7b379863b4
11 changed files with 190 additions and 12 deletions
+3
View File
@@ -5196,6 +5196,9 @@ const docTemplate = `{
"model.Upload": {
"type": "object",
"properties": {
"access_mode": {
"type": "integer"
},
"created_at": {
"type": "string"
},
+3
View File
@@ -5189,6 +5189,9 @@
"model.Upload": {
"type": "object",
"properties": {
"access_mode": {
"type": "integer"
},
"created_at": {
"type": "string"
},
+2
View File
@@ -420,6 +420,8 @@ definitions:
type: object
model.Upload:
properties:
access_mode:
type: integer
created_at:
type: string
extension:
@@ -46,7 +46,8 @@ export class UploadService extends BaseService {
static async uploadFile(
file: File,
type: string = 'generic',
metadata?: Record<string, unknown>
metadata?: Record<string, unknown>,
accessMode?: number
): Promise<Upload> {
const formData = new FormData()
formData.append('file', file)
@@ -54,6 +55,9 @@ export class UploadService extends BaseService {
if (metadata) {
formData.append('metadata', JSON.stringify(metadata))
}
if (accessMode !== undefined) {
formData.append('access_mode', String(accessMode))
}
return this.post<Upload>('', formData, {
headers: { 'Content-Type': 'multipart/form-data' },
@@ -113,13 +117,14 @@ export class UploadService extends BaseService {
static async uploadBase64Image(
base64: string,
type: string = 'generic',
filename: string = 'image.png'
filename: string = 'image.png',
accessMode?: number
): Promise<UploadImageResponse> {
const response = await fetch(base64)
const blob = await response.blob()
const mimeType = base64.match(/data:([^;]+);/)?.[1] || 'image/png'
const file = new File([blob], filename, { type: mimeType })
const result = await this.uploadFile(file, type)
const result = await this.uploadFile(file, type, undefined, accessMode)
return { id: result.id }
}
}
+32 -6
View File
@@ -23,6 +23,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
@@ -56,7 +57,7 @@ func ServeFileByID(c *gin.Context) {
}
// 校验业务白名单与访问权限
if err := checkFileAccessPermission(c, upload.Type); err != nil {
if err := checkFileAccessPermission(c, upload); err != nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
return
}
@@ -309,13 +310,38 @@ func isFilePublic(ctx context.Context, uploadType string) bool {
return false
}
// checkFileAccessPermission 校验文件是否可以被当前请求访问
func checkFileAccessPermission(c *gin.Context, uploadType string) error {
if !isFilePublic(c.Request.Context(), uploadType) {
// 必须进行鉴权
if _, err := oauth.GetUserFromRequest(c); err != nil {
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 {
currUser = u
} else {
currUser, err = oauth.GetUserFromRequest(c)
if err != nil {
return err
}
}
if currUser.ID != ownerID {
return errors.New("forbidden: cross-user access denied")
}
return nil
}
// checkFileAccessPermission 校验文件是否可以被当前请求访问
func checkFileAccessPermission(c *gin.Context, upload *model.Upload) error {
// 1. 私有文件校验(优先级高于当前白名单逻辑)
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 _, err := oauth.GetUserFromRequest(c); err != nil {
return err
}
}
}
return nil
}
+3
View File
@@ -65,6 +65,7 @@ func TestServeFileByIDAccessControl(t *testing.T) {
StorageDriver: "local",
Type: "avatar",
Status: model.UploadStatusUsed,
AccessMode: 1,
}
attachmentFile := model.Upload{
ID: 8002,
@@ -77,6 +78,7 @@ func TestServeFileByIDAccessControl(t *testing.T) {
StorageDriver: "local",
Type: "attachment",
Status: model.UploadStatusUsed,
AccessMode: 1,
}
_ = os.MkdirAll("uploads", 0755)
@@ -254,6 +256,7 @@ func TestImageCompression(t *testing.T) {
StorageDriver: "local",
Type: "avatar", // Whitelisted by default
Status: model.UploadStatusUsed,
AccessMode: 1,
}
dbConn.Create(&uploadRecord)
+37 -3
View File
@@ -25,6 +25,7 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
@@ -117,8 +118,27 @@ func UploadFile(c *gin.Context) {
return
}
uploadType := c.DefaultPostForm("type", "generic")
accessModeStr := c.PostForm("access_mode")
var accessMode int
if accessModeStr == "" {
if uploadType == "avatar" {
accessMode = 1
} else {
accessMode = 0
}
} else {
var err error
accessMode, err = strconv.Atoi(accessModeStr)
if err != nil || (accessMode != 0 && accessMode != 1) {
c.JSON(http.StatusOK, util.Err("无效的 access_mode 参数"))
return
}
}
// 6. 秒传匹配校验:校验数据库中是否存在相同 Hash 且大小一致的可用文件
handled, lookupErr := tryInstantUpload(ctx, c, currUser, fileHash, size, mimeType, ext, origName)
handled, lookupErr := tryInstantUpload(ctx, c, currUser, fileHash, size, mimeType, ext, origName, accessMode)
if handled {
return
}
@@ -155,8 +175,9 @@ func UploadFile(c *gin.Context) {
Extension: ext,
Hash: fileHash,
StorageDriver: storageDriver,
Type: c.DefaultPostForm("type", "generic"),
Type: uploadType,
Status: model.UploadStatusUsed,
AccessMode: accessMode,
Metadata: meta,
}
@@ -196,6 +217,12 @@ func DownloadFile(c *gin.Context) {
return
}
// 校验文件访问权限
if err := checkFileAccessPermission(c, upload); err != nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error_msg": common.UnAuthorized, "data": nil})
return
}
fileName := upload.FileName
quality := normalizeImageQuality(c.Query("quality"))
isImage := strings.HasPrefix(strings.ToLower(upload.MimeType), "image/") || isImageExtension(strings.ToLower(upload.Extension))
@@ -270,6 +297,12 @@ func BatchDownloadFiles(c *gin.Context) {
usedNames := make(map[string]int)
for _, upload := range uploads {
// 校验文件访问权限
if err := 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 {
@@ -460,7 +493,7 @@ func validateUploadExtension(ctx context.Context, ext string) string {
}
// tryInstantUpload 尝试秒传:若数据库已存在相同 Hash 且大小一致的可用文件,直接生成新记录
func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User, fileHash string, size int64, mimeType, ext, origName string) (bool, error) {
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 {
@@ -480,6 +513,7 @@ func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User,
StorageDriver: existing.StorageDriver,
Type: c.DefaultPostForm("type", "generic"),
Status: model.UploadStatusUsed,
AccessMode: accessMode,
Metadata: existing.Metadata,
}
+89
View File
@@ -14,6 +14,7 @@ import (
"net/http"
"net/http/httptest"
"os"
"strconv"
"strings"
"testing"
@@ -621,3 +622,91 @@ func TestBatchDownloadFiles(t *testing.T) {
t.Logf("Successfully unzipped batch. Extracted files: %+v", extracted)
})
}
func TestUploadAccessModeAccessControl(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }()
user1 := &model.User{ID: 1001, Username: "user1"}
user2 := &model.User{ID: 1002, Username: "user2"}
// Seed user1
if err := dbConn.Create(user1).Error; err != nil {
t.Fatalf("create test user1 failed: %v", err)
}
// Seed user2
if err := dbConn.Create(user2).Error; err != nil {
t.Fatalf("create test user2 failed: %v", err)
}
router := setupTestRouter(user1)
// 1. Upload private file for user1 (explicitly specifying access_mode = 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")
contentType, body := createMultipartRequest(t, "file", "private.png", imgContent, map[string]string{
"type": "generic",
"access_mode": "0",
})
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("Upload failed: %d, %s", w.Code, w.Body.String())
}
t.Logf("Raw upload response: %s", w.Body.String())
var resp1 testResponse
_ = json.Unmarshal(w.Body.Bytes(), &resp1)
var upload1 model.Upload
_ = json.Unmarshal(resp1.Data, &upload1)
if upload1.AccessMode != 0 {
t.Errorf("expected access_mode 0, got %d", upload1.AccessMode)
}
// 2. Upload public file for user1 (type avatar, should default to public 1)
contentType2, body2 := createMultipartRequest(t, "file", "public.png", []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"), 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)
var resp2 testResponse
_ = json.Unmarshal(w2.Body.Bytes(), &resp2)
var upload2 model.Upload
_ = json.Unmarshal(resp2.Data, &upload2)
if upload2.AccessMode != 1 {
t.Errorf("expected access_mode 1 (public) for avatar, got %d", upload2.AccessMode)
}
// 3. Verify accessing private file as user1 (owner) succeeds
wAccessOwner := httptest.NewRecorder()
reqAccessOwner, _ := http.NewRequest("GET", "/api/v1/upload/download/"+strconv.FormatUint(upload1.ID, 10), nil)
router.ServeHTTP(wAccessOwner, reqAccessOwner)
if wAccessOwner.Code != http.StatusOK {
t.Errorf("owner should be allowed to download private file, got status %d", wAccessOwner.Code)
}
// 4. Verify accessing private file as user2 (non-owner) fails
routerUser2 := setupTestRouter(user2)
wAccessOther := httptest.NewRecorder()
reqAccessOther, _ := http.NewRequest("GET", "/api/v1/upload/download/"+strconv.FormatUint(upload1.ID, 10), nil)
routerUser2.ServeHTTP(wAccessOther, reqAccessOther)
if wAccessOther.Code != http.StatusUnauthorized {
t.Errorf("non-owner should be denied download of private file, got status %d, want 401", wAccessOther.Code)
}
// 5. Verify accessing public file as user2 (non-owner) succeeds
wAccessPublic := httptest.NewRecorder()
reqAccessPublic, _ := http.NewRequest("GET", "/api/v1/upload/download/"+strconv.FormatUint(upload2.ID, 10), nil)
routerUser2.ServeHTTP(wAccessPublic, reqAccessPublic)
if wAccessPublic.Code != http.StatusOK {
t.Errorf("any logged-in user should be allowed to download public file, got status %d", wAccessPublic.Code)
}
}
@@ -0,0 +1,6 @@
-- +goose Up
ALTER TABLE w_uploads ADD COLUMN access_mode INTEGER NOT NULL DEFAULT 0;
UPDATE w_uploads SET access_mode = 1 WHERE type = 'avatar';
-- +goose Down
ALTER TABLE w_uploads DROP COLUMN access_mode;
@@ -0,0 +1,6 @@
-- +goose Up
ALTER TABLE w_uploads ADD COLUMN access_mode INTEGER NOT NULL DEFAULT 0;
UPDATE w_uploads SET access_mode = 1 WHERE type = 'avatar';
-- +goose Down
ALTER TABLE w_uploads DROP COLUMN access_mode;
+1
View File
@@ -43,6 +43,7 @@ type Upload struct {
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"` // 状态
AccessMode int `json:"access_mode" gorm:"column:access_mode;not null;default:0"`
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"`