fix(backend): harden write registrar boundaries
This commit is contained in:
@@ -29,7 +29,7 @@ func main() {
|
|||||||
tagHandler := projects.NewTagHandler(projectService)
|
tagHandler := projects.NewTagHandler(projectService)
|
||||||
cronHandler := projects.NewCronHandler(projectService)
|
cronHandler := projects.NewCronHandler(projectService)
|
||||||
taskHandler := tasks.NewHandler(taskService)
|
taskHandler := tasks.NewHandler(taskService)
|
||||||
fileHandler := files.NewHandler(fileService)
|
fileHandler := files.NewHandler(fileService, cfg.MaxUploadBytes)
|
||||||
inboxHandler := inbox.NewHandler(inbox.NewService(inbox.StaticAnalyzer{}))
|
inboxHandler := inbox.NewHandler(inbox.NewService(inbox.StaticAnalyzer{}))
|
||||||
authHandler := auth.NewHandler(authService)
|
authHandler := auth.NewHandler(authService)
|
||||||
// search 与 AI 当前只有 service,尚无 HTTP registrar;后续功能任务应在实现真实契约后从此处注入,不能注册伪端点。
|
// search 与 AI 当前只有 service,尚无 HTTP registrar;后续功能任务应在实现真实契约后从此处注入,不能注册伪端点。
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ env: development
|
|||||||
port: "9150"
|
port: "9150"
|
||||||
dsn: "postgres://postgres:postgres@localhost:5432/agent_dev?sslmode=disable"
|
dsn: "postgres://postgres:postgres@localhost:5432/agent_dev?sslmode=disable"
|
||||||
storage_dir: "./data/files"
|
storage_dir: "./data/files"
|
||||||
|
max_upload_bytes: 33554432
|
||||||
auth_secret: "development-auth-secret-change-me"
|
auth_secret: "development-auth-secret-change-me"
|
||||||
system_ai_key: ""
|
system_ai_key: ""
|
||||||
ai_key_encryption_secret: "development-ai-key-secret-change-me"
|
ai_key_encryption_secret: "development-ai-key-secret-change-me"
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ type Config struct {
|
|||||||
Port string `yaml:"port"`
|
Port string `yaml:"port"`
|
||||||
DSN string `yaml:"dsn"`
|
DSN string `yaml:"dsn"`
|
||||||
StorageDir string `yaml:"storage_dir"`
|
StorageDir string `yaml:"storage_dir"`
|
||||||
|
MaxUploadBytes int64 `yaml:"max_upload_bytes"`
|
||||||
AuthSecret string `yaml:"auth_secret"`
|
AuthSecret string `yaml:"auth_secret"`
|
||||||
SystemAIKey string `yaml:"system_ai_key"`
|
SystemAIKey string `yaml:"system_ai_key"`
|
||||||
AIKeyEncryptionSecret string `yaml:"ai_key_encryption_secret"`
|
AIKeyEncryptionSecret string `yaml:"ai_key_encryption_secret"`
|
||||||
@@ -41,5 +42,8 @@ func LoadFromDir(configDir string) (Config, error) {
|
|||||||
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
||||||
return Config{}, err
|
return Config{}, err
|
||||||
}
|
}
|
||||||
|
if cfg.MaxUploadBytes <= 0 {
|
||||||
|
cfg.MaxUploadBytes = 32 << 20
|
||||||
|
}
|
||||||
return cfg, nil
|
return cfg, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ func TestLoadFromDirDefaultsToDevYAML(t *testing.T) {
|
|||||||
require.Equal(t, "18080", cfg.Port)
|
require.Equal(t, "18080", cfg.Port)
|
||||||
require.Equal(t, "postgres://dev", cfg.DSN)
|
require.Equal(t, "postgres://dev", cfg.DSN)
|
||||||
require.Equal(t, "./dev-files", cfg.StorageDir)
|
require.Equal(t, "./dev-files", cfg.StorageDir)
|
||||||
|
require.Equal(t, int64(32<<20), cfg.MaxUploadBytes)
|
||||||
require.Equal(t, "dev-auth", cfg.AuthSecret)
|
require.Equal(t, "dev-auth", cfg.AuthSecret)
|
||||||
require.Equal(t, "dev-system", cfg.SystemAIKey)
|
require.Equal(t, "dev-system", cfg.SystemAIKey)
|
||||||
require.Equal(t, "dev-ai", cfg.AIKeyEncryptionSecret)
|
require.Equal(t, "dev-ai", cfg.AIKeyEncryptionSecret)
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package files
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
|
"log"
|
||||||
"net/http"
|
"net/http"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -26,12 +27,20 @@ type SourceDTO struct {
|
|||||||
|
|
||||||
// Handler 负责 multipart 边界,文件名清理和本地路径构造始终委托给 Service。
|
// Handler 负责 multipart 边界,文件名清理和本地路径构造始终委托给 Service。
|
||||||
type Handler struct {
|
type Handler struct {
|
||||||
service *Service
|
service *Service
|
||||||
|
maxUploadBytes int64
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// DefaultMaxUploadBytes 是未显式配置时的请求体上限,覆盖完整 multipart 内容。
|
||||||
|
const DefaultMaxUploadBytes int64 = 32 << 20
|
||||||
|
|
||||||
// NewHandler 创建文件资料 HTTP registrar。
|
// NewHandler 创建文件资料 HTTP registrar。
|
||||||
func NewHandler(service *Service) *Handler {
|
func NewHandler(service *Service, limits ...int64) *Handler {
|
||||||
return &Handler{service: service}
|
limit := DefaultMaxUploadBytes
|
||||||
|
if len(limits) > 0 && limits[0] > 0 {
|
||||||
|
limit = limits[0]
|
||||||
|
}
|
||||||
|
return &Handler{service: service, maxUploadBytes: limit}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Register 将文件资料上传接口注册到上层提供的 /api/v1 路由组。
|
// Register 将文件资料上传接口注册到上层提供的 /api/v1 路由组。
|
||||||
@@ -55,8 +64,15 @@ func (h *Handler) upload(c *gin.Context) {
|
|||||||
writeSourceError(c, err)
|
writeSourceError(c, err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
// MaxBytesReader 必须在任何 multipart 解析前安装,限制字段、边界和文件内容的总请求量。
|
||||||
|
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, h.maxUploadBytes)
|
||||||
fileHeader, err := c.FormFile("file")
|
fileHeader, err := c.FormFile("file")
|
||||||
if err != nil {
|
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", "请选择要上传的文件")
|
httpx.Error(c, http.StatusBadRequest, "invalid_request", "请选择要上传的文件")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -75,6 +91,10 @@ func (h *Handler) upload(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
source, err := h.service.CreateSource(userID, project, c.PostForm("title"), stored)
|
source, err := h.service.CreateSource(userID, project, c.PostForm("title"), stored)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
if cleanupErr := h.service.Remove(stored); cleanupErr != nil {
|
||||||
|
// 清理错误仅写服务端日志,响应继续使用稳定中文错误,不暴露 storage root。
|
||||||
|
log.Printf("source file cleanup failed: %v", cleanupErr)
|
||||||
|
}
|
||||||
writeSourceError(c, err)
|
writeSourceError(c, err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,7 +3,9 @@ package files
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"log"
|
||||||
"mime/multipart"
|
"mime/multipart"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
@@ -76,7 +78,102 @@ func TestFileRegistrarChecksOwnershipBeforeWritingFile(t *testing.T) {
|
|||||||
require.Empty(t, entries)
|
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) {
|
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()
|
t.Helper()
|
||||||
database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{TranslateError: true})
|
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, 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)}
|
project := models.SenlinAgentProject{OwnerID: ownerID, Name: "Files", Identifier: fmt.Sprintf("FILES-%d", ownerID)}
|
||||||
require.NoError(t, database.Create(&project).Error)
|
require.NoError(t, database.Create(&project).Error)
|
||||||
storageRoot := t.TempDir()
|
storageRoot := t.TempDir()
|
||||||
|
service := NewService(storageRoot, database)
|
||||||
router := httpx.NewProtectedRouter(
|
router := httpx.NewProtectedRouter(
|
||||||
config.Config{Env: "test"},
|
config.Config{Env: "test"},
|
||||||
func(string) (uint, error) { return currentUserID, nil },
|
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 {
|
type Service struct {
|
||||||
root string
|
root string
|
||||||
db *gorm.DB
|
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 {
|
type StoredFile struct {
|
||||||
@@ -30,7 +39,12 @@ func NewService(root string, databases ...*gorm.DB) *Service {
|
|||||||
if len(databases) > 0 {
|
if len(databases) > 0 {
|
||||||
database = 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。
|
// Save 清理客户端文件名,并只返回供持久化的相对路径;绝对路径不得进入 API DTO。
|
||||||
@@ -40,21 +54,90 @@ func (s *Service) Save(projectID uint, originalName string, content io.Reader) (
|
|||||||
cleanName = "upload.bin"
|
cleanName = "upload.bin"
|
||||||
}
|
}
|
||||||
relative := filepath.ToSlash(filepath.Join("projects", fmt.Sprint(projectID), fmt.Sprintf("%d-%s", time.Now().UnixNano(), cleanName)))
|
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))
|
absolute, err := s.absolutePath(relative)
|
||||||
if err := os.MkdirAll(filepath.Dir(absolute), 0o755); err != nil {
|
if err != nil {
|
||||||
return StoredFile{}, err
|
return StoredFile{}, err
|
||||||
}
|
}
|
||||||
file, err := os.Create(absolute)
|
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 {
|
if err != nil {
|
||||||
return StoredFile{}, err
|
return StoredFile{}, err
|
||||||
}
|
}
|
||||||
defer file.Close()
|
|
||||||
if _, err := io.Copy(file, content); err != nil {
|
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{}, err
|
||||||
}
|
}
|
||||||
return StoredFile{OriginalName: cleanName, RelativePath: relative, AbsolutePath: absolute}, nil
|
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 {
|
func (s *Service) database() *gorm.DB {
|
||||||
if s.db != nil {
|
if s.db != nil {
|
||||||
return s.db
|
return s.db
|
||||||
|
|||||||
@@ -1,12 +1,37 @@
|
|||||||
package files
|
package files
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
|
"io/fs"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
"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) {
|
func TestSaveStoresFileUnderProjectDirectory(t *testing.T) {
|
||||||
service := NewService(t.TempDir())
|
service := NewService(t.TempDir())
|
||||||
|
|
||||||
@@ -27,3 +52,57 @@ func TestSaveNeutralizesPathTraversal(t *testing.T) {
|
|||||||
require.Equal(t, "secret.txt", stored.OriginalName)
|
require.Equal(t, "secret.txt", stored.OriginalName)
|
||||||
require.Contains(t, stored.RelativePath, "projects/12/")
|
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
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,11 +1,13 @@
|
|||||||
package projects
|
package projects
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
"gorm.io/gorm"
|
||||||
"senlinai-agent/backend/internal/httpx"
|
"senlinai-agent/backend/internal/httpx"
|
||||||
"senlinai-agent/backend/internal/models"
|
"senlinai-agent/backend/internal/models"
|
||||||
)
|
)
|
||||||
@@ -62,12 +64,23 @@ func (h *CronHandler) create(c *gin.Context) {
|
|||||||
Title: input.Title, Schedule: input.Schedule, Enabled: input.Enabled, NextRunAt: nextRunAt,
|
Title: input.Title, Schedule: input.Schedule, Enabled: input.Enabled, NextRunAt: nextRunAt,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
httpx.Error(c, http.StatusBadRequest, "invalid_request", "计划任务参数无效")
|
writeCronError(c, err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
c.JSON(http.StatusCreated, cronPlanDTO(*plan))
|
c.JSON(http.StatusCreated, cronPlanDTO(*plan))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func writeCronError(c *gin.Context, err error) {
|
||||||
|
switch {
|
||||||
|
case errors.Is(err, ErrCronTitleRequired), errors.Is(err, ErrCronScheduleRequired):
|
||||||
|
httpx.Error(c, http.StatusBadRequest, "invalid_request", "计划任务参数无效")
|
||||||
|
case errors.Is(err, gorm.ErrRecordNotFound):
|
||||||
|
httpx.Error(c, http.StatusNotFound, "not_found", "项目不存在")
|
||||||
|
default:
|
||||||
|
httpx.Error(c, http.StatusInternalServerError, "internal_error", "计划任务创建失败")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func cronPlanDTO(plan models.SenlinAgentCronPlan) CronPlanDTO {
|
func cronPlanDTO(plan models.SenlinAgentCronPlan) CronPlanDTO {
|
||||||
return CronPlanDTO{
|
return CronPlanDTO{
|
||||||
ID: plan.Identity, ProjectID: plan.ProjectIdentity, Title: plan.Title, Schedule: plan.Schedule,
|
ID: plan.Identity, ProjectID: plan.ProjectIdentity, Title: plan.Title, Schedule: plan.Schedule,
|
||||||
|
|||||||
@@ -8,17 +8,23 @@ import (
|
|||||||
"senlinai-agent/backend/internal/models"
|
"senlinai-agent/backend/internal/models"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
ErrTagNameRequired = errors.New("tag name is required")
|
||||||
|
ErrCronTitleRequired = errors.New("cron plan title is required")
|
||||||
|
ErrCronScheduleRequired = errors.New("cron schedule is required")
|
||||||
|
)
|
||||||
|
|
||||||
func (s *Service) CreateCronPlan(ownerID uint, projectID uint, input CreateCronPlanInput) (*models.SenlinAgentCronPlan, error) {
|
func (s *Service) CreateCronPlan(ownerID uint, projectID uint, input CreateCronPlanInput) (*models.SenlinAgentCronPlan, error) {
|
||||||
if err := ensureProjectOwner(ownerID, projectID); err != nil {
|
if err := ensureProjectOwner(ownerID, projectID); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
title := strings.TrimSpace(input.Title)
|
title := strings.TrimSpace(input.Title)
|
||||||
if title == "" {
|
if title == "" {
|
||||||
return nil, errors.New("cron plan title is required")
|
return nil, ErrCronTitleRequired
|
||||||
}
|
}
|
||||||
schedule := strings.TrimSpace(input.Schedule)
|
schedule := strings.TrimSpace(input.Schedule)
|
||||||
if schedule == "" {
|
if schedule == "" {
|
||||||
return nil, errors.New("cron schedule is required")
|
return nil, ErrCronScheduleRequired
|
||||||
}
|
}
|
||||||
plan := &models.SenlinAgentCronPlan{
|
plan := &models.SenlinAgentCronPlan{
|
||||||
ProjectID: projectID, CreatedBy: ownerID, Title: title, Schedule: schedule,
|
ProjectID: projectID, CreatedBy: ownerID, Title: title, Schedule: schedule,
|
||||||
@@ -37,7 +43,7 @@ func (s *Service) CreateProjectTag(ownerID uint, projectID uint, name string) (*
|
|||||||
func (s *Service) CreateTag(projectID uint, name string) (*models.SenlinAgentTag, error) {
|
func (s *Service) CreateTag(projectID uint, name string) (*models.SenlinAgentTag, error) {
|
||||||
name = strings.TrimSpace(name)
|
name = strings.TrimSpace(name)
|
||||||
if name == "" {
|
if name == "" {
|
||||||
return nil, errors.New("tag name is required")
|
return nil, ErrTagNameRequired
|
||||||
}
|
}
|
||||||
var existing models.SenlinAgentTag
|
var existing models.SenlinAgentTag
|
||||||
err := models.DBService.Where("project_id = ? AND name = ?", projectID, name).First(&existing).Error
|
err := models.DBService.Where("project_id = ? AND name = ?", projectID, name).First(&existing).Error
|
||||||
|
|||||||
@@ -1,9 +1,11 @@
|
|||||||
package projects
|
package projects
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
"gorm.io/gorm"
|
||||||
"senlinai-agent/backend/internal/httpx"
|
"senlinai-agent/backend/internal/httpx"
|
||||||
"senlinai-agent/backend/internal/logic/auth"
|
"senlinai-agent/backend/internal/logic/auth"
|
||||||
"senlinai-agent/backend/internal/models"
|
"senlinai-agent/backend/internal/models"
|
||||||
@@ -56,12 +58,23 @@ func (h *TagHandler) create(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
tag, err := h.service.CreateProjectTag(userID, project.ID, input.Name)
|
tag, err := h.service.CreateProjectTag(userID, project.ID, input.Name)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
httpx.Error(c, http.StatusBadRequest, "invalid_request", "标签名称不能为空")
|
writeTagError(c, err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
c.JSON(http.StatusCreated, WorkspaceTagDTO{ID: tag.Identity, Name: tag.Name})
|
c.JSON(http.StatusCreated, WorkspaceTagDTO{ID: tag.Identity, Name: tag.Name})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func writeTagError(c *gin.Context, err error) {
|
||||||
|
switch {
|
||||||
|
case errors.Is(err, ErrTagNameRequired):
|
||||||
|
httpx.Error(c, http.StatusBadRequest, "invalid_request", "标签名称不能为空")
|
||||||
|
case errors.Is(err, gorm.ErrRecordNotFound):
|
||||||
|
httpx.Error(c, http.StatusNotFound, "not_found", "项目不存在")
|
||||||
|
default:
|
||||||
|
httpx.Error(c, http.StatusInternalServerError, "internal_error", "标签创建失败")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func ownedProjectFromRequest(c *gin.Context) (uint, *models.SenlinAgentProject, bool) {
|
func ownedProjectFromRequest(c *gin.Context) (uint, *models.SenlinAgentProject, bool) {
|
||||||
userID, ok := auth.CurrentUserID(c)
|
userID, ok := auth.CurrentUserID(c)
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|||||||
@@ -3,12 +3,14 @@ package projects
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/gorm"
|
||||||
"senlinai-agent/backend/internal/config"
|
"senlinai-agent/backend/internal/config"
|
||||||
"senlinai-agent/backend/internal/httpx"
|
"senlinai-agent/backend/internal/httpx"
|
||||||
"senlinai-agent/backend/internal/models"
|
"senlinai-agent/backend/internal/models"
|
||||||
@@ -62,6 +64,75 @@ func TestTagRegistrarRejectsProjectOwnedByAnotherUser(t *testing.T) {
|
|||||||
router.ServeHTTP(rec, req)
|
router.ServeHTTP(rec, req)
|
||||||
|
|
||||||
require.Equal(t, http.StatusNotFound, rec.Code, rec.Body.String())
|
require.Equal(t, http.StatusNotFound, rec.Code, rec.Body.String())
|
||||||
|
var payload httpx.ErrorEnvelope
|
||||||
|
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
|
||||||
|
require.Equal(t, "not_found", payload.Error.Code)
|
||||||
|
require.Equal(t, "项目不存在", payload.Error.Message)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTagAndCronServicesReturnTypedValidationErrors(t *testing.T) {
|
||||||
|
database := newTestDB(t)
|
||||||
|
require.NoError(t, database.Create(&models.SenlinAgentUser{Email: "owner@example.com", DisplayName: "Owner", PasswordHash: "hash"}).Error)
|
||||||
|
project := models.SenlinAgentProject{OwnerID: 1, Name: "Alpha", Identifier: "ALPHA"}
|
||||||
|
require.NoError(t, database.Create(&project).Error)
|
||||||
|
service := NewService()
|
||||||
|
|
||||||
|
_, err := service.CreateProjectTag(1, project.ID, " ")
|
||||||
|
require.ErrorIs(t, err, ErrTagNameRequired)
|
||||||
|
_, err = service.CreateCronPlan(1, project.ID, CreateCronPlanInput{Title: " ", Schedule: "0 9 * * *"})
|
||||||
|
require.ErrorIs(t, err, ErrCronTitleRequired)
|
||||||
|
_, err = service.CreateCronPlan(1, project.ID, CreateCronPlanInput{Title: "Daily", Schedule: " "})
|
||||||
|
require.ErrorIs(t, err, ErrCronScheduleRequired)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTagRegistrarMapsDatabaseFailureToInternalError(t *testing.T) {
|
||||||
|
database := newTestDB(t)
|
||||||
|
require.NoError(t, database.Create(&models.SenlinAgentUser{Email: "owner@example.com", DisplayName: "Owner", PasswordHash: "hash"}).Error)
|
||||||
|
project := models.SenlinAgentProject{OwnerID: 1, Name: "Alpha", Identifier: "ALPHA"}
|
||||||
|
require.NoError(t, database.Create(&project).Error)
|
||||||
|
require.NoError(t, database.Callback().Create().Before("gorm:create").Register("test:fail_tag_create", func(tx *gorm.DB) {
|
||||||
|
if tx.Statement.Schema != nil && tx.Statement.Schema.Table == (models.SenlinAgentTag{}).TableName() {
|
||||||
|
tx.AddError(errors.New("forced tag create failure"))
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
router := newProjectWriteHandlerRouter(t, 1, NewTagHandler(NewService()))
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/tags", bytes.NewBufferString(`{"name":"设计"}`))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Authorization", "Bearer test-token")
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusInternalServerError, rec.Code, rec.Body.String())
|
||||||
|
var payload httpx.ErrorEnvelope
|
||||||
|
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
|
||||||
|
require.Equal(t, "internal_error", payload.Error.Code)
|
||||||
|
require.Equal(t, "标签创建失败", payload.Error.Message)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCronRegistrarMapsDatabaseFailureToInternalError(t *testing.T) {
|
||||||
|
database := newTestDB(t)
|
||||||
|
require.NoError(t, database.Create(&models.SenlinAgentUser{Email: "owner@example.com", DisplayName: "Owner", PasswordHash: "hash"}).Error)
|
||||||
|
project := models.SenlinAgentProject{OwnerID: 1, Name: "Alpha", Identifier: "ALPHA"}
|
||||||
|
require.NoError(t, database.Create(&project).Error)
|
||||||
|
require.NoError(t, database.Callback().Create().Before("gorm:create").Register("test:fail_cron_create", func(tx *gorm.DB) {
|
||||||
|
if tx.Statement.Schema != nil && tx.Statement.Schema.Table == (models.SenlinAgentCronPlan{}).TableName() {
|
||||||
|
tx.AddError(errors.New("forced cron create failure"))
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
router := newProjectWriteHandlerRouter(t, 1, NewCronHandler(NewService()))
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/cron-plans", bytes.NewBufferString(`{"title":"每日整理","schedule":"0 9 * * *"}`))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Authorization", "Bearer test-token")
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusInternalServerError, rec.Code, rec.Body.String())
|
||||||
|
var payload httpx.ErrorEnvelope
|
||||||
|
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
|
||||||
|
require.Equal(t, "internal_error", payload.Error.Code)
|
||||||
|
require.Equal(t, "计划任务创建失败", payload.Error.Message)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCronRegistrarCreatesCamelCaseIdentityDTO(t *testing.T) {
|
func TestCronRegistrarCreatesCamelCaseIdentityDTO(t *testing.T) {
|
||||||
|
|||||||
94
backend/internal/logic/tasks/concurrency_postgres_test.go
Normal file
94
backend/internal/logic/tasks/concurrency_postgres_test.go
Normal file
@@ -0,0 +1,94 @@
|
|||||||
|
package tasks
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/driver/postgres"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
"senlinai-agent/backend/internal/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPostgresMovePreventsOldProjectShareFromBeingInsertedConcurrently(t *testing.T) {
|
||||||
|
dsn := os.Getenv("DATABASE_URL")
|
||||||
|
if dsn == "" {
|
||||||
|
t.Skip("DATABASE_URL is not configured; skipping PostgreSQL row-lock concurrency test")
|
||||||
|
}
|
||||||
|
database, err := gorm.Open(postgres.Open(dsn), &gorm.Config{TranslateError: true})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, models.AutoMigrate(database))
|
||||||
|
suffix := fmt.Sprint(time.Now().UnixNano())
|
||||||
|
owner := models.SenlinAgentUser{Email: "lock-" + suffix + "@example.com", DisplayName: "Lock Owner", PasswordHash: "hash"}
|
||||||
|
require.NoError(t, database.Create(&owner).Error)
|
||||||
|
first := models.SenlinAgentProject{OwnerID: owner.ID, Name: "First", Identifier: "FIRST-" + suffix}
|
||||||
|
second := models.SenlinAgentProject{OwnerID: owner.ID, Name: "Second", Identifier: "SECOND-" + suffix}
|
||||||
|
require.NoError(t, database.Create(&first).Error)
|
||||||
|
require.NoError(t, database.Create(&second).Error)
|
||||||
|
task := models.SenlinAgentTask{ProjectID: first.ID, CreatedBy: owner.ID, Title: "Move", Status: "open"}
|
||||||
|
note := models.SenlinAgentNote{ProjectID: first.ID, CreatedBy: owner.ID, Title: "Old context", Markdown: "private"}
|
||||||
|
require.NoError(t, database.Create(&task).Error)
|
||||||
|
require.NoError(t, database.Create(¬e).Error)
|
||||||
|
t.Cleanup(func() {
|
||||||
|
database.Where("task_id = ?", task.ID).Delete(&models.SenlinAgentTaskShare{})
|
||||||
|
database.Where("entity_type = ? AND entity_id = ?", "task", task.ID).Delete(&models.SenlinAgentProjectEvent{})
|
||||||
|
database.Delete(¬e)
|
||||||
|
database.Delete(&task)
|
||||||
|
database.Delete(&first)
|
||||||
|
database.Delete(&second)
|
||||||
|
database.Delete(&owner)
|
||||||
|
})
|
||||||
|
|
||||||
|
locked := make(chan struct{})
|
||||||
|
release := make(chan struct{})
|
||||||
|
var once sync.Once
|
||||||
|
var releaseOnce sync.Once
|
||||||
|
releaseLock := func() { releaseOnce.Do(func() { close(release) }) }
|
||||||
|
t.Cleanup(releaseLock)
|
||||||
|
callbackName := "test:pause_first_task_lock_" + suffix
|
||||||
|
require.NoError(t, database.Callback().Query().After("gorm:query").Register(callbackName, func(tx *gorm.DB) {
|
||||||
|
if tx.Statement.Schema == nil || tx.Statement.Schema.Table != (models.SenlinAgentTask{}).TableName() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if _, ok := tx.Statement.Clauses["FOR"]; !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
once.Do(func() {
|
||||||
|
close(locked)
|
||||||
|
<-release
|
||||||
|
})
|
||||||
|
}))
|
||||||
|
t.Cleanup(func() { database.Callback().Query().Remove(callbackName) })
|
||||||
|
|
||||||
|
service := NewService(database)
|
||||||
|
moveResult := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
_, err := service.Update(owner.ID, first.Identity, task.Identity, UpdateTaskInput{
|
||||||
|
Title: task.Title, NextProjectIdentity: second.Identity,
|
||||||
|
})
|
||||||
|
moveResult <- err
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-locked:
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("timed out waiting for Update to acquire the task row lock")
|
||||||
|
}
|
||||||
|
|
||||||
|
shareResult := make(chan error, 1)
|
||||||
|
go func() { shareResult <- service.ShareObject(task.ID, "note", note.ID) }()
|
||||||
|
select {
|
||||||
|
case err := <-shareResult:
|
||||||
|
t.Fatalf("ShareObject returned before the moving transaction released its task lock: %v", err)
|
||||||
|
case <-time.After(150 * time.Millisecond):
|
||||||
|
}
|
||||||
|
releaseLock()
|
||||||
|
require.NoError(t, <-moveResult)
|
||||||
|
require.ErrorContains(t, <-shareResult, "shared object not found in task project")
|
||||||
|
|
||||||
|
var shareCount int64
|
||||||
|
require.NoError(t, database.Model(&models.SenlinAgentTaskShare{}).Where("task_id = ?", task.ID).Count(&shareCount).Error)
|
||||||
|
require.Zero(t, shareCount)
|
||||||
|
}
|
||||||
@@ -93,15 +93,18 @@ func (h *Handler) update(c *gin.Context) {
|
|||||||
httpx.Error(c, http.StatusBadRequest, "invalid_request", "请求参数无效")
|
httpx.Error(c, http.StatusBadRequest, "invalid_request", "请求参数无效")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
nextProjectIdentity := ""
|
||||||
if strings.TrimSpace(input.NextProjectID) != "" {
|
if strings.TrimSpace(input.NextProjectID) != "" {
|
||||||
if _, ok := parseIdentity(input.NextProjectID); !ok {
|
var valid bool
|
||||||
|
nextProjectIdentity, valid = parseIdentity(input.NextProjectID)
|
||||||
|
if !valid {
|
||||||
httpx.Error(c, http.StatusBadRequest, "invalid_identity", "nextProjectId 必须是 UUIDv7")
|
httpx.Error(c, http.StatusBadRequest, "invalid_identity", "nextProjectId 必须是 UUIDv7")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
task, err := h.service.Update(userID, projectIdentity, taskIdentity, UpdateTaskInput{
|
task, err := h.service.Update(userID, projectIdentity, taskIdentity, UpdateTaskInput{
|
||||||
Title: input.Title, Description: input.Description, Status: input.Status, Completed: input.Completed,
|
Title: input.Title, Description: input.Description, Status: input.Status, Completed: input.Completed,
|
||||||
NextProjectIdentity: input.NextProjectID, Tag: input.Tag,
|
NextProjectIdentity: nextProjectIdentity, Tag: input.Tag,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
writeTaskError(c, err)
|
writeTaskError(c, err)
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
@@ -52,7 +53,7 @@ func TestTaskRegistrarMovesTaskByIdentityAndClearsForeignProjectTag(t *testing.T
|
|||||||
note := models.SenlinAgentNote{ProjectID: first.ID, CreatedBy: 1, Title: "旧项目资料", Markdown: "仅可在 Alpha 分享"}
|
note := models.SenlinAgentNote{ProjectID: first.ID, CreatedBy: 1, Title: "旧项目资料", Markdown: "仅可在 Alpha 分享"}
|
||||||
require.NoError(t, database.Create(¬e).Error)
|
require.NoError(t, database.Create(¬e).Error)
|
||||||
require.NoError(t, NewService(database).ShareObject(task.ID, "note", note.ID))
|
require.NoError(t, NewService(database).ShareObject(task.ID, "note", note.ID))
|
||||||
body := bytes.NewBufferString(fmt.Sprintf(`{"title":"迁移任务","description":"已移动","completed":false,"nextProjectId":%q}`, second.Identity))
|
body := bytes.NewBufferString(fmt.Sprintf(`{"title":"迁移任务","description":"已移动","completed":false,"nextProjectId":%q}`, strings.ToUpper(second.Identity)))
|
||||||
req := httptest.NewRequest(http.MethodPatch, "/api/v1/projects/"+first.Identity+"/tasks/"+task.Identity, body)
|
req := httptest.NewRequest(http.MethodPatch, "/api/v1/projects/"+first.Identity+"/tasks/"+task.Identity, body)
|
||||||
req.Header.Set("Content-Type", "application/json")
|
req.Header.Set("Content-Type", "application/json")
|
||||||
req.Header.Set("Authorization", "Bearer test-token")
|
req.Header.Set("Authorization", "Bearer test-token")
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
"gorm.io/gorm/clause"
|
||||||
"senlinai-agent/backend/internal/models"
|
"senlinai-agent/backend/internal/models"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -15,8 +16,8 @@ type Service struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type LinkedObject struct {
|
type LinkedObject struct {
|
||||||
ObjectType string `json:"object_type"`
|
ObjectType string `json:"objectType"`
|
||||||
ObjectID uint `json:"object_id"`
|
ObjectID string `json:"objectId"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewService 接受数据库或上层事务,确保任务写入、标签调整和分享检查使用同一依赖。
|
// NewService 接受数据库或上层事务,确保任务写入、标签调整和分享检查使用同一依赖。
|
||||||
@@ -35,13 +36,20 @@ func (s *Service) database() *gorm.DB {
|
|||||||
return models.DBService
|
return models.DBService
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Service) Assign(taskID uint, assigneeID uint) error {
|
// Assign 在同一事务内由用户 identity 解析内部主键,并同步两个指派字段,避免 DTO 返回陈旧 identity。
|
||||||
|
func (s *Service) Assign(taskID uint, assigneeIdentity string) error {
|
||||||
return s.database().Transaction(func(tx *gorm.DB) error {
|
return s.database().Transaction(func(tx *gorm.DB) error {
|
||||||
var task models.SenlinAgentTask
|
var task models.SenlinAgentTask
|
||||||
if err := tx.First(&task, taskID).Error; err != nil {
|
if err := tx.First(&task, taskID).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if err := tx.Model(&task).Update("assignee_id", assigneeID).Error; err != nil {
|
var assignee models.SenlinAgentUser
|
||||||
|
if err := tx.Where("identity = ?", assigneeIdentity).First(&assignee).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := tx.Model(&task).Updates(map[string]any{
|
||||||
|
"assignee_id": assignee.ID, "assignee_identity": assignee.Identity,
|
||||||
|
}).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return tx.Create(&models.SenlinAgentProjectEvent{
|
return tx.Create(&models.SenlinAgentProjectEvent{
|
||||||
@@ -50,7 +58,7 @@ func (s *Service) Assign(taskID uint, assigneeID uint) error {
|
|||||||
EventType: "task_assigned",
|
EventType: "task_assigned",
|
||||||
EntityType: "task",
|
EntityType: "task",
|
||||||
EntityID: task.ID,
|
EntityID: task.ID,
|
||||||
Summary: fmt.Sprintf("Task assigned to user %d", assigneeID),
|
Summary: fmt.Sprintf("Task assigned to user %s", assignee.Identity),
|
||||||
}).Error
|
}).Error
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -61,8 +69,8 @@ func (s *Service) ShareObject(taskID uint, objectType string, objectID uint) err
|
|||||||
return errors.New("unsupported shared object type")
|
return errors.New("unsupported shared object type")
|
||||||
}
|
}
|
||||||
return s.database().Transaction(func(tx *gorm.DB) error {
|
return s.database().Transaction(func(tx *gorm.DB) error {
|
||||||
var task models.SenlinAgentTask
|
task, err := lockTaskByID(tx, taskID)
|
||||||
if err := tx.First(&task, taskID).Error; err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if err := ensureSharedObjectInProject(tx, task.ProjectID, objectType, objectID); err != nil {
|
if err := ensureSharedObjectInProject(tx, task.ProjectID, objectType, objectID); err != nil {
|
||||||
@@ -97,7 +105,7 @@ func (s *Service) VisibleLinkedObjects(taskID uint, viewerID uint) ([]LinkedObje
|
|||||||
}
|
}
|
||||||
objects := make([]LinkedObject, 0, len(shares))
|
objects := make([]LinkedObject, 0, len(shares))
|
||||||
for _, share := range shares {
|
for _, share := range shares {
|
||||||
objects = append(objects, LinkedObject{ObjectType: share.ObjectType, ObjectID: share.ObjectID})
|
objects = append(objects, LinkedObject{ObjectType: share.ObjectType, ObjectID: share.ObjectIdentity})
|
||||||
}
|
}
|
||||||
return objects, nil
|
return objects, nil
|
||||||
}
|
}
|
||||||
@@ -141,13 +149,16 @@ func (s *Service) Create(ownerID uint, projectIdentity string, input CreateTaskI
|
|||||||
func (s *Service) Update(ownerID uint, projectIdentity, taskIdentity string, input UpdateTaskInput) (TaskDTO, error) {
|
func (s *Service) Update(ownerID uint, projectIdentity, taskIdentity string, input UpdateTaskInput) (TaskDTO, error) {
|
||||||
var result TaskDTO
|
var result TaskDTO
|
||||||
err := s.database().Transaction(func(tx *gorm.DB) error {
|
err := s.database().Transaction(func(tx *gorm.DB) error {
|
||||||
|
task, err := lockTaskByIdentity(tx, taskIdentity)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
currentProject, err := findOwnedProject(tx, ownerID, projectIdentity)
|
currentProject, err := findOwnedProject(tx, ownerID, projectIdentity)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
var task models.SenlinAgentTask
|
if task.ProjectID != currentProject.ID {
|
||||||
if err := tx.Where("identity = ? AND project_id = ?", taskIdentity, currentProject.ID).First(&task).Error; err != nil {
|
return gorm.ErrRecordNotFound
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
targetProject := currentProject
|
targetProject := currentProject
|
||||||
if strings.TrimSpace(input.NextProjectIdentity) != "" && input.NextProjectIdentity != currentProject.Identity {
|
if strings.TrimSpace(input.NextProjectIdentity) != "" && input.NextProjectIdentity != currentProject.Identity {
|
||||||
@@ -187,15 +198,31 @@ func (s *Service) Update(ownerID uint, projectIdentity, taskIdentity string, inp
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := tx.Save(&task).Error; err != nil {
|
if err := tx.Save(task).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
result = makeTaskDTO(task, targetProject.Identity, tagIdentity, tagName)
|
result = makeTaskDTO(*task, targetProject.Identity, tagIdentity, tagName)
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func lockTaskByID(tx *gorm.DB, taskID uint) (*models.SenlinAgentTask, error) {
|
||||||
|
var task models.SenlinAgentTask
|
||||||
|
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&task, taskID).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &task, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func lockTaskByIdentity(tx *gorm.DB, taskIdentity string) (*models.SenlinAgentTask, error) {
|
||||||
|
var task models.SenlinAgentTask
|
||||||
|
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("identity = ?", taskIdentity).First(&task).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &task, nil
|
||||||
|
}
|
||||||
|
|
||||||
func findOwnedProject(tx *gorm.DB, ownerID uint, identity string) (*models.SenlinAgentProject, error) {
|
func findOwnedProject(tx *gorm.DB, ownerID uint, identity string) (*models.SenlinAgentProject, error) {
|
||||||
var project models.SenlinAgentProject
|
var project models.SenlinAgentProject
|
||||||
if err := tx.Where("owner_id = ? AND identity = ?", ownerID, identity).First(&project).Error; err != nil {
|
if err := tx.Where("owner_id = ? AND identity = ?", ownerID, identity).First(&project).Error; err != nil {
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package tasks
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
@@ -10,6 +11,30 @@ import (
|
|||||||
"senlinai-agent/backend/internal/models"
|
"senlinai-agent/backend/internal/models"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func TestUpdateLocksTaskBeforeProjectValidation(t *testing.T) {
|
||||||
|
database, owner, project, task := newTaskLockFixture(t)
|
||||||
|
queries := captureTaskQueryOrder(t, database)
|
||||||
|
|
||||||
|
_, err := NewService(database).Update(owner.ID, project.Identity, task.Identity, UpdateTaskInput{Title: task.Title})
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotEmpty(t, *queries)
|
||||||
|
require.Equal(t, "task:locked-in-tx", (*queries)[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestShareObjectLocksTaskBeforeObjectValidation(t *testing.T) {
|
||||||
|
database, _, project, task := newTaskLockFixture(t)
|
||||||
|
note := models.SenlinAgentNote{ProjectID: project.ID, CreatedBy: task.CreatedBy, Title: "Context", Markdown: "Private"}
|
||||||
|
require.NoError(t, database.Create(¬e).Error)
|
||||||
|
queries := captureTaskQueryOrder(t, database)
|
||||||
|
|
||||||
|
err := NewService(database).ShareObject(task.ID, "note", note.ID)
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotEmpty(t, *queries)
|
||||||
|
require.Equal(t, "task:locked-in-tx", (*queries)[0])
|
||||||
|
}
|
||||||
|
|
||||||
func TestAssigneeOnlySeesExplicitlySharedObjects(t *testing.T) {
|
func TestAssigneeOnlySeesExplicitlySharedObjects(t *testing.T) {
|
||||||
database := newTestDB(t)
|
database := newTestDB(t)
|
||||||
assigneeID := uint(2)
|
assigneeID := uint(2)
|
||||||
@@ -28,7 +53,7 @@ func TestAssigneeOnlySeesExplicitlySharedObjects(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Len(t, after, 1)
|
require.Len(t, after, 1)
|
||||||
require.Equal(t, "note", after[0].ObjectType)
|
require.Equal(t, "note", after[0].ObjectType)
|
||||||
require.Equal(t, note.ID, after[0].ObjectID)
|
require.Equal(t, note.Identity, after[0].ObjectID)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestShareObjectRejectsUnsupportedType(t *testing.T) {
|
func TestShareObjectRejectsUnsupportedType(t *testing.T) {
|
||||||
@@ -53,16 +78,32 @@ func TestShareObjectRejectsObjectFromAnotherProject(t *testing.T) {
|
|||||||
require.ErrorContains(t, err, "shared object not found in task project")
|
require.ErrorContains(t, err, "shared object not found in task project")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAssignRecordsProjectEvent(t *testing.T) {
|
func TestAssignResolvesIdentityAndTaskDTOReflectsReassignment(t *testing.T) {
|
||||||
database := newTestDB(t)
|
database := newTestDB(t)
|
||||||
task := models.SenlinAgentTask{ProjectID: 7, CreatedBy: 1, Title: "安排评审"}
|
owner := models.SenlinAgentUser{Email: "owner@example.com", DisplayName: "Owner", PasswordHash: "hash"}
|
||||||
|
first := models.SenlinAgentUser{Email: "first@example.com", DisplayName: "First", PasswordHash: "hash"}
|
||||||
|
second := models.SenlinAgentUser{Email: "second@example.com", DisplayName: "Second", PasswordHash: "hash"}
|
||||||
|
require.NoError(t, database.Create(&owner).Error)
|
||||||
|
require.NoError(t, database.Create(&first).Error)
|
||||||
|
require.NoError(t, database.Create(&second).Error)
|
||||||
|
project := models.SenlinAgentProject{OwnerID: owner.ID, Name: "Alpha", Identifier: "ALPHA"}
|
||||||
|
require.NoError(t, database.Create(&project).Error)
|
||||||
|
task := models.SenlinAgentTask{ProjectID: project.ID, CreatedBy: owner.ID, Title: "安排评审", Status: "open"}
|
||||||
require.NoError(t, database.Create(&task).Error)
|
require.NoError(t, database.Create(&task).Error)
|
||||||
service := NewService()
|
service := NewService(database)
|
||||||
|
|
||||||
require.NoError(t, service.Assign(task.ID, 2))
|
require.NoError(t, service.Assign(task.ID, first.Identity))
|
||||||
|
assigned, err := service.Update(owner.ID, project.Identity, task.Identity, UpdateTaskInput{Title: task.Title})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, first.Identity, *assigned.AssigneeID)
|
||||||
|
|
||||||
|
require.NoError(t, service.Assign(task.ID, second.Identity))
|
||||||
|
reassigned, err := service.Update(owner.ID, project.Identity, task.Identity, UpdateTaskInput{Title: task.Title})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, second.Identity, *reassigned.AssigneeID)
|
||||||
|
|
||||||
var event models.SenlinAgentProjectEvent
|
var event models.SenlinAgentProjectEvent
|
||||||
require.NoError(t, database.Where("project_id = ? AND entity_type = ? AND entity_id = ?", 7, "task", task.ID).First(&event).Error)
|
require.NoError(t, database.Where("project_id = ? AND entity_type = ? AND entity_id = ?", project.ID, "task", task.ID).Order("id desc").First(&event).Error)
|
||||||
require.Equal(t, "task_assigned", event.EventType)
|
require.Equal(t, "task_assigned", event.EventType)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -74,3 +115,45 @@ func newTestDB(t *testing.T) *gorm.DB {
|
|||||||
models.DBService = database
|
models.DBService = database
|
||||||
return database
|
return database
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func newTaskLockFixture(t *testing.T) (*gorm.DB, models.SenlinAgentUser, models.SenlinAgentProject, models.SenlinAgentTask) {
|
||||||
|
t.Helper()
|
||||||
|
database := newTestDB(t)
|
||||||
|
owner := models.SenlinAgentUser{Email: "lock-owner@example.com", DisplayName: "Owner", PasswordHash: "hash"}
|
||||||
|
require.NoError(t, database.Create(&owner).Error)
|
||||||
|
project := models.SenlinAgentProject{OwnerID: owner.ID, Name: "Lock", Identifier: "LOCK"}
|
||||||
|
require.NoError(t, database.Create(&project).Error)
|
||||||
|
task := models.SenlinAgentTask{ProjectID: project.ID, CreatedBy: owner.ID, Title: "Lock me", Status: "open"}
|
||||||
|
require.NoError(t, database.Create(&task).Error)
|
||||||
|
return database, owner, project, task
|
||||||
|
}
|
||||||
|
|
||||||
|
func captureTaskQueryOrder(t *testing.T, database *gorm.DB) *[]string {
|
||||||
|
t.Helper()
|
||||||
|
var mutex sync.Mutex
|
||||||
|
queries := make([]string, 0, 4)
|
||||||
|
callbackName := "test:capture_task_lock_" + t.Name()
|
||||||
|
require.NoError(t, database.Callback().Query().Before("gorm:query").Register(callbackName, func(tx *gorm.DB) {
|
||||||
|
table := tx.Statement.Table
|
||||||
|
if tx.Statement.Schema != nil {
|
||||||
|
table = tx.Statement.Schema.Table
|
||||||
|
}
|
||||||
|
if table != (models.SenlinAgentTask{}).TableName() && table != (models.SenlinAgentProject{}).TableName() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
entry := "project"
|
||||||
|
if table == (models.SenlinAgentTask{}).TableName() {
|
||||||
|
entry = "task"
|
||||||
|
if _, ok := tx.Statement.Clauses["FOR"]; ok {
|
||||||
|
entry += ":locked"
|
||||||
|
if _, ok := tx.Statement.ConnPool.(gorm.TxCommitter); ok {
|
||||||
|
entry += "-in-tx"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
mutex.Lock()
|
||||||
|
queries = append(queries, entry)
|
||||||
|
mutex.Unlock()
|
||||||
|
}))
|
||||||
|
return &queries
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user