mirror of
https://github.com/MHSanaei/3x-ui.git
synced 2026-09-27 01:32:08 +03:00
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).
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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{}{}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user