ソースを参照

fix(hwid): serialize the device-limit write with its trim (#6591)

setClientLimitHwidByEmail wrote clients.limit_hwid and then trimmed client_hwids as two independent statements. A traffic-cycle Save that read the record before the limit changed could write the stale value back after it, and a failed trim committed the new limit anyway.

Both halves now run inside runSerializedTx, the transaction the traffic writer already owns. setClientLimitHwidByEmailTx and clearClientHwidsBySubIDTx refuse a handle that is not that transaction (errClientHwidWriteNotSerialized) instead of falling back to the shared handle. Client delete moves onto the same writer, and BulkCreate withdraws a re-created client's tombstone before applying its optional HWID limit.

TestSetClientLimitHwidIsSerializedWithSyncInbound holds a stale traffic-cycle Save open across the limit change and fails without the serialization (limit_hwid = 5, want 1).
n0ctal 9 時間 前
親
コミット
c0c0136037

+ 5 - 5
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
 }
 

+ 3 - 4
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

+ 14 - 5
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{}{}

+ 1 - 1
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

+ 218 - 0
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)
+					}
+				}
+			}
+		})
+	}
+}

+ 14 - 1
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 {