memory_test.go 3.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151
  1. package repository
  2. import (
  3. "sync"
  4. "testing"
  5. "time"
  6. "job-cheng-xing/model"
  7. )
  8. func strPtr(s string) *string { return &s }
  9. func TestMemoryRepo_CreateAndFind(t *testing.T) {
  10. repo := NewMemoryRepo()
  11. order := &model.Order{
  12. ID: "ord_001",
  13. Status: model.StatusPending,
  14. ServiceTime: time.Now(),
  15. Duration: 120,
  16. Address: "北京朝阳",
  17. CreatedAt: time.Now(),
  18. UpdatedAt: time.Now(),
  19. }
  20. err := repo.Create(order)
  21. if err != nil {
  22. t.Fatalf("Create failed: %v", err)
  23. }
  24. found, err := repo.FindByID("ord_001")
  25. if err != nil {
  26. t.Fatalf("FindByID failed: %v", err)
  27. }
  28. if found.ID != "ord_001" {
  29. t.Errorf("expected ord_001, got %s", found.ID)
  30. }
  31. }
  32. func TestMemoryRepo_FindByID_NotFound(t *testing.T) {
  33. repo := NewMemoryRepo()
  34. _, err := repo.FindByID("nonexistent")
  35. if err != model.ErrNotFound {
  36. t.Errorf("expected ErrNotFound, got %v", err)
  37. }
  38. }
  39. func TestMemoryRepo_Update_CAS_Success(t *testing.T) {
  40. repo := NewMemoryRepo()
  41. order := &model.Order{
  42. ID: "ord_001",
  43. Status: model.StatusPending,
  44. ServiceTime: time.Now(),
  45. Duration: 120,
  46. Address: "北京朝阳",
  47. CreatedAt: time.Now(),
  48. UpdatedAt: time.Now(),
  49. }
  50. repo.Create(order)
  51. newOrder := &model.Order{
  52. Status: model.StatusAccepted,
  53. ProviderID: strPtr("prov_001"),
  54. UpdatedAt: time.Now(),
  55. }
  56. err := repo.Update("ord_001", model.StatusPending, newOrder)
  57. if err != nil {
  58. t.Fatalf("Update failed: %v", err)
  59. }
  60. found, _ := repo.FindByID("ord_001")
  61. if found.Status != model.StatusAccepted {
  62. t.Errorf("expected accepted, got %s", found.Status)
  63. }
  64. if *found.ProviderID != "prov_001" {
  65. t.Errorf("expected prov_001, got %s", *found.ProviderID)
  66. }
  67. }
  68. func TestMemoryRepo_Update_CAS_Fail_WrongStatus(t *testing.T) {
  69. repo := NewMemoryRepo()
  70. order := &model.Order{
  71. ID: "ord_001",
  72. Status: model.StatusPending,
  73. }
  74. repo.Create(order)
  75. newOrder := &model.Order{Status: model.StatusCanceled}
  76. err := repo.Update("ord_001", model.StatusAccepted, newOrder)
  77. if err != model.ErrStatusConflict {
  78. t.Errorf("expected ErrStatusConflict, got %v", err)
  79. }
  80. }
  81. func TestMemoryRepo_ConcurrentAccept(t *testing.T) {
  82. repo := NewMemoryRepo()
  83. order := &model.Order{
  84. ID: "ord_001",
  85. Status: model.StatusPending,
  86. ServiceTime: time.Now(),
  87. Duration: 120,
  88. Address: "北京朝阳",
  89. CreatedAt: time.Now(),
  90. UpdatedAt: time.Now(),
  91. }
  92. repo.Create(order)
  93. const numGoroutines = 50
  94. var wg sync.WaitGroup
  95. results := make(chan error, numGoroutines)
  96. for i := 0; i < numGoroutines; i++ {
  97. wg.Add(1)
  98. go func(idx int) {
  99. defer wg.Done()
  100. pid := "prov_" + string(rune('a'+idx%26)) + string(rune('0'+idx/26))
  101. newOrder := &model.Order{
  102. Status: model.StatusAccepted,
  103. ProviderID: &pid,
  104. UpdatedAt: time.Now(),
  105. }
  106. results <- repo.Update("ord_001", model.StatusPending, newOrder)
  107. }(i)
  108. }
  109. wg.Wait()
  110. close(results)
  111. successCount := 0
  112. failCount := 0
  113. for err := range results {
  114. if err == nil {
  115. successCount++
  116. } else if err == model.ErrStatusConflict {
  117. failCount++
  118. } else {
  119. t.Errorf("unexpected error: %v", err)
  120. }
  121. }
  122. if successCount != 1 {
  123. t.Errorf("expected exactly 1 success, got %d", successCount)
  124. }
  125. if failCount != numGoroutines-1 {
  126. t.Errorf("expected %d failures, got %d", numGoroutines-1, failCount)
  127. }
  128. found, _ := repo.FindByID("ord_001")
  129. if found.ProviderID == nil {
  130. t.Fatal("expected ProviderID to be set")
  131. }
  132. t.Logf("winner provider: %s", *found.ProviderID)
  133. }