diff --git a/internal/web/service/client_bulk.go b/internal/web/service/client_bulk.go index 8bc44564e..aece9b9db 100644 --- a/internal/web/service/client_bulk.go +++ b/internal/web/service/client_bulk.go @@ -553,7 +553,7 @@ func (s *ClientService) BulkAdjust(inboundSvc *InboundService, emails []string, } } if adjustHwid { - if err := s.setClientLimitHwidByEmail(db, email, *limitHwid); err != nil { + if err := s.setClientLimitHwidByEmail(email, *limitHwid); err != nil { if _, already := skippedReasons[email]; !already { skippedReasons[email] = err.Error() } @@ -1487,22 +1487,22 @@ func (s *ClientService) BulkCreate(inboundSvc *InboundService, payloads []Client } } - createdEmails := make([]string, 0, len(prep)) for idx := range prep { if failed[idx] { skip(prep[idx].client.Email, reason[idx]) continue } - if err := s.setClientLimitHwidByEmail(nil, prep[idx].client.Email, prep[idx].limitHwid); err != nil { + // The client is already live after fanout; never leave a stale delete + // tombstone merely because applying its optional HWID limit failed. + withdrawClientTombstones(prep[idx].client.Email) + if err := s.setClientLimitHwidByEmail(prep[idx].client.Email, prep[idx].limitHwid); err != nil { skip(prep[idx].client.Email, err.Error()) continue } - createdEmails = append(createdEmails, prep[idx].client.Email) result.Created++ } // A re-created email is a live identity again: a delete tombstone left // standing makes the next node merge prune the new client's inbound links. - withdrawClientTombstones(createdEmails...) return result, needRestart, nil } diff --git a/internal/web/service/client_crud.go b/internal/web/service/client_crud.go index 0a5eaa40d..cb1885102 100644 --- a/internal/web/service/client_crud.go +++ b/internal/web/service/client_crud.go @@ -237,7 +237,7 @@ func (s *ClientService) Create(inboundSvc *InboundService, payload *ClientCreate // A re-created email is a live identity again: a delete tombstone left // standing makes the next node merge prune the new client's inbound links. withdrawClientTombstones(client.Email) - return needRestart, s.setClientLimitHwidByEmail(nil, client.Email, payload.LimitHwid) + return needRestart, s.setClientLimitHwidByEmail(client.Email, payload.LimitHwid) } // inboundFanoutConcurrency caps how many inbounds one client op applies at @@ -803,7 +803,7 @@ func (s *ClientService) Update(inboundSvc *InboundService, id int, updated model return needRestart, err } - if err := s.setClientLimitHwidByEmail(nil, updated.Email, limitHwid); err != nil { + if err := s.setClientLimitHwidByEmail(updated.Email, limitHwid); err != nil { return needRestart, err } @@ -868,8 +868,7 @@ func (s *ClientService) Delete(inboundSvc *InboundService, id int, keepTraffic b return needRestart, errors.Join(delErrs...) } - db := database.GetDB() - if err := db.Transaction(func(tx *gorm.DB) error { + if err := runSerializedTx(func(tx *gorm.DB) error { if existing.Email != "" { if err := adjustGroupBaselinesForRemovedTraffic(tx, []string{existing.Email}); err != nil { return err diff --git a/internal/web/service/client_hwid.go b/internal/web/service/client_hwid.go index 2a780b8f0..606ea6ac3 100644 --- a/internal/web/service/client_hwid.go +++ b/internal/web/service/client_hwid.go @@ -48,6 +48,8 @@ const ( hwidFingerprintLength = 12 ) +var errClientHwidWriteNotSerialized = errors.New("client HWID write requires the serialized transaction") + type ClientHwidInfo struct { Id int `json:"id"` FirstSeen int64 `json:"firstSeen"` @@ -300,9 +302,16 @@ func (s *ClientService) DeleteClientHwid(email string, id int) error { return nil } -func (s *ClientService) setClientLimitHwidByEmail(tx *gorm.DB, email string, limit int) error { - if tx == nil { - tx = database.GetDB() +// Serialize the limit write and trim with SyncInbound and client deletion. +func (s *ClientService) setClientLimitHwidByEmail(email string, limit int) error { + return runSerializedTx(func(tx *gorm.DB) error { + return s.setClientLimitHwidByEmailTx(tx, email, limit) + }) +} + +func (s *ClientService) setClientLimitHwidByEmailTx(tx *gorm.DB, email string, limit int) error { + if !isSerializedTx(tx) { + return errClientHwidWriteNotSerialized } if limit < 0 { limit = 0 @@ -346,8 +355,8 @@ func trimClientHwidsForSubID(tx *gorm.DB, subID string, limit int) error { } func clearClientHwidsBySubIDTx(tx *gorm.DB, subIDs ...string) error { - if tx == nil { - tx = database.GetDB() + if !isSerializedTx(tx) { + return errClientHwidWriteNotSerialized } clean := make([]string, 0, len(subIDs)) seen := map[string]struct{}{} diff --git a/internal/web/service/client_hwid_test.go b/internal/web/service/client_hwid_test.go index 1d0b6eca9..7c8792886 100644 --- a/internal/web/service/client_hwid_test.go +++ b/internal/web/service/client_hwid_test.go @@ -153,7 +153,7 @@ func TestClientHwidGateRegistersAndBlocks(t *testing.T) { t.Fatalf("updated HWID metadata missing: %#v", list) } - if err := svc.setClientLimitHwidByEmail(nil, rec.Email, 1); err != nil { + if err := svc.setClientLimitHwidByEmail(rec.Email, 1); err != nil { t.Fatalf("lower limit: %v", err) } var count int64 diff --git a/internal/web/service/client_hwid_tx_test.go b/internal/web/service/client_hwid_tx_test.go new file mode 100644 index 000000000..e6d33050d --- /dev/null +++ b/internal/web/service/client_hwid_tx_test.go @@ -0,0 +1,218 @@ +package service + +import ( + "errors" + "fmt" + "path/filepath" + "testing" + "time" + + "github.com/mhsanaei/3x-ui/v3/internal/database" + "github.com/mhsanaei/3x-ui/v3/internal/database/model" + + "gorm.io/gorm" +) + +var errInjectedHwidDelete = errors.New("injected client_hwids delete failure") + +func failHwidDeletes(t *testing.T, db *gorm.DB) { + t.Helper() + if err := db.Callback().Delete().Before("gorm:delete").Register("t:hwid:fail", func(tx *gorm.DB) { + if tx.Statement != nil && tx.Statement.Table == "client_hwids" { + _ = tx.AddError(errInjectedHwidDelete) + } + }); err != nil { + t.Fatalf("register callback: %v", err) + } + t.Cleanup(func() { + if err := db.Callback().Delete().Remove("t:hwid:fail"); err != nil { + t.Fatalf("remove callback: %v", err) + } + }) +} + +func seedHwids(t *testing.T, db *gorm.DB, subID string, n int) { + t.Helper() + now := time.Now().UnixMilli() + rows := make([]model.ClientHwid, 0, n) + for i := range n { + rows = append(rows, model.ClientHwid{ + SubID: subID, HwidHash: fmt.Sprintf("%s-hash-%d", subID, i), + FirstSeen: now, LastSeen: now + int64(i), + }) + } + if err := db.Create(&rows).Error; err != nil { + t.Fatalf("seed client_hwids for %q: %v", subID, err) + } +} + +func assertHwidState(t *testing.T, db *gorm.DB, email string, limit, devices int) { + t.Helper() + var rec model.ClientRecord + if err := db.Where("email = ?", email).First(&rec).Error; err != nil { + t.Fatalf("reload client: %v", err) + } + if rec.LimitHwid != limit { + t.Fatalf("limit_hwid = %d, want %d", rec.LimitHwid, limit) + } + var n int64 + if err := db.Model(&model.ClientHwid{}).Where("sub_id = ?", rec.SubID).Count(&n).Error; err != nil { + t.Fatalf("count client_hwids: %v", err) + } + if n != int64(devices) { + t.Fatalf("client_hwids = %d, want %d", n, devices) + } +} + +func TestSetClientLimitHwidRollsBackFailedTrim(t *testing.T) { + initClientHwidTestDB(t) + db := database.GetDB() + rec := seedHwidClient(t, 5) + seedHwids(t, db, rec.SubID, 3) + failHwidDeletes(t, db) + + err := (&ClientService{}).setClientLimitHwidByEmail(rec.Email, 1) + if !errors.Is(err, errInjectedHwidDelete) { + t.Fatalf("want errInjectedHwidDelete, got: %v", err) + } + assertHwidState(t, db, rec.Email, 5, 3) +} + +func TestClientHwidTxRejectsUnserializedHandle(t *testing.T) { + initClientHwidTestDB(t) + db := database.GetDB() + rec := seedHwidClient(t, 5) + svc := &ClientService{} + + if err := svc.setClientLimitHwidByEmailTx(db, rec.Email, 1); !errors.Is(err, errClientHwidWriteNotSerialized) { + t.Fatalf("bare handle error = %v, want errClientHwidWriteNotSerialized", err) + } + if err := runSerializedTx(func(tx *gorm.DB) error { + return svc.setClientLimitHwidByEmailTx(tx, rec.Email, 1) + }); err != nil { + t.Fatalf("serialized update: %v", err) + } + assertHwidState(t, db, rec.Email, 1, 0) +} + +func TestBulkAdjustHwidRollsBackFailedTrim(t *testing.T) { + setupBulkDB(t) + db := database.GetDB() + rec := &model.ClientRecord{Email: "bulk-hwid@x", SubID: "bulk-sub", Enable: true, LimitHwid: 5} + if err := db.Create(rec).Error; err != nil { + t.Fatalf("seed client: %v", err) + } + seedHwids(t, db, rec.SubID, 3) + failHwidDeletes(t, db) + + limit := 1 + res, _, err := (&ClientService{}).BulkAdjust(&InboundService{}, []string{rec.Email}, 0, 0, "", &limit, "") + if err != nil { + t.Fatalf("BulkAdjust: %v", err) + } + if len(res.Skipped) != 1 || res.Skipped[0].Reason != errInjectedHwidDelete.Error() { + t.Fatalf("skipped = %+v, want injected failure", res.Skipped) + } + assertHwidState(t, db, rec.Email, 5, 3) +} + +func TestBulkCreateWithdrawsTombstoneWhenHwidTrimFails(t *testing.T) { + setupBulkDB(t) + StartTrafficWriter() + t.Cleanup(StopTrafficWriter) + db := database.GetDB() + const email = "reborn-bulk@x" + const subID = "reborn-bulk-sub" + tombstoneClientEmail(email) + t.Cleanup(func() { withdrawClientTombstones(email) }) + seedHwids(t, db, subID, 3) + failHwidDeletes(t, db) + ib := mkInbound(t, 30441, model.VLESS, `{"clients":[]}`) + + res, _, err := (&ClientService{}).BulkCreate(&InboundService{}, []ClientCreatePayload{{ + Client: model.Client{ + Email: email, SubID: subID, ID: "aaaaaaaa-bbbb-4ccc-8ddd-eeeeeeeeeeee", Enable: true, + }, + InboundIds: []int{ib.Id}, LimitHwid: 1, + }}) + if err != nil { + t.Fatalf("BulkCreate: %v", err) + } + if len(res.Skipped) != 1 || res.Skipped[0].Reason != errInjectedHwidDelete.Error() { + t.Fatalf("skipped = %+v, want injected HWID failure", res.Skipped) + } + if isClientEmailTombstoned(email) { + t.Fatal("live bulk-created client retained a delete tombstone") + } +} + +func TestSetClientLimitHwidIsSerializedWithSyncInbound(t *testing.T) { + db := durablePostgresDB(t) + if err := db.Exec("TRUNCATE client_hwids, clients RESTART IDENTITY CASCADE").Error; err != nil { + t.Fatalf("reset tables: %v", err) + } + rec := seedHwidClient(t, 5) + seedHwids(t, db, rec.SubID, 3) + StartTrafficWriter() + t.Cleanup(StopTrafficWriter) + + read := make(chan struct{}) + release := make(chan struct{}) + staleDone := make(chan error, 1) + go func() { + staleDone <- runSerializedTx(func(tx *gorm.DB) error { + var stale model.ClientRecord + if err := tx.Where("email = ?", rec.Email).First(&stale).Error; err != nil { + return err + } + close(read) + <-release + return tx.Save(&stale).Error + }) + }() + <-read + + limitDone := make(chan error, 1) + go func() { limitDone <- (&ClientService{}).setClientLimitHwidByEmail(rec.Email, 1) }() + time.Sleep(100 * time.Millisecond) + close(release) + if err := <-staleDone; err != nil { + t.Fatalf("stale SyncInbound write: %v", err) + } + if err := <-limitDone; err != nil { + t.Fatalf("set limit: %v", err) + } + assertHwidState(t, db, rec.Email, 1, 1) +} + +func BenchmarkSetClientLimitHwidSerialized(b *testing.B) { + dbDir := b.TempDir() + b.Setenv("XUI_DB_FOLDER", dbDir) + if err := database.InitDB(filepath.Join(dbDir, "x-ui.db")); err != nil { + b.Fatalf("InitDB: %v", err) + } + b.Cleanup(func() { _ = database.CloseDB() }) + StartTrafficWriter() + b.Cleanup(StopTrafficWriter) + db := database.GetDB() + emails := make([]string, 100) + for i := range emails { + emails[i] = fmt.Sprintf("bench-%03d@x", i) + rec := &model.ClientRecord{Email: emails[i], SubID: fmt.Sprintf("bench-sub-%03d", i), Enable: true} + if err := db.Create(rec).Error; err != nil { + b.Fatalf("seed client: %v", err) + } + } + svc := &ClientService{} + for _, count := range []int{1, 100} { + b.Run(fmt.Sprintf("clients_%d", count), func(b *testing.B) { + for range b.N { + for i := range count { + if err := svc.setClientLimitHwidByEmail(emails[i], 2); err != nil { + b.Fatal(err) + } + } + } + }) + } +} diff --git a/internal/web/service/traffic_writer.go b/internal/web/service/traffic_writer.go index be0cea8bd..4cf4c3466 100644 --- a/internal/web/service/traffic_writer.go +++ b/internal/web/service/traffic_writer.go @@ -23,6 +23,8 @@ type trafficWriteRequest struct { done chan error } +type serializedTxContextKey struct{} + var ( twMu sync.Mutex twQueue chan *trafficWriteRequest @@ -125,10 +127,21 @@ func runTrafficWriter(ctx context.Context, queue chan *trafficWriteRequest, done // timeout. Apply runtime changes after this returns. func runSerializedTx(fn func(tx *gorm.DB) error) error { return submitTrafficWrite(func() error { - return database.GetDB().Transaction(fn) + return database.GetDB().Transaction(func(tx *gorm.DB) error { + ctx := context.WithValue(tx.Statement.Context, serializedTxContextKey{}, true) + return fn(tx.WithContext(ctx)) + }) }) } +func isSerializedTx(tx *gorm.DB) bool { + if tx == nil || tx.Statement == nil || tx.Statement.Context == nil { + return false + } + active, _ := tx.Statement.Context.Value(serializedTxContextKey{}).(bool) + return active +} + func safeApply(fn func() error) (err error) { defer func() { if r := recover(); r != nil {