diff --git a/backend/cmd/api/main.go b/backend/cmd/api/main.go index 096c073..2de0405 100644 --- a/backend/cmd/api/main.go +++ b/backend/cmd/api/main.go @@ -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;后续功能任务应在实现真实契约后从此处注入,不能注册伪端点。 diff --git a/backend/etc/agent.dev.yaml b/backend/etc/agent.dev.yaml index ab8be8b..bb89521 100644 --- a/backend/etc/agent.dev.yaml +++ b/backend/etc/agent.dev.yaml @@ -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" diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index c3e2d91..51bcaaf 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -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 } diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index c9b9efd..ad90a9c 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -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) diff --git a/backend/internal/logic/files/handlers.go b/backend/internal/logic/files/handlers.go index c0fc50b..244e659 100644 --- a/backend/internal/logic/files/handlers.go +++ b/backend/internal/logic/files/handlers.go @@ -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 } diff --git a/backend/internal/logic/files/handlers_test.go b/backend/internal/logic/files/handlers_test.go index b7351f2..d91e61f 100644 --- a/backend/internal/logic/files/handlers_test.go +++ b/backend/internal/logic/files/handlers_test.go @@ -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 } diff --git a/backend/internal/logic/files/service.go b/backend/internal/logic/files/service.go index 8bb827a..13727ad 100644 --- a/backend/internal/logic/files/service.go +++ b/backend/internal/logic/files/service.go @@ -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 diff --git a/backend/internal/logic/files/service_test.go b/backend/internal/logic/files/service_test.go index df05fa6..d205d87 100644 --- a/backend/internal/logic/files/service_test.go +++ b/backend/internal/logic/files/service_test.go @@ -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 + })) +} diff --git a/backend/internal/logic/projects/cron_handlers.go b/backend/internal/logic/projects/cron_handlers.go index 07798b9..6d76da0 100644 --- a/backend/internal/logic/projects/cron_handlers.go +++ b/backend/internal/logic/projects/cron_handlers.go @@ -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, diff --git a/backend/internal/logic/projects/object_crud.go b/backend/internal/logic/projects/object_crud.go index a042a17..397558c 100644 --- a/backend/internal/logic/projects/object_crud.go +++ b/backend/internal/logic/projects/object_crud.go @@ -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 diff --git a/backend/internal/logic/projects/tag_handlers.go b/backend/internal/logic/projects/tag_handlers.go index c6c9b65..4169ede 100644 --- a/backend/internal/logic/projects/tag_handlers.go +++ b/backend/internal/logic/projects/tag_handlers.go @@ -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 { diff --git a/backend/internal/logic/projects/write_handlers_test.go b/backend/internal/logic/projects/write_handlers_test.go index 8bc1fdb..6046e68 100644 --- a/backend/internal/logic/projects/write_handlers_test.go +++ b/backend/internal/logic/projects/write_handlers_test.go @@ -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) { diff --git a/backend/internal/logic/tasks/concurrency_postgres_test.go b/backend/internal/logic/tasks/concurrency_postgres_test.go new file mode 100644 index 0000000..78a8074 --- /dev/null +++ b/backend/internal/logic/tasks/concurrency_postgres_test.go @@ -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) +} diff --git a/backend/internal/logic/tasks/handlers.go b/backend/internal/logic/tasks/handlers.go index d1ba379..8eee16b 100644 --- a/backend/internal/logic/tasks/handlers.go +++ b/backend/internal/logic/tasks/handlers.go @@ -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) diff --git a/backend/internal/logic/tasks/handlers_test.go b/backend/internal/logic/tasks/handlers_test.go index 2cc765c..ebb7e69 100644 --- a/backend/internal/logic/tasks/handlers_test.go +++ b/backend/internal/logic/tasks/handlers_test.go @@ -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") diff --git a/backend/internal/logic/tasks/service.go b/backend/internal/logic/tasks/service.go index 812c77c..27a6b56 100644 --- a/backend/internal/logic/tasks/service.go +++ b/backend/internal/logic/tasks/service.go @@ -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 { diff --git a/backend/internal/logic/tasks/service_test.go b/backend/internal/logic/tasks/service_test.go index 74de8c7..52aa619 100644 --- a/backend/internal/logic/tasks/service_test.go +++ b/backend/internal/logic/tasks/service_test.go @@ -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 +}