1
0

client_hwid_tx_test.go 6.4 KB

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