fix(backend): harden write registrar boundaries

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

View File

@@ -29,7 +29,7 @@ func main() {
tagHandler := projects.NewTagHandler(projectService)
cronHandler := projects.NewCronHandler(projectService)
taskHandler := tasks.NewHandler(taskService)
fileHandler := files.NewHandler(fileService)
fileHandler := files.NewHandler(fileService, cfg.MaxUploadBytes)
inboxHandler := inbox.NewHandler(inbox.NewService(inbox.StaticAnalyzer{}))
authHandler := auth.NewHandler(authService)
// search 与 AI 当前只有 service尚无 HTTP registrar后续功能任务应在实现真实契约后从此处注入不能注册伪端点。

View File

@@ -2,6 +2,7 @@ env: development
port: "9150"
dsn: "postgres://postgres:postgres@localhost:5432/agent_dev?sslmode=disable"
storage_dir: "./data/files"
max_upload_bytes: 33554432
auth_secret: "development-auth-secret-change-me"
system_ai_key: ""
ai_key_encryption_secret: "development-ai-key-secret-change-me"

View File

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

View File

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

View File

@@ -2,6 +2,7 @@ package files
import (
"errors"
"log"
"net/http"
"time"
@@ -26,12 +27,20 @@ type SourceDTO struct {
// Handler 负责 multipart 边界,文件名清理和本地路径构造始终委托给 Service。
type Handler struct {
service *Service
service *Service
maxUploadBytes int64
}
// DefaultMaxUploadBytes 是未显式配置时的请求体上限,覆盖完整 multipart 内容。
const DefaultMaxUploadBytes int64 = 32 << 20
// NewHandler 创建文件资料 HTTP registrar。
func NewHandler(service *Service) *Handler {
return &Handler{service: service}
func NewHandler(service *Service, limits ...int64) *Handler {
limit := DefaultMaxUploadBytes
if len(limits) > 0 && limits[0] > 0 {
limit = limits[0]
}
return &Handler{service: service, maxUploadBytes: limit}
}
// Register 将文件资料上传接口注册到上层提供的 /api/v1 路由组。
@@ -55,8 +64,15 @@ func (h *Handler) upload(c *gin.Context) {
writeSourceError(c, err)
return
}
// MaxBytesReader 必须在任何 multipart 解析前安装,限制字段、边界和文件内容的总请求量。
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, h.maxUploadBytes)
fileHeader, err := c.FormFile("file")
if err != nil {
var maxBytesError *http.MaxBytesError
if errors.As(err, &maxBytesError) {
httpx.Error(c, http.StatusRequestEntityTooLarge, "payload_too_large", "上传内容超过大小限制")
return
}
httpx.Error(c, http.StatusBadRequest, "invalid_request", "请选择要上传的文件")
return
}
@@ -75,6 +91,10 @@ func (h *Handler) upload(c *gin.Context) {
}
source, err := h.service.CreateSource(userID, project, c.PostForm("title"), stored)
if err != nil {
if cleanupErr := h.service.Remove(stored); cleanupErr != nil {
// 清理错误仅写服务端日志,响应继续使用稳定中文错误,不暴露 storage root。
log.Printf("source file cleanup failed: %v", cleanupErr)
}
writeSourceError(c, err)
return
}

View File

@@ -3,7 +3,9 @@ package files
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"log"
"mime/multipart"
"net/http"
"net/http/httptest"
@@ -76,7 +78,102 @@ func TestFileRegistrarChecksOwnershipBeforeWritingFile(t *testing.T) {
require.Empty(t, entries)
}
func TestFileRegistrarRemovesStoredFileWhenSourceCreateFails(t *testing.T) {
router, database, project, storageRoot := newFileHandlerTestRouter(t, 1, 1)
require.NoError(t, database.Callback().Create().Before("gorm:create").Register("test:fail_source_create", func(tx *gorm.DB) {
if tx.Statement.Schema != nil && tx.Statement.Schema.Table == (models.SenlinAgentSource{}).TableName() {
tx.AddError(errors.New("forced source create failure"))
}
}))
body := &bytes.Buffer{}
writer := multipart.NewWriter(body)
part, err := writer.CreateFormFile("file", "orphan.txt")
require.NoError(t, err)
_, err = part.Write([]byte("must be compensated"))
require.NoError(t, err)
require.NoError(t, writer.Close())
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/sources", body)
req.Header.Set("Content-Type", writer.FormDataContentType())
req.Header.Set("Authorization", "Bearer test-token")
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusInternalServerError, rec.Code, rec.Body.String())
require.NotContains(t, rec.Body.String(), storageRoot)
entries, err := os.ReadDir(storageRoot)
require.NoError(t, err)
require.Empty(t, entries, "source 元数据失败后存储目录应为空")
}
func TestFileRegistrarRejectsRequestOverConfiguredUploadLimit(t *testing.T) {
router, _, project, storageRoot := newFileHandlerTestRouterWithLimit(t, 1, 1, 256)
body := &bytes.Buffer{}
writer := multipart.NewWriter(body)
part, err := writer.CreateFormFile("file", "oversize.bin")
require.NoError(t, err)
_, err = part.Write(bytes.Repeat([]byte("x"), 1024))
require.NoError(t, err)
require.NoError(t, writer.Close())
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/sources", body)
req.Header.Set("Content-Type", writer.FormDataContentType())
req.Header.Set("Authorization", "Bearer test-token")
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusRequestEntityTooLarge, rec.Code, rec.Body.String())
var payload httpx.ErrorEnvelope
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
require.Equal(t, "payload_too_large", payload.Error.Code)
require.Equal(t, "上传内容超过大小限制", payload.Error.Message)
entries, err := os.ReadDir(storageRoot)
require.NoError(t, err)
require.Empty(t, entries)
}
func TestFileRegistrarLogsCleanupFailureWithoutLeakingStorageRoot(t *testing.T) {
router, database, project, storageRoot, service := newFileHandlerTestRouterWithService(t, 1, 1, DefaultMaxUploadBytes)
require.NoError(t, database.Callback().Create().Before("gorm:create").Register("test:fail_source_create_for_cleanup_log", func(tx *gorm.DB) {
if tx.Statement.Schema != nil && tx.Statement.Schema.Table == (models.SenlinAgentSource{}).TableName() {
tx.AddError(errors.New("forced source create failure"))
}
}))
service.remove = func(path string) error { return fmt.Errorf("cleanup blocked for %s", path) }
var serverLog bytes.Buffer
previousLogOutput := log.Writer()
log.SetOutput(&serverLog)
t.Cleanup(func() { log.SetOutput(previousLogOutput) })
body := &bytes.Buffer{}
writer := multipart.NewWriter(body)
part, err := writer.CreateFormFile("file", "cleanup-failure.txt")
require.NoError(t, err)
_, err = part.Write([]byte("content"))
require.NoError(t, err)
require.NoError(t, writer.Close())
req := httptest.NewRequest(http.MethodPost, "/api/v1/projects/"+project.Identity+"/sources", body)
req.Header.Set("Content-Type", writer.FormDataContentType())
req.Header.Set("Authorization", "Bearer test-token")
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusInternalServerError, rec.Code, rec.Body.String())
require.NotContains(t, rec.Body.String(), storageRoot)
require.Contains(t, serverLog.String(), "source file cleanup failed")
require.Contains(t, serverLog.String(), storageRoot)
}
func newFileHandlerTestRouter(t *testing.T, currentUserID, ownerID uint) (*gin.Engine, *gorm.DB, models.SenlinAgentProject, string) {
return newFileHandlerTestRouterWithLimit(t, currentUserID, ownerID, DefaultMaxUploadBytes)
}
func newFileHandlerTestRouterWithLimit(t *testing.T, currentUserID, ownerID uint, maxUploadBytes int64) (*gin.Engine, *gorm.DB, models.SenlinAgentProject, string) {
router, database, project, storageRoot, _ := newFileHandlerTestRouterWithService(t, currentUserID, ownerID, maxUploadBytes)
return router, database, project, storageRoot
}
func newFileHandlerTestRouterWithService(t *testing.T, currentUserID, ownerID uint, maxUploadBytes int64) (*gin.Engine, *gorm.DB, models.SenlinAgentProject, string, *Service) {
t.Helper()
database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{TranslateError: true})
require.NoError(t, err)
@@ -87,10 +184,11 @@ func newFileHandlerTestRouter(t *testing.T, currentUserID, ownerID uint) (*gin.E
project := models.SenlinAgentProject{OwnerID: ownerID, Name: "Files", Identifier: fmt.Sprintf("FILES-%d", ownerID)}
require.NoError(t, database.Create(&project).Error)
storageRoot := t.TempDir()
service := NewService(storageRoot, database)
router := httpx.NewProtectedRouter(
config.Config{Env: "test"},
func(string) (uint, error) { return currentUserID, nil },
NewHandler(NewService(storageRoot, database)),
NewHandler(service, maxUploadBytes),
)
return router, database, project, storageRoot
return router, database, project, storageRoot, service
}

View File

@@ -14,8 +14,17 @@ import (
)
type Service struct {
root string
db *gorm.DB
root string
db *gorm.DB
createTemp func(string, string) (stagedFile, error)
rename func(string, string) error
remove func(string) error
}
type stagedFile interface {
io.Writer
Close() error
Name() string
}
type StoredFile struct {
@@ -30,7 +39,12 @@ func NewService(root string, databases ...*gorm.DB) *Service {
if len(databases) > 0 {
database = databases[0]
}
return &Service{root: root, db: database}
return &Service{
root: root, db: database,
createTemp: func(directory, pattern string) (stagedFile, error) { return os.CreateTemp(directory, pattern) },
rename: os.Rename,
remove: os.Remove,
}
}
// Save 清理客户端文件名,并只返回供持久化的相对路径;绝对路径不得进入 API DTO。
@@ -40,21 +54,90 @@ func (s *Service) Save(projectID uint, originalName string, content io.Reader) (
cleanName = "upload.bin"
}
relative := filepath.ToSlash(filepath.Join("projects", fmt.Sprint(projectID), fmt.Sprintf("%d-%s", time.Now().UnixNano(), cleanName)))
absolute := filepath.Join(s.root, filepath.FromSlash(relative))
if err := os.MkdirAll(filepath.Dir(absolute), 0o755); err != nil {
return StoredFile{}, err
}
file, err := os.Create(absolute)
absolute, err := s.absolutePath(relative)
if err != nil {
return StoredFile{}, err
}
if err := os.MkdirAll(filepath.Dir(absolute), 0o755); err != nil {
return StoredFile{}, err
}
// 临时文件与最终文件位于同一目录Close 成功后再原子替换,避免暴露半写入内容。
file, err := s.createTemp(filepath.Dir(absolute), ".upload-*")
if err != nil {
return StoredFile{}, err
}
defer file.Close()
if _, err := io.Copy(file, content); err != nil {
_ = file.Close()
s.cleanupFailedSave(file.Name(), absolute)
return StoredFile{}, err
}
if err := file.Close(); err != nil {
s.cleanupFailedSave(file.Name(), absolute)
return StoredFile{}, err
}
if err := s.rename(file.Name(), absolute); err != nil {
s.cleanupFailedSave(file.Name(), absolute)
return StoredFile{}, err
}
return StoredFile{OriginalName: cleanName, RelativePath: relative, AbsolutePath: absolute}, nil
}
// Remove 只依据 Service 生成的相对路径定位文件,并尽量移除直至存储根目录的空父目录。
func (s *Service) Remove(stored StoredFile) error {
absolute, err := s.absolutePath(stored.RelativePath)
if err != nil {
return err
}
if err := s.remove(absolute); err != nil && !errors.Is(err, os.ErrNotExist) {
return err
}
root, err := filepath.Abs(s.root)
if err != nil {
return err
}
for directory := filepath.Dir(absolute); directory != root; directory = filepath.Dir(directory) {
entries, err := os.ReadDir(directory)
if errors.Is(err, os.ErrNotExist) {
continue
}
if err != nil {
return err
}
if len(entries) > 0 {
return nil
}
if err := s.remove(directory); err != nil && !errors.Is(err, os.ErrNotExist) {
return err
}
}
return nil
}
func (s *Service) absolutePath(relativePath string) (string, error) {
cleaned := filepath.Clean(filepath.FromSlash(strings.TrimSpace(relativePath)))
if cleaned == "." || filepath.IsAbs(cleaned) || cleaned == ".." || strings.HasPrefix(cleaned, ".."+string(filepath.Separator)) {
return "", ErrSourcePathRequired
}
root, err := filepath.Abs(s.root)
if err != nil {
return "", err
}
absolute, err := filepath.Abs(filepath.Join(root, cleaned))
if err != nil {
return "", err
}
relativeToRoot, err := filepath.Rel(root, absolute)
if err != nil || relativeToRoot == ".." || strings.HasPrefix(relativeToRoot, ".."+string(filepath.Separator)) {
return "", ErrSourcePathRequired
}
return absolute, nil
}
func (s *Service) cleanupFailedSave(tempPath, finalPath string) {
_ = s.remove(tempPath)
_ = s.remove(finalPath)
}
func (s *Service) database() *gorm.DB {
if s.db != nil {
return s.db

View File

@@ -1,12 +1,37 @@
package files
import (
"errors"
"io/fs"
"os"
"path/filepath"
"strings"
"testing"
"github.com/stretchr/testify/require"
)
type failingReader struct {
sent bool
}
func (r *failingReader) Read(buffer []byte) (int, error) {
if !r.sent {
r.sent = true
return copy(buffer, "partial"), nil
}
return 0, errors.New("copy failed")
}
type closeFailingFile struct {
*os.File
}
func (f *closeFailingFile) Close() error {
_ = f.File.Close()
return errors.New("close failed")
}
func TestSaveStoresFileUnderProjectDirectory(t *testing.T) {
service := NewService(t.TempDir())
@@ -27,3 +52,57 @@ func TestSaveNeutralizesPathTraversal(t *testing.T) {
require.Equal(t, "secret.txt", stored.OriginalName)
require.Contains(t, stored.RelativePath, "projects/12/")
}
func TestSaveRemovesTemporaryAndPartialFilesWhenCopyFails(t *testing.T) {
storageRoot := t.TempDir()
service := NewService(storageRoot)
_, err := service.Save(12, "broken.bin", &failingReader{})
require.ErrorContains(t, err, "copy failed")
requireNoStoredFiles(t, storageRoot)
}
func TestSaveRemovesTemporaryFileWhenCloseFails(t *testing.T) {
storageRoot := t.TempDir()
service := NewService(storageRoot)
service.createTemp = func(directory, pattern string) (stagedFile, error) {
file, err := os.CreateTemp(directory, pattern)
if err != nil {
return nil, err
}
return &closeFailingFile{File: file}, nil
}
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
require.ErrorContains(t, err, "close failed")
requireNoStoredFiles(t, storageRoot)
}
func TestSaveRemovesTemporaryAndPartialFilesWhenRenameFails(t *testing.T) {
storageRoot := t.TempDir()
service := NewService(storageRoot)
service.rename = func(oldPath, newPath string) error {
require.NoError(t, os.WriteFile(newPath, []byte("partial final"), 0o600))
return errors.New("rename failed")
}
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
require.ErrorContains(t, err, "rename failed")
requireNoStoredFiles(t, storageRoot)
}
func requireNoStoredFiles(t *testing.T, storageRoot string) {
t.Helper()
require.NoError(t, filepath.WalkDir(storageRoot, func(path string, entry fs.DirEntry, err error) error {
if err != nil {
return err
}
if path != storageRoot && !entry.IsDir() {
t.Fatalf("unexpected stored file after failure: %s", path)
}
return nil
}))
}

View File

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

View File

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

View File

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

View File

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

View 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(&note).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(&note)
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)
}

View File

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

View File

@@ -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(&note).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")

View File

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

View File

@@ -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(&note).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
}