client_hwid_tx_test.go 6.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216
  1. package service
  2. import (
  3. "errors"
  4. "fmt"
  5. "path/filepath"
  6. "testing"
  7. "time"
  8. "github.com/mhsanaei/3x-ui/v3/internal/database"
  9. "github.com/mhsanaei/3x-ui/v3/internal/database/dbtest"
  10. "github.com/mhsanaei/3x-ui/v3/internal/database/model"
  11. "gorm.io/gorm"
  12. )
  13. var errInjectedHwidDelete = errors.New("injected client_hwids delete failure")
  14. func failHwidDeletes(t *testing.T, db *gorm.DB) {
  15. t.Helper()
  16. if err := db.Callback().Delete().Before("gorm:delete").Register("t:hwid:fail", func(tx *gorm.DB) {
  17. if tx.Statement != nil && tx.Statement.Table == "client_hwids" {
  18. _ = tx.AddError(errInjectedHwidDelete)
  19. }
  20. }); err != nil {
  21. t.Fatalf("register callback: %v", err)
  22. }
  23. t.Cleanup(func() {
  24. if err := db.Callback().Delete().Remove("t:hwid:fail"); err != nil {
  25. t.Fatalf("remove callback: %v", err)
  26. }
  27. })
  28. }
  29. func seedHwids(t *testing.T, db *gorm.DB, subID string, n int) {
  30. t.Helper()
  31. now := time.Now().UnixMilli()
  32. rows := make([]model.ClientHwid, 0, n)
  33. for i := range n {
  34. rows = append(rows, model.ClientHwid{
  35. SubID: subID, HwidHash: fmt.Sprintf("%s-hash-%d", subID, i),
  36. FirstSeen: now, LastSeen: now + int64(i),
  37. })
  38. }
  39. if err := db.Create(&rows).Error; err != nil {
  40. t.Fatalf("seed client_hwids for %q: %v", subID, err)
  41. }
  42. }
  43. func assertHwidState(t *testing.T, db *gorm.DB, email string, limit, devices int) {
  44. t.Helper()
  45. var rec model.ClientRecord
  46. if err := db.Where("email = ?", email).First(&rec).Error; err != nil {
  47. t.Fatalf("reload client: %v", err)
  48. }
  49. if rec.LimitHwid != limit {
  50. t.Fatalf("limit_hwid = %d, want %d", rec.LimitHwid, limit)
  51. }
  52. var n int64
  53. if err := db.Model(&model.ClientHwid{}).Where("sub_id = ?", rec.SubID).Count(&n).Error; err != nil {
  54. t.Fatalf("count client_hwids: %v", err)
  55. }
  56. if n != int64(devices) {
  57. t.Fatalf("client_hwids = %d, want %d", n, devices)
  58. }
  59. }
  60. func TestSetClientLimitHwidRollsBackFailedTrim(t *testing.T) {
  61. initClientHwidTestDB(t)
  62. db := database.GetDB()
  63. rec := seedHwidClient(t, 5)
  64. seedHwids(t, db, rec.SubID, 3)
  65. failHwidDeletes(t, db)
  66. err := (&ClientService{}).setClientLimitHwidByEmail(rec.Email, 1)
  67. if !errors.Is(err, errInjectedHwidDelete) {
  68. t.Fatalf("want errInjectedHwidDelete, got: %v", err)
  69. }
  70. assertHwidState(t, db, rec.Email, 5, 3)
  71. }
  72. func TestClientHwidTxRejectsUnserializedHandle(t *testing.T) {
  73. initClientHwidTestDB(t)
  74. db := database.GetDB()
  75. rec := seedHwidClient(t, 5)
  76. svc := &ClientService{}
  77. if err := svc.setClientLimitHwidByEmailTx(db, rec.Email, 1); !errors.Is(err, errClientHwidWriteNotSerialized) {
  78. t.Fatalf("bare handle error = %v, want errClientHwidWriteNotSerialized", err)
  79. }
  80. if err := runSerializedTx(func(tx *gorm.DB) error {
  81. return svc.setClientLimitHwidByEmailTx(tx, rec.Email, 1)
  82. }); err != nil {
  83. t.Fatalf("serialized update: %v", err)
  84. }
  85. assertHwidState(t, db, rec.Email, 1, 0)
  86. }
  87. func TestBulkAdjustHwidRollsBackFailedTrim(t *testing.T) {
  88. setupBulkDB(t)
  89. db := database.GetDB()
  90. rec := &model.ClientRecord{Email: "bulk-hwid@x", SubID: "bulk-sub", Enable: true, LimitHwid: 5}
  91. if err := db.Create(rec).Error; err != nil {
  92. t.Fatalf("seed client: %v", err)
  93. }
  94. seedHwids(t, db, rec.SubID, 3)
  95. failHwidDeletes(t, db)
  96. limit := 1
  97. res, _, err := (&ClientService{}).BulkAdjust(&InboundService{}, []string{rec.Email}, 0, 0, "", &limit, "")
  98. if err != nil {
  99. t.Fatalf("BulkAdjust: %v", err)
  100. }
  101. if len(res.Skipped) != 1 || res.Skipped[0].Reason != errInjectedHwidDelete.Error() {
  102. t.Fatalf("skipped = %+v, want injected failure", res.Skipped)
  103. }
  104. assertHwidState(t, db, rec.Email, 5, 3)
  105. }
  106. func TestBulkCreateWithdrawsTombstoneWhenHwidTrimFails(t *testing.T) {
  107. setupBulkDB(t)
  108. StartTrafficWriter()
  109. t.Cleanup(StopTrafficWriter)
  110. db := database.GetDB()
  111. const email = "reborn-bulk@x"
  112. const subID = "reborn-bulk-sub"
  113. tombstoneClientEmail(email)
  114. t.Cleanup(func() { withdrawClientTombstones(email) })
  115. seedHwids(t, db, subID, 3)
  116. failHwidDeletes(t, db)
  117. ib := mkInbound(t, 30441, model.VLESS, `{"clients":[]}`)
  118. res, _, err := (&ClientService{}).BulkCreate(&InboundService{}, []ClientCreatePayload{{
  119. Client: model.Client{
  120. Email: email, SubID: subID, ID: "aaaaaaaa-bbbb-4ccc-8ddd-eeeeeeeeeeee", Enable: true,
  121. },
  122. InboundIds: []int{ib.Id}, LimitHwid: 1,
  123. }})
  124. if err != nil {
  125. t.Fatalf("BulkCreate: %v", err)
  126. }
  127. if len(res.Skipped) != 1 || res.Skipped[0].Reason != errInjectedHwidDelete.Error() {
  128. t.Fatalf("skipped = %+v, want injected HWID failure", res.Skipped)
  129. }
  130. if isClientEmailTombstoned(email) {
  131. t.Fatal("live bulk-created client retained a delete tombstone")
  132. }
  133. }
  134. func TestSetClientLimitHwidIsSerializedWithSyncInbound(t *testing.T) {
  135. db := durablePostgresDB(t)
  136. if err := db.Exec("TRUNCATE client_hwids, clients RESTART IDENTITY CASCADE").Error; err != nil {
  137. t.Fatalf("reset tables: %v", err)
  138. }
  139. rec := seedHwidClient(t, 5)
  140. seedHwids(t, db, rec.SubID, 3)
  141. StartTrafficWriter()
  142. t.Cleanup(StopTrafficWriter)
  143. read := make(chan struct{})
  144. release := make(chan struct{})
  145. staleDone := make(chan error, 1)
  146. go func() {
  147. staleDone <- runSerializedTx(func(tx *gorm.DB) error {
  148. var stale model.ClientRecord
  149. if err := tx.Where("email = ?", rec.Email).First(&stale).Error; err != nil {
  150. return err
  151. }
  152. close(read)
  153. <-release
  154. return tx.Save(&stale).Error
  155. })
  156. }()
  157. <-read
  158. limitDone := make(chan error, 1)
  159. go func() { limitDone <- (&ClientService{}).setClientLimitHwidByEmail(rec.Email, 1) }()
  160. time.Sleep(100 * time.Millisecond)
  161. close(release)
  162. if err := <-staleDone; err != nil {
  163. t.Fatalf("stale SyncInbound write: %v", err)
  164. }
  165. if err := <-limitDone; err != nil {
  166. t.Fatalf("set limit: %v", err)
  167. }
  168. assertHwidState(t, db, rec.Email, 1, 1)
  169. }
  170. func BenchmarkSetClientLimitHwidSerialized(b *testing.B) {
  171. dbDir := b.TempDir()
  172. b.Setenv("XUI_DB_FOLDER", dbDir)
  173. dbtest.InitDB(b, filepath.Join(dbDir, "x-ui.db"))
  174. StartTrafficWriter()
  175. b.Cleanup(StopTrafficWriter)
  176. db := database.GetDB()
  177. emails := make([]string, 100)
  178. for i := range emails {
  179. emails[i] = fmt.Sprintf("bench-%03d@x", i)
  180. rec := &model.ClientRecord{Email: emails[i], SubID: fmt.Sprintf("bench-sub-%03d", i), Enable: true}
  181. if err := db.Create(rec).Error; err != nil {
  182. b.Fatalf("seed client: %v", err)
  183. }
  184. }
  185. svc := &ClientService{}
  186. for _, count := range []int{1, 100} {
  187. b.Run(fmt.Sprintf("clients_%d", count), func(b *testing.B) {
  188. for range b.N {
  189. for i := range count {
  190. if err := svc.setClientLimitHwidByEmail(emails[i], 2); err != nil {
  191. b.Fatal(err)
  192. }
  193. }
  194. }
  195. })
  196. }
  197. }