修复(#12): 优化用户列表限速计算性能 (#14)

Co-authored-by: multica-agent <github@multica.ai>
This commit is contained in:
2026-06-07 23:29:15 -07:00
committed by GitHub
parent b98d718f3c
commit 24caf58987
2 changed files with 517 additions and 17 deletions
+255 -17
View File
@@ -19,6 +19,7 @@ import (
"github.com/perfect-panel/server/pkg/tool"
"github.com/perfect-panel/server/pkg/uuidx"
"github.com/perfect-panel/server/pkg/xerr"
"github.com/redis/go-redis/v9"
)
type GetServerUserListLogic struct {
@@ -27,6 +28,46 @@ type GetServerUserListLogic struct {
svcCtx *svc.ServiceContext
}
type serverUserListPerfStats struct {
serverID int64
protocol string
cacheHit bool
nodesCount int
subsCount int
usersCount int
speedLimitZeroCount int
speedLimitPositiveCount int
trafficCalcMS int64
startedAt time.Time
}
type serverUserSpeedLimitCandidate struct {
userID int64
userSubscribeID int64
baseSpeed int64
trafficLimit string
}
type serverUserTrafficWindow struct {
statType string
statValue int64
start time.Time
end time.Time
}
type serverUserTrafficUsageKey struct {
userID int64
userSubscribeID int64
statType string
statValue int64
}
type serverUserTrafficUsage struct {
userID int64
userSubscribeID int64
usedGB float64
}
// NewGetServerUserListLogic Get user list
func NewGetServerUserListLogic(ctx *gin.Context, svcCtx *svc.ServiceContext) *GetServerUserListLogic {
return &GetServerUserListLogic{
@@ -37,15 +78,27 @@ func NewGetServerUserListLogic(ctx *gin.Context, svcCtx *svc.ServiceContext) *Ge
}
func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListRequest) (resp *types.GetServerUserListResponse, err error) {
startedAt := time.Now()
protocolRequest := normalizeServerUserListProtocol(req.Protocol)
stats := serverUserListPerfStats{
serverID: req.ServerId,
protocol: protocolRequest,
startedAt: startedAt,
cacheHit: false,
}
cacheKey := fmt.Sprintf("%s%d:%s", node.ServerUserListCacheKey, req.ServerId, protocolRequest)
cache, err := l.svcCtx.Redis.Get(l.ctx, cacheKey).Result()
if err != nil && err != redis.Nil {
l.Errorw("[ServerUserListCacheKey] redis get error", logger.Field("error", err.Error()))
}
if cache != "" {
stats.cacheHit = true
etag := tool.GenerateETag([]byte(cache))
resp = &types.GetServerUserListResponse{}
// Check If-None-Match header
if match := l.ctx.GetHeader("If-None-Match"); match == etag {
l.logServerUserListPerf(stats)
return nil, xerr.StatusNotModified
}
l.ctx.Header("ETag", etag)
@@ -54,6 +107,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
l.Errorw("[ServerUserListCacheKey] json unmarshal error", logger.Field("error", err.Error()))
return nil, err
}
stats.usersCount = len(resp.Users)
stats.recordSpeedLimits(resp.Users)
l.logServerUserListPerf(stats)
return resp, nil
}
server, err := l.svcCtx.NodeModel.FindOneServer(l.ctx, req.ServerId)
@@ -72,12 +128,16 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
l.Errorw("FilterNodeList error", logger.Field("error", err.Error()))
return nil, err
}
stats.nodesCount = len(nodes)
if len(nodes) == 0 {
l.Errorw("[ServerUserList] fallback: no nodes matched server+protocol, returning placeholder without cache",
logger.Field("server_id", req.ServerId),
logger.Field("protocol", req.Protocol),
)
stats.usersCount = 1
stats.speedLimitZeroCount = 1
l.logServerUserListPerf(stats)
return &types.GetServerUserListResponse{
Users: []types.ServerUser{
{
@@ -149,6 +209,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
logger.Field("server_id", req.ServerId),
logger.Field("protocol", req.Protocol),
)
stats.usersCount = 1
stats.speedLimitZeroCount = 1
l.logServerUserListPerf(stats)
return &types.GetServerUserListResponse{
Users: []types.ServerUser{
{
@@ -158,7 +221,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
},
}, nil
}
stats.subsCount = len(subs)
users := make([]types.ServerUser, 0)
speedCandidates := make(map[int64]serverUserSpeedLimitCandidate)
for _, sub := range subs {
data, err := l.svcCtx.UserModel.FindUsersSubscribeBySubscribeId(l.ctx, sub.Id)
if err != nil {
@@ -169,17 +234,29 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
continue
}
// 计算该用户的实际限速值(考虑按量限速规则)
effectiveSpeedLimit := l.calculateEffectiveSpeedLimit(sub, datum)
baseSpeed, trafficLimit := serverUserSpeedLimitInputs(sub, datum)
speedCandidates[datum.Id] = serverUserSpeedLimitCandidate{
userID: datum.UserId,
userSubscribeID: datum.Id,
baseSpeed: baseSpeed,
trafficLimit: trafficLimit,
}
users = append(users, types.ServerUser{
Id: datum.Id,
UUID: datum.UUID,
SpeedLimit: effectiveSpeedLimit,
SpeedLimit: baseSpeed,
DeviceLimit: sub.DeviceLimit,
})
}
}
trafficCalcStartedAt := time.Now()
speedLimits := l.calculateServerUserSpeedLimits(speedCandidates)
stats.trafficCalcMS = time.Since(trafficCalcStartedAt).Milliseconds()
for i := range users {
if speedLimit, ok := speedLimits[users[i].Id]; ok {
users[i].SpeedLimit = speedLimit
}
}
// 处理过期订阅用户:如果当前节点属于过期节点组,添加符合条件的过期用户
if len(nodeGroupIds) > 0 {
@@ -197,6 +274,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
logger.Field("server_id", req.ServerId),
logger.Field("protocol", req.Protocol),
)
stats.usersCount = 1
stats.speedLimitZeroCount = 1
l.logServerUserListPerf(stats)
return &types.GetServerUserListResponse{
Users: []types.ServerUser{
{
@@ -209,6 +289,8 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
resp = &types.GetServerUserListResponse{
Users: users,
}
stats.usersCount = len(users)
stats.recordSpeedLimits(users)
val, _ := json.Marshal(resp)
etag := tool.GenerateETag(val)
l.ctx.Header("ETag", etag)
@@ -218,8 +300,10 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
}
// Check If-None-Match header
if match := l.ctx.GetHeader("If-None-Match"); match == etag {
l.logServerUserListPerf(stats)
return nil, xerr.StatusNotModified
}
l.logServerUserListPerf(stats)
return resp, nil
}
@@ -334,8 +418,7 @@ func (l *GetServerUserListLogic) canUseExpiredNodeGroup(userSub *user.Subscribe,
return true
}
// calculateEffectiveSpeedLimit 计算用户的实际限速值(考虑按量限速规则)
func (l *GetServerUserListLogic) calculateEffectiveSpeedLimit(sub *subscribe.Subscribe, userSub *user.Subscribe) int64 {
func serverUserSpeedLimitInputs(sub *subscribe.Subscribe, userSub *user.Subscribe) (int64, string) {
baseSpeed := sub.SpeedLimit
if userSub.SpeedLimit > 0 {
baseSpeed = userSub.SpeedLimit
@@ -346,15 +429,170 @@ func (l *GetServerUserListLogic) calculateEffectiveSpeedLimit(sub *subscribe.Sub
trafficLimit = *userSub.TrafficLimit
}
result := speedlimit.CalculateWithCache(
l.ctx.Request.Context(),
l.svcCtx.Redis,
l.svcCtx.DB,
userSub.UserId,
userSub.Id,
baseSpeed,
trafficLimit,
30*time.Second,
)
return result.EffectiveSpeed
return baseSpeed, trafficLimit
}
func (l *GetServerUserListLogic) calculateServerUserSpeedLimits(candidates map[int64]serverUserSpeedLimitCandidate) map[int64]int64 {
speedLimits := make(map[int64]int64, len(candidates))
rulesByUserSubscribeID := make(map[int64][]speedlimit.TrafficLimitRule)
windowByKey := make(map[string]serverUserTrafficWindow)
now := time.Now()
for userSubscribeID, candidate := range candidates {
speedLimits[userSubscribeID] = candidate.baseSpeed
if candidate.trafficLimit == "" {
continue
}
var rules []speedlimit.TrafficLimitRule
if err := json.Unmarshal([]byte(candidate.trafficLimit), &rules); err != nil || len(rules) == 0 {
continue
}
rulesByUserSubscribeID[userSubscribeID] = rules
for _, rule := range rules {
window, ok := serverUserTrafficRuleWindow(rule, now)
if !ok {
continue
}
windowByKey[serverUserTrafficWindowKey(rule.StatType, rule.StatValue)] = window
}
}
if len(rulesByUserSubscribeID) == 0 || len(windowByKey) == 0 {
return speedLimits
}
usageByKey := make(map[serverUserTrafficUsageKey]float64)
for _, window := range windowByKey {
trafficUsage, err := l.queryServerUserTrafficUsage(candidates, window)
if err != nil {
l.Errorw("[ServerUserList] batch traffic usage query failed",
logger.Field("error", err.Error()),
logger.Field("stat_type", window.statType),
logger.Field("stat_value", window.statValue),
)
continue
}
for _, usage := range trafficUsage {
usageByKey[serverUserTrafficUsageKey{
userID: usage.userID,
userSubscribeID: usage.userSubscribeID,
statType: window.statType,
statValue: window.statValue,
}] = usage.usedGB
}
}
for userSubscribeID, rules := range rulesByUserSubscribeID {
candidate := candidates[userSubscribeID]
for _, rule := range rules {
if rule.SpeedLimit <= 0 {
continue
}
if _, ok := serverUserTrafficRuleWindow(rule, now); !ok {
continue
}
usedGB := usageByKey[serverUserTrafficUsageKey{
userID: candidate.userID,
userSubscribeID: candidate.userSubscribeID,
statType: rule.StatType,
statValue: rule.StatValue,
}]
if usedGB < float64(rule.TrafficUsage) {
continue
}
current := speedLimits[userSubscribeID]
if current == 0 || rule.SpeedLimit < current {
speedLimits[userSubscribeID] = rule.SpeedLimit
}
}
}
return speedLimits
}
func (l *GetServerUserListLogic) queryServerUserTrafficUsage(
candidates map[int64]serverUserSpeedLimitCandidate,
window serverUserTrafficWindow,
) ([]serverUserTrafficUsage, error) {
userSubscribeIDs := make([]int64, 0, len(candidates))
for userSubscribeID := range candidates {
userSubscribeIDs = append(userSubscribeIDs, userSubscribeID)
}
var rows []struct {
UserID int64
UserSubscribeID int64
Upload int64
Download int64
}
err := l.svcCtx.DB.WithContext(l.ctx.Request.Context()).
Table("traffic_log").
Select("user_id, subscribe_id AS user_subscribe_id, COALESCE(SUM(upload), 0) AS upload, COALESCE(SUM(download), 0) AS download").
Where("subscribe_id IN ? AND timestamp >= ? AND timestamp < ?", userSubscribeIDs, window.start, window.end).
Group("user_id, subscribe_id").
Scan(&rows).Error
if err != nil {
return nil, err
}
trafficUsage := make([]serverUserTrafficUsage, 0, len(rows))
for _, row := range rows {
trafficUsage = append(trafficUsage, serverUserTrafficUsage{
userID: row.UserID,
userSubscribeID: row.UserSubscribeID,
usedGB: float64(row.Upload+row.Download) / (1024 * 1024 * 1024),
})
}
return trafficUsage, nil
}
func serverUserTrafficRuleWindow(rule speedlimit.TrafficLimitRule, now time.Time) (serverUserTrafficWindow, bool) {
if rule.StatValue <= 0 {
return serverUserTrafficWindow{}, false
}
window := serverUserTrafficWindow{
statType: rule.StatType,
statValue: rule.StatValue,
end: now,
}
switch rule.StatType {
case "hour":
window.start = now.Add(-time.Duration(rule.StatValue) * time.Hour)
case "day":
window.start = now.AddDate(0, 0, -int(rule.StatValue))
default:
return serverUserTrafficWindow{}, false
}
return window, true
}
func serverUserTrafficWindowKey(statType string, statValue int64) string {
return fmt.Sprintf("%s:%d", statType, statValue)
}
func (s *serverUserListPerfStats) recordSpeedLimits(users []types.ServerUser) {
for _, user := range users {
if user.SpeedLimit > 0 {
s.speedLimitPositiveCount++
continue
}
s.speedLimitZeroCount++
}
}
func (l *GetServerUserListLogic) logServerUserListPerf(stats serverUserListPerfStats) {
l.Infow("[ServerUserList] performance",
logger.Field("server_id", stats.serverID),
logger.Field("protocol", stats.protocol),
logger.Field("cache_hit", stats.cacheHit),
logger.Field("nodes_count", stats.nodesCount),
logger.Field("subs_count", stats.subsCount),
logger.Field("users_count", stats.usersCount),
logger.Field("speed_limit_0_count", stats.speedLimitZeroCount),
logger.Field("speed_limit_positive_count", stats.speedLimitPositiveCount),
logger.Field("traffic_calc_ms", stats.trafficCalcMS),
logger.Field("total_ms", time.Since(stats.startedAt).Milliseconds()),
)
}
@@ -1,11 +1,24 @@
package server
import (
"context"
"fmt"
"net/http"
"regexp"
"slices"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"github.com/gin-gonic/gin"
"github.com/perfect-panel/server/internal/model/node"
"github.com/perfect-panel/server/internal/model/subscribe"
"github.com/perfect-panel/server/internal/model/user"
"github.com/perfect-panel/server/internal/svc"
"github.com/perfect-panel/server/pkg/logger"
"github.com/perfect-panel/server/pkg/speedlimit"
"gorm.io/driver/mysql"
"gorm.io/gorm"
)
// TestNormalizeServerUserListProtocol 验证客户端 hysteria2 兼容字段被映射回
@@ -76,3 +89,252 @@ func TestServerUserListCacheKey_ShapeMatchesDelPath(t *testing.T) {
t.Fatalf("write-path key %q not present in Del-path enumeration %v", writeShape, delKeys)
}
}
func TestServerUserSpeedLimitInputs(t *testing.T) {
planTrafficLimit := `[{"stat_type":"day","stat_value":1,"traffic_usage":10,"speed_limit":5}]`
userTrafficLimit := `[{"stat_type":"day","stat_value":1,"traffic_usage":1,"speed_limit":3}]`
cases := []struct {
name string
sub *subscribe.Subscribe
userSub *user.Subscribe
wantSpeedLimit int64
wantTrafficLimit string
}{
{
name: "未设置用户覆盖时返回套餐基础限速",
sub: &subscribe.Subscribe{
SpeedLimit: 100,
TrafficLimit: planTrafficLimit,
},
userSub: &user.Subscribe{},
wantSpeedLimit: 100,
wantTrafficLimit: planTrafficLimit,
},
{
name: "用户 speed_limit 覆盖套餐基础限速",
sub: &subscribe.Subscribe{
SpeedLimit: 100,
TrafficLimit: planTrafficLimit,
},
userSub: &user.Subscribe{
SpeedLimit: 30,
},
wantSpeedLimit: 30,
wantTrafficLimit: planTrafficLimit,
},
{
name: "用户 traffic_limit 覆盖套餐规则",
sub: &subscribe.Subscribe{
SpeedLimit: 100,
TrafficLimit: planTrafficLimit,
},
userSub: &user.Subscribe{
TrafficLimit: &userTrafficLimit,
},
wantSpeedLimit: 100,
wantTrafficLimit: userTrafficLimit,
},
{
name: "用户 traffic_limit 为空时保留套餐规则",
sub: &subscribe.Subscribe{
SpeedLimit: 100,
TrafficLimit: planTrafficLimit,
},
userSub: &user.Subscribe{
TrafficLimit: ptrString(""),
},
wantSpeedLimit: 100,
wantTrafficLimit: planTrafficLimit,
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
gotSpeedLimit, gotTrafficLimit := serverUserSpeedLimitInputs(tc.sub, tc.userSub)
if gotSpeedLimit != tc.wantSpeedLimit {
t.Fatalf("speedLimit = %d, want %d", gotSpeedLimit, tc.wantSpeedLimit)
}
if gotTrafficLimit != tc.wantTrafficLimit {
t.Fatalf("trafficLimit = %q, want %q", gotTrafficLimit, tc.wantTrafficLimit)
}
})
}
}
func TestCalculateServerUserSpeedLimits_BatchTrafficRules(t *testing.T) {
db, mock, cleanup := newServerUserListTestDB(t)
defer cleanup()
ctx, _ := gin.CreateTestContext(nil)
ctx.Request = httptestNewRequest()
logic := &GetServerUserListLogic{
Logger: logger.WithContext(ctx.Request.Context()),
ctx: ctx,
svcCtx: &svc.ServiceContext{DB: db},
}
trafficLimit := `[{"stat_type":"day","stat_value":1,"traffic_usage":10,"speed_limit":5}]`
candidates := map[int64]serverUserSpeedLimitCandidate{
101: {userID: 1001, userSubscribeID: 101, baseSpeed: 0, trafficLimit: trafficLimit},
102: {userID: 1002, userSubscribeID: 102, baseSpeed: 0, trafficLimit: trafficLimit},
103: {userID: 1003, userSubscribeID: 103, baseSpeed: 20, trafficLimit: trafficLimit},
}
mock.ExpectQuery("FROM `traffic_log`").
WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg()).
WillReturnRows(sqlmock.NewRows([]string{"user_id", "user_subscribe_id", "upload", "download"}).
AddRow(1001, 101, gb(4), gb(5)).
AddRow(1002, 102, gb(4), gb(6)).
AddRow(1003, 103, gb(12), int64(0)))
got := logic.calculateServerUserSpeedLimits(candidates)
assertSpeedLimit(t, got, 101, 0)
assertSpeedLimit(t, got, 102, 5)
assertSpeedLimit(t, got, 103, 5)
assertServerUserListExpectations(t, mock)
}
func TestCalculateServerUserSpeedLimits_UserTrafficLimitOverride(t *testing.T) {
db, mock, cleanup := newServerUserListTestDB(t)
defer cleanup()
ctx, _ := gin.CreateTestContext(nil)
ctx.Request = httptestNewRequest()
logic := &GetServerUserListLogic{
Logger: logger.WithContext(ctx.Request.Context()),
ctx: ctx,
svcCtx: &svc.ServiceContext{DB: db},
}
candidates := map[int64]serverUserSpeedLimitCandidate{
201: {
userID: 2001,
userSubscribeID: 201,
baseSpeed: 50,
trafficLimit: `[{"stat_type":"day","stat_value":1,"traffic_usage":2,"speed_limit":3}]`,
},
202: {
userID: 2002,
userSubscribeID: 202,
baseSpeed: 50,
trafficLimit: "",
},
}
mock.ExpectQuery("FROM `traffic_log`").
WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg()).
WillReturnRows(sqlmock.NewRows([]string{"user_id", "user_subscribe_id", "upload", "download"}).
AddRow(2001, 201, gb(2), int64(0)))
got := logic.calculateServerUserSpeedLimits(candidates)
assertSpeedLimit(t, got, 201, 3)
assertSpeedLimit(t, got, 202, 50)
assertServerUserListExpectations(t, mock)
}
func TestServerUserTrafficRuleWindow(t *testing.T) {
now := time.Date(2026, 6, 8, 12, 0, 0, 0, time.UTC)
cases := []struct {
name string
rule speedlimit.TrafficLimitRule
want time.Time
ok bool
}{
{
name: "hour window",
rule: speedlimit.TrafficLimitRule{StatType: "hour", StatValue: 2},
want: now.Add(-2 * time.Hour),
ok: true,
},
{
name: "day window",
rule: speedlimit.TrafficLimitRule{StatType: "day", StatValue: 1},
want: now.AddDate(0, 0, -1),
ok: true,
},
{
name: "invalid stat type",
rule: speedlimit.TrafficLimitRule{StatType: "week", StatValue: 1},
ok: false,
},
{
name: "invalid stat value",
rule: speedlimit.TrafficLimitRule{StatType: "day", StatValue: 0},
ok: false,
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
got, ok := serverUserTrafficRuleWindow(tc.rule, now)
if ok != tc.ok {
t.Fatalf("ok = %v, want %v", ok, tc.ok)
}
if !ok {
return
}
if !got.start.Equal(tc.want) || !got.end.Equal(now) {
t.Fatalf("window = [%s, %s), want [%s, %s)", got.start, got.end, tc.want, now)
}
})
}
}
func newServerUserListTestDB(t *testing.T) (*gorm.DB, sqlmock.Sqlmock, func()) {
t.Helper()
sqlDB, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherFunc(func(expectedSQL, actualSQL string) error {
if matched, _ := regexp.MatchString(expectedSQL, actualSQL); matched {
return nil
}
t.Fatalf("unexpected SQL\nexpected pattern: %s\nactual: %s", expectedSQL, actualSQL)
return nil
})))
if err != nil {
t.Fatalf("create sqlmock: %v", err)
}
db, err := gorm.Open(mysql.New(mysql.Config{
Conn: sqlDB,
SkipInitializeWithVersion: true,
}), &gorm.Config{})
if err != nil {
_ = sqlDB.Close()
t.Fatalf("open gorm db: %v", err)
}
return db, mock, func() {
_ = sqlDB.Close()
}
}
func assertServerUserListExpectations(t *testing.T, mock sqlmock.Sqlmock) {
t.Helper()
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatalf("unmet sql expectations: %v", err)
}
}
func assertSpeedLimit(t *testing.T, got map[int64]int64, userSubscribeID int64, want int64) {
t.Helper()
if got[userSubscribeID] != want {
t.Fatalf("speed limit for user_subscribe_id=%d = %d, want %d", userSubscribeID, got[userSubscribeID], want)
}
}
func gb(value int64) int64 {
return value * 1024 * 1024 * 1024
}
func ptrString(value string) *string {
return &value
}
func httptestNewRequest() *http.Request {
req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, "/v1/server/user", nil)
return req
}