fix(backend): harden write registrar boundaries
This commit is contained in:
@@ -2,6 +2,7 @@ package files
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
@@ -26,12 +27,20 @@ type SourceDTO struct {
|
||||
|
||||
// Handler 负责 multipart 边界,文件名清理和本地路径构造始终委托给 Service。
|
||||
type Handler struct {
|
||||
service *Service
|
||||
service *Service
|
||||
maxUploadBytes int64
|
||||
}
|
||||
|
||||
// DefaultMaxUploadBytes 是未显式配置时的请求体上限,覆盖完整 multipart 内容。
|
||||
const DefaultMaxUploadBytes int64 = 32 << 20
|
||||
|
||||
// NewHandler 创建文件资料 HTTP registrar。
|
||||
func NewHandler(service *Service) *Handler {
|
||||
return &Handler{service: service}
|
||||
func NewHandler(service *Service, limits ...int64) *Handler {
|
||||
limit := DefaultMaxUploadBytes
|
||||
if len(limits) > 0 && limits[0] > 0 {
|
||||
limit = limits[0]
|
||||
}
|
||||
return &Handler{service: service, maxUploadBytes: limit}
|
||||
}
|
||||
|
||||
// Register 将文件资料上传接口注册到上层提供的 /api/v1 路由组。
|
||||
@@ -55,8 +64,15 @@ func (h *Handler) upload(c *gin.Context) {
|
||||
writeSourceError(c, err)
|
||||
return
|
||||
}
|
||||
// MaxBytesReader 必须在任何 multipart 解析前安装,限制字段、边界和文件内容的总请求量。
|
||||
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, h.maxUploadBytes)
|
||||
fileHeader, err := c.FormFile("file")
|
||||
if err != nil {
|
||||
var maxBytesError *http.MaxBytesError
|
||||
if errors.As(err, &maxBytesError) {
|
||||
httpx.Error(c, http.StatusRequestEntityTooLarge, "payload_too_large", "上传内容超过大小限制")
|
||||
return
|
||||
}
|
||||
httpx.Error(c, http.StatusBadRequest, "invalid_request", "请选择要上传的文件")
|
||||
return
|
||||
}
|
||||
@@ -75,6 +91,10 @@ func (h *Handler) upload(c *gin.Context) {
|
||||
}
|
||||
source, err := h.service.CreateSource(userID, project, c.PostForm("title"), stored)
|
||||
if err != nil {
|
||||
if cleanupErr := h.service.Remove(stored); cleanupErr != nil {
|
||||
// 清理错误仅写服务端日志,响应继续使用稳定中文错误,不暴露 storage root。
|
||||
log.Printf("source file cleanup failed: %v", cleanupErr)
|
||||
}
|
||||
writeSourceError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -3,7 +3,9 @@ package files
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -76,7 +78,102 @@ func TestFileRegistrarChecksOwnershipBeforeWritingFile(t *testing.T) {
|
||||
require.Empty(t, entries)
|
||||
}
|
||||
|
||||
func TestFileRegistrarRemovesStoredFileWhenSourceCreateFails(t *testing.T) {
|
||||
router, database, project, storageRoot := newFileHandlerTestRouter(t, 1, 1)
|
||||
require.NoError(t, database.Callback().Create().Before("gorm:create").Register("test:fail_source_create", func(tx *gorm.DB) {
|
||||
if tx.Statement.Schema != nil && tx.Statement.Schema.Table == (models.SenlinAgentSource{}).TableName() {
|
||||
tx.AddError(errors.New("forced source create failure"))
|
||||
}
|
||||
}))
|
||||
body := &bytes.Buffer{}
|
||||
writer := multipart.NewWriter(body)
|
||||
part, err := writer.CreateFormFile("file", "orphan.txt")
|
||||
require.NoError(t, err)
|
||||
_, err = part.Write([]byte("must be compensated"))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, writer.Close())
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/sources", body)
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer test-token")
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusInternalServerError, rec.Code, rec.Body.String())
|
||||
require.NotContains(t, rec.Body.String(), storageRoot)
|
||||
entries, err := os.ReadDir(storageRoot)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, entries, "source 元数据失败后存储目录应为空")
|
||||
}
|
||||
|
||||
func TestFileRegistrarRejectsRequestOverConfiguredUploadLimit(t *testing.T) {
|
||||
router, _, project, storageRoot := newFileHandlerTestRouterWithLimit(t, 1, 1, 256)
|
||||
body := &bytes.Buffer{}
|
||||
writer := multipart.NewWriter(body)
|
||||
part, err := writer.CreateFormFile("file", "oversize.bin")
|
||||
require.NoError(t, err)
|
||||
_, err = part.Write(bytes.Repeat([]byte("x"), 1024))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, writer.Close())
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/sources", body)
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer test-token")
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusRequestEntityTooLarge, rec.Code, rec.Body.String())
|
||||
var payload httpx.ErrorEnvelope
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
|
||||
require.Equal(t, "payload_too_large", payload.Error.Code)
|
||||
require.Equal(t, "上传内容超过大小限制", payload.Error.Message)
|
||||
entries, err := os.ReadDir(storageRoot)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, entries)
|
||||
}
|
||||
|
||||
func TestFileRegistrarLogsCleanupFailureWithoutLeakingStorageRoot(t *testing.T) {
|
||||
router, database, project, storageRoot, service := newFileHandlerTestRouterWithService(t, 1, 1, DefaultMaxUploadBytes)
|
||||
require.NoError(t, database.Callback().Create().Before("gorm:create").Register("test:fail_source_create_for_cleanup_log", func(tx *gorm.DB) {
|
||||
if tx.Statement.Schema != nil && tx.Statement.Schema.Table == (models.SenlinAgentSource{}).TableName() {
|
||||
tx.AddError(errors.New("forced source create failure"))
|
||||
}
|
||||
}))
|
||||
service.remove = func(path string) error { return fmt.Errorf("cleanup blocked for %s", path) }
|
||||
var serverLog bytes.Buffer
|
||||
previousLogOutput := log.Writer()
|
||||
log.SetOutput(&serverLog)
|
||||
t.Cleanup(func() { log.SetOutput(previousLogOutput) })
|
||||
body := &bytes.Buffer{}
|
||||
writer := multipart.NewWriter(body)
|
||||
part, err := writer.CreateFormFile("file", "cleanup-failure.txt")
|
||||
require.NoError(t, err)
|
||||
_, err = part.Write([]byte("content"))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, writer.Close())
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/sources", body)
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer test-token")
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusInternalServerError, rec.Code, rec.Body.String())
|
||||
require.NotContains(t, rec.Body.String(), storageRoot)
|
||||
require.Contains(t, serverLog.String(), "source file cleanup failed")
|
||||
require.Contains(t, serverLog.String(), storageRoot)
|
||||
}
|
||||
|
||||
func newFileHandlerTestRouter(t *testing.T, currentUserID, ownerID uint) (*gin.Engine, *gorm.DB, models.SenlinAgentProject, string) {
|
||||
return newFileHandlerTestRouterWithLimit(t, currentUserID, ownerID, DefaultMaxUploadBytes)
|
||||
}
|
||||
|
||||
func newFileHandlerTestRouterWithLimit(t *testing.T, currentUserID, ownerID uint, maxUploadBytes int64) (*gin.Engine, *gorm.DB, models.SenlinAgentProject, string) {
|
||||
router, database, project, storageRoot, _ := newFileHandlerTestRouterWithService(t, currentUserID, ownerID, maxUploadBytes)
|
||||
return router, database, project, storageRoot
|
||||
}
|
||||
|
||||
func newFileHandlerTestRouterWithService(t *testing.T, currentUserID, ownerID uint, maxUploadBytes int64) (*gin.Engine, *gorm.DB, models.SenlinAgentProject, string, *Service) {
|
||||
t.Helper()
|
||||
database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{TranslateError: true})
|
||||
require.NoError(t, err)
|
||||
@@ -87,10 +184,11 @@ func newFileHandlerTestRouter(t *testing.T, currentUserID, ownerID uint) (*gin.E
|
||||
project := models.SenlinAgentProject{OwnerID: ownerID, Name: "Files", Identifier: fmt.Sprintf("FILES-%d", ownerID)}
|
||||
require.NoError(t, database.Create(&project).Error)
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot, database)
|
||||
router := httpx.NewProtectedRouter(
|
||||
config.Config{Env: "test"},
|
||||
func(string) (uint, error) { return currentUserID, nil },
|
||||
NewHandler(NewService(storageRoot, database)),
|
||||
NewHandler(service, maxUploadBytes),
|
||||
)
|
||||
return router, database, project, storageRoot
|
||||
return router, database, project, storageRoot, service
|
||||
}
|
||||
|
||||
@@ -14,8 +14,17 @@ import (
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
root string
|
||||
db *gorm.DB
|
||||
root string
|
||||
db *gorm.DB
|
||||
createTemp func(string, string) (stagedFile, error)
|
||||
rename func(string, string) error
|
||||
remove func(string) error
|
||||
}
|
||||
|
||||
type stagedFile interface {
|
||||
io.Writer
|
||||
Close() error
|
||||
Name() string
|
||||
}
|
||||
|
||||
type StoredFile struct {
|
||||
@@ -30,7 +39,12 @@ func NewService(root string, databases ...*gorm.DB) *Service {
|
||||
if len(databases) > 0 {
|
||||
database = databases[0]
|
||||
}
|
||||
return &Service{root: root, db: database}
|
||||
return &Service{
|
||||
root: root, db: database,
|
||||
createTemp: func(directory, pattern string) (stagedFile, error) { return os.CreateTemp(directory, pattern) },
|
||||
rename: os.Rename,
|
||||
remove: os.Remove,
|
||||
}
|
||||
}
|
||||
|
||||
// Save 清理客户端文件名,并只返回供持久化的相对路径;绝对路径不得进入 API DTO。
|
||||
@@ -40,21 +54,90 @@ func (s *Service) Save(projectID uint, originalName string, content io.Reader) (
|
||||
cleanName = "upload.bin"
|
||||
}
|
||||
relative := filepath.ToSlash(filepath.Join("projects", fmt.Sprint(projectID), fmt.Sprintf("%d-%s", time.Now().UnixNano(), cleanName)))
|
||||
absolute := filepath.Join(s.root, filepath.FromSlash(relative))
|
||||
if err := os.MkdirAll(filepath.Dir(absolute), 0o755); err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
file, err := os.Create(absolute)
|
||||
absolute, err := s.absolutePath(relative)
|
||||
if err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(absolute), 0o755); err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
// 临时文件与最终文件位于同一目录,Close 成功后再原子替换,避免暴露半写入内容。
|
||||
file, err := s.createTemp(filepath.Dir(absolute), ".upload-*")
|
||||
if err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
defer file.Close()
|
||||
if _, err := io.Copy(file, content); err != nil {
|
||||
_ = file.Close()
|
||||
s.cleanupFailedSave(file.Name(), absolute)
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := file.Close(); err != nil {
|
||||
s.cleanupFailedSave(file.Name(), absolute)
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := s.rename(file.Name(), absolute); err != nil {
|
||||
s.cleanupFailedSave(file.Name(), absolute)
|
||||
return StoredFile{}, err
|
||||
}
|
||||
return StoredFile{OriginalName: cleanName, RelativePath: relative, AbsolutePath: absolute}, nil
|
||||
}
|
||||
|
||||
// Remove 只依据 Service 生成的相对路径定位文件,并尽量移除直至存储根目录的空父目录。
|
||||
func (s *Service) Remove(stored StoredFile) error {
|
||||
absolute, err := s.absolutePath(stored.RelativePath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.remove(absolute); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
return err
|
||||
}
|
||||
root, err := filepath.Abs(s.root)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for directory := filepath.Dir(absolute); directory != root; directory = filepath.Dir(directory) {
|
||||
entries, err := os.ReadDir(directory)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(entries) > 0 {
|
||||
return nil
|
||||
}
|
||||
if err := s.remove(directory); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) absolutePath(relativePath string) (string, error) {
|
||||
cleaned := filepath.Clean(filepath.FromSlash(strings.TrimSpace(relativePath)))
|
||||
if cleaned == "." || filepath.IsAbs(cleaned) || cleaned == ".." || strings.HasPrefix(cleaned, ".."+string(filepath.Separator)) {
|
||||
return "", ErrSourcePathRequired
|
||||
}
|
||||
root, err := filepath.Abs(s.root)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
absolute, err := filepath.Abs(filepath.Join(root, cleaned))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
relativeToRoot, err := filepath.Rel(root, absolute)
|
||||
if err != nil || relativeToRoot == ".." || strings.HasPrefix(relativeToRoot, ".."+string(filepath.Separator)) {
|
||||
return "", ErrSourcePathRequired
|
||||
}
|
||||
return absolute, nil
|
||||
}
|
||||
|
||||
func (s *Service) cleanupFailedSave(tempPath, finalPath string) {
|
||||
_ = s.remove(tempPath)
|
||||
_ = s.remove(finalPath)
|
||||
}
|
||||
|
||||
func (s *Service) database() *gorm.DB {
|
||||
if s.db != nil {
|
||||
return s.db
|
||||
|
||||
@@ -1,12 +1,37 @@
|
||||
package files
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type failingReader struct {
|
||||
sent bool
|
||||
}
|
||||
|
||||
func (r *failingReader) Read(buffer []byte) (int, error) {
|
||||
if !r.sent {
|
||||
r.sent = true
|
||||
return copy(buffer, "partial"), nil
|
||||
}
|
||||
return 0, errors.New("copy failed")
|
||||
}
|
||||
|
||||
type closeFailingFile struct {
|
||||
*os.File
|
||||
}
|
||||
|
||||
func (f *closeFailingFile) Close() error {
|
||||
_ = f.File.Close()
|
||||
return errors.New("close failed")
|
||||
}
|
||||
|
||||
func TestSaveStoresFileUnderProjectDirectory(t *testing.T) {
|
||||
service := NewService(t.TempDir())
|
||||
|
||||
@@ -27,3 +52,57 @@ func TestSaveNeutralizesPathTraversal(t *testing.T) {
|
||||
require.Equal(t, "secret.txt", stored.OriginalName)
|
||||
require.Contains(t, stored.RelativePath, "projects/12/")
|
||||
}
|
||||
|
||||
func TestSaveRemovesTemporaryAndPartialFilesWhenCopyFails(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
|
||||
_, err := service.Save(12, "broken.bin", &failingReader{})
|
||||
|
||||
require.ErrorContains(t, err, "copy failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
}
|
||||
|
||||
func TestSaveRemovesTemporaryFileWhenCloseFails(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
service.createTemp = func(directory, pattern string) (stagedFile, error) {
|
||||
file, err := os.CreateTemp(directory, pattern)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &closeFailingFile{File: file}, nil
|
||||
}
|
||||
|
||||
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
|
||||
|
||||
require.ErrorContains(t, err, "close failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
}
|
||||
|
||||
func TestSaveRemovesTemporaryAndPartialFilesWhenRenameFails(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
service.rename = func(oldPath, newPath string) error {
|
||||
require.NoError(t, os.WriteFile(newPath, []byte("partial final"), 0o600))
|
||||
return errors.New("rename failed")
|
||||
}
|
||||
|
||||
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
|
||||
|
||||
require.ErrorContains(t, err, "rename failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
}
|
||||
|
||||
func requireNoStoredFiles(t *testing.T, storageRoot string) {
|
||||
t.Helper()
|
||||
require.NoError(t, filepath.WalkDir(storageRoot, func(path string, entry fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if path != storageRoot && !entry.IsDir() {
|
||||
t.Fatalf("unexpected stored file after failure: %s", path)
|
||||
}
|
||||
return nil
|
||||
}))
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user