mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-06 18:06:36 +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 forwardRuntimeLimiters struct {
|
||||
TrafficLimiter string
|
||||
ConnLimiter string
|
||||
}
|
||||
|
||||
type forwardLimiterConfig struct {
|
||||
Name string
|
||||
Limits []string
|
||||
}
|
||||
|
||||
type diagnosisTarget struct {
|
||||
Address string
|
||||
IP string
|
||||
@@ -264,6 +274,13 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
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)
|
||||
|
||||
user, err := h.repo.GetUserByID(forward.UserID)
|
||||
@@ -271,19 +288,31 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var cLimiterName string
|
||||
var maxConnToSet int
|
||||
|
||||
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)
|
||||
userMaxConn := 0
|
||||
if user != nil && user.MaxConn > 0 {
|
||||
userMaxConn = user.MaxConn
|
||||
}
|
||||
connLimiterConfig := buildConnLimiterConfig(forward, userMaxConn)
|
||||
|
||||
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 the limiter push fails because the node is offline, skip it with a warning
|
||||
if isNodeOfflineOrTimeoutError(err) {
|
||||
@@ -299,8 +328,8 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
}
|
||||
}
|
||||
|
||||
if cLimiterName != "" {
|
||||
if err := h.ensureConnLimiterOnNode(fp.NodeID, cLimiterName, maxConnToSet); err != nil {
|
||||
if connLimiterConfig.Name != "" {
|
||||
if err := h.ensureConnLimiterOnNode(fp.NodeID, connLimiterConfig); err != nil {
|
||||
warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err))
|
||||
}
|
||||
}
|
||||
@@ -309,7 +338,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
if err != nil {
|
||||
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)
|
||||
if err != nil && allowFallbackAdd && method == "UpdateService" {
|
||||
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) {
|
||||
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 != "" {
|
||||
warnings = append(warnings, warning)
|
||||
}
|
||||
@@ -350,7 +379,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
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 {
|
||||
return "", errors.New("invalid bind fallback context")
|
||||
}
|
||||
@@ -367,7 +396,7 @@ func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunne
|
||||
}
|
||||
|
||||
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 {
|
||||
return "", err
|
||||
}
|
||||
@@ -1659,7 +1688,7 @@ func compactErrorMessage(msg string) string {
|
||||
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"}
|
||||
services := make([]map[string]interface{}, 0, 2)
|
||||
targets := splitRemoteTargets(forward.RemoteAddr)
|
||||
@@ -1702,8 +1731,11 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
},
|
||||
},
|
||||
}
|
||||
if cLimiterName != "" {
|
||||
service["climiter"] = cLimiterName
|
||||
if runtimeLimiters.ConnLimiter != "" {
|
||||
service["climiter"] = runtimeLimiters.ConnLimiter
|
||||
}
|
||||
if runtimeLimiters.TrafficLimiter != "" {
|
||||
service["limiter"] = runtimeLimiters.TrafficLimiter
|
||||
}
|
||||
if forward.ProxyProtocol > 0 {
|
||||
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
|
||||
}
|
||||
if limiterID != nil && *limiterID > 0 {
|
||||
service["limiter"] = strconv.FormatInt(*limiterID, 10)
|
||||
}
|
||||
services = append(services, service)
|
||||
}
|
||||
|
||||
@@ -1830,22 +1859,16 @@ func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxConn int) error {
|
||||
limitStr := fmt.Sprintf("$ %d", maxConn)
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"name": limiterName,
|
||||
"limits": []string{limitStr},
|
||||
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, cfg forwardLimiterConfig) error {
|
||||
if cfg.Name == "" || len(cfg.Limits) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
payload := map[string]interface{}{"name": cfg.Name, "limits": cfg.Limits}
|
||||
if _, err := h.sendNodeCommand(nodeID, "AddCLimiters", payload, false, false); err != nil {
|
||||
if !isAlreadyExistsMessage(err.Error()) {
|
||||
return fmt.Errorf("连接限制器下发失败: %w", err)
|
||||
}
|
||||
updatePayload := map[string]interface{}{
|
||||
"limiter": limiterName,
|
||||
"data": payload,
|
||||
}
|
||||
updatePayload := map[string]interface{}{"limiter": cfg.Name, "data": payload}
|
||||
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil {
|
||||
return fmt.Errorf("连接限制器更新失败: %w", updateErr)
|
||||
}
|
||||
@@ -1853,14 +1876,51 @@ func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxC
|
||||
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
|
||||
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)
|
||||
|
||||
return name, map[string]interface{}{
|
||||
"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
|
||||
}
|
||||
|
||||
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) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
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 {
|
||||
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) {
|
||||
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, "", nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -409,7 +409,7 @@ func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.
|
||||
func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
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 {
|
||||
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) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
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 {
|
||||
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) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
@@ -25,7 +25,7 @@ func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing
|
||||
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 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user