fix(backend): harden write registrar boundaries
This commit is contained in:
@@ -13,6 +13,7 @@ type Config struct {
|
||||
Port string `yaml:"port"`
|
||||
DSN string `yaml:"dsn"`
|
||||
StorageDir string `yaml:"storage_dir"`
|
||||
MaxUploadBytes int64 `yaml:"max_upload_bytes"`
|
||||
AuthSecret string `yaml:"auth_secret"`
|
||||
SystemAIKey string `yaml:"system_ai_key"`
|
||||
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 {
|
||||
return Config{}, err
|
||||
}
|
||||
if cfg.MaxUploadBytes <= 0 {
|
||||
cfg.MaxUploadBytes = 32 << 20
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
@@ -21,6 +21,7 @@ func TestLoadFromDirDefaultsToDevYAML(t *testing.T) {
|
||||
require.Equal(t, "18080", cfg.Port)
|
||||
require.Equal(t, "postgres://dev", cfg.DSN)
|
||||
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-system", cfg.SystemAIKey)
|
||||
require.Equal(t, "dev-ai", cfg.AIKeyEncryptionSecret)
|
||||
|
||||
@@ -2,6 +2,7 @@ package files
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
@@ -26,12 +27,20 @@ type SourceDTO struct {
|
||||
|
||||
// Handler 负责 multipart 边界,文件名清理和本地路径构造始终委托给 Service。
|
||||
type Handler struct {
|
||||
service *Service
|
||||
service *Service
|
||||
maxUploadBytes int64
|
||||
}
|
||||
|
||||
// DefaultMaxUploadBytes 是未显式配置时的请求体上限,覆盖完整 multipart 内容。
|
||||
const DefaultMaxUploadBytes int64 = 32 << 20
|
||||
|
||||
// NewHandler 创建文件资料 HTTP registrar。
|
||||
func NewHandler(service *Service) *Handler {
|
||||
return &Handler{service: service}
|
||||
func NewHandler(service *Service, limits ...int64) *Handler {
|
||||
limit := DefaultMaxUploadBytes
|
||||
if len(limits) > 0 && limits[0] > 0 {
|
||||
limit = limits[0]
|
||||
}
|
||||
return &Handler{service: service, maxUploadBytes: limit}
|
||||
}
|
||||
|
||||
// Register 将文件资料上传接口注册到上层提供的 /api/v1 路由组。
|
||||
@@ -55,8 +64,15 @@ func (h *Handler) upload(c *gin.Context) {
|
||||
writeSourceError(c, err)
|
||||
return
|
||||
}
|
||||
// MaxBytesReader 必须在任何 multipart 解析前安装,限制字段、边界和文件内容的总请求量。
|
||||
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, h.maxUploadBytes)
|
||||
fileHeader, err := c.FormFile("file")
|
||||
if err != nil {
|
||||
var maxBytesError *http.MaxBytesError
|
||||
if errors.As(err, &maxBytesError) {
|
||||
httpx.Error(c, http.StatusRequestEntityTooLarge, "payload_too_large", "上传内容超过大小限制")
|
||||
return
|
||||
}
|
||||
httpx.Error(c, http.StatusBadRequest, "invalid_request", "请选择要上传的文件")
|
||||
return
|
||||
}
|
||||
@@ -75,6 +91,10 @@ func (h *Handler) upload(c *gin.Context) {
|
||||
}
|
||||
source, err := h.service.CreateSource(userID, project, c.PostForm("title"), stored)
|
||||
if err != nil {
|
||||
if cleanupErr := h.service.Remove(stored); cleanupErr != nil {
|
||||
// 清理错误仅写服务端日志,响应继续使用稳定中文错误,不暴露 storage root。
|
||||
log.Printf("source file cleanup failed: %v", cleanupErr)
|
||||
}
|
||||
writeSourceError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -3,7 +3,9 @@ package files
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -76,7 +78,102 @@ func TestFileRegistrarChecksOwnershipBeforeWritingFile(t *testing.T) {
|
||||
require.Empty(t, entries)
|
||||
}
|
||||
|
||||
func TestFileRegistrarRemovesStoredFileWhenSourceCreateFails(t *testing.T) {
|
||||
router, database, project, storageRoot := newFileHandlerTestRouter(t, 1, 1)
|
||||
require.NoError(t, database.Callback().Create().Before("gorm:create").Register("test:fail_source_create", func(tx *gorm.DB) {
|
||||
if tx.Statement.Schema != nil && tx.Statement.Schema.Table == (models.SenlinAgentSource{}).TableName() {
|
||||
tx.AddError(errors.New("forced source create failure"))
|
||||
}
|
||||
}))
|
||||
body := &bytes.Buffer{}
|
||||
writer := multipart.NewWriter(body)
|
||||
part, err := writer.CreateFormFile("file", "orphan.txt")
|
||||
require.NoError(t, err)
|
||||
_, err = part.Write([]byte("must be compensated"))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, writer.Close())
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/sources", body)
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer test-token")
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusInternalServerError, rec.Code, rec.Body.String())
|
||||
require.NotContains(t, rec.Body.String(), storageRoot)
|
||||
entries, err := os.ReadDir(storageRoot)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, entries, "source 元数据失败后存储目录应为空")
|
||||
}
|
||||
|
||||
func TestFileRegistrarRejectsRequestOverConfiguredUploadLimit(t *testing.T) {
|
||||
router, _, project, storageRoot := newFileHandlerTestRouterWithLimit(t, 1, 1, 256)
|
||||
body := &bytes.Buffer{}
|
||||
writer := multipart.NewWriter(body)
|
||||
part, err := writer.CreateFormFile("file", "oversize.bin")
|
||||
require.NoError(t, err)
|
||||
_, err = part.Write(bytes.Repeat([]byte("x"), 1024))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, writer.Close())
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/sources", body)
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer test-token")
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusRequestEntityTooLarge, rec.Code, rec.Body.String())
|
||||
var payload httpx.ErrorEnvelope
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
|
||||
require.Equal(t, "payload_too_large", payload.Error.Code)
|
||||
require.Equal(t, "上传内容超过大小限制", payload.Error.Message)
|
||||
entries, err := os.ReadDir(storageRoot)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, entries)
|
||||
}
|
||||
|
||||
func TestFileRegistrarLogsCleanupFailureWithoutLeakingStorageRoot(t *testing.T) {
|
||||
router, database, project, storageRoot, service := newFileHandlerTestRouterWithService(t, 1, 1, DefaultMaxUploadBytes)
|
||||
require.NoError(t, database.Callback().Create().Before("gorm:create").Register("test:fail_source_create_for_cleanup_log", func(tx *gorm.DB) {
|
||||
if tx.Statement.Schema != nil && tx.Statement.Schema.Table == (models.SenlinAgentSource{}).TableName() {
|
||||
tx.AddError(errors.New("forced source create failure"))
|
||||
}
|
||||
}))
|
||||
service.remove = func(path string) error { return fmt.Errorf("cleanup blocked for %s", path) }
|
||||
var serverLog bytes.Buffer
|
||||
previousLogOutput := log.Writer()
|
||||
log.SetOutput(&serverLog)
|
||||
t.Cleanup(func() { log.SetOutput(previousLogOutput) })
|
||||
body := &bytes.Buffer{}
|
||||
writer := multipart.NewWriter(body)
|
||||
part, err := writer.CreateFormFile("file", "cleanup-failure.txt")
|
||||
require.NoError(t, err)
|
||||
_, err = part.Write([]byte("content"))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, writer.Close())
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/sources", body)
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer test-token")
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusInternalServerError, rec.Code, rec.Body.String())
|
||||
require.NotContains(t, rec.Body.String(), storageRoot)
|
||||
require.Contains(t, serverLog.String(), "source file cleanup failed")
|
||||
require.Contains(t, serverLog.String(), storageRoot)
|
||||
}
|
||||
|
||||
func newFileHandlerTestRouter(t *testing.T, currentUserID, ownerID uint) (*gin.Engine, *gorm.DB, models.SenlinAgentProject, string) {
|
||||
return newFileHandlerTestRouterWithLimit(t, currentUserID, ownerID, DefaultMaxUploadBytes)
|
||||
}
|
||||
|
||||
func newFileHandlerTestRouterWithLimit(t *testing.T, currentUserID, ownerID uint, maxUploadBytes int64) (*gin.Engine, *gorm.DB, models.SenlinAgentProject, string) {
|
||||
router, database, project, storageRoot, _ := newFileHandlerTestRouterWithService(t, currentUserID, ownerID, maxUploadBytes)
|
||||
return router, database, project, storageRoot
|
||||
}
|
||||
|
||||
func newFileHandlerTestRouterWithService(t *testing.T, currentUserID, ownerID uint, maxUploadBytes int64) (*gin.Engine, *gorm.DB, models.SenlinAgentProject, string, *Service) {
|
||||
t.Helper()
|
||||
database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{TranslateError: true})
|
||||
require.NoError(t, err)
|
||||
@@ -87,10 +184,11 @@ func newFileHandlerTestRouter(t *testing.T, currentUserID, ownerID uint) (*gin.E
|
||||
project := models.SenlinAgentProject{OwnerID: ownerID, Name: "Files", Identifier: fmt.Sprintf("FILES-%d", ownerID)}
|
||||
require.NoError(t, database.Create(&project).Error)
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot, database)
|
||||
router := httpx.NewProtectedRouter(
|
||||
config.Config{Env: "test"},
|
||||
func(string) (uint, error) { return currentUserID, nil },
|
||||
NewHandler(NewService(storageRoot, database)),
|
||||
NewHandler(service, maxUploadBytes),
|
||||
)
|
||||
return router, database, project, storageRoot
|
||||
return router, database, project, storageRoot, service
|
||||
}
|
||||
|
||||
@@ -14,8 +14,17 @@ import (
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
root string
|
||||
db *gorm.DB
|
||||
root string
|
||||
db *gorm.DB
|
||||
createTemp func(string, string) (stagedFile, error)
|
||||
rename func(string, string) error
|
||||
remove func(string) error
|
||||
}
|
||||
|
||||
type stagedFile interface {
|
||||
io.Writer
|
||||
Close() error
|
||||
Name() string
|
||||
}
|
||||
|
||||
type StoredFile struct {
|
||||
@@ -30,7 +39,12 @@ func NewService(root string, databases ...*gorm.DB) *Service {
|
||||
if len(databases) > 0 {
|
||||
database = databases[0]
|
||||
}
|
||||
return &Service{root: root, db: database}
|
||||
return &Service{
|
||||
root: root, db: database,
|
||||
createTemp: func(directory, pattern string) (stagedFile, error) { return os.CreateTemp(directory, pattern) },
|
||||
rename: os.Rename,
|
||||
remove: os.Remove,
|
||||
}
|
||||
}
|
||||
|
||||
// Save 清理客户端文件名,并只返回供持久化的相对路径;绝对路径不得进入 API DTO。
|
||||
@@ -40,21 +54,90 @@ func (s *Service) Save(projectID uint, originalName string, content io.Reader) (
|
||||
cleanName = "upload.bin"
|
||||
}
|
||||
relative := filepath.ToSlash(filepath.Join("projects", fmt.Sprint(projectID), fmt.Sprintf("%d-%s", time.Now().UnixNano(), cleanName)))
|
||||
absolute := filepath.Join(s.root, filepath.FromSlash(relative))
|
||||
if err := os.MkdirAll(filepath.Dir(absolute), 0o755); err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
file, err := os.Create(absolute)
|
||||
absolute, err := s.absolutePath(relative)
|
||||
if err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(absolute), 0o755); err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
// 临时文件与最终文件位于同一目录,Close 成功后再原子替换,避免暴露半写入内容。
|
||||
file, err := s.createTemp(filepath.Dir(absolute), ".upload-*")
|
||||
if err != nil {
|
||||
return StoredFile{}, err
|
||||
}
|
||||
defer file.Close()
|
||||
if _, err := io.Copy(file, content); err != nil {
|
||||
_ = file.Close()
|
||||
s.cleanupFailedSave(file.Name(), absolute)
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := file.Close(); err != nil {
|
||||
s.cleanupFailedSave(file.Name(), absolute)
|
||||
return StoredFile{}, err
|
||||
}
|
||||
if err := s.rename(file.Name(), absolute); err != nil {
|
||||
s.cleanupFailedSave(file.Name(), absolute)
|
||||
return StoredFile{}, err
|
||||
}
|
||||
return StoredFile{OriginalName: cleanName, RelativePath: relative, AbsolutePath: absolute}, nil
|
||||
}
|
||||
|
||||
// Remove 只依据 Service 生成的相对路径定位文件,并尽量移除直至存储根目录的空父目录。
|
||||
func (s *Service) Remove(stored StoredFile) error {
|
||||
absolute, err := s.absolutePath(stored.RelativePath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.remove(absolute); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
return err
|
||||
}
|
||||
root, err := filepath.Abs(s.root)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for directory := filepath.Dir(absolute); directory != root; directory = filepath.Dir(directory) {
|
||||
entries, err := os.ReadDir(directory)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(entries) > 0 {
|
||||
return nil
|
||||
}
|
||||
if err := s.remove(directory); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) absolutePath(relativePath string) (string, error) {
|
||||
cleaned := filepath.Clean(filepath.FromSlash(strings.TrimSpace(relativePath)))
|
||||
if cleaned == "." || filepath.IsAbs(cleaned) || cleaned == ".." || strings.HasPrefix(cleaned, ".."+string(filepath.Separator)) {
|
||||
return "", ErrSourcePathRequired
|
||||
}
|
||||
root, err := filepath.Abs(s.root)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
absolute, err := filepath.Abs(filepath.Join(root, cleaned))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
relativeToRoot, err := filepath.Rel(root, absolute)
|
||||
if err != nil || relativeToRoot == ".." || strings.HasPrefix(relativeToRoot, ".."+string(filepath.Separator)) {
|
||||
return "", ErrSourcePathRequired
|
||||
}
|
||||
return absolute, nil
|
||||
}
|
||||
|
||||
func (s *Service) cleanupFailedSave(tempPath, finalPath string) {
|
||||
_ = s.remove(tempPath)
|
||||
_ = s.remove(finalPath)
|
||||
}
|
||||
|
||||
func (s *Service) database() *gorm.DB {
|
||||
if s.db != nil {
|
||||
return s.db
|
||||
|
||||
@@ -1,12 +1,37 @@
|
||||
package files
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type failingReader struct {
|
||||
sent bool
|
||||
}
|
||||
|
||||
func (r *failingReader) Read(buffer []byte) (int, error) {
|
||||
if !r.sent {
|
||||
r.sent = true
|
||||
return copy(buffer, "partial"), nil
|
||||
}
|
||||
return 0, errors.New("copy failed")
|
||||
}
|
||||
|
||||
type closeFailingFile struct {
|
||||
*os.File
|
||||
}
|
||||
|
||||
func (f *closeFailingFile) Close() error {
|
||||
_ = f.File.Close()
|
||||
return errors.New("close failed")
|
||||
}
|
||||
|
||||
func TestSaveStoresFileUnderProjectDirectory(t *testing.T) {
|
||||
service := NewService(t.TempDir())
|
||||
|
||||
@@ -27,3 +52,57 @@ func TestSaveNeutralizesPathTraversal(t *testing.T) {
|
||||
require.Equal(t, "secret.txt", stored.OriginalName)
|
||||
require.Contains(t, stored.RelativePath, "projects/12/")
|
||||
}
|
||||
|
||||
func TestSaveRemovesTemporaryAndPartialFilesWhenCopyFails(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
|
||||
_, err := service.Save(12, "broken.bin", &failingReader{})
|
||||
|
||||
require.ErrorContains(t, err, "copy failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
}
|
||||
|
||||
func TestSaveRemovesTemporaryFileWhenCloseFails(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
service.createTemp = func(directory, pattern string) (stagedFile, error) {
|
||||
file, err := os.CreateTemp(directory, pattern)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &closeFailingFile{File: file}, nil
|
||||
}
|
||||
|
||||
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
|
||||
|
||||
require.ErrorContains(t, err, "close failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
}
|
||||
|
||||
func TestSaveRemovesTemporaryAndPartialFilesWhenRenameFails(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
service.rename = func(oldPath, newPath string) error {
|
||||
require.NoError(t, os.WriteFile(newPath, []byte("partial final"), 0o600))
|
||||
return errors.New("rename failed")
|
||||
}
|
||||
|
||||
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
|
||||
|
||||
require.ErrorContains(t, err, "rename failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
}
|
||||
|
||||
func requireNoStoredFiles(t *testing.T, storageRoot string) {
|
||||
t.Helper()
|
||||
require.NoError(t, filepath.WalkDir(storageRoot, func(path string, entry fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if path != storageRoot && !entry.IsDir() {
|
||||
t.Fatalf("unexpected stored file after failure: %s", path)
|
||||
}
|
||||
return nil
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -1,11 +1,13 @@
|
||||
package projects
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
"senlinai-agent/backend/internal/httpx"
|
||||
"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,
|
||||
})
|
||||
if err != nil {
|
||||
httpx.Error(c, http.StatusBadRequest, "invalid_request", "计划任务参数无效")
|
||||
writeCronError(c, err)
|
||||
return
|
||||
}
|
||||
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 {
|
||||
return CronPlanDTO{
|
||||
ID: plan.Identity, ProjectID: plan.ProjectIdentity, Title: plan.Title, Schedule: plan.Schedule,
|
||||
|
||||
@@ -8,17 +8,23 @@ import (
|
||||
"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) {
|
||||
if err := ensureProjectOwner(ownerID, projectID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
title := strings.TrimSpace(input.Title)
|
||||
if title == "" {
|
||||
return nil, errors.New("cron plan title is required")
|
||||
return nil, ErrCronTitleRequired
|
||||
}
|
||||
schedule := strings.TrimSpace(input.Schedule)
|
||||
if schedule == "" {
|
||||
return nil, errors.New("cron schedule is required")
|
||||
return nil, ErrCronScheduleRequired
|
||||
}
|
||||
plan := &models.SenlinAgentCronPlan{
|
||||
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) {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return nil, errors.New("tag name is required")
|
||||
return nil, ErrTagNameRequired
|
||||
}
|
||||
var existing models.SenlinAgentTag
|
||||
err := models.DBService.Where("project_id = ? AND name = ?", projectID, name).First(&existing).Error
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
package projects
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
"senlinai-agent/backend/internal/httpx"
|
||||
"senlinai-agent/backend/internal/logic/auth"
|
||||
"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)
|
||||
if err != nil {
|
||||
httpx.Error(c, http.StatusBadRequest, "invalid_request", "标签名称不能为空")
|
||||
writeTagError(c, err)
|
||||
return
|
||||
}
|
||||
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) {
|
||||
userID, ok := auth.CurrentUserID(c)
|
||||
if !ok {
|
||||
|
||||
@@ -3,12 +3,14 @@ package projects
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
"senlinai-agent/backend/internal/config"
|
||||
"senlinai-agent/backend/internal/httpx"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
@@ -62,6 +64,75 @@ func TestTagRegistrarRejectsProjectOwnedByAnotherUser(t *testing.T) {
|
||||
router.ServeHTTP(rec, req)
|
||||
|
||||
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) {
|
||||
|
||||
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", "请求参数无效")
|
||||
return
|
||||
}
|
||||
nextProjectIdentity := ""
|
||||
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")
|
||||
return
|
||||
}
|
||||
}
|
||||
task, err := h.service.Update(userID, projectIdentity, taskIdentity, UpdateTaskInput{
|
||||
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 {
|
||||
writeTaskError(c, err)
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"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 分享"}
|
||||
require.NoError(t, database.Create(¬e).Error)
|
||||
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.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer test-token")
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
@@ -15,8 +16,8 @@ type Service struct {
|
||||
}
|
||||
|
||||
type LinkedObject struct {
|
||||
ObjectType string `json:"object_type"`
|
||||
ObjectID uint `json:"object_id"`
|
||||
ObjectType string `json:"objectType"`
|
||||
ObjectID string `json:"objectId"`
|
||||
}
|
||||
|
||||
// NewService 接受数据库或上层事务,确保任务写入、标签调整和分享检查使用同一依赖。
|
||||
@@ -35,13 +36,20 @@ func (s *Service) database() *gorm.DB {
|
||||
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 {
|
||||
var task models.SenlinAgentTask
|
||||
if err := tx.First(&task, taskID).Error; err != nil {
|
||||
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 tx.Create(&models.SenlinAgentProjectEvent{
|
||||
@@ -50,7 +58,7 @@ func (s *Service) Assign(taskID uint, assigneeID uint) error {
|
||||
EventType: "task_assigned",
|
||||
EntityType: "task",
|
||||
EntityID: task.ID,
|
||||
Summary: fmt.Sprintf("Task assigned to user %d", assigneeID),
|
||||
Summary: fmt.Sprintf("Task assigned to user %s", assignee.Identity),
|
||||
}).Error
|
||||
})
|
||||
}
|
||||
@@ -61,8 +69,8 @@ func (s *Service) ShareObject(taskID uint, objectType string, objectID uint) err
|
||||
return errors.New("unsupported shared object type")
|
||||
}
|
||||
return s.database().Transaction(func(tx *gorm.DB) error {
|
||||
var task models.SenlinAgentTask
|
||||
if err := tx.First(&task, taskID).Error; err != nil {
|
||||
task, err := lockTaskByID(tx, taskID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
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))
|
||||
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
|
||||
}
|
||||
@@ -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) {
|
||||
var result TaskDTO
|
||||
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)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var task models.SenlinAgentTask
|
||||
if err := tx.Where("identity = ? AND project_id = ?", taskIdentity, currentProject.ID).First(&task).Error; err != nil {
|
||||
return err
|
||||
if task.ProjectID != currentProject.ID {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
targetProject := currentProject
|
||||
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
|
||||
}
|
||||
}
|
||||
if err := tx.Save(&task).Error; err != nil {
|
||||
if err := tx.Save(task).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
result = makeTaskDTO(task, targetProject.Identity, tagIdentity, tagName)
|
||||
result = makeTaskDTO(*task, targetProject.Identity, tagIdentity, tagName)
|
||||
return nil
|
||||
})
|
||||
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) {
|
||||
var project models.SenlinAgentProject
|
||||
if err := tx.Where("owner_id = ? AND identity = ?", ownerID, identity).First(&project).Error; err != nil {
|
||||
|
||||
@@ -2,6 +2,7 @@ package tasks
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
@@ -10,6 +11,30 @@ import (
|
||||
"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) {
|
||||
database := newTestDB(t)
|
||||
assigneeID := uint(2)
|
||||
@@ -28,7 +53,7 @@ func TestAssigneeOnlySeesExplicitlySharedObjects(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.Len(t, after, 1)
|
||||
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) {
|
||||
@@ -53,16 +78,32 @@ func TestShareObjectRejectsObjectFromAnotherProject(t *testing.T) {
|
||||
require.ErrorContains(t, err, "shared object not found in task project")
|
||||
}
|
||||
|
||||
func TestAssignRecordsProjectEvent(t *testing.T) {
|
||||
func TestAssignResolvesIdentityAndTaskDTOReflectsReassignment(t *testing.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)
|
||||
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
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -74,3 +115,45 @@ func newTestDB(t *testing.T) *gorm.DB {
|
||||
models.DBService = 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