node_reset_replay_test.go 2.7 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980
  1. package job
  2. import (
  3. "net/http"
  4. "net/http/httptest"
  5. "path/filepath"
  6. "slices"
  7. "strconv"
  8. "strings"
  9. "sync"
  10. "testing"
  11. "github.com/op/go-logging"
  12. "github.com/mhsanaei/3x-ui/v3/internal/database"
  13. "github.com/mhsanaei/3x-ui/v3/internal/database/dbtest"
  14. "github.com/mhsanaei/3x-ui/v3/internal/database/model"
  15. xuilogger "github.com/mhsanaei/3x-ui/v3/internal/logger"
  16. "github.com/mhsanaei/3x-ui/v3/internal/web/runtime"
  17. "github.com/mhsanaei/3x-ui/v3/internal/web/service"
  18. )
  19. // A reset the node missed is replayed by the next sync, ahead of the snapshot
  20. // fetch so the merge already sees the zeroed counters.
  21. func TestNodeTrafficSyncReplaysOwedResetBeforeSnapshot(t *testing.T) {
  22. xuilogger.InitLogger(logging.ERROR)
  23. dbtest.InitDB(t, filepath.Join(t.TempDir(), "x-ui.db"))
  24. service.StartTrafficWriter()
  25. t.Cleanup(service.StopTrafficWriter)
  26. runtime.SetManager(runtime.NewManager(runtime.LocalDeps{APIPort: func() int { return 0 }, SetNeedRestart: func() {}}))
  27. t.Cleanup(func() { runtime.SetManager(nil) })
  28. var mu sync.Mutex
  29. var calls []string
  30. srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  31. mu.Lock()
  32. switch {
  33. case strings.Contains(r.URL.Path, "clients/resetTraffic/"):
  34. calls = append(calls, "reset:"+r.URL.Path[strings.LastIndex(r.URL.Path, "/")+1:])
  35. case strings.HasSuffix(r.URL.Path, "inbounds/list"):
  36. calls = append(calls, "snapshot")
  37. }
  38. mu.Unlock()
  39. w.Header().Set("Content-Type", "application/json")
  40. if strings.HasSuffix(r.URL.Path, "inbounds/list") {
  41. _, _ = w.Write([]byte(`{"success":true,"obj":[]}`))
  42. return
  43. }
  44. _, _ = w.Write([]byte(`{"success":true}`))
  45. }))
  46. t.Cleanup(srv.Close)
  47. host, port, _ := strings.Cut(strings.TrimPrefix(srv.URL, "http://"), ":")
  48. portNum, _ := strconv.Atoi(port)
  49. node := &model.Node{
  50. Name: "owes-reset", Scheme: "http", Address: host, Port: portNum, BasePath: "/", ApiToken: "tok",
  51. Enable: true, Status: "online", AllowPrivateAddress: true, TlsVerifyMode: "verify",
  52. }
  53. if err := database.GetDB().Create(node).Error; err != nil {
  54. t.Fatalf("create node: %v", err)
  55. }
  56. if err := database.GetDB().Create(&model.NodePendingReset{NodeId: node.Id, Email: "owed@node", QueuedAt: 1}).Error; err != nil {
  57. t.Fatalf("seed pending reset: %v", err)
  58. }
  59. NewNodeTrafficSyncJob().Run()
  60. mu.Lock()
  61. got := slices.Clone(calls)
  62. mu.Unlock()
  63. if len(got) < 2 || got[0] != "reset:owed@node" || !slices.Contains(got, "snapshot") {
  64. t.Fatalf("node calls %v, want the owed reset first, then the snapshot", got)
  65. }
  66. var left int64
  67. if err := database.GetDB().Model(&model.NodePendingReset{}).Count(&left).Error; err != nil {
  68. t.Fatalf("count pending: %v", err)
  69. }
  70. if left != 0 {
  71. t.Fatalf("replayed reset still queued (%d rows)", left)
  72. }
  73. }