fix(backend): harden write registrar boundaries
This commit is contained in:
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