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

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