diff --git a/server/app/goauto/device/diagnostics_mysql_test.go b/server/app/goauto/device/diagnostics_mysql_test.go index 5e78676..8212ff2 100644 --- a/server/app/goauto/device/diagnostics_mysql_test.go +++ b/server/app/goauto/device/diagnostics_mysql_test.go @@ -112,48 +112,31 @@ func TestMySQLAgentDiagnostics(t *testing.T) { }) return s, d, token } + t.Run("v4", func(t *testing.T) { runOfflineV4MySQLTests(t, db, seed) }) - t.Run("scan_revalidates_heartbeat_after_lock_wait", func(t *testing.T) { - s, d, _ := seed(t) - lock := db.Begin() - if lock.Error != nil { - t.Fatal("begin failed") - } - defer lock.Rollback() - var device models.AgentDevice - if lock.Clauses(clause.Locking{Strength: "UPDATE"}).First(&device, d.DeviceID).Error != nil { - t.Fatal("lock failed") - } - done := make(chan struct { - n int64 - err error - }, 1) - go func() { - n, e := s.MarkStaleDevicesOffline(context.Background(), DefaultOfflineThreshold) - done <- struct { - n int64 - err error - }{n, e} - }() - select { - case <-done: - t.Fatal("scan did not wait for device lock") - case <-time.After(100 * time.Millisecond): - } - if lock.Model(&models.AgentDevice{}).Where("id = ?", d.DeviceID).Update("last_heartbeat_at", s.Now()).Error != nil { - t.Fatal("refresh failed") - } - if lock.Commit().Error != nil { - t.Fatal("commit failed") - } - select { - case result := <-done: - if result.err != nil || result.n != 0 { - t.Fatalf("fresh device falsely marked offline: count=%d error=%v", result.n, result.err != nil) + t.Run("scan_revalidates_real_heartbeat_after_candidate_read", func(t *testing.T) { + s, d, token := seed(t) + type scanMarker struct{} + callback := "v4_heartbeat_before_scan_lock" + if db.Callback().Query().Before("gorm:query").Register(callback, func(tx *gorm.DB) { + if tx.Statement.Context.Value(scanMarker{}) != true || tx.Statement.Table != "agent_device" { + return } - case <-time.After(5 * time.Second): - t.Fatal("scan stuck") + if _, ok := tx.Statement.Clauses["FOR"]; !ok { + return + } + if _, err := s.Heartbeat(context.Background(), HeartbeatRequest{RequestID: uuid.NewString()}, token); err != nil { + tx.AddError(errors.New("concurrent heartbeat failed")) + } + }) != nil { + t.Fatal("register heartbeat barrier failed") } + defer db.Callback().Query().Remove(callback) + n, err := s.MarkStaleDevicesOffline(context.WithValue(context.Background(), scanMarker{}, true), DefaultOfflineThreshold) + if err != nil || n != 0 { + t.Fatal("fresh heartbeat falsely marked offline") + } + requireOfflineDevice(t, db, d.DeviceID, models.DeviceStatusOnline, 0) }) t.Run("concurrent_scans_and_resume_are_unique", func(t *testing.T) { s, d, testDeviceToken := seed(t) @@ -317,12 +300,8 @@ func TestMySQLAgentDiagnostics(t *testing.T) { } t.Fatal("reset lock failed") } - var mysqlErr *driver.MySQLError - if scanErr == nil || !errors.As(scanErr, &mysqlErr) || mysqlErr.Number != 3572 { - if mysqlErr != nil { - t.Fatalf("scan must decline busy task without deadlock; mysql_errno=%d", mysqlErr.Number) - } - t.Fatal("scan must decline busy task") + if scanErr != nil { + t.Fatal("scan must skip busy task without returning contention") } var eventCount int64 db.Model(&models.AgentDeviceStatusEvent{}).Where("device_id = ?", d.DeviceID).Count(&eventCount) diff --git a/server/app/goauto/device/heartbeat.go b/server/app/goauto/device/heartbeat.go index 14b92c8..08cfa78 100644 --- a/server/app/goauto/device/heartbeat.go +++ b/server/app/goauto/device/heartbeat.go @@ -9,9 +9,12 @@ import ( "go-admin/app/goauto/models" + log "github.com/go-admin-team/go-admin-core/logger" + mysqlDriver "github.com/go-sql-driver/mysql" "github.com/google/uuid" "gorm.io/gorm" "gorm.io/gorm/clause" + "gorm.io/gorm/logger" ) const ( @@ -137,8 +140,8 @@ func tokenInvalidError() error { return &ServiceError{Code: CodeTokenInvalid, Message: "设备 Token 无效", Retryable: false} } -// MarkStaleDevicesOffline atomically marks stale devices offline and fails any -// running task held by those devices. It returns the number of devices changed. +// MarkStaleDevicesOffline commits each device, its running tasks and its event +// independently. The count always describes committed device transitions. func (service *Service) MarkStaleDevicesOffline(ctx context.Context, threshold time.Duration) (int64, error) { if threshold <= 0 { return 0, fmt.Errorf("offline threshold must be positive") @@ -146,7 +149,6 @@ func (service *Service) MarkStaleDevicesOffline(ctx context.Context, threshold t now := service.Now() cutoff := now.Add(-threshold) var changed int64 - var events []models.AgentDeviceStatusEvent var deviceIDs []uint64 if err := service.DB.WithContext(ctx).Model(&models.AgentDevice{}). Where("status = ? AND COALESCE(last_heartbeat_at, created_at) < ?", models.DeviceStatusOnline, cutoff). @@ -156,15 +158,20 @@ func (service *Service) MarkStaleDevicesOffline(ctx context.Context, threshold t if len(deviceIDs) == 0 { return 0, nil } - err := service.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { - // Lock primary keys in deterministic order, not the mutable status index: - // a status-range FOR UPDATE can deadlock concurrent scans on gap locks. - // Candidate selection is only a hint; revalidate the current locked row. - for _, deviceID := range deviceIDs { + for _, deviceID := range deviceIDs { + if err := ctx.Err(); err != nil { + return changed, err + } + var event *models.AgentDeviceStatusEvent + // Expected NOWAIT contention must not be printed by GORM with SQL or + // parameters. The fixed warning below is the only contention log. + err := service.DB.WithContext(ctx).Session(&gorm.Session{Logger: logger.Default.LogMode(logger.Silent)}).Transaction(func(tx *gorm.DB) error { var device models.AgentDevice - if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&device, deviceID).Error; err != nil { + // The unlocked candidate list is only a hint. Lock the exact primary + // key without waiting, then revalidate its current state and heartbeat. + if err := tx.Clauses(clause.Locking{Strength: "UPDATE", Options: "NOWAIT"}).First(&device, deviceID).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { - continue + return nil } return internalError(err) } @@ -173,77 +180,103 @@ func (service *Service) MarkStaleDevicesOffline(ctx context.Context, threshold t last = *device.LastHeartbeatAt } if device.Status != models.DeviceStatusOnline || !last.Before(cutoff) { - continue + return nil } - // Existing Admin reset paths lock task -> device, while Agent paths - // lock device -> task. Never wait for a task while holding this device: - // contention aborts this atomic scan and the next scan retries it. - if err := lockOfflineTasksNowait(tx, device.ID); err != nil { - return err - } - result := tx.Session(&gorm.Session{SkipHooks: true}).Model(&models.CollectionTask{}). - Where("device_id = ? AND status = ?", device.ID, models.TaskStatusRunning). - Updates(map[string]any{ - "status": models.TaskStatusFailed, "active_slot": gorm.Expr("NULL"), "device_run_slot": gorm.Expr("NULL"), - "lease_expires_at": gorm.Expr("NULL"), "error_code": "DEVICE_OFFLINE", - "error_message": "设备心跳超时,任务失败且不会自动重试", "finished_at": now, - }) - if result.Error != nil { - return internalError(result.Error) - } - failed, unknown, err := markOfflinePurchaseTasks(tx, []uint64{device.ID}, now) + // Existing Admin reset paths lock task -> device, while collection + // Claim locks device -> task. Never wait for a task holding this device: + // contention rolls back only this device; a later scan retries it. + tasks, err := lockOfflineTasksNowait(tx, device.ID) if err != nil { return err } - event := models.AgentDeviceStatusEvent{DeviceID: device.ID, OccurredAt: now, FromStatus: device.Status, ToStatus: models.DeviceStatusOffline, Reason: "heartbeat_timeout", LastHeartbeatAt: device.LastHeartbeatAt, FailedTaskCount: result.RowsAffected + failed, OrderResultUnknownCount: unknown} - result = tx.Model(&models.AgentDevice{}).Where("id = ? AND status = ?", device.ID, models.DeviceStatusOnline). + var collectionFailed int64 + if len(tasks.CollectionIDs) > 0 { + result := tx.Session(&gorm.Session{SkipHooks: true}).Model(&models.CollectionTask{}). + Where("id IN ? AND status = ?", tasks.CollectionIDs, models.TaskStatusRunning). + Updates(map[string]any{ + "status": models.TaskStatusFailed, "active_slot": gorm.Expr("NULL"), "device_run_slot": gorm.Expr("NULL"), + "lease_expires_at": gorm.Expr("NULL"), "error_code": "DEVICE_OFFLINE", + "error_message": "设备心跳超时,任务失败且不会自动重试", "finished_at": now, + }) + if result.Error != nil { + return internalError(result.Error) + } + collectionFailed = result.RowsAffected + } + failed, unknown, err := markOfflinePurchaseTasks(tx, tasks, now) + if err != nil { + return err + } + event = &models.AgentDeviceStatusEvent{DeviceID: device.ID, OccurredAt: now, FromStatus: device.Status, ToStatus: models.DeviceStatusOffline, Reason: "heartbeat_timeout", LastHeartbeatAt: device.LastHeartbeatAt, FailedTaskCount: collectionFailed + failed, OrderResultUnknownCount: unknown} + result := tx.Model(&models.AgentDevice{}).Where("id = ? AND status = ?", device.ID, models.DeviceStatusOnline). Update("status", models.DeviceStatusOffline) if result.Error != nil { return internalError(result.Error) } - if err := insertDeviceStatusEvent(tx, &event); err != nil { - return err + return insertDeviceStatusEvent(tx, event) + }) + if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return changed, ctxErr } - events = append(events, event) - changed += result.RowsAffected + var mysqlErr *mysqlDriver.MySQLError + if errors.As(err, &mysqlErr) && mysqlErr.Number == 3572 { + log.Warnf("agent offline scan skipped: device_id=%d classification=lock_busy", deviceID) + continue + } + return changed, err + } + if event != nil { + changed++ + logDeviceStatusEvent(*event) } - return nil - }) - if err != nil { - return 0, err } - for _, event := range events { - logDeviceStatusEvent(event) - } - return changed, err + return changed, nil } -func lockOfflineTasksNowait(tx *gorm.DB, deviceID uint64) error { - var collectionIDs, purchaseIDs []uint64 +type offlineTasks struct { + CollectionIDs []uint64 + PurchaseRunningIDs []uint64 + PurchaseSubmitStartedIDs []uint64 +} + +func lockOfflineTasksNowait(tx *gorm.DB, deviceID uint64) (offlineTasks, error) { + var tasks offlineTasks locking := clause.Locking{Strength: "UPDATE", Options: "NOWAIT"} if err := tx.Model(&models.CollectionTask{}).Clauses(locking). Where("device_id = ? AND status = ?", deviceID, models.TaskStatusRunning). - Order("id").Pluck("id", &collectionIDs).Error; err != nil { - return internalError(err) + Order("id").Pluck("id", &tasks.CollectionIDs).Error; err != nil { + return tasks, internalError(err) } - if err := tx.Model(&models.PurchaseTask{}).Clauses(locking). + var purchases []struct { + ID uint64 + Status string + } + if err := tx.Model(&models.PurchaseTask{}).Select("id, status").Clauses(locking). Where("device_id = ? AND status IN ?", deviceID, []string{models.PurchaseTaskStatusRunning, models.PurchaseTaskStatusOrderSubmitStarted}). - Order("id").Pluck("id", &purchaseIDs).Error; err != nil { - return internalError(err) + Order("id").Find(&purchases).Error; err != nil { + return tasks, internalError(err) } - return nil + for _, task := range purchases { + if task.Status == models.PurchaseTaskStatusRunning { + tasks.PurchaseRunningIDs = append(tasks.PurchaseRunningIDs, task.ID) + } else { + tasks.PurchaseSubmitStartedIDs = append(tasks.PurchaseSubmitStartedIDs, task.ID) + } + } + return tasks, nil } -func markOfflinePurchaseTasks(tx *gorm.DB, deviceIDs []uint64, now time.Time) (int64, int64, error) { +func markOfflinePurchaseTasks(tx *gorm.DB, tasks offlineTasks, now time.Time) (int64, int64, error) { failed, err := markOfflinePurchaseStatus( - tx, deviceIDs, models.PurchaseTaskStatusRunning, models.PurchaseTaskStatusFailed, + tx, tasks.PurchaseRunningIDs, models.PurchaseTaskStatusRunning, models.PurchaseTaskStatusFailed, "failed", "设备心跳超时,采购任务失败且不会自动重试", now, ) if err != nil { return 0, 0, err } unknown, err := markOfflinePurchaseStatus( - tx, deviceIDs, models.PurchaseTaskStatusOrderSubmitStarted, models.PurchaseTaskStatusOrderResultUnknown, + tx, tasks.PurchaseSubmitStartedIDs, models.PurchaseTaskStatusOrderSubmitStarted, models.PurchaseTaskStatusOrderResultUnknown, "order_result_unknown", "设备心跳超时,订单结果未知,请人工核对,禁止自动重试", now, ) return failed, unknown, err @@ -251,19 +284,15 @@ func markOfflinePurchaseTasks(tx *gorm.DB, deviceIDs []uint64, now time.Time) (i func markOfflinePurchaseStatus( tx *gorm.DB, - deviceIDs []uint64, + taskIDs []uint64, fromStatus string, toStatus string, resultType string, errorMessage string, now time.Time, ) (int64, error) { - var taskIDs []uint64 - if err := tx.Model(&models.PurchaseTask{}). - Where("device_id IN ? AND status = ?", deviceIDs, fromStatus). - Pluck("id", &taskIDs).Error; err != nil { - return 0, internalError(err) - } + // Reuse the IDs and source states from the locking current read. An + // ordinary SELECT here could see an older REPEATABLE READ snapshot. if len(taskIDs) == 0 { return 0, nil } diff --git a/server/app/goauto/device/offline_mysql_test.go b/server/app/goauto/device/offline_mysql_test.go new file mode 100644 index 0000000..c0885bf --- /dev/null +++ b/server/app/goauto/device/offline_mysql_test.go @@ -0,0 +1,432 @@ +package device_test + +import ( + "context" + "errors" + "fmt" + "strings" + "sync" + "testing" + "time" + + "go-admin/app/goauto/device" + "go-admin/app/goauto/models" + "go-admin/app/goauto/purchase" + "go-admin/app/goauto/task" + + driver "github.com/go-sql-driver/mysql" + "github.com/google/uuid" + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +type offlineSeed func(*testing.T) (*device.Service, device.RegisterResponse, string) + +func offlineCollection(t *testing.T, db *gorm.DB, deviceID uint64, status string) models.CollectionTask { + t.Helper() + product := models.PDDProduct{GoodsID: strings.ReplaceAll(uuid.NewString(), "-", ""), URL: "https://example.invalid/test"} + rule := models.CollectionRule{Name: "test", ContentJSON: `{"steps":[]}`} + if err := db.Create(&product).Error; err != nil { + t.Fatal("synthetic product failed") + } + if err := db.Create(&rule).Error; err != nil { + t.Fatal("synthetic rule failed") + } + task := models.CollectionTask{PDDProductID: &product.ID, RuleID: rule.ID, DeviceID: &deviceID, Status: status, URLSnapshot: product.URL, GoodsIDSnapshot: product.GoodsID, RuleSnapshot: rule.ContentJSON} + if err := db.Create(&task).Error; err != nil { + t.Fatal("synthetic collection task failed") + } + return task +} + +func offlinePurchase(t *testing.T, db *gorm.DB, deviceID uint64, status string) models.PurchaseTask { + t.Helper() + product := models.PDDProduct{GoodsID: strings.ReplaceAll(uuid.NewString(), "-", ""), URL: "https://example.invalid/test"} + if err := db.Create(&product).Error; err != nil { + t.Fatal("synthetic product failed") + } + task := models.PurchaseTask{PDDProductID: product.ID, DeviceID: &deviceID, ExecutionMode: models.PurchaseExecutionModeRehearsal, Status: status, PDDURLSnapshot: product.URL, PDDGoodsIDSnapshot: product.GoodsID, Quantity: 1, ReferenceUnitPriceCent: 100, MinUnitPriceCent: 20, MaxUnitPriceCent: 150, Currency: "CNY", RuleType: "pddPurchase", RuleSchemaVersion: 1, RequiredCapabilitiesJSON: `[]`, RuleSnapshot: `{}`, CreateRequestID: uuid.NewString()} + if err := db.Create(&task).Error; err != nil { + t.Fatal("synthetic purchase task failed") + } + return task +} + +func requireOfflineDevice(t *testing.T, db *gorm.DB, id uint64, status string, events int64) { + t.Helper() + var row models.AgentDevice + if db.First(&row, id).Error != nil || row.Status != status { + t.Fatalf("device %d status=%s want=%s", id, row.Status, status) + } + var count int64 + if db.Model(&models.AgentDeviceStatusEvent{}).Where("device_id = ?", id).Count(&count).Error != nil || count != events { + t.Fatalf("device %d events=%d want=%d", id, count, events) + } +} + +func runOfflineV4MySQLTests(t *testing.T, db *gorm.DB, seed offlineSeed) { + for _, busyKind := range []string{"device", "collection", "purchase"} { + t.Run("busy_"+busyKind+"_does_not_block_other_devices", func(t *testing.T) { + s, first, _ := seed(t) + _, busy, _ := seed(t) + _, last, _ := seed(t) + var lockedModel any + var lockedID uint64 + switch busyKind { + case "device": + lockedModel = &models.AgentDevice{} + lockedID = busy.DeviceID + case "collection": + task := offlineCollection(t, db, busy.DeviceID, models.TaskStatusRunning) + lockedModel = &models.CollectionTask{} + lockedID = task.ID + case "purchase": + task := offlinePurchase(t, db, busy.DeviceID, models.PurchaseTaskStatusRunning) + lockedModel = &models.PurchaseTask{} + lockedID = task.ID + } + holder := db.Begin() + if holder.Error != nil { + t.Fatal("begin failed") + } + defer holder.Rollback() + if holder.Clauses(clause.Locking{Strength: "UPDATE"}).First(lockedModel, lockedID).Error != nil { + t.Fatal("synthetic lock failed") + } + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + n, err := s.MarkStaleDevicesOffline(ctx, device.DefaultOfflineThreshold) + if err != nil || n != 2 { + t.Fatalf("busy %s prevented unrelated devices: changed=%d failed=%v", busyKind, n, err != nil) + } + requireOfflineDevice(t, db, first.DeviceID, models.DeviceStatusOffline, 1) + requireOfflineDevice(t, db, busy.DeviceID, models.DeviceStatusOnline, 0) + requireOfflineDevice(t, db, last.DeviceID, models.DeviceStatusOffline, 1) + if holder.Rollback().Error != nil { + t.Fatal("release failed") + } + n, err = s.MarkStaleDevicesOffline(context.Background(), device.DefaultOfflineThreshold) + if err != nil || n != 1 { + t.Fatal("busy device not retried on next scan") + } + }) + } + + for _, target := range []string{models.PurchaseTaskStatusRunning, models.PurchaseTaskStatusOrderSubmitStarted} { + t.Run("second_device_current_read_"+target, func(t *testing.T) { + s, first, _ := seed(t) + _, second, _ := seed(t) + offlinePurchase(t, db, first.DeviceID, models.PurchaseTaskStatusRunning) + task := offlinePurchase(t, db, second.DeviceID, models.PurchaseTaskStatusPending) + attempt := models.PurchaseTaskAttempt{TaskID: task.ID, AttemptID: uuid.NewString(), AttemptNumber: 1, Phase: models.PurchaseAttemptPhasePurchase, Status: models.PurchaseAttemptStatusRunning, DeviceID: &second.DeviceID, RuleSnapshotHash: "test", SpecDecisionSnapshot: `{}`} + if db.Create(&attempt).Error != nil { + t.Fatal("synthetic attempt failed") + } + callback := "v4_advance_second_device" + if db.Callback().Create().After("gorm:create").Register(callback, func(tx *gorm.DB) { + event, ok := tx.Statement.Dest.(*models.AgentDeviceStatusEvent) + if !ok || event.DeviceID != first.DeviceID { + return + } + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + if err := db.WithContext(ctx).Session(&gorm.Session{SkipHooks: true}).Model(&models.PurchaseTask{}).Where("id = ?", task.ID).Updates(map[string]any{"status": target, "device_run_slot": 1}).Error; err != nil { + tx.AddError(errors.New("synthetic second device advancement failed")) + } + }) != nil { + t.Fatal("register advancement failed") + } + defer db.Callback().Create().Remove(callback) + n, err := s.MarkStaleDevicesOffline(context.Background(), device.DefaultOfflineThreshold) + if err != nil || n != 2 { + t.Fatalf("scan changed=%d failed=%v", n, err != nil) + } + var stored models.PurchaseTask + var storedAttempt models.PurchaseTaskAttempt + var event models.AgentDeviceStatusEvent + if db.First(&stored, task.ID).Error != nil || db.First(&storedAttempt, attempt.ID).Error != nil || db.Where("device_id = ?", second.DeviceID).First(&event).Error != nil { + t.Fatal("read outcome failed") + } + expected := models.PurchaseTaskStatusFailed + failed, unknown := int64(1), int64(0) + if target == models.PurchaseTaskStatusOrderSubmitStarted { + expected = models.PurchaseTaskStatusOrderResultUnknown + failed, unknown = 0, 1 + } + if stored.Status != expected || storedAttempt.Status != models.PurchaseAttemptStatusFailed || event.FailedTaskCount != failed || event.OrderResultUnknownCount != unknown { + t.Fatalf("stale task snapshot: status=%s attempt=%s failed=%d unknown=%d", stored.Status, storedAttempt.Status, event.FailedTaskCount, event.OrderResultUnknownCount) + } + }) + } + + t.Run("event_failure_preserves_prior_device_commit", func(t *testing.T) { + s, first, _ := seed(t) + _, second, _ := seed(t) + _, last, _ := seed(t) + task := offlineCollection(t, db, second.DeviceID, models.TaskStatusRunning) + statement := fmt.Sprintf("CREATE TRIGGER v4_reject_event BEFORE INSERT ON agent_device_status_event FOR EACH ROW BEGIN IF NEW.device_id = %d THEN SIGNAL SQLSTATE '45000' SET MESSAGE_TEXT = 'synthetic_failure'; END IF; END", second.DeviceID) + if db.Exec(statement).Error != nil { + t.Fatal("test trigger failed") + } + defer db.Exec("DROP TRIGGER IF EXISTS v4_reject_event") + n, err := s.MarkStaleDevicesOffline(context.Background(), device.DefaultOfflineThreshold) + if err == nil || n != 1 { + t.Fatalf("must return earlier committed count on event failure: changed=%d failed=%v", n, err != nil) + } + requireOfflineDevice(t, db, first.DeviceID, models.DeviceStatusOffline, 1) + requireOfflineDevice(t, db, second.DeviceID, models.DeviceStatusOnline, 0) + requireOfflineDevice(t, db, last.DeviceID, models.DeviceStatusOnline, 0) + var stored models.CollectionTask + if db.First(&stored, task.ID).Error != nil || stored.Status != models.TaskStatusRunning { + t.Fatal("current task escaped rollback") + } + if db.Exec("DROP TRIGGER v4_reject_event").Error != nil { + t.Fatal("remove trigger failed") + } + n, err = s.MarkStaleDevicesOffline(context.Background(), device.DefaultOfflineThreshold) + if err != nil || n != 2 { + t.Fatal("next scan retry failed") + } + }) + + for _, kind := range []string{"1213", "1205", "text_3572", "cancel", "cancel_with_3572"} { + t.Run("non_contention_error_"+kind, func(t *testing.T) { + s, first, _ := seed(t) + _, second, _ := seed(t) + _, last, _ := seed(t) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + callback := "v4_inject_non_contention" + if db.Callback().Query().After("gorm:query").Register(callback, func(tx *gorm.DB) { + row, ok := tx.Statement.Dest.(*models.AgentDevice) + if !ok || row.ID != second.DeviceID { + return + } + if _, locking := tx.Statement.Clauses["FOR"]; !locking { + return + } + switch kind { + case "1213": + tx.AddError(&driver.MySQLError{Number: 1213, Message: "synthetic"}) + case "1205": + tx.AddError(&driver.MySQLError{Number: 1205, Message: "synthetic"}) + case "text_3572": + tx.AddError(errors.New("synthetic 3572")) + case "cancel": + cancel() + tx.AddError(context.Canceled) + case "cancel_with_3572": + cancel() + tx.AddError(&driver.MySQLError{Number: 3572, Message: "synthetic"}) + } + }) != nil { + t.Fatal("register injection failed") + } + defer db.Callback().Query().Remove(callback) + n, err := s.MarkStaleDevicesOffline(ctx, device.DefaultOfflineThreshold) + if err == nil || n != 1 { + t.Fatalf("non-contention error swallowed or lost prior commit: changed=%d failed=%v", n, err != nil) + } + if (kind == "cancel" || kind == "cancel_with_3572") && !errors.Is(err, context.Canceled) { + t.Fatal("cancellation not preserved") + } + requireOfflineDevice(t, db, first.DeviceID, models.DeviceStatusOffline, 1) + requireOfflineDevice(t, db, second.DeviceID, models.DeviceStatusOnline, 0) + requireOfflineDevice(t, db, last.DeviceID, models.DeviceStatusOnline, 0) + }) + } + + for _, domain := range []string{"collection_claim", "purchase_replayed_claim"} { + t.Run("actual_"+domain+"_does_not_deadlock_scan", func(t *testing.T) { + s, d, token := seed(t) + _, other, _ := seed(t) + type claimMarker struct{} + locked, release := make(chan struct{}), make(chan struct{}) + var once sync.Once + callback := "v4_actual_claim_barrier" + if db.Callback().Query().After("gorm:query").Register(callback, func(tx *gorm.DB) { + if tx.Statement.Context.Value(claimMarker{}) != true { + return + } + lockTable := "agent_device" + if domain == "purchase_replayed_claim" { + lockTable = "purchase_task" + } + if tx.Statement.Table != lockTable { + return + } + if _, ok := tx.Statement.Clauses["FOR"]; !ok { + return + } + once.Do(func() { close(locked); <-release }) + }) != nil { + t.Fatal("register actual claim barrier failed") + } + defer db.Callback().Query().Remove(callback) + ctx, cancel := context.WithTimeout(context.WithValue(context.Background(), claimMarker{}, true), 5*time.Second) + defer cancel() + done := make(chan error, 1) + var taskID uint64 + if domain == "collection_claim" { + record := offlineCollection(t, db, d.DeviceID, models.TaskStatusPending) + taskID = record.ID + service := task.NewService(db) + service.Now = s.Now + go func() { + _, err := service.Claim(ctx, record.ID, task.ActionRequest{RequestID: uuid.NewString()}, token) + done <- err + }() + } else { + record := offlinePurchase(t, db, d.DeviceID, models.PurchaseTaskStatusRunning) + taskID = record.ID + requestID := uuid.NewString() + if db.Session(&gorm.Session{SkipHooks: true}).Model(&models.PurchaseTask{}).Where("id = ?", record.ID).Update("claim_request_id", requestID).Error != nil { + t.Fatal("replay fixture failed") + } + service := purchase.NewService(db) + service.Now = s.Now + go func() { + _, err := service.Claim(ctx, record.ID, purchase.ActionRequest{RequestID: requestID}, token) + done <- err + }() + } + select { + case <-locked: + case <-ctx.Done(): + close(release) + t.Fatal("actual claim lock not reached") + } + scanCtx, stopScan := context.WithTimeout(context.Background(), 2*time.Second) + n, err := s.MarkStaleDevicesOffline(scanCtx, device.DefaultOfflineThreshold) + stopScan() + close(release) + claimErr := <-done + if err != nil || claimErr != nil || n != 1 { + t.Fatalf("claim/scan conflict: changed=%d scanFailed=%v claimFailed=%v", n, err != nil, claimErr != nil) + } + requireOfflineDevice(t, db, d.DeviceID, models.DeviceStatusOnline, 0) + requireOfflineDevice(t, db, other.DeviceID, models.DeviceStatusOffline, 1) + if domain == "collection_claim" { + var stored models.CollectionTask + if db.First(&stored, taskID).Error != nil || stored.Status != models.TaskStatusPending || stored.LeaseExpiresAt == nil || stored.ClaimRequestID == nil { + t.Fatal("claim partially persisted") + } + } + }) + } + + t.Run("real_heartbeat_waiting_for_scan_resumes_once", func(t *testing.T) { + s, d, token := seed(t) + type scanMarker struct{} + type heartbeatMarker struct{} + locked, release, heartbeatStarted := make(chan struct{}), make(chan struct{}), make(chan struct{}) + scanCallback, heartbeatCallback := "v4_scan_heartbeat_barrier", "v4_heartbeat_scan_barrier" + if db.Callback().Query().After("gorm:query").Register(scanCallback, func(tx *gorm.DB) { + if tx.Statement.Context.Value(scanMarker{}) == true && tx.Statement.Table == "agent_device" { + if _, ok := tx.Statement.Clauses["FOR"]; ok { + close(locked) + <-release + } + } + }) != nil { + t.Fatal("register scan barrier failed") + } + defer db.Callback().Query().Remove(scanCallback) + if db.Callback().Query().Before("gorm:query").Register(heartbeatCallback, func(tx *gorm.DB) { + if tx.Statement.Context.Value(heartbeatMarker{}) == true && tx.Statement.Table == "agent_device" { + if _, ok := tx.Statement.Clauses["FOR"]; ok { + close(heartbeatStarted) + } + } + }) != nil { + t.Fatal("register heartbeat barrier failed") + } + defer db.Callback().Query().Remove(heartbeatCallback) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + scanDone, heartbeatDone := make(chan error, 1), make(chan error, 1) + go func() { + n, err := s.MarkStaleDevicesOffline(context.WithValue(ctx, scanMarker{}, true), device.DefaultOfflineThreshold) + if err == nil && n != 1 { + err = errors.New("wrong scan count") + } + scanDone <- err + }() + select { + case <-locked: + case <-ctx.Done(): + close(release) + t.Fatal("scan lock not reached") + } + go func() { + _, err := s.Heartbeat(context.WithValue(ctx, heartbeatMarker{}, true), device.HeartbeatRequest{RequestID: uuid.NewString()}, token) + heartbeatDone <- err + }() + select { + case <-heartbeatStarted: + case <-ctx.Done(): + close(release) + t.Fatal("heartbeat not started") + } + close(release) + if err := <-scanDone; err != nil { + t.Fatal("scan failed") + } + if err := <-heartbeatDone; err != nil { + t.Fatal("heartbeat failed") + } + requireOfflineDevice(t, db, d.DeviceID, models.DeviceStatusOnline, 2) + }) + + t.Run("explain_actual_lock_queries", func(t *testing.T) { + s, d, _ := seed(t) + offlineCollection(t, db, d.DeviceID, models.TaskStatusRunning) + offlinePurchase(t, db, d.DeviceID, models.PurchaseTaskStatusRunning) + type captured struct { + sql string + vars []any + } + queries := map[string]captured{} + callback := "v4_capture_lock_sql" + if db.Callback().Query().After("gorm:query").Register(callback, func(tx *gorm.DB) { + if _, ok := tx.Statement.Clauses["FOR"]; ok { + queries[tx.Statement.Table] = captured{tx.Statement.SQL.String(), append([]any(nil), tx.Statement.Vars...)} + } + }) != nil { + t.Fatal("capture SQL failed") + } + n, err := s.MarkStaleDevicesOffline(context.Background(), device.DefaultOfflineThreshold) + db.Callback().Query().Remove(callback) + if err != nil || n != 1 { + t.Fatal("capture scan failed") + } + for _, table := range []string{"agent_device", "collection_task", "purchase_task"} { + query, ok := queries[table] + if !ok { + t.Fatalf("missing production lock query %s", table) + } + var plan []struct { + Type string + Key *string + PossibleKeys *string + Rows int64 + Extra string + } + if db.Raw("EXPLAIN "+query.sql, query.vars...).Scan(&plan).Error != nil { + t.Fatal("EXPLAIN failed") + } + t.Logf("synthetic fixture only table=%s sql=%s", table, query.sql) + for _, row := range plan { + key := "NULL" + if row.Key != nil { + key = *row.Key + } + possible := "NULL" + if row.PossibleKeys != nil { + possible = *row.PossibleKeys + } + t.Logf("type=%s key=%s possible_keys=%s estimated_rows=%d extra=%s", row.Type, key, possible, row.Rows, row.Extra) + } + } + }) +}