fix(backend): harden write registrar boundaries

This commit is contained in:
2026-07-21 16:12:20 +08:00
parent 1fcbb31301
commit 80bec26839
17 changed files with 639 additions and 42 deletions

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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

View File

@@ -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
}))
}