node_bulk_dispatch_test.go 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392
  1. package service
  2. import (
  3. "context"
  4. "errors"
  5. "fmt"
  6. "sync/atomic"
  7. "testing"
  8. "github.com/google/uuid"
  9. "gorm.io/gorm"
  10. "github.com/mhsanaei/3x-ui/v3/internal/database"
  11. "github.com/mhsanaei/3x-ui/v3/internal/database/model"
  12. "github.com/mhsanaei/3x-ui/v3/internal/web/runtime"
  13. )
  14. // fakeNodeRuntime is a runtime.Runtime stub that counts the per-client dispatch
  15. // calls so a test can assert a bulk op does NOT stream one RPC per client.
  16. type fakeNodeRuntime struct {
  17. addInbound atomic.Int32
  18. delInbound atomic.Int32
  19. addClient atomic.Int32
  20. deleteClient atomic.Int32
  21. deleteUser atomic.Int32
  22. updateInbound atomic.Int32
  23. updateUser atomic.Int32
  24. }
  25. func (f *fakeNodeRuntime) Name() string { return "fake-node" }
  26. func (f *fakeNodeRuntime) AddInbound(context.Context, *model.Inbound) error {
  27. f.addInbound.Add(1)
  28. return nil
  29. }
  30. func (f *fakeNodeRuntime) DelInbound(context.Context, *model.Inbound) error {
  31. f.delInbound.Add(1)
  32. return nil
  33. }
  34. func (f *fakeNodeRuntime) UpdateInbound(context.Context, *model.Inbound, *model.Inbound) error {
  35. f.updateInbound.Add(1)
  36. return nil
  37. }
  38. func (f *fakeNodeRuntime) AddUser(context.Context, *model.Inbound, map[string]any) error { return nil }
  39. func (f *fakeNodeRuntime) RemoveUser(context.Context, *model.Inbound, string) error { return nil }
  40. func (f *fakeNodeRuntime) UpdateUser(context.Context, *model.Inbound, string, model.Client) error {
  41. f.updateUser.Add(1)
  42. return nil
  43. }
  44. func (f *fakeNodeRuntime) DeleteUser(context.Context, *model.Inbound, string) error {
  45. f.deleteUser.Add(1)
  46. return nil
  47. }
  48. func (f *fakeNodeRuntime) DeleteClient(context.Context, string) error {
  49. f.deleteClient.Add(1)
  50. return nil
  51. }
  52. func (f *fakeNodeRuntime) AddClient(context.Context, *model.Inbound, model.Client) error {
  53. f.addClient.Add(1)
  54. return nil
  55. }
  56. func (f *fakeNodeRuntime) RestartXray(context.Context) error { return nil }
  57. func (f *fakeNodeRuntime) ResetClientTraffic(context.Context, *model.Inbound, string) error {
  58. return nil
  59. }
  60. func (f *fakeNodeRuntime) ResetInboundTraffic(context.Context, *model.Inbound) error { return nil }
  61. func (f *fakeNodeRuntime) ResetAllTraffics(context.Context) error { return nil }
  62. // setupNodeRuntime wires an online node + a fake runtime override and returns the
  63. // node id and the fake so a test can drive the service node-dispatch path without
  64. // a network node.
  65. func setupNodeRuntime(t *testing.T) (int, *fakeNodeRuntime) {
  66. t.Helper()
  67. prev := runtime.GetManager()
  68. mgr := runtime.NewManager(runtime.LocalDeps{APIPort: func() int { return 0 }, SetNeedRestart: func() {}})
  69. runtime.SetManager(mgr)
  70. t.Cleanup(func() { runtime.SetManager(prev) })
  71. node := &model.Node{Name: "n1-" + t.Name(), Address: "127.0.0.1", Port: 2096, ApiToken: "tok", Enable: true, Status: "online"}
  72. if err := database.GetDB().Create(node).Error; err != nil {
  73. t.Fatalf("create node: %v", err)
  74. }
  75. t.Cleanup(func() {
  76. _ = database.GetDB().Where("id = ?", node.Id).Delete(&model.Node{}).Error
  77. })
  78. fake := &fakeNodeRuntime{}
  79. mgr.SetRuntimeOverride(node.Id, fake)
  80. return node.Id, fake
  81. }
  82. func nodeInbound(t *testing.T, nodeID, port int, clients []model.Client) *model.Inbound {
  83. t.Helper()
  84. if clients == nil {
  85. clients = []model.Client{}
  86. }
  87. ib := &model.Inbound{
  88. UserId: 1, NodeID: &nodeID, Tag: fmt.Sprintf("in-%d", port), Enable: true,
  89. Port: port, Protocol: model.VLESS, Settings: clientsSettings(t, clients),
  90. }
  91. if err := database.GetDB().Create(ib).Error; err != nil {
  92. t.Fatalf("create node inbound: %v", err)
  93. }
  94. if err := (&ClientService{}).SyncInbound(nil, ib.Id, clients); err != nil {
  95. t.Fatalf("seed SyncInbound: %v", err)
  96. }
  97. return ib
  98. }
  99. func makeNodeClients(n int) []model.Client {
  100. out := make([]model.Client, n)
  101. for i := range n {
  102. out[i] = model.Client{ID: uuid.NewString(), Email: fmt.Sprintf("nu-%05d@x", i), Enable: true}
  103. }
  104. return out
  105. }
  106. // TestNodeBulk_LargeAddFoldsToDirty: adding more than the threshold of clients to
  107. // an online node inbound must NOT stream one AddClient RPC per client; it marks
  108. // the node dirty so a single reconcile push converges it instead.
  109. func TestNodeBulk_LargeAddFoldsToDirty(t *testing.T) {
  110. setupBulkDB(t)
  111. nodeID, fake := setupNodeRuntime(t)
  112. ib := nodeInbound(t, nodeID, 30001, nil)
  113. svc := &ClientService{}
  114. inboundSvc := &InboundService{}
  115. add := makeNodeClients(nodeBulkPushThreshold + 10)
  116. if _, err := svc.AddInboundClient(inboundSvc, &model.Inbound{Id: ib.Id, Protocol: model.VLESS, Settings: clientsSettings(t, add)}); err != nil {
  117. t.Fatalf("AddInboundClient: %v", err)
  118. }
  119. if got := fake.addClient.Load(); got != 0 {
  120. t.Fatalf("large add streamed %d AddClient RPCs, want 0 (should fold to dirty)", got)
  121. }
  122. if _, _, dirty, _, err := (&NodeService{}).NodeSyncState(nodeID); err != nil {
  123. t.Fatalf("NodeSyncState: %v", err)
  124. } else if !dirty {
  125. t.Fatal("large add must mark the node dirty")
  126. }
  127. }
  128. // TestNodeBulk_SmallAddPushesLive: a small add stays on the live per-client path.
  129. func TestNodeBulk_SmallAddPushesLive(t *testing.T) {
  130. setupBulkDB(t)
  131. nodeID, fake := setupNodeRuntime(t)
  132. ib := nodeInbound(t, nodeID, 30002, nil)
  133. svc := &ClientService{}
  134. inboundSvc := &InboundService{}
  135. const small = 3
  136. add := makeNodeClients(small)
  137. if _, err := svc.AddInboundClient(inboundSvc, &model.Inbound{Id: ib.Id, Protocol: model.VLESS, Settings: clientsSettings(t, add)}); err != nil {
  138. t.Fatalf("AddInboundClient: %v", err)
  139. }
  140. if got := fake.addClient.Load(); got != int32(small) {
  141. t.Fatalf("small add streamed %d AddClient RPCs, want %d", got, small)
  142. }
  143. }
  144. func TestNodeBulkAdjustDoesNotPushBeforeFailedCommit(t *testing.T) {
  145. setupBulkDB(t)
  146. nodeID, fake := setupNodeRuntime(t)
  147. client := model.Client{
  148. ID: uuid.NewString(),
  149. Email: "txfail-adjust@x",
  150. Enable: true,
  151. ExpiryTime: 1_900_000_000_000,
  152. }
  153. nodeInbound(t, nodeID, 30022, []model.Client{client})
  154. db := database.GetDB()
  155. const callbackName = "bulk-adjust:fail-inbound-update"
  156. if err := db.Callback().Update().After("gorm:update").Register(callbackName, func(tx *gorm.DB) {
  157. if tx.Statement != nil && tx.Statement.Table == "inbounds" {
  158. tx.AddError(errors.New("injected bulk-adjust transaction failure"))
  159. }
  160. }); err != nil {
  161. t.Fatalf("register callback: %v", err)
  162. }
  163. t.Cleanup(func() { _ = db.Callback().Update().Remove(callbackName) })
  164. result, _, err := (&ClientService{}).BulkAdjust(&InboundService{}, []string{client.Email}, 1, 0, "")
  165. if err != nil {
  166. t.Fatalf("BulkAdjust: %v", err)
  167. }
  168. if result.Adjusted != 0 || len(result.Skipped) != 1 {
  169. t.Fatalf("BulkAdjust result = %+v, want one skipped client after injected failure", result)
  170. }
  171. if got := fake.updateUser.Load(); got != 0 {
  172. t.Fatalf("failed transaction pushed %d UpdateUser call(s) to the node, want 0", got)
  173. }
  174. }
  175. func TestNodeBulkDeleteDoesNotPushBeforeFailedCommit(t *testing.T) {
  176. setupBulkDB(t)
  177. nodeID, fake := setupNodeRuntime(t)
  178. client := model.Client{ID: uuid.NewString(), Email: "txfail-delete@x", Enable: true}
  179. nodeInbound(t, nodeID, 30023, []model.Client{client})
  180. db := database.GetDB()
  181. const callbackName = "bulk-delete:fail-inbound-update"
  182. if err := db.Callback().Update().After("gorm:update").Register(callbackName, func(tx *gorm.DB) {
  183. if tx.Statement != nil && tx.Statement.Table == "inbounds" {
  184. tx.AddError(errors.New("injected bulk-delete transaction failure"))
  185. }
  186. }); err != nil {
  187. t.Fatalf("register callback: %v", err)
  188. }
  189. t.Cleanup(func() { _ = db.Callback().Update().Remove(callbackName) })
  190. result, _, err := (&ClientService{}).BulkDelete(&InboundService{}, []string{client.Email}, true)
  191. if err != nil {
  192. t.Fatalf("BulkDelete: %v", err)
  193. }
  194. if result.Deleted != 0 || len(result.Skipped) != 1 {
  195. t.Fatalf("BulkDelete result = %+v, want one skipped client after injected failure", result)
  196. }
  197. if got := fake.deleteClient.Load() + fake.deleteUser.Load(); got != 0 {
  198. t.Fatalf("failed transaction pushed %d delete call(s) to the node, want 0", got)
  199. }
  200. }
  201. func TestNodeBulkSmallDeleteRemovesWholeRemoteClient(t *testing.T) {
  202. setupBulkDB(t)
  203. nodeID, fake := setupNodeRuntime(t)
  204. client := model.Client{ID: uuid.NewString(), Email: "full-delete@x", Enable: true}
  205. nodeInbound(t, nodeID, 30024, []model.Client{client})
  206. result, _, err := (&ClientService{}).BulkDelete(&InboundService{}, []string{client.Email}, true)
  207. if err != nil {
  208. t.Fatalf("BulkDelete: %v", err)
  209. }
  210. if result.Deleted != 1 || len(result.Skipped) != 0 {
  211. t.Fatalf("BulkDelete result = %+v, want one deleted client", result)
  212. }
  213. if got := fake.deleteClient.Load(); got != 1 {
  214. t.Fatalf("remote DeleteClient calls = %d, want 1", got)
  215. }
  216. if got := fake.deleteUser.Load(); got != 0 {
  217. t.Fatalf("remote DeleteUser detach calls = %d, want 0 for full deletion", got)
  218. }
  219. }
  220. func TestNodeUpdateInboundClientNoopSkipsRuntimeAndDirty(t *testing.T) {
  221. setupBulkDB(t)
  222. nodeID, fake := setupNodeRuntime(t)
  223. client := model.Client{
  224. ID: uuid.NewString(),
  225. Email: "noop@x",
  226. SubID: "sub-noop",
  227. Enable: true,
  228. CreatedAt: 111,
  229. UpdatedAt: 222,
  230. }
  231. ib := nodeInbound(t, nodeID, 30020, []model.Client{client})
  232. svc := &ClientService{}
  233. inboundSvc := &InboundService{}
  234. if _, err := svc.UpdateInboundClient(inboundSvc, &model.Inbound{
  235. Id: ib.Id,
  236. Protocol: model.VLESS,
  237. Settings: clientsSettings(t, []model.Client{client}),
  238. }, client.Email); err != nil {
  239. t.Fatalf("UpdateInboundClient: %v", err)
  240. }
  241. if got := fake.updateUser.Load(); got != 0 {
  242. t.Fatalf("no-op update streamed %d UpdateUser RPCs, want 0", got)
  243. }
  244. if _, _, dirty, _, err := (&NodeService{}).NodeSyncState(nodeID); err != nil {
  245. t.Fatalf("NodeSyncState: %v", err)
  246. } else if dirty {
  247. t.Fatal("no-op update must not mark the node dirty")
  248. }
  249. reloaded, err := inboundSvc.GetInbound(ib.Id)
  250. if err != nil {
  251. t.Fatalf("GetInbound: %v", err)
  252. }
  253. if reloaded.Settings != ib.Settings {
  254. t.Fatal("no-op update rewrote inbound settings")
  255. }
  256. }
  257. func TestNodeUpdateInboundClientLivePushKeepsDirtyBackup(t *testing.T) {
  258. setupBulkDB(t)
  259. nodeID, fake := setupNodeRuntime(t)
  260. client := model.Client{
  261. ID: uuid.NewString(),
  262. Email: "edit@x",
  263. SubID: "sub-edit",
  264. Enable: true,
  265. CreatedAt: 111,
  266. UpdatedAt: 222,
  267. }
  268. ib := nodeInbound(t, nodeID, 30021, []model.Client{client})
  269. edited := client
  270. edited.Comment = "changed"
  271. svc := &ClientService{}
  272. inboundSvc := &InboundService{}
  273. if _, err := svc.UpdateInboundClient(inboundSvc, &model.Inbound{
  274. Id: ib.Id,
  275. Protocol: model.VLESS,
  276. Settings: clientsSettings(t, []model.Client{edited}),
  277. }, client.Email); err != nil {
  278. t.Fatalf("UpdateInboundClient: %v", err)
  279. }
  280. if got := fake.updateUser.Load(); got != 1 {
  281. t.Fatalf("edit streamed %d UpdateUser RPCs, want 1", got)
  282. }
  283. if _, _, dirty, _, err := (&NodeService{}).NodeSyncState(nodeID); err != nil {
  284. t.Fatalf("NodeSyncState: %v", err)
  285. } else if !dirty {
  286. t.Fatal("successful live update should keep node dirty as reconcile backup")
  287. }
  288. }
  289. // TestNodeBulk_LargeDeleteFoldsToDirty: deleting more than the threshold from an
  290. // online node inbound must fold into a reconcile rather than per-client deletes.
  291. func TestNodeBulk_LargeDeleteFoldsToDirty(t *testing.T) {
  292. setupBulkDB(t)
  293. nodeID, fake := setupNodeRuntime(t)
  294. seed := makeNodeClients(nodeBulkPushThreshold + 10)
  295. nodeInbound(t, nodeID, 30003, seed)
  296. svc := &ClientService{}
  297. inboundSvc := &InboundService{}
  298. emails := make([]string, len(seed))
  299. for i := range seed {
  300. emails[i] = seed[i].Email
  301. }
  302. if _, _, err := svc.BulkDelete(inboundSvc, emails, false); err != nil {
  303. t.Fatalf("BulkDelete: %v", err)
  304. }
  305. if got := fake.deleteUser.Load(); got != 0 {
  306. t.Fatalf("large delete streamed %d DeleteUser RPCs, want 0 (should fold to dirty)", got)
  307. }
  308. if _, _, dirty, _, err := (&NodeService{}).NodeSyncState(nodeID); err != nil {
  309. t.Fatalf("NodeSyncState: %v", err)
  310. } else if !dirty {
  311. t.Fatal("large delete must mark the node dirty")
  312. }
  313. }
  314. func TestDelInbound_NodeSelectedModeDeletesRemoteImmediately(t *testing.T) {
  315. setupBulkDB(t)
  316. nodeID, fake := setupNodeRuntime(t)
  317. if err := database.GetDB().Model(&model.Node{}).Where("id = ?", nodeID).
  318. Updates(map[string]any{
  319. "inbound_sync_mode": "selected",
  320. "inbound_tags": []string{"other-tag"},
  321. }).Error; err != nil {
  322. t.Fatalf("set selected mode: %v", err)
  323. }
  324. ib := nodeInbound(t, nodeID, 30004, makeNodeClients(1))
  325. needRestart, err := (&InboundService{}).DelInbound(ib.Id)
  326. if err != nil {
  327. t.Fatalf("DelInbound: %v", err)
  328. }
  329. if needRestart {
  330. t.Fatal("node-owned delete should not request local restart")
  331. }
  332. if got := fake.delInbound.Load(); got != 1 {
  333. t.Fatalf("node-owned delete streamed %d DelInbound RPCs, want 1", got)
  334. }
  335. var count int64
  336. if err := database.GetDB().Model(&model.Inbound{}).Where("id = ?", ib.Id).Count(&count).Error; err != nil {
  337. t.Fatalf("count inbound: %v", err)
  338. }
  339. if count != 0 {
  340. t.Fatalf("deleted inbound row count = %d, want 0", count)
  341. }
  342. if _, _, dirty, _, err := (&NodeService{}).NodeSyncState(nodeID); err != nil {
  343. t.Fatalf("NodeSyncState: %v", err)
  344. } else if !dirty {
  345. t.Fatal("node-owned delete should still mark the node dirty as reconcile backup")
  346. }
  347. }