195 lines
8.3 KiB
Go
195 lines
8.3 KiB
Go
package files
|
|
|
|
import (
|
|
"bytes"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"log"
|
|
"mime/multipart"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"os"
|
|
"path/filepath"
|
|
"testing"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/glebarez/sqlite"
|
|
"github.com/stretchr/testify/require"
|
|
"gorm.io/gorm"
|
|
"senlinai-agent/backend/internal/config"
|
|
"senlinai-agent/backend/internal/httpx"
|
|
"senlinai-agent/backend/internal/models"
|
|
)
|
|
|
|
func TestFileRegistrarSavesBeforeCreatingIdentitySourceDTO(t *testing.T) {
|
|
router, database, project, storageRoot := newFileHandlerTestRouter(t, 1, 1)
|
|
body := &bytes.Buffer{}
|
|
writer := multipart.NewWriter(body)
|
|
require.NoError(t, writer.WriteField("title", "客户资料"))
|
|
part, err := writer.CreateFormFile("file", "客户资料.txt")
|
|
require.NoError(t, err)
|
|
_, err = part.Write([]byte("private 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.StatusCreated, rec.Code, rec.Body.String())
|
|
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)))
|
|
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.NotContains(t, payload, "AbsolutePath")
|
|
require.NotContains(t, rec.Body.String(), storageRoot)
|
|
}
|
|
|
|
func TestFileRegistrarChecksOwnershipBeforeWritingFile(t *testing.T) {
|
|
router, database, project, storageRoot := newFileHandlerTestRouter(t, 1, 2)
|
|
body := &bytes.Buffer{}
|
|
writer := multipart.NewWriter(body)
|
|
part, err := writer.CreateFormFile("file", "hidden.txt")
|
|
require.NoError(t, err)
|
|
_, err = part.Write([]byte("must not be stored"))
|
|
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.StatusNotFound, rec.Code, rec.Body.String())
|
|
var count int64
|
|
require.NoError(t, database.Model(&models.SenlinAgentSource{}).Count(&count).Error)
|
|
require.Zero(t, count)
|
|
entries, err := os.ReadDir(storageRoot)
|
|
require.NoError(t, err)
|
|
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)
|
|
require.NoError(t, models.AutoMigrate(database))
|
|
models.DBService = database
|
|
require.NoError(t, database.Create(&models.SenlinAgentUser{Email: "owner@example.com", DisplayName: "Owner", PasswordHash: "hash"}).Error)
|
|
require.NoError(t, database.Create(&models.SenlinAgentUser{Email: "other@example.com", DisplayName: "Other", PasswordHash: "hash"}).Error)
|
|
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(service, maxUploadBytes),
|
|
)
|
|
return router, database, project, storageRoot, service
|
|
}
|