Co-authored-by: multica-agent <github@multica.ai>
This commit is contained in:
@@ -19,6 +19,7 @@ import (
|
|||||||
"github.com/perfect-panel/server/pkg/tool"
|
"github.com/perfect-panel/server/pkg/tool"
|
||||||
"github.com/perfect-panel/server/pkg/uuidx"
|
"github.com/perfect-panel/server/pkg/uuidx"
|
||||||
"github.com/perfect-panel/server/pkg/xerr"
|
"github.com/perfect-panel/server/pkg/xerr"
|
||||||
|
"github.com/redis/go-redis/v9"
|
||||||
)
|
)
|
||||||
|
|
||||||
type GetServerUserListLogic struct {
|
type GetServerUserListLogic struct {
|
||||||
@@ -27,6 +28,46 @@ type GetServerUserListLogic struct {
|
|||||||
svcCtx *svc.ServiceContext
|
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
|
// NewGetServerUserListLogic Get user list
|
||||||
func NewGetServerUserListLogic(ctx *gin.Context, svcCtx *svc.ServiceContext) *GetServerUserListLogic {
|
func NewGetServerUserListLogic(ctx *gin.Context, svcCtx *svc.ServiceContext) *GetServerUserListLogic {
|
||||||
return &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) {
|
func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListRequest) (resp *types.GetServerUserListResponse, err error) {
|
||||||
|
startedAt := time.Now()
|
||||||
protocolRequest := normalizeServerUserListProtocol(req.Protocol)
|
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)
|
cacheKey := fmt.Sprintf("%s%d:%s", node.ServerUserListCacheKey, req.ServerId, protocolRequest)
|
||||||
cache, err := l.svcCtx.Redis.Get(l.ctx, cacheKey).Result()
|
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 != "" {
|
if cache != "" {
|
||||||
|
stats.cacheHit = true
|
||||||
etag := tool.GenerateETag([]byte(cache))
|
etag := tool.GenerateETag([]byte(cache))
|
||||||
resp = &types.GetServerUserListResponse{}
|
resp = &types.GetServerUserListResponse{}
|
||||||
// Check If-None-Match header
|
// Check If-None-Match header
|
||||||
if match := l.ctx.GetHeader("If-None-Match"); match == etag {
|
if match := l.ctx.GetHeader("If-None-Match"); match == etag {
|
||||||
|
l.logServerUserListPerf(stats)
|
||||||
return nil, xerr.StatusNotModified
|
return nil, xerr.StatusNotModified
|
||||||
}
|
}
|
||||||
l.ctx.Header("ETag", etag)
|
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()))
|
l.Errorw("[ServerUserListCacheKey] json unmarshal error", logger.Field("error", err.Error()))
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
stats.usersCount = len(resp.Users)
|
||||||
|
stats.recordSpeedLimits(resp.Users)
|
||||||
|
l.logServerUserListPerf(stats)
|
||||||
return resp, nil
|
return resp, nil
|
||||||
}
|
}
|
||||||
server, err := l.svcCtx.NodeModel.FindOneServer(l.ctx, req.ServerId)
|
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()))
|
l.Errorw("FilterNodeList error", logger.Field("error", err.Error()))
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
stats.nodesCount = len(nodes)
|
||||||
|
|
||||||
if len(nodes) == 0 {
|
if len(nodes) == 0 {
|
||||||
l.Errorw("[ServerUserList] fallback: no nodes matched server+protocol, returning placeholder without cache",
|
l.Errorw("[ServerUserList] fallback: no nodes matched server+protocol, returning placeholder without cache",
|
||||||
logger.Field("server_id", req.ServerId),
|
logger.Field("server_id", req.ServerId),
|
||||||
logger.Field("protocol", req.Protocol),
|
logger.Field("protocol", req.Protocol),
|
||||||
)
|
)
|
||||||
|
stats.usersCount = 1
|
||||||
|
stats.speedLimitZeroCount = 1
|
||||||
|
l.logServerUserListPerf(stats)
|
||||||
return &types.GetServerUserListResponse{
|
return &types.GetServerUserListResponse{
|
||||||
Users: []types.ServerUser{
|
Users: []types.ServerUser{
|
||||||
{
|
{
|
||||||
@@ -149,6 +209,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
|
|||||||
logger.Field("server_id", req.ServerId),
|
logger.Field("server_id", req.ServerId),
|
||||||
logger.Field("protocol", req.Protocol),
|
logger.Field("protocol", req.Protocol),
|
||||||
)
|
)
|
||||||
|
stats.usersCount = 1
|
||||||
|
stats.speedLimitZeroCount = 1
|
||||||
|
l.logServerUserListPerf(stats)
|
||||||
return &types.GetServerUserListResponse{
|
return &types.GetServerUserListResponse{
|
||||||
Users: []types.ServerUser{
|
Users: []types.ServerUser{
|
||||||
{
|
{
|
||||||
@@ -158,7 +221,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
|
|||||||
},
|
},
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
stats.subsCount = len(subs)
|
||||||
users := make([]types.ServerUser, 0)
|
users := make([]types.ServerUser, 0)
|
||||||
|
speedCandidates := make(map[int64]serverUserSpeedLimitCandidate)
|
||||||
for _, sub := range subs {
|
for _, sub := range subs {
|
||||||
data, err := l.svcCtx.UserModel.FindUsersSubscribeBySubscribeId(l.ctx, sub.Id)
|
data, err := l.svcCtx.UserModel.FindUsersSubscribeBySubscribeId(l.ctx, sub.Id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -169,17 +234,29 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// 计算该用户的实际限速值(考虑按量限速规则)
|
baseSpeed, trafficLimit := serverUserSpeedLimitInputs(sub, datum)
|
||||||
effectiveSpeedLimit := l.calculateEffectiveSpeedLimit(sub, datum)
|
speedCandidates[datum.Id] = serverUserSpeedLimitCandidate{
|
||||||
|
userID: datum.UserId,
|
||||||
|
userSubscribeID: datum.Id,
|
||||||
|
baseSpeed: baseSpeed,
|
||||||
|
trafficLimit: trafficLimit,
|
||||||
|
}
|
||||||
users = append(users, types.ServerUser{
|
users = append(users, types.ServerUser{
|
||||||
Id: datum.Id,
|
Id: datum.Id,
|
||||||
UUID: datum.UUID,
|
UUID: datum.UUID,
|
||||||
SpeedLimit: effectiveSpeedLimit,
|
SpeedLimit: baseSpeed,
|
||||||
DeviceLimit: sub.DeviceLimit,
|
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 {
|
if len(nodeGroupIds) > 0 {
|
||||||
@@ -197,6 +274,9 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
|
|||||||
logger.Field("server_id", req.ServerId),
|
logger.Field("server_id", req.ServerId),
|
||||||
logger.Field("protocol", req.Protocol),
|
logger.Field("protocol", req.Protocol),
|
||||||
)
|
)
|
||||||
|
stats.usersCount = 1
|
||||||
|
stats.speedLimitZeroCount = 1
|
||||||
|
l.logServerUserListPerf(stats)
|
||||||
return &types.GetServerUserListResponse{
|
return &types.GetServerUserListResponse{
|
||||||
Users: []types.ServerUser{
|
Users: []types.ServerUser{
|
||||||
{
|
{
|
||||||
@@ -209,6 +289,8 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
|
|||||||
resp = &types.GetServerUserListResponse{
|
resp = &types.GetServerUserListResponse{
|
||||||
Users: users,
|
Users: users,
|
||||||
}
|
}
|
||||||
|
stats.usersCount = len(users)
|
||||||
|
stats.recordSpeedLimits(users)
|
||||||
val, _ := json.Marshal(resp)
|
val, _ := json.Marshal(resp)
|
||||||
etag := tool.GenerateETag(val)
|
etag := tool.GenerateETag(val)
|
||||||
l.ctx.Header("ETag", etag)
|
l.ctx.Header("ETag", etag)
|
||||||
@@ -218,8 +300,10 @@ func (l *GetServerUserListLogic) GetServerUserList(req *types.GetServerUserListR
|
|||||||
}
|
}
|
||||||
// Check If-None-Match header
|
// Check If-None-Match header
|
||||||
if match := l.ctx.GetHeader("If-None-Match"); match == etag {
|
if match := l.ctx.GetHeader("If-None-Match"); match == etag {
|
||||||
|
l.logServerUserListPerf(stats)
|
||||||
return nil, xerr.StatusNotModified
|
return nil, xerr.StatusNotModified
|
||||||
}
|
}
|
||||||
|
l.logServerUserListPerf(stats)
|
||||||
return resp, nil
|
return resp, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -334,8 +418,7 @@ func (l *GetServerUserListLogic) canUseExpiredNodeGroup(userSub *user.Subscribe,
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// calculateEffectiveSpeedLimit 计算用户的实际限速值(考虑按量限速规则)
|
func serverUserSpeedLimitInputs(sub *subscribe.Subscribe, userSub *user.Subscribe) (int64, string) {
|
||||||
func (l *GetServerUserListLogic) calculateEffectiveSpeedLimit(sub *subscribe.Subscribe, userSub *user.Subscribe) int64 {
|
|
||||||
baseSpeed := sub.SpeedLimit
|
baseSpeed := sub.SpeedLimit
|
||||||
if userSub.SpeedLimit > 0 {
|
if userSub.SpeedLimit > 0 {
|
||||||
baseSpeed = userSub.SpeedLimit
|
baseSpeed = userSub.SpeedLimit
|
||||||
@@ -346,15 +429,170 @@ func (l *GetServerUserListLogic) calculateEffectiveSpeedLimit(sub *subscribe.Sub
|
|||||||
trafficLimit = *userSub.TrafficLimit
|
trafficLimit = *userSub.TrafficLimit
|
||||||
}
|
}
|
||||||
|
|
||||||
result := speedlimit.CalculateWithCache(
|
return baseSpeed, trafficLimit
|
||||||
l.ctx.Request.Context(),
|
}
|
||||||
l.svcCtx.Redis,
|
|
||||||
l.svcCtx.DB,
|
func (l *GetServerUserListLogic) calculateServerUserSpeedLimits(candidates map[int64]serverUserSpeedLimitCandidate) map[int64]int64 {
|
||||||
userSub.UserId,
|
speedLimits := make(map[int64]int64, len(candidates))
|
||||||
userSub.Id,
|
rulesByUserSubscribeID := make(map[int64][]speedlimit.TrafficLimitRule)
|
||||||
baseSpeed,
|
windowByKey := make(map[string]serverUserTrafficWindow)
|
||||||
trafficLimit,
|
now := time.Now()
|
||||||
30*time.Second,
|
|
||||||
)
|
for userSubscribeID, candidate := range candidates {
|
||||||
return result.EffectiveSpeed
|
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
|
package server
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"regexp"
|
||||||
"slices"
|
"slices"
|
||||||
"testing"
|
"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/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 兼容字段被映射回
|
// 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)
|
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
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user