mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-11 11:46:37 +08:00
feat: sync per-IP runtime limiters
This commit is contained in:
@@ -27,6 +27,16 @@ type nodeRecord = model.NodeRecord
|
|||||||
|
|
||||||
type chainNodeRecord = model.ChainNodeRecord
|
type chainNodeRecord = model.ChainNodeRecord
|
||||||
|
|
||||||
|
type forwardRuntimeLimiters struct {
|
||||||
|
TrafficLimiter string
|
||||||
|
ConnLimiter string
|
||||||
|
}
|
||||||
|
|
||||||
|
type forwardLimiterConfig struct {
|
||||||
|
Name string
|
||||||
|
Limits []string
|
||||||
|
}
|
||||||
|
|
||||||
type diagnosisTarget struct {
|
type diagnosisTarget struct {
|
||||||
Address string
|
Address string
|
||||||
IP string
|
IP string
|
||||||
@@ -264,6 +274,13 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
|||||||
speed = utSpeed
|
speed = utSpeed
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var ipSpeed *int
|
||||||
|
if forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0 {
|
||||||
|
if speedVal, err := h.repo.GetSpeedLimitSpeed(forward.IPSpeedID.Int64); err == nil && speedVal > 0 {
|
||||||
|
ipSpeed = &speedVal
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID)
|
serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID)
|
||||||
|
|
||||||
user, err := h.repo.GetUserByID(forward.UserID)
|
user, err := h.repo.GetUserByID(forward.UserID)
|
||||||
@@ -271,19 +288,31 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
var cLimiterName string
|
userMaxConn := 0
|
||||||
var maxConnToSet int
|
if user != nil && user.MaxConn > 0 {
|
||||||
|
userMaxConn = user.MaxConn
|
||||||
if forward.MaxConn > 0 {
|
|
||||||
maxConnToSet = forward.MaxConn
|
|
||||||
cLimiterName = fmt.Sprintf("rule_conn_limit_%d", forward.ID)
|
|
||||||
} else if user != nil && user.MaxConn > 0 {
|
|
||||||
maxConnToSet = user.MaxConn
|
|
||||||
cLimiterName = fmt.Sprintf("user_conn_limit_%d", user.ID)
|
|
||||||
}
|
}
|
||||||
|
connLimiterConfig := buildConnLimiterConfig(forward, userMaxConn)
|
||||||
|
|
||||||
for _, fp := range ports {
|
for _, fp := range ports {
|
||||||
if limiterID != nil && speed != nil {
|
runtimeLimiters := forwardRuntimeLimiters{ConnLimiter: connLimiterConfig.Name}
|
||||||
|
if ipSpeed != nil {
|
||||||
|
runtimeLimiters.TrafficLimiter = fmt.Sprintf("rule_traffic_limit_%d", forward.ID)
|
||||||
|
if err := h.ensureTrafficLimiterOnNode(fp.NodeID, runtimeLimiters.TrafficLimiter, speed, ipSpeed); err != nil {
|
||||||
|
// If the limiter push fails because the node is offline, skip it with a warning
|
||||||
|
if isNodeOfflineOrTimeoutError(err) {
|
||||||
|
node, _ := h.getNodeRecord(fp.NodeID)
|
||||||
|
nodeName := fmt.Sprintf("%d", fp.NodeID)
|
||||||
|
if node != nil && strings.TrimSpace(node.Name) != "" {
|
||||||
|
nodeName = strings.TrimSpace(node.Name)
|
||||||
|
}
|
||||||
|
warnings = append(warnings, fmt.Sprintf("节点 %s 不在线,已跳过下发", nodeName))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
} else if limiterID != nil && speed != nil {
|
||||||
|
runtimeLimiters.TrafficLimiter = strconv.FormatInt(*limiterID, 10)
|
||||||
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
|
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
|
||||||
// If the limiter push fails because the node is offline, skip it with a warning
|
// If the limiter push fails because the node is offline, skip it with a warning
|
||||||
if isNodeOfflineOrTimeoutError(err) {
|
if isNodeOfflineOrTimeoutError(err) {
|
||||||
@@ -299,8 +328,8 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if cLimiterName != "" {
|
if connLimiterConfig.Name != "" {
|
||||||
if err := h.ensureConnLimiterOnNode(fp.NodeID, cLimiterName, maxConnToSet); err != nil {
|
if err := h.ensureConnLimiterOnNode(fp.NodeID, connLimiterConfig); err != nil {
|
||||||
warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err))
|
warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -309,7 +338,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID, cLimiterName)
|
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), runtimeLimiters)
|
||||||
_, err = h.sendNodeCommand(node.ID, method, services, true, false)
|
_, err = h.sendNodeCommand(node.ID, method, services, true, false)
|
||||||
if err != nil && allowFallbackAdd && method == "UpdateService" {
|
if err != nil && allowFallbackAdd && method == "UpdateService" {
|
||||||
if isNotFoundError(err) {
|
if isNotFoundError(err) {
|
||||||
@@ -324,7 +353,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
|||||||
}
|
}
|
||||||
if err != nil && strings.EqualFold(strings.TrimSpace(method), "UpdateService") && isCannotAssignRequestedAddressError(err) {
|
if err != nil && strings.EqualFold(strings.TrimSpace(method), "UpdateService") && isCannotAssignRequestedAddressError(err) {
|
||||||
var warning string
|
var warning string
|
||||||
warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, limiterID, cLimiterName)
|
warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, runtimeLimiters)
|
||||||
if err == nil && warning != "" {
|
if err == nil && warning != "" {
|
||||||
warnings = append(warnings, warning)
|
warnings = append(warnings, warning)
|
||||||
}
|
}
|
||||||
@@ -350,7 +379,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
|||||||
return warnings, nil
|
return warnings, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, limiterID *int64, cLimiterName string) (string, error) {
|
func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, runtimeLimiters forwardRuntimeLimiters) (string, error) {
|
||||||
if h == nil || forward == nil || tunnel == nil || node == nil {
|
if h == nil || forward == nil || tunnel == nil || node == nil {
|
||||||
return "", errors.New("invalid bind fallback context")
|
return "", errors.New("invalid bind fallback context")
|
||||||
}
|
}
|
||||||
@@ -367,7 +396,7 @@ func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunne
|
|||||||
}
|
}
|
||||||
|
|
||||||
time.Sleep(150 * time.Millisecond)
|
time.Sleep(150 * time.Millisecond)
|
||||||
defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID, cLimiterName)
|
defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", runtimeLimiters)
|
||||||
if _, err := h.sendNodeCommand(node.ID, "AddService", defaultServices, true, false); err != nil {
|
if _, err := h.sendNodeCommand(node.ID, "AddService", defaultServices, true, false); err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
@@ -1659,7 +1688,7 @@ func compactErrorMessage(msg string) string {
|
|||||||
return strings.Join(strings.Fields(strings.ToLower(msg)), "")
|
return strings.Join(strings.Fields(strings.ToLower(msg)), "")
|
||||||
}
|
}
|
||||||
|
|
||||||
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64, cLimiterName string) []map[string]interface{} {
|
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, runtimeLimiters forwardRuntimeLimiters) []map[string]interface{} {
|
||||||
protocols := []string{"tcp", "udp"}
|
protocols := []string{"tcp", "udp"}
|
||||||
services := make([]map[string]interface{}, 0, 2)
|
services := make([]map[string]interface{}, 0, 2)
|
||||||
targets := splitRemoteTargets(forward.RemoteAddr)
|
targets := splitRemoteTargets(forward.RemoteAddr)
|
||||||
@@ -1702,8 +1731,11 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
if cLimiterName != "" {
|
if runtimeLimiters.ConnLimiter != "" {
|
||||||
service["climiter"] = cLimiterName
|
service["climiter"] = runtimeLimiters.ConnLimiter
|
||||||
|
}
|
||||||
|
if runtimeLimiters.TrafficLimiter != "" {
|
||||||
|
service["limiter"] = runtimeLimiters.TrafficLimiter
|
||||||
}
|
}
|
||||||
if forward.ProxyProtocol > 0 {
|
if forward.ProxyProtocol > 0 {
|
||||||
handlerConfig := service["handler"].(map[string]interface{})
|
handlerConfig := service["handler"].(map[string]interface{})
|
||||||
@@ -1728,9 +1760,6 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
|||||||
}
|
}
|
||||||
service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName
|
service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName
|
||||||
}
|
}
|
||||||
if limiterID != nil && *limiterID > 0 {
|
|
||||||
service["limiter"] = strconv.FormatInt(*limiterID, 10)
|
|
||||||
}
|
|
||||||
services = append(services, service)
|
services = append(services, service)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1830,22 +1859,16 @@ func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int)
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxConn int) error {
|
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, cfg forwardLimiterConfig) error {
|
||||||
limitStr := fmt.Sprintf("$ %d", maxConn)
|
if cfg.Name == "" || len(cfg.Limits) == 0 {
|
||||||
|
return nil
|
||||||
payload := map[string]interface{}{
|
|
||||||
"name": limiterName,
|
|
||||||
"limits": []string{limitStr},
|
|
||||||
}
|
}
|
||||||
|
payload := map[string]interface{}{"name": cfg.Name, "limits": cfg.Limits}
|
||||||
if _, err := h.sendNodeCommand(nodeID, "AddCLimiters", payload, false, false); err != nil {
|
if _, err := h.sendNodeCommand(nodeID, "AddCLimiters", payload, false, false); err != nil {
|
||||||
if !isAlreadyExistsMessage(err.Error()) {
|
if !isAlreadyExistsMessage(err.Error()) {
|
||||||
return fmt.Errorf("连接限制器下发失败: %w", err)
|
return fmt.Errorf("连接限制器下发失败: %w", err)
|
||||||
}
|
}
|
||||||
updatePayload := map[string]interface{}{
|
updatePayload := map[string]interface{}{"limiter": cfg.Name, "data": payload}
|
||||||
"limiter": limiterName,
|
|
||||||
"data": payload,
|
|
||||||
}
|
|
||||||
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil {
|
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil {
|
||||||
return fmt.Errorf("连接限制器更新失败: %w", updateErr)
|
return fmt.Errorf("连接限制器更新失败: %w", updateErr)
|
||||||
}
|
}
|
||||||
@@ -1853,14 +1876,51 @@ func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxC
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func buildLimiterAddPayload(limiterID int64, speed int) (string, map[string]interface{}) {
|
func buildConnLimiterConfig(forward *forwardRecord, userMaxConn int) forwardLimiterConfig {
|
||||||
|
if forward == nil {
|
||||||
|
return forwardLimiterConfig{}
|
||||||
|
}
|
||||||
|
limits := make([]string, 0, 2)
|
||||||
|
if forward.MaxConn > 0 {
|
||||||
|
limits = append(limits, fmt.Sprintf("$ %d", forward.MaxConn))
|
||||||
|
} else if userMaxConn > 0 {
|
||||||
|
limits = append(limits, fmt.Sprintf("$ %d", userMaxConn))
|
||||||
|
}
|
||||||
|
if forward.IPMaxConn > 0 {
|
||||||
|
limits = append(limits, fmt.Sprintf("$$ %d", forward.IPMaxConn))
|
||||||
|
}
|
||||||
|
if len(limits) == 0 {
|
||||||
|
return forwardLimiterConfig{}
|
||||||
|
}
|
||||||
|
name := fmt.Sprintf("user_conn_limit_%d", forward.UserID)
|
||||||
|
if forward.MaxConn > 0 || forward.IPMaxConn > 0 {
|
||||||
|
name = fmt.Sprintf("rule_conn_limit_%d", forward.ID)
|
||||||
|
}
|
||||||
|
return forwardLimiterConfig{Name: name, Limits: limits}
|
||||||
|
}
|
||||||
|
|
||||||
|
func speedToLimitLine(key string, speed int) string {
|
||||||
rate := float64(speed) / 8.0
|
rate := float64(speed) / 8.0
|
||||||
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
|
return fmt.Sprintf("%s %.1fMB %.1fMB", key, rate, rate)
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildTrafficLimiterPayload(name string, totalSpeed *int, ipSpeed *int) map[string]interface{} {
|
||||||
|
limits := make([]string, 0, 3)
|
||||||
|
if totalSpeed != nil && *totalSpeed > 0 {
|
||||||
|
limits = append(limits, speedToLimitLine("$", *totalSpeed))
|
||||||
|
}
|
||||||
|
if ipSpeed != nil && *ipSpeed > 0 {
|
||||||
|
limits = append(limits, speedToLimitLine("0.0.0.0/0", *ipSpeed), speedToLimitLine("::/0", *ipSpeed))
|
||||||
|
}
|
||||||
|
return map[string]interface{}{"name": name, "limits": limits}
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildLimiterAddPayload(limiterID int64, speed int) (string, map[string]interface{}) {
|
||||||
name := strconv.FormatInt(limiterID, 10)
|
name := strconv.FormatInt(limiterID, 10)
|
||||||
|
|
||||||
return name, map[string]interface{}{
|
return name, map[string]interface{}{
|
||||||
"name": name,
|
"name": name,
|
||||||
"limits": []string{limitStr},
|
"limits": []string{speedToLimitLine("$", speed)},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1888,3 +1948,20 @@ func (h *Handler) upsertLimiterOnNode(nodeID int64, limiterID int64, speed int)
|
|||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (h *Handler) ensureTrafficLimiterOnNode(nodeID int64, name string, totalSpeed *int, ipSpeed *int) error {
|
||||||
|
payload := buildTrafficLimiterPayload(name, totalSpeed, ipSpeed)
|
||||||
|
limits, _ := payload["limits"].([]string)
|
||||||
|
if name == "" || len(limits) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil {
|
||||||
|
if !isAlreadyExistsMessage(err.Error()) {
|
||||||
|
return fmt.Errorf("限速规则下发失败: %w", err)
|
||||||
|
}
|
||||||
|
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateLimiters", buildLimiterUpdatePayload(name, payload), false, false); updateErr != nil {
|
||||||
|
return fmt.Errorf("限速规则更新失败: %w", updateErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -378,7 +378,7 @@ func TestRetryTunnelServiceAddWithCleanupReturnsCleanupError(t *testing.T) {
|
|||||||
func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
||||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", nil, "")
|
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", forwardRuntimeLimiters{})
|
||||||
if len(services) != 2 {
|
if len(services) != 2 {
|
||||||
t.Fatalf("expected 2 services, got %d", len(services))
|
t.Fatalf("expected 2 services, got %d", len(services))
|
||||||
}
|
}
|
||||||
@@ -393,7 +393,7 @@ func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
|||||||
func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.T) {
|
func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.T) {
|
||||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||||
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
|
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
|
||||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", nil, "")
|
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{})
|
||||||
if len(services) != 2 {
|
if len(services) != 2 {
|
||||||
t.Fatalf("expected 2 services, got %d", len(services))
|
t.Fatalf("expected 2 services, got %d", len(services))
|
||||||
}
|
}
|
||||||
@@ -409,7 +409,7 @@ func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.
|
|||||||
func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(t *testing.T) {
|
func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(t *testing.T) {
|
||||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", nil, "")
|
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", forwardRuntimeLimiters{})
|
||||||
if len(services) != 2 {
|
if len(services) != 2 {
|
||||||
t.Fatalf("expected 2 services, got %d", len(services))
|
t.Fatalf("expected 2 services, got %d", len(services))
|
||||||
}
|
}
|
||||||
@@ -464,7 +464,7 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
|
|||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, nil, "")
|
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, forwardRuntimeLimiters{})
|
||||||
if len(services) != 2 {
|
if len(services) != 2 {
|
||||||
t.Fatalf("expected 2 services, got %d", len(services))
|
t.Fatalf("expected 2 services, got %d", len(services))
|
||||||
}
|
}
|
||||||
@@ -478,6 +478,52 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestBuildConnLimiterConfigCombinesTotalAndPerIP(t *testing.T) {
|
||||||
|
cfg := buildConnLimiterConfig(&forwardRecord{ID: 42, UserID: 9, MaxConn: 100, IPMaxConn: 5}, 37)
|
||||||
|
want := forwardLimiterConfig{Name: "rule_conn_limit_42", Limits: []string{"$ 100", "$$ 5"}}
|
||||||
|
if !reflect.DeepEqual(cfg, want) {
|
||||||
|
t.Fatalf("expected %+v, got %+v", want, cfg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildConnLimiterConfigUsesUserTotalWithRulePerIP(t *testing.T) {
|
||||||
|
cfg := buildConnLimiterConfig(&forwardRecord{ID: 42, UserID: 9, IPMaxConn: 5}, 37)
|
||||||
|
want := forwardLimiterConfig{Name: "rule_conn_limit_42", Limits: []string{"$ 37", "$$ 5"}}
|
||||||
|
if !reflect.DeepEqual(cfg, want) {
|
||||||
|
t.Fatalf("expected %+v, got %+v", want, cfg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildTrafficLimiterPayloadCombinesTotalAndPerIP(t *testing.T) {
|
||||||
|
payload := buildTrafficLimiterPayload("rule_traffic_limit_42", intPtr(80), intPtr(40))
|
||||||
|
wantLimits := []string{"$ 10.0MB 10.0MB", "0.0.0.0/0 5.0MB 5.0MB", "::/0 5.0MB 5.0MB"}
|
||||||
|
if payload["name"] != "rule_traffic_limit_42" {
|
||||||
|
t.Fatalf("expected name rule_traffic_limit_42, got %v", payload["name"])
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(payload["limits"], wantLimits) {
|
||||||
|
t.Fatalf("expected limits %v, got %v", wantLimits, payload["limits"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildForwardServiceConfigsUsesRuntimeLimiterNames(t *testing.T) {
|
||||||
|
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||||
|
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
|
||||||
|
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{TrafficLimiter: "rule_traffic_limit_42", ConnLimiter: "rule_conn_limit_42"})
|
||||||
|
if len(services) != 2 {
|
||||||
|
t.Fatalf("expected 2 services, got %d", len(services))
|
||||||
|
}
|
||||||
|
for _, service := range services {
|
||||||
|
if service["limiter"] != "rule_traffic_limit_42" {
|
||||||
|
t.Fatalf("expected traffic limiter rule_traffic_limit_42, got %v", service["limiter"])
|
||||||
|
}
|
||||||
|
if service["climiter"] != "rule_conn_limit_42" {
|
||||||
|
t.Fatalf("expected conn limiter rule_conn_limit_42, got %v", service["climiter"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func intPtr(v int) *int { return &v }
|
||||||
|
|
||||||
func TestProcessServerAddress_StripsURLSchemeAndPath(t *testing.T) {
|
func TestProcessServerAddress_StripsURLSchemeAndPath(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing
|
|||||||
UDPListenAddr: "0.0.0.0",
|
UDPListenAddr: "0.0.0.0",
|
||||||
}
|
}
|
||||||
|
|
||||||
services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", nil, "")
|
services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", forwardRuntimeLimiters{})
|
||||||
if len(services) != 2 {
|
if len(services) != 2 {
|
||||||
t.Fatalf("expected 2 services, got %d", len(services))
|
t.Fatalf("expected 2 services, got %d", len(services))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -97,6 +97,7 @@ func TestMaxConnLimit(t *testing.T) {
|
|||||||
"remoteAddr": "1.1.1.1:443",
|
"remoteAddr": "1.1.1.1:443",
|
||||||
"strategy": "fifo",
|
"strategy": "fifo",
|
||||||
"maxConn": 42,
|
"maxConn": 42,
|
||||||
|
"ipMaxConn": 7,
|
||||||
"proxyProtocol": 2,
|
"proxyProtocol": 2,
|
||||||
}
|
}
|
||||||
body, err := json.Marshal(payload)
|
body, err := json.Marshal(payload)
|
||||||
@@ -195,8 +196,8 @@ func TestMaxConnLimit(t *testing.T) {
|
|||||||
t.Fatalf("expected limiter name %s, got %v", expectedName, addData["name"])
|
t.Fatalf("expected limiter name %s, got %v", expectedName, addData["name"])
|
||||||
}
|
}
|
||||||
if limits, ok := addData["limits"].([]interface{}); ok {
|
if limits, ok := addData["limits"].([]interface{}); ok {
|
||||||
if len(limits) != 1 || limits[0] != "$ 42" {
|
if len(limits) != 2 || limits[0] != "$ 42" || limits[1] != "$$ 7" {
|
||||||
t.Fatalf("expected limits to contain '$ 42', got %v", limits)
|
t.Fatalf("expected limits to contain '$ 42' and '$$ 7', got %v", limits)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
t.Fatalf("invalid limits type in AddCLimiters data: %v", addData)
|
t.Fatalf("invalid limits type in AddCLimiters data: %v", addData)
|
||||||
@@ -218,8 +219,8 @@ func TestMaxConnLimit(t *testing.T) {
|
|||||||
t.Fatalf("expected nested name %s, got %v", expectedName, nestedData["name"])
|
t.Fatalf("expected nested name %s, got %v", expectedName, nestedData["name"])
|
||||||
}
|
}
|
||||||
if nestedLimits, ok := nestedData["limits"].([]interface{}); ok {
|
if nestedLimits, ok := nestedData["limits"].([]interface{}); ok {
|
||||||
if len(nestedLimits) != 1 || nestedLimits[0] != "$ 42" {
|
if len(nestedLimits) != 2 || nestedLimits[0] != "$ 42" || nestedLimits[1] != "$$ 7" {
|
||||||
t.Fatalf("expected nested limits to contain '$ 42', got %v", nestedLimits)
|
t.Fatalf("expected nested limits to contain '$ 42' and '$$ 7', got %v", nestedLimits)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
t.Fatalf("invalid limits type in UpdateCLimiters nested data: %v", nestedData)
|
t.Fatalf("invalid limits type in UpdateCLimiters nested data: %v", nestedData)
|
||||||
|
|||||||
@@ -0,0 +1,168 @@
|
|||||||
|
package contract_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"reflect"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"go-backend/internal/auth"
|
||||||
|
"go-backend/internal/http/response"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPerIPSpeedLimitRuntimePayload(t *testing.T) {
|
||||||
|
secret := "contract-jwt-secret"
|
||||||
|
router, r := setupContractRouter(t, secret)
|
||||||
|
server := httptest.NewServer(router)
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("generate admin token: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
if err := r.DB().Exec(`
|
||||||
|
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
|
`, "per-ip-speed-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||||
|
t.Fatalf("insert tunnel: %v", err)
|
||||||
|
}
|
||||||
|
var tunnelID int64
|
||||||
|
if err := r.DB().Raw("SELECT id FROM tunnel WHERE name = ?", "per-ip-speed-tunnel").Scan(&tunnelID).Error; err != nil {
|
||||||
|
t.Fatalf("get tunnel ID: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.DB().Exec(`
|
||||||
|
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||||
|
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
|
`, "per-ip-speed-node", "per-ip-speed-secret", "10.22.0.1", "10.22.0.1", "", "32200-32210", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||||
|
t.Fatalf("insert node: %v", err)
|
||||||
|
}
|
||||||
|
var nodeID int64
|
||||||
|
if err := r.DB().Raw("SELECT id FROM node WHERE name = ?", "per-ip-speed-node").Scan(&nodeID).Error; err != nil {
|
||||||
|
t.Fatalf("get node ID: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := r.DB().Exec(`
|
||||||
|
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||||
|
VALUES(?, 1, ?, 32201, 'round', 1, 'tls')
|
||||||
|
`, tunnelID, nodeID).Error; err != nil {
|
||||||
|
t.Fatalf("insert chain_tunnel: %v", err)
|
||||||
|
}
|
||||||
|
if err := r.DB().Exec(`
|
||||||
|
INSERT INTO user_tunnel(user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||||
|
VALUES(1, ?, 10, 99999, 0, 0, 1, ?, 1)
|
||||||
|
`, tunnelID, now+365*24*3600*1000).Error; err != nil {
|
||||||
|
t.Fatalf("insert user_tunnel: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
totalSpeedID, err := r.CreateSpeedLimit("per-ip-total-speed", 80, now, 1)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("create total speed limit: %v", err)
|
||||||
|
}
|
||||||
|
ipSpeedID, err := r.CreateSpeedLimit("per-ip-client-speed", 40, now, 1)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("create per-ip speed limit: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var commandMu sync.Mutex
|
||||||
|
receivedCommands := make([]string, 0)
|
||||||
|
var addLimitersData json.RawMessage
|
||||||
|
var updateServiceData json.RawMessage
|
||||||
|
|
||||||
|
stopNode := startMockSessionForMaxConn(t, server.URL, "per-ip-speed-secret", func(cmdType string, data json.RawMessage) (bool, string) {
|
||||||
|
commandMu.Lock()
|
||||||
|
defer commandMu.Unlock()
|
||||||
|
receivedCommands = append(receivedCommands, cmdType)
|
||||||
|
if cmdType == "AddLimiters" {
|
||||||
|
addLimitersData = append([]byte(nil), data...)
|
||||||
|
}
|
||||||
|
if cmdType == "UpdateService" {
|
||||||
|
updateServiceData = append([]byte(nil), data...)
|
||||||
|
}
|
||||||
|
return false, ""
|
||||||
|
})
|
||||||
|
defer stopNode()
|
||||||
|
|
||||||
|
waitNodeStatus(t, r, nodeID, 1)
|
||||||
|
|
||||||
|
payload := map[string]interface{}{
|
||||||
|
"name": "per-ip-speed-forward",
|
||||||
|
"tunnelId": tunnelID,
|
||||||
|
"remoteAddr": "1.1.1.1:443",
|
||||||
|
"strategy": "fifo",
|
||||||
|
"speedId": totalSpeedID,
|
||||||
|
"ipSpeedId": ipSpeedID,
|
||||||
|
}
|
||||||
|
body, err := json.Marshal(payload)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal payload: %v", err)
|
||||||
|
}
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Authorization", adminToken)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
res := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(res, req)
|
||||||
|
|
||||||
|
var out response.R
|
||||||
|
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||||
|
t.Fatalf("decode response: %v", err)
|
||||||
|
}
|
||||||
|
if out.Code != 0 {
|
||||||
|
t.Fatalf("expected create success, got code=%d msg=%s", out.Code, out.Msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
var forwardID int64
|
||||||
|
if err := r.DB().Raw("SELECT id FROM forward WHERE name = ?", "per-ip-speed-forward").Scan(&forwardID).Error; err != nil {
|
||||||
|
t.Fatalf("get forward ID: %v", err)
|
||||||
|
}
|
||||||
|
expectedName := fmt.Sprintf("rule_traffic_limit_%d", forwardID)
|
||||||
|
expectedLimits := []string{"$ 10.0MB 10.0MB", "0.0.0.0/0 5.0MB 5.0MB", "::/0 5.0MB 5.0MB"}
|
||||||
|
|
||||||
|
commandMu.Lock()
|
||||||
|
defer commandMu.Unlock()
|
||||||
|
if addLimitersData == nil {
|
||||||
|
t.Fatalf("expected AddLimiters to be sent. Received: %v", receivedCommands)
|
||||||
|
}
|
||||||
|
if updateServiceData == nil {
|
||||||
|
t.Fatalf("expected UpdateService to be sent. Received: %v", receivedCommands)
|
||||||
|
}
|
||||||
|
|
||||||
|
var addData map[string]interface{}
|
||||||
|
if err := json.Unmarshal(addLimitersData, &addData); err != nil {
|
||||||
|
t.Fatalf("unmarshal AddLimiters data: %v", err)
|
||||||
|
}
|
||||||
|
if addData["name"] != expectedName {
|
||||||
|
t.Fatalf("expected limiter name %s, got %v", expectedName, addData["name"])
|
||||||
|
}
|
||||||
|
limits, ok := addData["limits"].([]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected limits array, got %T", addData["limits"])
|
||||||
|
}
|
||||||
|
gotLimits := make([]string, 0, len(limits))
|
||||||
|
for _, limit := range limits {
|
||||||
|
gotLimits = append(gotLimits, fmt.Sprint(limit))
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(gotLimits, expectedLimits) {
|
||||||
|
t.Fatalf("expected limits %v, got %v", expectedLimits, gotLimits)
|
||||||
|
}
|
||||||
|
|
||||||
|
var services []map[string]interface{}
|
||||||
|
if err := json.Unmarshal(updateServiceData, &services); err != nil {
|
||||||
|
t.Fatalf("unmarshal UpdateService data: %v", err)
|
||||||
|
}
|
||||||
|
if len(services) == 0 {
|
||||||
|
t.Fatalf("expected services in UpdateService")
|
||||||
|
}
|
||||||
|
for _, service := range services {
|
||||||
|
if service["limiter"] != expectedName {
|
||||||
|
t.Fatalf("expected service limiter %s, got %v", expectedName, service["limiter"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user