fix(backend): harden final MVP invariants
This commit is contained in:
@@ -4,6 +4,7 @@ import (
|
||||
"errors"
|
||||
"log"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -14,15 +15,15 @@ import (
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
// SourceDTO 是文件资料写接口的稳定响应,仅包含可公开的相对存储路径。
|
||||
// SourceDTO 是文件资料写接口的稳定响应,仅返回 opaque storage key,不暴露物理目录布局。
|
||||
type SourceDTO struct {
|
||||
ID string `json:"id"`
|
||||
ProjectID string `json:"projectId"`
|
||||
Kind string `json:"kind"`
|
||||
Title string `json:"title"`
|
||||
FilePath string `json:"filePath"`
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
UpdatedAt time.Time `json:"updatedAt"`
|
||||
ID string `json:"id"`
|
||||
ProjectID string `json:"projectId"`
|
||||
Kind string `json:"kind"`
|
||||
Title string `json:"title"`
|
||||
StorageKey string `json:"storageKey"`
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
UpdatedAt time.Time `json:"updatedAt"`
|
||||
}
|
||||
|
||||
// Handler 负责 multipart 边界,文件名清理和本地路径构造始终委托给 Service。
|
||||
@@ -84,7 +85,7 @@ func (h *Handler) upload(c *gin.Context) {
|
||||
defer file.Close()
|
||||
|
||||
// 先由文件服务保存内容,再用其返回的相对路径创建资料记录;handler 不拼接任何本地路径。
|
||||
stored, err := h.service.Save(project.ID, fileHeader.Filename, file)
|
||||
stored, err := h.service.Save(project.Identity, fileHeader.Filename, file)
|
||||
if err != nil {
|
||||
httpx.Error(c, http.StatusInternalServerError, "internal_error", "文件保存失败")
|
||||
return
|
||||
@@ -104,7 +105,7 @@ func (h *Handler) upload(c *gin.Context) {
|
||||
func sourceDTO(source models.SenlinAgentSource) SourceDTO {
|
||||
return SourceDTO{
|
||||
ID: source.Identity, ProjectID: source.ProjectIdentity, Kind: source.Kind,
|
||||
Title: source.Title, FilePath: source.FilePath,
|
||||
Title: source.Title, StorageKey: filepath.Base(filepath.FromSlash(source.FilePath)),
|
||||
CreatedAt: source.CreatedAt.UTC(), UpdatedAt: source.UpdatedAt.UTC(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -43,14 +44,17 @@ func TestFileRegistrarSavesBeforeCreatingIdentitySourceDTO(t *testing.T) {
|
||||
var source models.SenlinAgentSource
|
||||
require.NoError(t, database.Where("project_id = ?", project.ID).First(&source).Error)
|
||||
require.FileExists(t, filepath.Join(storageRoot, filepath.FromSlash(source.FilePath)))
|
||||
require.Contains(t, filepath.ToSlash(source.FilePath), "projects/"+project.Identity+"/")
|
||||
require.NotContains(t, filepath.ToSlash(source.FilePath), fmt.Sprintf("projects/%d/", project.ID))
|
||||
var payload map[string]any
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
|
||||
require.Equal(t, source.Identity, payload["id"])
|
||||
require.Equal(t, project.Identity, payload["projectId"])
|
||||
require.Equal(t, filepath.ToSlash(source.FilePath), payload["filePath"])
|
||||
require.NotContains(t, payload["filePath"], storageRoot)
|
||||
require.Equal(t, filepath.Base(source.FilePath), payload["storageKey"])
|
||||
require.NotContains(t, payload, "filePath")
|
||||
require.NotContains(t, payload, "AbsolutePath")
|
||||
require.NotContains(t, rec.Body.String(), storageRoot)
|
||||
require.NotContains(t, rec.Body.String(), fmt.Sprintf("projects/%d/", project.ID))
|
||||
}
|
||||
|
||||
func TestFileRegistrarChecksOwnershipBeforeWritingFile(t *testing.T) {
|
||||
@@ -139,7 +143,12 @@ func TestFileRegistrarLogsCleanupFailureWithoutLeakingStorageRoot(t *testing.T)
|
||||
tx.AddError(errors.New("forced source create failure"))
|
||||
}
|
||||
}))
|
||||
service.remove = func(path string) error { return fmt.Errorf("cleanup blocked for %s", path) }
|
||||
service.remove = func(path string) error {
|
||||
if strings.HasPrefix(filepath.Base(path), ".upload-") {
|
||||
return os.Remove(path)
|
||||
}
|
||||
return fmt.Errorf("cleanup blocked for %s", path)
|
||||
}
|
||||
var serverLog bytes.Buffer
|
||||
previousLogOutput := log.Writer()
|
||||
log.SetOutput(&serverLog)
|
||||
|
||||
@@ -1,24 +1,25 @@
|
||||
package files
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
root string
|
||||
db *gorm.DB
|
||||
createTemp func(string, string) (stagedFile, error)
|
||||
rename func(string, string) error
|
||||
remove func(string) error
|
||||
root string
|
||||
db *gorm.DB
|
||||
createTemp func(string, string) (stagedFile, error)
|
||||
publish func(string, string) error
|
||||
remove func(string) error
|
||||
newStorageKey func() (string, error)
|
||||
}
|
||||
|
||||
type stagedFile interface {
|
||||
@@ -29,6 +30,7 @@ type stagedFile interface {
|
||||
|
||||
type StoredFile struct {
|
||||
OriginalName string
|
||||
StorageKey string
|
||||
RelativePath string
|
||||
AbsolutePath string
|
||||
}
|
||||
@@ -41,45 +43,73 @@ func NewService(root string, databases ...*gorm.DB) *Service {
|
||||
}
|
||||
return &Service{
|
||||
root: root, db: database,
|
||||
createTemp: func(directory, pattern string) (stagedFile, error) { return os.CreateTemp(directory, pattern) },
|
||||
rename: os.Rename,
|
||||
remove: os.Remove,
|
||||
createTemp: func(directory, pattern string) (stagedFile, error) { return os.CreateTemp(directory, pattern) },
|
||||
publish: os.Link,
|
||||
remove: os.Remove,
|
||||
newStorageKey: randomStorageKey,
|
||||
}
|
||||
}
|
||||
|
||||
// Save 清理客户端文件名,并只返回供持久化的相对路径;绝对路径不得进入 API DTO。
|
||||
func (s *Service) Save(projectID uint, originalName string, content io.Reader) (StoredFile, error) {
|
||||
// Save 使用项目公开 identity 和随机 opaque key 定位文件;硬链接发布提供跨请求的排他创建语义。
|
||||
func (s *Service) Save(projectIdentity string, originalName string, content io.Reader) (StoredFile, error) {
|
||||
cleanName := filepath.Base(strings.ReplaceAll(strings.TrimSpace(originalName), "\\", "/"))
|
||||
if cleanName == "." || cleanName == "" {
|
||||
cleanName = "upload.bin"
|
||||
}
|
||||
relative := filepath.ToSlash(filepath.Join("projects", fmt.Sprint(projectID), fmt.Sprintf("%d-%s", time.Now().UnixNano(), cleanName)))
|
||||
absolute, err := s.absolutePath(relative)
|
||||
projectSegment := filepath.Base(strings.ReplaceAll(strings.TrimSpace(projectIdentity), "\\", "/"))
|
||||
if projectSegment == "." || projectSegment == "" || projectSegment != strings.TrimSpace(projectIdentity) {
|
||||
return StoredFile{}, ErrSourcePathRequired
|
||||
}
|
||||
directoryRelative := filepath.ToSlash(filepath.Join("projects", projectSegment))
|
||||
directoryAbsolute, err := s.absolutePath(directoryRelative)
|
||||
if err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(absolute), 0o755); err != nil {
|
||||
if err := os.MkdirAll(directoryAbsolute, 0o755); err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
// 临时文件与最终文件位于同一目录,Close 成功后再原子替换,避免暴露半写入内容。
|
||||
file, err := s.createTemp(filepath.Dir(absolute), ".upload-*")
|
||||
// 临时文件与最终文件位于同一目录,完整关闭后再排他发布,避免暴露半写入内容。
|
||||
file, err := s.createTemp(directoryAbsolute, ".upload-*")
|
||||
if err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if _, err := io.Copy(file, content); err != nil {
|
||||
_ = file.Close()
|
||||
s.cleanupFailedSave(file.Name(), absolute)
|
||||
s.cleanupOwnedFiles(file.Name())
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := file.Close(); err != nil {
|
||||
s.cleanupFailedSave(file.Name(), absolute)
|
||||
s.cleanupOwnedFiles(file.Name())
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := s.rename(file.Name(), absolute); err != nil {
|
||||
s.cleanupFailedSave(file.Name(), absolute)
|
||||
return StoredFile{}, err
|
||||
for attempt := 0; attempt < storageKeyAttempts; attempt++ {
|
||||
storageKey, err := s.newStorageKey()
|
||||
if err != nil {
|
||||
s.cleanupOwnedFiles(file.Name())
|
||||
return StoredFile{}, err
|
||||
}
|
||||
relative := filepath.ToSlash(filepath.Join(directoryRelative, storageKey))
|
||||
absolute, err := s.absolutePath(relative)
|
||||
if err != nil {
|
||||
s.cleanupOwnedFiles(file.Name())
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := s.publish(file.Name(), absolute); err != nil {
|
||||
if errors.Is(err, os.ErrExist) {
|
||||
continue
|
||||
}
|
||||
s.cleanupOwnedFiles(file.Name())
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := s.remove(file.Name()); err != nil {
|
||||
// final path 由本次排他发布创建,因此这里只清理本次请求拥有的两个路径。
|
||||
s.cleanupOwnedFiles(file.Name(), absolute)
|
||||
return StoredFile{}, err
|
||||
}
|
||||
return StoredFile{OriginalName: cleanName, StorageKey: storageKey, RelativePath: relative, AbsolutePath: absolute}, nil
|
||||
}
|
||||
return StoredFile{OriginalName: cleanName, RelativePath: relative, AbsolutePath: absolute}, nil
|
||||
s.cleanupOwnedFiles(file.Name())
|
||||
return StoredFile{}, ErrStorageKeyCollision
|
||||
}
|
||||
|
||||
// Remove 只依据 Service 生成的相对路径定位文件,并尽量移除直至存储根目录的空父目录。
|
||||
@@ -133,9 +163,10 @@ func (s *Service) absolutePath(relativePath string) (string, error) {
|
||||
return absolute, nil
|
||||
}
|
||||
|
||||
func (s *Service) cleanupFailedSave(tempPath, finalPath string) {
|
||||
_ = s.remove(tempPath)
|
||||
_ = s.remove(finalPath)
|
||||
func (s *Service) cleanupOwnedFiles(paths ...string) {
|
||||
for _, path := range paths {
|
||||
_ = s.remove(path)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) database() *gorm.DB {
|
||||
@@ -148,8 +179,19 @@ func (s *Service) database() *gorm.DB {
|
||||
var (
|
||||
ErrSourceTitleRequired = errors.New("source title is required")
|
||||
ErrSourcePathRequired = errors.New("source file path is required")
|
||||
ErrStorageKeyCollision = errors.New("unable to allocate unique storage key")
|
||||
)
|
||||
|
||||
const storageKeyAttempts = 8
|
||||
|
||||
func randomStorageKey() (string, error) {
|
||||
buffer := make([]byte, 16)
|
||||
if _, err := rand.Read(buffer); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(buffer), nil
|
||||
}
|
||||
|
||||
// CreateSource 只持久化 Save 产生的相对路径,不接受 handler 自行拼接本地路径。
|
||||
func (s *Service) CreateSource(ownerID uint, project *models.SenlinAgentProject, title string, stored StoredFile) (*models.SenlinAgentSource, error) {
|
||||
title = strings.TrimSpace(title)
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -32,32 +33,36 @@ func (f *closeFailingFile) Close() error {
|
||||
return errors.New("close failed")
|
||||
}
|
||||
|
||||
func TestSaveStoresFileUnderProjectDirectory(t *testing.T) {
|
||||
func TestSaveStoresFileUnderProjectIdentityWithOpaqueKey(t *testing.T) {
|
||||
service := NewService(t.TempDir())
|
||||
projectIdentity := "019b0000-0000-7000-8000-000000000012"
|
||||
|
||||
stored, err := service.Save(12, "brief.md", strings.NewReader("hello"))
|
||||
stored, err := service.Save(projectIdentity, "brief.md", strings.NewReader("hello"))
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "brief.md", stored.OriginalName)
|
||||
require.Contains(t, stored.RelativePath, "projects/12/")
|
||||
require.Contains(t, stored.RelativePath, "projects/"+projectIdentity+"/")
|
||||
require.NotContains(t, stored.RelativePath, "projects/12/")
|
||||
require.Regexp(t, `^[a-f0-9]{32}$`, stored.StorageKey)
|
||||
require.Equal(t, stored.StorageKey, filepath.Base(stored.RelativePath))
|
||||
require.FileExists(t, stored.AbsolutePath)
|
||||
}
|
||||
|
||||
func TestSaveNeutralizesPathTraversal(t *testing.T) {
|
||||
service := NewService(t.TempDir())
|
||||
|
||||
stored, err := service.Save(12, "..\\..\\secret.txt", strings.NewReader("hello"))
|
||||
stored, err := service.Save("019b0000-0000-7000-8000-000000000012", "..\\..\\secret.txt", strings.NewReader("hello"))
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "secret.txt", stored.OriginalName)
|
||||
require.Contains(t, stored.RelativePath, "projects/12/")
|
||||
require.NotContains(t, stored.RelativePath, "secret.txt")
|
||||
}
|
||||
|
||||
func TestSaveRemovesTemporaryAndPartialFilesWhenCopyFails(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
|
||||
_, err := service.Save(12, "broken.bin", &failingReader{})
|
||||
_, err := service.Save("019b0000-0000-7000-8000-000000000012", "broken.bin", &failingReader{})
|
||||
|
||||
require.ErrorContains(t, err, "copy failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
@@ -74,24 +79,63 @@ func TestSaveRemovesTemporaryFileWhenCloseFails(t *testing.T) {
|
||||
return &closeFailingFile{File: file}, nil
|
||||
}
|
||||
|
||||
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
|
||||
_, err := service.Save("019b0000-0000-7000-8000-000000000012", "broken.bin", strings.NewReader("content"))
|
||||
|
||||
require.ErrorContains(t, err, "close failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
}
|
||||
|
||||
func TestSaveRemovesTemporaryAndPartialFilesWhenRenameFails(t *testing.T) {
|
||||
func TestSaveCollisionNeverOverwritesOrDeletesExistingFile(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")
|
||||
service.newStorageKey = func() (string, error) { return "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", nil }
|
||||
projectIdentity := "019b0000-0000-7000-8000-000000000012"
|
||||
first, err := service.Save(projectIdentity, "first.bin", strings.NewReader("first owner"))
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = service.Save(projectIdentity, "second.bin", strings.NewReader("second owner"))
|
||||
|
||||
require.ErrorIs(t, err, ErrStorageKeyCollision)
|
||||
content, readErr := os.ReadFile(first.AbsolutePath)
|
||||
require.NoError(t, readErr)
|
||||
require.Equal(t, "first owner", string(content))
|
||||
requireOnlyStoredFile(t, storageRoot, first.AbsolutePath)
|
||||
}
|
||||
|
||||
func TestConcurrentSaveWithSameKeyPublishesExactlyOneOwner(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
service.newStorageKey = func() (string, error) { return "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb", nil }
|
||||
projectIdentity := "019b0000-0000-7000-8000-000000000012"
|
||||
start := make(chan struct{})
|
||||
type result struct {
|
||||
stored StoredFile
|
||||
err error
|
||||
}
|
||||
results := make(chan result, 2)
|
||||
var ready sync.WaitGroup
|
||||
ready.Add(2)
|
||||
for _, body := range []string{"alpha", "beta"} {
|
||||
go func(content string) {
|
||||
ready.Done()
|
||||
<-start
|
||||
stored, err := service.Save(projectIdentity, content+".bin", strings.NewReader(content))
|
||||
results <- result{stored: stored, err: err}
|
||||
}(body)
|
||||
}
|
||||
ready.Wait()
|
||||
close(start)
|
||||
firstResult, secondResult := <-results, <-results
|
||||
|
||||
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
|
||||
|
||||
require.ErrorContains(t, err, "rename failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
errors := []error{firstResult.err, secondResult.err}
|
||||
require.Equal(t, 1, countNilErrors(errors))
|
||||
require.Equal(t, 1, countMatchingErrors(errors, ErrStorageKeyCollision))
|
||||
winner := firstResult.stored
|
||||
if firstResult.err != nil {
|
||||
winner = secondResult.stored
|
||||
}
|
||||
require.FileExists(t, winner.AbsolutePath)
|
||||
requireOnlyStoredFile(t, storageRoot, winner.AbsolutePath)
|
||||
}
|
||||
|
||||
func requireNoStoredFiles(t *testing.T, storageRoot string) {
|
||||
@@ -106,3 +150,38 @@ func requireNoStoredFiles(t *testing.T, storageRoot string) {
|
||||
return nil
|
||||
}))
|
||||
}
|
||||
|
||||
func requireOnlyStoredFile(t *testing.T, storageRoot, expectedPath string) {
|
||||
t.Helper()
|
||||
files := []string{}
|
||||
require.NoError(t, filepath.WalkDir(storageRoot, func(path string, entry fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !entry.IsDir() {
|
||||
files = append(files, path)
|
||||
}
|
||||
return nil
|
||||
}))
|
||||
require.Equal(t, []string{expectedPath}, files)
|
||||
}
|
||||
|
||||
func countNilErrors(values []error) int {
|
||||
count := 0
|
||||
for _, err := range values {
|
||||
if err == nil {
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func countMatchingErrors(values []error, target error) int {
|
||||
count := 0
|
||||
for _, err := range values {
|
||||
if errors.Is(err, target) {
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user