diff --git a/docs/architecture.md b/docs/architecture.md index e31848a24..92bc8525f 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -367,6 +367,8 @@ merged with GUID-based baselines to avoid double counting after resets. `job/xray_traffic_job.go`, `job/node_traffic_sync_job.go`, `service/inbound_node.go` (`SetRemoteTraffic` / `upsertNodeBaseline`), models `xray.ClientTraffic`, `model.NodeClientTraffic`, `model.ClientGlobalTraffic` (cross-master totals). +A client reset is queued per hosting node in `model.NodePendingReset` (`service/node_reset_queue.go`) +and replayed by the node sync until the node accepts it. Periodic resets: `job/periodic_traffic_reset_job.go` (keyed off `Inbound.TrafficReset`). ### 5.4 Background jobs (cron) @@ -465,6 +467,7 @@ for AutoMigrate in `internal/database/db.go`. | `Host` | Subscription host overrides (per inbound) | `Address`, `Port`, `Sni`, `Path`, `Security`, `Fingerprint`, `SortOrder`, visibility/exclusion flags | | `Node` | A managed child panel | `Guid`, `Address`, `Status`, `TlsVerifyMode`, `PinnedCertSha256`, `ConfigDirty`, version/heartbeat/metric fields | | `NodeClientTraffic` | Per-node client traffic baseline | cross-node merge (anti-double-count) | +| `NodePendingReset` | Client resets a node has not confirmed | `NodeId`, `Email`, `QueuedAt`; replayed by the node sync, freezes that client's node verdict until delivered | | `NodeClientIp` | Per-node client IP attribution | `NodeGuid`, `Email`, `Ips` | | `ClientGlobalTraffic` | Cross-master usage totals | `MasterGuid`, `Email`, `Up`, `Down` | | `xray.ClientTraffic` | Per-client counters (`client_traffics`) | `Email`, `Up`, `Down`, `Total`, `ExpiryTime`, `LastOnline` | diff --git a/internal/database/db.go b/internal/database/db.go index 16608c55b..f469a81d9 100644 --- a/internal/database/db.go +++ b/internal/database/db.go @@ -83,6 +83,7 @@ func allModels() []any { &model.NodeClientTraffic{}, &model.NodeClientIp{}, &model.ClientGlobalTraffic{}, + &model.NodePendingReset{}, &model.OutboundSubscription{}, &model.SubBalancer{}, } diff --git a/internal/database/migrate_data.go b/internal/database/migrate_data.go index 3a423c932..38d343bb7 100644 --- a/internal/database/migrate_data.go +++ b/internal/database/migrate_data.go @@ -56,6 +56,7 @@ func migrationModels() []any { &model.NodeClientTraffic{}, &model.NodeClientIp{}, &model.ClientGlobalTraffic{}, + &model.NodePendingReset{}, &model.OutboundSubscription{}, &model.SubBalancer{}, } diff --git a/internal/database/model/node_pending_reset.go b/internal/database/model/node_pending_reset.go new file mode 100644 index 000000000..4dd182727 --- /dev/null +++ b/internal/database/model/node_pending_reset.go @@ -0,0 +1,11 @@ +package model + +// NodePendingReset is a client traffic reset a hosting node has not confirmed; +// until it lands the node still counts pre-reset usage, so every sync replays it. +type NodePendingReset struct { + Id int `json:"id" gorm:"primaryKey;autoIncrement"` + NodeId int `json:"nodeId" gorm:"uniqueIndex:idx_node_pending_reset,priority:1;not null"` + Email string `json:"email" gorm:"uniqueIndex:idx_node_pending_reset,priority:2;not null"` + // QueuedAt (ns) tells a delivery apart from a reset re-queued while it ran. + QueuedAt int64 `json:"queuedAt"` +} diff --git a/internal/web/job/node_reset_replay_test.go b/internal/web/job/node_reset_replay_test.go new file mode 100644 index 000000000..961b67969 --- /dev/null +++ b/internal/web/job/node_reset_replay_test.go @@ -0,0 +1,80 @@ +package job + +import ( + "net/http" + "net/http/httptest" + "path/filepath" + "slices" + "strconv" + "strings" + "sync" + "testing" + + "github.com/op/go-logging" + + "github.com/mhsanaei/3x-ui/v3/internal/database" + "github.com/mhsanaei/3x-ui/v3/internal/database/dbtest" + "github.com/mhsanaei/3x-ui/v3/internal/database/model" + xuilogger "github.com/mhsanaei/3x-ui/v3/internal/logger" + "github.com/mhsanaei/3x-ui/v3/internal/web/runtime" + "github.com/mhsanaei/3x-ui/v3/internal/web/service" +) + +// A reset the node missed is replayed by the next sync, ahead of the snapshot +// fetch so the merge already sees the zeroed counters. +func TestNodeTrafficSyncReplaysOwedResetBeforeSnapshot(t *testing.T) { + xuilogger.InitLogger(logging.ERROR) + dbtest.InitDB(t, filepath.Join(t.TempDir(), "x-ui.db")) + service.StartTrafficWriter() + t.Cleanup(service.StopTrafficWriter) + runtime.SetManager(runtime.NewManager(runtime.LocalDeps{APIPort: func() int { return 0 }, SetNeedRestart: func() {}})) + t.Cleanup(func() { runtime.SetManager(nil) }) + + var mu sync.Mutex + var calls []string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + switch { + case strings.Contains(r.URL.Path, "clients/resetTraffic/"): + calls = append(calls, "reset:"+r.URL.Path[strings.LastIndex(r.URL.Path, "/")+1:]) + case strings.HasSuffix(r.URL.Path, "inbounds/list"): + calls = append(calls, "snapshot") + } + mu.Unlock() + w.Header().Set("Content-Type", "application/json") + if strings.HasSuffix(r.URL.Path, "inbounds/list") { + _, _ = w.Write([]byte(`{"success":true,"obj":[]}`)) + return + } + _, _ = w.Write([]byte(`{"success":true}`)) + })) + t.Cleanup(srv.Close) + host, port, _ := strings.Cut(strings.TrimPrefix(srv.URL, "http://"), ":") + portNum, _ := strconv.Atoi(port) + node := &model.Node{ + Name: "owes-reset", Scheme: "http", Address: host, Port: portNum, BasePath: "/", ApiToken: "tok", + Enable: true, Status: "online", AllowPrivateAddress: true, TlsVerifyMode: "verify", + } + if err := database.GetDB().Create(node).Error; err != nil { + t.Fatalf("create node: %v", err) + } + if err := database.GetDB().Create(&model.NodePendingReset{NodeId: node.Id, Email: "owed@node", QueuedAt: 1}).Error; err != nil { + t.Fatalf("seed pending reset: %v", err) + } + + NewNodeTrafficSyncJob().Run() + + mu.Lock() + got := slices.Clone(calls) + mu.Unlock() + if len(got) < 2 || got[0] != "reset:owed@node" || !slices.Contains(got, "snapshot") { + t.Fatalf("node calls %v, want the owed reset first, then the snapshot", got) + } + var left int64 + if err := database.GetDB().Model(&model.NodePendingReset{}).Count(&left).Error; err != nil { + t.Fatalf("count pending: %v", err) + } + if left != 0 { + t.Fatalf("replayed reset still queued (%d rows)", left) + } +} diff --git a/internal/web/job/node_traffic_sync_job.go b/internal/web/job/node_traffic_sync_job.go index 71931b797..aaa3b22b0 100644 --- a/internal/web/job/node_traffic_sync_job.go +++ b/internal/web/job/node_traffic_sync_job.go @@ -387,6 +387,13 @@ func (j *NodeTrafficSyncJob) syncOne(mgr *runtime.Manager, n *model.Node, doIpSy } } + // Before the snapshot, so counters a reset just zeroed are what gets merged. + resetCtx, resetCancel := context.WithTimeout(context.Background(), nodeTrafficSyncRequestTimeout) + if resetErr := j.inboundService.DeliverNodeResets(resetCtx, n.Id, rt); resetErr != nil { + logger.Warningf("node traffic sync: reset delivery to %s failed, retrying next tick: %v", n.Name, resetErr) + } + resetCancel() + ctx, cancel := context.WithTimeout(context.Background(), nodeTrafficSyncRequestTimeout) defer cancel() diff --git a/internal/web/runtime/remote.go b/internal/web/runtime/remote.go index 5c8942d77..99587f28b 100644 --- a/internal/web/runtime/remote.go +++ b/internal/web/runtime/remote.go @@ -701,6 +701,12 @@ func (r *Remote) ResetClientTraffic(ctx context.Context, _ *model.Inbound, email return err } +// ResetClientTraffics zeroes many clients on the node in one request. +func (r *Remote) ResetClientTraffics(ctx context.Context, emails []string) error { + _, err := r.do(ctx, http.MethodPost, "panel/api/clients/bulkResetTraffic", map[string]any{"emails": emails}) + return err +} + func (r *Remote) ResetAllTraffics(ctx context.Context) error { _, err := r.do(ctx, http.MethodPost, "panel/api/inbounds/resetAllTraffics", nil) return err diff --git a/internal/web/runtime/remote_reset_test.go b/internal/web/runtime/remote_reset_test.go new file mode 100644 index 000000000..4eb196036 --- /dev/null +++ b/internal/web/runtime/remote_reset_test.go @@ -0,0 +1,33 @@ +package runtime + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "slices" + "testing" +) + +// The master replays a node's reset backlog through the node's bulk endpoint. +func TestRemoteResetClientTrafficsPostsEmailsToBulkEndpoint(t *testing.T) { + var path string + var body struct { + Emails []string `json:"emails"` + } + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + path = r.URL.Path + _ = json.NewDecoder(r.Body).Decode(&body) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"success":true}`)) + })) + t.Cleanup(srv.Close) + + r := NewRemote(nodeForPlainServer(t, srv, "verify", "tok"), nil) + if err := r.ResetClientTraffics(context.Background(), []string{"a@x", "b@x"}); err != nil { + t.Fatalf("ResetClientTraffics: %v", err) + } + if path != "/panel/api/clients/bulkResetTraffic" || !slices.Equal(body.Emails, []string{"a@x", "b@x"}) { + t.Fatalf("node got %s %v, want /panel/api/clients/bulkResetTraffic [a@x b@x]", path, body.Emails) + } +} diff --git a/internal/web/service/client_traffic.go b/internal/web/service/client_traffic.go index 326a75552..5082d9525 100644 --- a/internal/web/service/client_traffic.go +++ b/internal/web/service/client_traffic.go @@ -74,6 +74,7 @@ func (s *ClientService) BulkResetTraffic(inboundSvc *InboundService, emails []st return 0, err } affected := 0 + var resetNodes []int err = submitTrafficWrite(func() error { db := database.GetDB() return db.Transaction(func(tx *gorm.DB) error { @@ -97,12 +98,16 @@ func (s *ClientService) BulkResetTraffic(inboundSvc *InboundService, emails []st return err } } - return nil + var qErr error + resetNodes, qErr = queueNodeResets(tx, cleanEmails) + return qErr }) }) if err != nil { return 0, err } + inboundSvc.resetMtprotoClientQuotas(cleanEmails) + inboundSvc.deliverNodeResetsNow(resetNodes) // After the zeroing, as in ResetTrafficByEmail: enabling a still-depleted // client first lets the next traffic tick switch it off again. for _, e := range cleanEmails { @@ -120,18 +125,23 @@ func (s *ClientService) BulkResetTraffic(inboundSvc *InboundService, emails []st } func (s *ClientService) ResetAllClientTraffics(inboundSvc *InboundService, id int) error { + var resetNodes []int err := submitTrafficWrite(func() error { - return s.resetAllClientTrafficsLocked(id) + var inner error + resetNodes, inner = s.resetAllClientTrafficsLocked(id) + return inner }) if err == nil { inboundSvc.resetAllMtprotoQuotas() + inboundSvc.deliverNodeResetsNow(resetNodes) } return err } -func (s *ClientService) resetAllClientTrafficsLocked(id int) error { +func (s *ClientService) resetAllClientTrafficsLocked(id int) ([]int, error) { db := database.GetDB() now := time.Now().Unix() * 1000 + var resetNodes []int if err := db.Transaction(func(tx *gorm.DB) error { // client_traffics.inbound_id is stale: it reflects the inbound the row was @@ -176,6 +186,10 @@ func (s *ClientService) resetAllClientTrafficsLocked(id int) error { return err } } + var qErr error + if resetNodes, qErr = queueNodeResets(tx, resetEmails); qErr != nil { + return qErr + } inboundWhereText := "id " if id == -1 { @@ -190,13 +204,14 @@ func (s *ClientService) resetAllClientTrafficsLocked(id int) error { return result.Error }); err != nil { - return err + return nil, err } - return nil + return resetNodes, nil } func (s *ClientService) ResetAllTraffics() (bool, error) { var affected int64 + var resetNodes []int err := submitTrafficWrite(func() error { return database.GetDB().Transaction(func(tx *gorm.DB) error { res := tx.Model(&xray.ClientTraffic{}). @@ -209,11 +224,19 @@ func (s *ClientService) ResetAllTraffics() (bool, error) { if err := tx.Where("1 = 1").Delete(&model.ClientGlobalTraffic{}).Error; err != nil { return err } - return tx.Where("1 = 1").Delete(&model.NodeClientTraffic{}).Error + if err := tx.Where("1 = 1").Delete(&model.NodeClientTraffic{}).Error; err != nil { + return err + } + var qErr error + resetNodes, qErr = queueNodeResets(tx, nil) + return qErr }) }) if err != nil { return false, err } + inbounds := &InboundService{} + inbounds.resetAllMtprotoQuotas() + inbounds.deliverNodeResetsNow(resetNodes) return affected > 0, nil } diff --git a/internal/web/service/inbound_mtproto.go b/internal/web/service/inbound_mtproto.go index 8fd5fc4f4..d4e5125fd 100644 --- a/internal/web/service/inbound_mtproto.go +++ b/internal/web/service/inbound_mtproto.go @@ -95,16 +95,47 @@ func (s *InboundService) applyLocalMtproto(inboundId int) { } func (s *InboundService) resetMtprotoClientQuota(email string) { + s.resetMtprotoClientQuotas([]string{email}) +} + +// resetMtprotoClientQuotas zeroes the sidecar's own quota counter for each local +// MTProto client in emails, or it keeps blocking a client the panel just reset. +func (s *InboundService) resetMtprotoClientQuotas(emails []string) { mgr := mtproto.GetManager() - if !mgr.HasRunning() { + if !mgr.HasRunning() || len(emails) == 0 { return } - id, ok := s.localMtprotoInboundIdForEmail(email) - if !ok { + var inbounds []*model.Inbound + if err := database.GetDB().Model(model.Inbound{}). + Where("protocol = ? AND node_id IS NULL", model.MTProto). + Find(&inbounds).Error; err != nil { return } - s.applyLocalMtproto(id) - mgr.ResetQuota(email) + want := make(map[string]struct{}, len(emails)) + for _, e := range emails { + want[e] = struct{}{} + } + var hit []string + for _, ib := range inbounds { + inst, ok := mtproto.InstanceFromInbound(ib) + if !ok { + continue + } + applied := false + for _, sec := range inst.Secrets { + if _, ok := want[sec.Name]; !ok { + continue + } + if !applied { + s.applyLocalMtproto(ib.Id) + applied = true + } + hit = append(hit, sec.Name) + } + } + for _, email := range hit { + mgr.ResetQuota(email) + } } func (s *InboundService) resetAllMtprotoQuotas() { @@ -123,25 +154,3 @@ func (s *InboundService) resetAllMtprotoQuotas() { } } } - -func (s *InboundService) localMtprotoInboundIdForEmail(email string) (int, bool) { - db := database.GetDB() - var inbounds []*model.Inbound - if err := db.Model(model.Inbound{}). - Where("protocol = ? AND node_id IS NULL", model.MTProto). - Find(&inbounds).Error; err != nil { - return 0, false - } - for _, ib := range inbounds { - inst, ok := mtproto.InstanceFromInbound(ib) - if !ok { - continue - } - for _, sec := range inst.Secrets { - if sec.Name == email { - return ib.Id, true - } - } - } - return 0, false -} diff --git a/internal/web/service/inbound_node.go b/internal/web/service/inbound_node.go index 24b907058..1feec0f46 100644 --- a/internal/web/service/inbound_node.go +++ b/internal/web/service/inbound_node.go @@ -547,6 +547,11 @@ func (s *InboundService) setRemoteTrafficLocked(nodeID int, snap *runtime.Traffi centralCSByEmail[centralClientStats[i].Email] = ¢ralClientStats[i] } + owedResets, err := pendingNodeResetEmails(db, nodeID) + if err != nil { + return false, err + } + nodeBaselines := make(map[string]nodeTrafficCounter) var baselineRows []model.NodeClientTraffic if err := db.Model(&model.NodeClientTraffic{}). @@ -925,6 +930,10 @@ func (s *InboundService) setRemoteTrafficLocked(nodeID int, snap *runtime.Traffi // Node-wide total, not this inbound's possibly-stale copy (#5274). canon := nodeEmailTotals[cs.Email] + // Until the node applies a reset it owes, its verdict rests on the + // pre-reset counters: only usage may move for this client. + _, owed := owedResets[cs.Email] + clientFrozen := lifecycleFrozen || owed base, seen := nodeBaselines[cs.Email] var deltaUp, deltaDown int64 @@ -986,18 +995,18 @@ func (s *InboundService) setRemoteTrafficLocked(nodeID int, snap *runtime.Traffi existing := centralCSByEmail[cs.Email] if existing != nil { - expiryChanged := !lifecycleFrozen && existing.ExpiryTime != mergeActivationExpiry(existing.ExpiryTime, cs.ExpiryTime) + expiryChanged := !clientFrozen && existing.ExpiryTime != mergeActivationExpiry(existing.ExpiryTime, cs.ExpiryTime) // Only a real latch to disabled is structural; one-way merge never // re-enables from the node. - enableChanged := !lifecycleFrozen && existing.Enable && !cs.Enable && + enableChanged := !clientFrozen && existing.Enable && !cs.Enable && !nodeDisableIsStale(existing, cs, now, deltaUp, deltaDown) - metaChanged := !lifecycleFrozen && (existing.Total != cs.Total || existing.Reset != cs.Reset || existing.ResetWeekday != cs.ResetWeekday) + metaChanged := !clientFrozen && (existing.Total != cs.Total || existing.Reset != cs.Reset || existing.ResetWeekday != cs.ResetWeekday) if enableChanged || metaChanged || expiryChanged { structuralChange = true } } - renewed := !lifecycleFrozen && seen && existing != nil && nodeClientRenewed(existing, cs, canon, base) + renewed := !clientFrozen && seen && existing != nil && nodeClientRenewed(existing, cs, canon, base) if renewed { // Reject when the node's own settings still carry the old absolute: // lagging ClientStats after a master shorten mimic a renew (#6228). @@ -1037,7 +1046,7 @@ func (s *InboundService) setRemoteTrafficLocked(nodeID int, snap *runtime.Traffi existing.ResetWeekday = cs.ResetWeekday existing.ResetCount = cs.ResetCount structuralChange = true - } else if lifecycleFrozen { + } else if clientFrozen { // Push pending or just landed: only counters may move, the master // keeps expiry/enable/total/reset. if err := tx.Exec( @@ -1096,7 +1105,7 @@ func (s *InboundService) setRemoteTrafficLocked(nodeID int, snap *runtime.Traffi } // A dip plus a lagging longer expiry mimics nodeClientRenewed and would // undo a master shorten once the freeze lifts (#6228). - if lifecycleFrozen && seen && (canon.Up < base.Up || canon.Down < base.Down) { + if clientFrozen && seen && (canon.Up < base.Up || canon.Down < base.Down) { continue } if err := s.upsertNodeBaseline(tx, nodeID, cs.Email, canon.Up, canon.Down); err != nil { diff --git a/internal/web/service/inbound_traffic.go b/internal/web/service/inbound_traffic.go index 507c31c94..904499b27 100644 --- a/internal/web/service/inbound_traffic.go +++ b/internal/web/service/inbound_traffic.go @@ -28,14 +28,16 @@ const depletedClientsClause = "reset = 0 and reset_day = 0 and reset_weekday = 0 func (s *InboundService) AddTraffic(inboundTraffics []*xray.Traffic, clientTraffics []*xray.ClientTraffic) (needRestart bool, clientsDisabled bool, err error) { var disabledNodeIDs []int var remotePlans []trafficInboundUpdatePlan + var renewed []string err = submitTrafficWrite(func() error { var inner error - needRestart, clientsDisabled, disabledNodeIDs, remotePlans, inner = s.addTrafficLocked(inboundTraffics, clientTraffics) + needRestart, clientsDisabled, disabledNodeIDs, remotePlans, renewed, inner = s.addTrafficLocked(inboundTraffics, clientTraffics) return inner }) if err != nil { return } + s.resetMtprotoClientQuotas(renewed) // Off the serial writer: a hanging node must not stall traffic accounting. needRestart = s.applyTrafficRemotePlans(remotePlans) || needRestart if len(disabledNodeIDs) > 0 { @@ -44,7 +46,7 @@ func (s *InboundService) AddTraffic(inboundTraffics []*xray.Traffic, clientTraff return } -func (s *InboundService) addTrafficLocked(inboundTraffics []*xray.Traffic, clientTraffics []*xray.ClientTraffic) (bool, bool, []int, []trafficInboundUpdatePlan, error) { +func (s *InboundService) addTrafficLocked(inboundTraffics []*xray.Traffic, clientTraffics []*xray.ClientTraffic) (bool, bool, []int, []trafficInboundUpdatePlan, []string, error) { db := database.GetDB() // Commit durable traffic before best-effort lifecycle maintenance so helper // failures cannot discard usage already reported by Xray. @@ -54,7 +56,7 @@ func (s *InboundService) addTrafficLocked(inboundTraffics []*xray.Traffic, clien } return s.addClientTraffic(tx, clientTraffics) }); err != nil { - return false, false, nil, nil, err + return false, false, nil, nil, nil, err } var ( @@ -99,10 +101,10 @@ func (s *InboundService) addTrafficLocked(inboundTraffics []*xray.Traffic, clien }) if err != nil { logger.Warning("traffic lifecycle maintenance failed after traffic commit:", err) - return false, false, nil, nil, nil + return false, false, nil, nil, nil, nil } needRestart = needRestart || s.applyTrafficMutationBatch(batch) - return needRestart, clientsDisabled, disabledNodeIDs, batch.remotePlans, nil + return needRestart, clientsDisabled, disabledNodeIDs, batch.remotePlans, batch.renewedEmails, nil } func (s *InboundService) addInboundTraffic(tx *gorm.DB, traffics []*xray.Traffic) error { @@ -515,6 +517,7 @@ func (s *InboundService) autoRenewClients(tx *gorm.DB, mutationBatch *trafficMut if err = clearGlobalTraffic(tx, renewedEmails...); err != nil { return false, 0, err } + mutationBatch.renewedEmails = append(mutationBatch.renewedEmails, renewedEmails...) for _, clientToAdd := range clientsToAdd { if clientToAdd.inbound.NodeID != nil { mutationBatch.addNode(*clientToAdd.inbound.NodeID) @@ -648,33 +651,23 @@ func (s *InboundService) ResetClientTrafficByEmail(clientEmail string) error { } func (s *InboundService) ResetClientTraffic(id int, clientEmail string) (needRestart bool, err error) { - var resetInbound *model.Inbound + var ownNode *int err = submitTrafficWrite(func() error { var inner error - needRestart, resetInbound, inner = s.resetClientTrafficLocked(id, clientEmail) + needRestart, ownNode, inner = s.resetClientTrafficLocked(id, clientEmail) return inner }) if err == nil { s.resetMtprotoClientQuota(clientEmail) - if resetInbound != nil && resetInbound.NodeID != nil { - // Attempted whatever the node's status: nothing replays a reset, so a - // node still serving after being marked offline must get it now. - if rt, rterr := s.runtimeFor(resetInbound); rterr != nil { - logger.Warning("ResetClientTraffic: runtime lookup failed:", rterr) - } else { - ctx, cancel := nodePushContext() - e := rt.ResetClientTraffic(ctx, resetInbound, clientEmail) - cancel() - if e != nil { - logger.Warning("ResetClientTraffic: remote propagation to", rt.Name(), "failed:", e) - } - } + // Siblings on other nodes are delivered by their own inbound's reset. + if ownNode != nil { + s.deliverNodeResetsNow([]int{*ownNode}) } } return } -func (s *InboundService) resetClientTrafficLocked(id int, clientEmail string) (bool, *model.Inbound, error) { +func (s *InboundService) resetClientTrafficLocked(id int, clientEmail string) (bool, *int, error) { needRestart := false var reenablePlan *trafficLocalApplyPlan var reenableNodeID *int @@ -747,6 +740,9 @@ func (s *InboundService) resetClientTrafficLocked(id int, clientEmail string) (b if err := tx.Where("email = ?", clientEmail).Delete(&model.NodeClientTraffic{}).Error; err != nil { return err } + if _, err := queueNodeResets(tx, []string{clientEmail}); err != nil { + return err + } if err := tx.Model(model.Inbound{}). Where("id = ?", id). Update("last_traffic_reset_time", now).Error; err != nil { @@ -775,7 +771,10 @@ func (s *InboundService) resetClientTrafficLocked(id int, clientEmail string) (b } } - return needRestart, inbound, nil + if inbound != nil { + return needRestart, inbound.NodeID, nil + } + return needRestart, nil, nil } func (s *InboundService) ResetAllTraffics() error { diff --git a/internal/web/service/inbound_traffic_apply.go b/internal/web/service/inbound_traffic_apply.go index 5bbeb8b18..3d962f30d 100644 --- a/internal/web/service/inbound_traffic_apply.go +++ b/internal/web/service/inbound_traffic_apply.go @@ -30,6 +30,8 @@ type trafficMutationBatch struct { localPlans []trafficLocalApplyPlan remotePlans []trafficInboundUpdatePlan nodeIDs map[int]struct{} + // renewedEmails get their MTProto sidecar quota zeroed once the tick commits. + renewedEmails []string } type trafficInboundUpdatePlan struct{ oldInbound, newInbound model.Inbound } diff --git a/internal/web/service/mtproto_fake_test.go b/internal/web/service/mtproto_fake_test.go index efc7bffae..b56f99842 100644 --- a/internal/web/service/mtproto_fake_test.go +++ b/internal/web/service/mtproto_fake_test.go @@ -2,8 +2,13 @@ package service import ( "fmt" + "net" + "net/http" + "net/url" "os" "path/filepath" + "regexp" + "slices" "strings" "testing" "time" @@ -46,9 +51,79 @@ func fakeMtgChildMain() { fmt.Fprintf(f, "%d\n", os.Getpid()) f.Close() } + if logPath := os.Getenv("MTG_FAKE_APILOG"); logPath != "" && len(os.Args) > 2 { + go serveFakeMtgAPI(os.Args[len(os.Args)-1], logPath) + } select {} } +// serveFakeMtgAPI answers the management API on the config's api-bind-to and +// logs each reset-quota call, so a test sees which sidecar quotas were zeroed. +func serveFakeMtgAPI(configPath, logPath string) { + cfg, err := os.ReadFile(configPath) + if err != nil { + return + } + m := regexp.MustCompile(`api-bind-to = "([^"]+)"`).FindSubmatch(cfg) + if m == nil { + return + } + ln, err := net.Listen("tcp", string(m[1])) + if err != nil { + return + } + appendFakeMtgLog(logPath, "ready") + _ = http.Serve(ln, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if name, ok := strings.CutSuffix(strings.TrimPrefix(r.URL.Path, "/secrets/"), "/reset-quota"); ok && r.Method == http.MethodPost { + if unescaped, err := url.PathUnescape(name); err == nil { + appendFakeMtgLog(logPath, "reset:"+unescaped) + } + } + _, _ = w.Write([]byte("{}")) + })) +} + +func appendFakeMtgLog(path, line string) { + if f, err := os.OpenFile(path, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0o644); err == nil { + fmt.Fprintln(f, line) + f.Close() + } +} + +// installFakeMtgAPI is installFakeMtg whose children also serve the management +// API; it returns the pid file and the API call log. +func installFakeMtgAPI(t *testing.T) (string, string) { + t.Helper() + pidFile := installFakeMtg(t) + logPath := filepath.Join(filepath.Dir(pidFile), "mtg-api.log") + t.Setenv("MTG_FAKE_APILOG", logPath) + return pidFile, logPath +} + +func fakeMtgLog(t *testing.T, logPath string) []string { + t.Helper() + data, err := os.ReadFile(logPath) + if os.IsNotExist(err) { + return nil + } + if err != nil { + t.Fatalf("read mtg api log: %v", err) + } + return strings.Fields(string(data)) +} + +// waitFakeMtgLog polls until the log holds want, failing on timeout. +func waitFakeMtgLog(t *testing.T, logPath, want string) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for !slices.Contains(fakeMtgLog(t, logPath), want) { + if time.Now().After(deadline) { + t.Fatalf("mtg api log never recorded %q: %v", want, fakeMtgLog(t, logPath)) + } + time.Sleep(20 * time.Millisecond) + } +} + // installFakeMtg points the mtproto manager at a copy of the running test // binary posing as mtg (via the MTG_FAKE_CHILD gate in TestMain) and returns // the pid file whose line count equals the number of processes spawned so far. diff --git a/internal/web/service/mtproto_quota_reset_test.go b/internal/web/service/mtproto_quota_reset_test.go new file mode 100644 index 000000000..0ae1e41a0 --- /dev/null +++ b/internal/web/service/mtproto_quota_reset_test.go @@ -0,0 +1,100 @@ +package service + +import ( + "slices" + "strings" + "testing" + "time" + + "github.com/mhsanaei/3x-ui/v3/internal/database" + "github.com/mhsanaei/3x-ui/v3/internal/database/model" + "github.com/mhsanaei/3x-ui/v3/internal/mtproto" + "github.com/mhsanaei/3x-ui/v3/internal/web/runtime" + "github.com/mhsanaei/3x-ui/v3/internal/xray" +) + +// startQuotaSidecar runs a local MTProto inbound for mtga and mtgb under the fake +// mtg and returns its API log once the sidecar answers. +func startQuotaSidecar(t *testing.T, port int, mtga model.Client) (*model.Inbound, string) { + t.Helper() + setupConflictDB(t) + pidFile, logPath := installFakeMtgAPI(t) + runtime.SetManager(runtime.NewManager(runtime.LocalDeps{APIPort: func() int { return 0 }, SetNeedRestart: func() {}})) + t.Cleanup(func() { runtime.SetManager(nil) }) + + mtga.Email, mtga.Secret = "mtga", mtprotoTestSecretA + clients := []model.Client{mtga, {Email: "mtgb", Secret: mtprotoTestSecretB, Enable: true}} + ib := &model.Inbound{Tag: "mt-quota", Enable: true, Port: port, Protocol: model.MTProto, Settings: clientsSettings(t, clients)} + if err := database.GetDB().Create(ib).Error; err != nil { + t.Fatalf("create inbound: %v", err) + } + if err := (&ClientService{}).SyncInbound(nil, ib.Id, clients); err != nil { + t.Fatalf("SyncInbound: %v", err) + } + for _, c := range clients { + row := xray.ClientTraffic{InboundId: ib.Id, Email: c.Email, Enable: true, Up: 5, Total: c.TotalGB, ExpiryTime: c.ExpiryTime, Reset: c.Reset} + if err := database.GetDB().Create(&row).Error; err != nil { + t.Fatalf("seed traffic: %v", err) + } + } + // A running sidecar needs a served client, so prime with the healthy set. + inst, ok := mtproto.InstanceFromInbound(&model.Inbound{ + Id: ib.Id, Tag: ib.Tag, Port: port, Protocol: model.MTProto, + Settings: clientsSettings(t, []model.Client{{Email: "mtgb", Secret: mtprotoTestSecretB, Enable: true}}), + }) + if !ok { + t.Fatal("seed inbound must produce an mtg instance") + } + if err := mtproto.GetManager().Ensure(inst); err != nil { + t.Fatalf("start mtg: %v", err) + } + t.Cleanup(func() { mtproto.GetManager().Remove(ib.Id) }) + waitForSpawns(t, pidFile, 1) + waitFakeMtgLog(t, logPath, "ready") + return ib, logPath +} + +func quotaResets(t *testing.T, logPath string) []string { + t.Helper() + var out []string + for _, line := range fakeMtgLog(t, logPath) { + if name, ok := strings.CutPrefix(line, "reset:"); ok { + out = append(out, name) + } + } + slices.Sort(out) + return out +} + +// Every path that zeroes a client's panel counters must zero the sidecar's own +// quota counter too, or the sidecar keeps refusing the client. +func TestPanelResetsZeroSidecarQuota(t *testing.T) { + t.Run("bulk reset", func(t *testing.T) { + _, logPath := startQuotaSidecar(t, 46201, model.Client{Enable: true}) + if _, err := (&ClientService{}).BulkResetTraffic(&InboundService{}, []string{"mtga"}); err != nil { + t.Fatalf("BulkResetTraffic: %v", err) + } + if got := quotaResets(t, logPath); !slices.Equal(got, []string{"mtga"}) { + t.Fatalf("sidecar quota resets %v, want [mtga]", got) + } + }) + t.Run("reset all", func(t *testing.T) { + _, logPath := startQuotaSidecar(t, 46202, model.Client{Enable: true}) + if _, err := (&ClientService{}).ResetAllTraffics(); err != nil { + t.Fatalf("ResetAllTraffics: %v", err) + } + if got := quotaResets(t, logPath); !slices.Equal(got, []string{"mtga", "mtgb"}) { + t.Fatalf("sidecar quota resets %v, want [mtga mtgb]", got) + } + }) + t.Run("auto renew", func(t *testing.T) { + expired := time.Now().Add(-time.Hour).UnixMilli() + _, logPath := startQuotaSidecar(t, 46203, model.Client{Enable: true, Reset: 30, ExpiryTime: expired}) + if _, _, err := (&InboundService{}).AddTraffic(nil, nil); err != nil { + t.Fatalf("AddTraffic: %v", err) + } + if got := quotaResets(t, logPath); !slices.Equal(got, []string{"mtga"}) { + t.Fatalf("sidecar quota resets %v, want [mtga]", got) + } + }) +} diff --git a/internal/web/service/node.go b/internal/web/service/node.go index 437b2ac47..6e014ec87 100644 --- a/internal/web/service/node.go +++ b/internal/web/service/node.go @@ -864,6 +864,9 @@ func (s *NodeService) Delete(id int) error { if err := tx.Where("node_id = ?", id).Delete(&model.NodeClientTraffic{}).Error; err != nil { return err } + if err := tx.Where("node_id = ?", id).Delete(&model.NodePendingReset{}).Error; err != nil { + return err + } guids := []string{synthNodeGuid(id)} if guid != "" { guids = append(guids, guid) diff --git a/internal/web/service/node_reset_queue.go b/internal/web/service/node_reset_queue.go new file mode 100644 index 000000000..5772a23db --- /dev/null +++ b/internal/web/service/node_reset_queue.go @@ -0,0 +1,152 @@ +package service + +import ( + "context" + "sync" + "time" + + "github.com/mhsanaei/3x-ui/v3/internal/database" + "github.com/mhsanaei/3x-ui/v3/internal/database/model" + "github.com/mhsanaei/3x-ui/v3/internal/logger" + "github.com/mhsanaei/3x-ui/v3/internal/web/runtime" + + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +// nodeBulkResetter is a node runtime that can zero many clients in one call. +type nodeBulkResetter interface { + ResetClientTraffics(ctx context.Context, emails []string) error +} + +type nodeEmail struct { + NodeId int `gorm:"column:node_id"` + Email string `gorm:"column:email"` +} + +// queueNodeResets records a reset for every node hosting one of emails (all +// node-hosted clients when emails is nil) and returns the nodes involved. +func queueNodeResets(tx *gorm.DB, emails []string) ([]int, error) { + base := func() *gorm.DB { + return tx.Table("clients"). + Select("DISTINCT inbounds.node_id AS node_id, clients.email AS email"). + Joins("JOIN client_inbounds ON client_inbounds.client_id = clients.id"). + Joins("JOIN inbounds ON inbounds.id = client_inbounds.inbound_id"). + Where("inbounds.node_id IS NOT NULL") + } + var pairs []nodeEmail + if emails == nil { + if err := base().Scan(&pairs).Error; err != nil { + return nil, err + } + } else { + for _, batch := range chunkStrings(uniqueNonEmptyStrings(emails), sqlInChunk) { + var page []nodeEmail + if err := base().Where("clients.email IN ?", batch).Scan(&page).Error; err != nil { + return nil, err + } + pairs = append(pairs, page...) + } + } + if len(pairs) == 0 { + return nil, nil + } + now := time.Now().UnixNano() + rows := make([]model.NodePendingReset, 0, len(pairs)) + nodes := make(map[int]struct{}) + for _, p := range pairs { + rows = append(rows, model.NodePendingReset{NodeId: p.NodeId, Email: p.Email, QueuedAt: now}) + nodes[p.NodeId] = struct{}{} + } + if err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "node_id"}, {Name: "email"}}, + DoUpdates: clause.AssignmentColumns([]string{"queued_at"}), + }).CreateInBatches(rows, 200).Error; err != nil { + return nil, err + } + ids := make([]int, 0, len(nodes)) + for id := range nodes { + ids = append(ids, id) + } + return ids, nil +} + +// pendingNodeResetEmails lists the clients whose reset the node still owes. +func pendingNodeResetEmails(tx *gorm.DB, nodeID int) (map[string]struct{}, error) { + var emails []string + if err := tx.Model(&model.NodePendingReset{}).Where("node_id = ?", nodeID).Pluck("email", &emails).Error; err != nil { + return nil, err + } + out := make(map[string]struct{}, len(emails)) + for _, e := range emails { + out[e] = struct{}{} + } + return out, nil +} + +var nodeResetDeliveryLocks sync.Map + +// DeliverNodeResets sends the node every reset it has not confirmed. A row is +// dropped only after the node accepted it and only if nothing re-queued it since. +func (s *InboundService) DeliverNodeResets(ctx context.Context, nodeID int, rt runtime.Runtime) error { + lock, _ := nodeResetDeliveryLocks.LoadOrStore(nodeID, &sync.Mutex{}) + lock.(*sync.Mutex).Lock() + defer lock.(*sync.Mutex).Unlock() + db := database.GetDB() + var rows []model.NodePendingReset + if err := db.Where("node_id = ?", nodeID).Order("id").Find(&rows).Error; err != nil { + return err + } + if len(rows) == 0 { + return nil + } + bulk, canBulk := rt.(nodeBulkResetter) + for start := 0; start < len(rows); start += sqlInChunk { + batch := rows[start:min(start+sqlInChunk, len(rows))] + emails := make([]string, len(batch)) + for i := range batch { + emails[i] = batch[i].Email + } + var err error + if canBulk && len(batch) > nodeBulkPushThreshold { + err = bulk.ResetClientTraffics(ctx, emails) + } else { + for _, email := range emails { + if err = rt.ResetClientTraffic(ctx, nil, email); err != nil { + break + } + } + } + if err != nil { + return err + } + for i := range batch { + if err := db.Where("id = ? AND queued_at = ?", batch[i].Id, batch[i].QueuedAt). + Delete(&model.NodePendingReset{}).Error; err != nil { + return err + } + } + } + return nil +} + +// deliverNodeResetsNow tries each node once right after a reset commits; what +// fails stays queued for the node sync job. +func (s *InboundService) deliverNodeResetsNow(nodeIDs []int) { + mgr := runtime.GetManager() + if mgr == nil || len(nodeIDs) == 0 { + return + } + fanoutInboundResults(nodeIDs, nodeFanoutConcurrency, func(i int) struct{} { + rt, err := mgr.RuntimeFor(&nodeIDs[i]) + if err != nil { + return struct{}{} + } + ctx, cancel := nodePushContext() + defer cancel() + if err := s.DeliverNodeResets(ctx, nodeIDs[i], rt); err != nil { + logger.Warning("reset delivery to", rt.Name(), "deferred to the next sync:", err) + } + return struct{}{} + }) +} diff --git a/internal/web/service/node_reset_undelivered_test.go b/internal/web/service/node_reset_undelivered_test.go new file mode 100644 index 000000000..988fd5b20 --- /dev/null +++ b/internal/web/service/node_reset_undelivered_test.go @@ -0,0 +1,275 @@ +package service + +import ( + "context" + "errors" + "fmt" + "slices" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/mhsanaei/3x-ui/v3/internal/database" + "github.com/mhsanaei/3x-ui/v3/internal/database/model" + "github.com/mhsanaei/3x-ui/v3/internal/xray" + + "gorm.io/gorm" +) + +const ( + resetLostOn = `{"clients":[{"email":"reset-lost","totalGB":100,"enable":true}]}` + resetLostOff = `{"clients":[{"email":"reset-lost","totalGB":100,"enable":false}]}` +) + +// seedLatchedNodeClient leaves reset-lost depleted and latched off on the +// master by its node's own usage, as a real node sync does. +func seedLatchedNodeClient(t *testing.T, svc *InboundService) (*gorm.DB, *model.Inbound) { + t.Helper() + db := initTrafficTestDB(t) + createNodeInboundWithClient(t, db, 1, "n1-in", 41901, "reset-lost") + syncNodeWithSettings(t, svc, 1, "n1-in", resetLostOn, xray.ClientTraffic{Email: "reset-lost", Up: 10, Down: 10, Total: 100, Enable: true}) + syncNodeWithSettings(t, svc, 1, "n1-in", resetLostOff, xray.ClientTraffic{Email: "reset-lost", Up: 60, Down: 60, Total: 100, Enable: false}) + if got := readTraffic(t, db, "reset-lost"); got.Enable { + t.Fatal("setup: the depleted client should be latched off") + } + var ib model.Inbound + if err := db.Where("tag = ?", "n1-in").First(&ib).Error; err != nil { + t.Fatalf("load inbound: %v", err) + } + return db, &ib +} + +// A reset the node never received leaves its old counters, so the node keeps +// switching the client off; the master must not adopt that verdict. +func TestNodeResetNotDeliveredDoesNotRedisableClient(t *testing.T) { + resets := []struct { + name string + run func(svc *InboundService, ib *model.Inbound) error + }{ + {"single", func(svc *InboundService, ib *model.Inbound) error { + _, err := svc.ResetClientTraffic(ib.Id, "reset-lost") + return err + }}, + {"bulk", func(svc *InboundService, _ *model.Inbound) error { + _, err := (&ClientService{}).BulkResetTraffic(svc, []string{"reset-lost"}) + return err + }}, + {"inbound", func(svc *InboundService, ib *model.Inbound) error { + return (&ClientService{}).ResetAllClientTraffics(svc, ib.Id) + }}, + {"all", func(*InboundService, *model.Inbound) error { + _, err := (&ClientService{}).ResetAllTraffics() + return err + }}, + } + for _, reset := range resets { + t.Run(reset.name, func(t *testing.T) { + svc := &InboundService{} + db, ib := seedLatchedNodeClient(t, svc) + if err := reset.run(svc, ib); err != nil { + t.Fatalf("reset: %v", err) + } + syncNodeWithSettings(t, svc, 1, "n1-in", resetLostOff, xray.ClientTraffic{Email: "reset-lost", Up: 60, Down: 60, Total: 100, Enable: false}) + got := readTraffic(t, db, "reset-lost") + if !got.Enable || got.Up+got.Down != 0 { + t.Fatalf("after reset: enable=%v used=%d, want enabled at 0 — the undelivered reset re-disabled it", got.Enable, got.Up+got.Down) + } + }) + } +} + +// resetRecordingRuntime is a node that accepts per-client resets unless failing. +type resetRecordingRuntime struct { + fakeNodeRuntime + mu sync.Mutex + fail bool + got []string +} + +func (r *resetRecordingRuntime) ResetClientTraffic(_ context.Context, _ *model.Inbound, email string) error { + r.mu.Lock() + defer r.mu.Unlock() + if r.fail { + return errors.New("node unreachable") + } + r.got = append(r.got, email) + return nil +} + +func (r *resetRecordingRuntime) delivered() []string { + r.mu.Lock() + defer r.mu.Unlock() + return slices.Clone(r.got) +} + +func pendingResetEmails(t *testing.T, nodeID int) []string { + t.Helper() + var emails []string + if err := database.GetDB().Model(&model.NodePendingReset{}).Where("node_id = ?", nodeID). + Order("email").Pluck("email", &emails).Error; err != nil { + t.Fatalf("read pending resets: %v", err) + } + return emails +} + +func setupRecordingNode(t *testing.T, fail bool) (int, *resetRecordingRuntime, *model.Inbound) { + t.Helper() + setupBulkDB(t) + mgr := useTestRuntimeManager(t) + node := &model.Node{Name: "reset-node", Address: "127.0.0.1", Port: 2096, ApiToken: "tok", Enable: true, Status: "online"} + if err := database.GetDB().Create(node).Error; err != nil { + t.Fatalf("create node: %v", err) + } + rec := &resetRecordingRuntime{fail: fail} + mgr.SetRuntimeOverride(node.Id, rec) + ib := nodeInbound(t, node.Id, 41911, []model.Client{{Email: "reset-lost", ID: "11111111-1111-1111-1111-1111111111aa", Enable: true}}) + if err := (&InboundService{}).AddClientStat(database.GetDB(), ib.Id, &model.Client{Email: "reset-lost", Enable: true}); err != nil { + t.Fatalf("AddClientStat: %v", err) + } + return node.Id, rec, ib +} + +// A reachable node gets the reset right after the master commits it. +func TestNodeResetDeliveredRightAway(t *testing.T) { + resets := []struct { + name string + run func(svc *InboundService, ib *model.Inbound) error + }{ + {"single", func(svc *InboundService, ib *model.Inbound) error { + _, err := svc.ResetClientTraffic(ib.Id, "reset-lost") + return err + }}, + {"bulk", func(svc *InboundService, _ *model.Inbound) error { + _, err := (&ClientService{}).BulkResetTraffic(svc, []string{"reset-lost"}) + return err + }}, + {"inbound", func(svc *InboundService, ib *model.Inbound) error { + return (&ClientService{}).ResetAllClientTraffics(svc, ib.Id) + }}, + {"all", func(*InboundService, *model.Inbound) error { + _, err := (&ClientService{}).ResetAllTraffics() + return err + }}, + } + for _, reset := range resets { + t.Run(reset.name, func(t *testing.T) { + nodeID, rec, ib := setupRecordingNode(t, false) + if err := reset.run(&InboundService{}, ib); err != nil { + t.Fatalf("reset: %v", err) + } + if got := rec.delivered(); !slices.Equal(got, []string{"reset-lost"}) { + t.Fatalf("node received resets %v, want [reset-lost]", got) + } + if left := pendingResetEmails(t, nodeID); len(left) != 0 { + t.Fatalf("delivered reset still queued: %v", left) + } + }) + } +} + +// bulkResetRuntime also takes a batch in one call. +type bulkResetRuntime struct { + resetRecordingRuntime + batches [][]string +} + +func (b *bulkResetRuntime) ResetClientTraffics(_ context.Context, emails []string) error { + b.mu.Lock() + defer b.mu.Unlock() + b.batches = append(b.batches, slices.Clone(emails)) + return nil +} + +// Above the per-client push threshold a backlog goes out as one bulk request, +// not one round-trip per client. +func TestNodeResetBacklogUsesBulkRequest(t *testing.T) { + setupBulkDB(t) + const nodeID = 7 + rows := make([]model.NodePendingReset, nodeBulkPushThreshold+1) + for i := range rows { + rows[i] = model.NodePendingReset{NodeId: nodeID, Email: fmt.Sprintf("owed-%02d", i), QueuedAt: 1} + } + if err := database.GetDB().Create(&rows).Error; err != nil { + t.Fatalf("seed pending resets: %v", err) + } + rt := &bulkResetRuntime{} + if err := (&InboundService{}).DeliverNodeResets(context.Background(), nodeID, rt); err != nil { + t.Fatalf("DeliverNodeResets: %v", err) + } + if len(rt.batches) != 1 || len(rt.batches[0]) != len(rows) || len(rt.delivered()) != 0 { + t.Fatalf("bulk batches %d (first %d emails), per-client calls %d; want one batch of %d", + len(rt.batches), len(rt.batches[0]), len(rt.delivered()), len(rows)) + } + if left := pendingResetEmails(t, nodeID); len(left) != 0 { + t.Fatalf("delivered backlog still queued: %d rows", len(left)) + } +} + +// An unreachable node keeps the reset queued until a later delivery lands. +func TestNodeResetReplayedAfterFailure(t *testing.T) { + nodeID, rec, ib := setupRecordingNode(t, true) + if _, err := (&InboundService{}).ResetClientTraffic(ib.Id, "reset-lost"); err != nil { + t.Fatalf("ResetClientTraffic: %v", err) + } + if left := pendingResetEmails(t, nodeID); !slices.Equal(left, []string{"reset-lost"}) { + t.Fatalf("pending after failed delivery = %v, want [reset-lost]", left) + } + + rec.mu.Lock() + rec.fail = false + rec.mu.Unlock() + if err := (&InboundService{}).DeliverNodeResets(context.Background(), nodeID, rec); err != nil { + t.Fatalf("DeliverNodeResets: %v", err) + } + if got := rec.delivered(); !slices.Equal(got, []string{"reset-lost"}) { + t.Fatalf("node received resets %v, want [reset-lost]", got) + } + if left := pendingResetEmails(t, nodeID); len(left) != 0 { + t.Fatalf("delivered reset still queued: %v", left) + } +} + +// slowResetRuntime holds each reset until a second one arrives or a short +// timeout passes, so two unserialized deliveries both reach the node. +type slowResetRuntime struct { + resetRecordingRuntime + calls atomic.Int32 + both chan struct{} +} + +func (r *slowResetRuntime) ResetClientTraffic(ctx context.Context, ib *model.Inbound, email string) error { + if r.calls.Add(1) == 2 { + close(r.both) + } + select { + case <-r.both: + case <-time.After(300 * time.Millisecond): + } + return r.resetRecordingRuntime.ResetClientTraffic(ctx, ib, email) +} + +// The sync job and a reset's own delivery can run at once; the node must still +// get each owed reset once, or usage made in between is wiped a second time. +func TestConcurrentNodeResetDeliveriesSendOnce(t *testing.T) { + setupBulkDB(t) + const nodeID = 9 + if err := database.GetDB().Create(&model.NodePendingReset{NodeId: nodeID, Email: "once", QueuedAt: 1}).Error; err != nil { + t.Fatalf("seed pending reset: %v", err) + } + rt := &slowResetRuntime{both: make(chan struct{})} + var wg sync.WaitGroup + for range 2 { + wg.Add(1) + go func() { + defer wg.Done() + if err := (&InboundService{}).DeliverNodeResets(context.Background(), nodeID, rt); err != nil { + t.Errorf("DeliverNodeResets: %v", err) + } + }() + } + wg.Wait() + if got := rt.delivered(); !slices.Equal(got, []string{"once"}) { + t.Fatalf("node received resets %v, want exactly [once]", got) + } +}