修复(#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()),
)
}