fix: preserve shared limiters with per-IP rules

This commit is contained in:
sagitchu
2026-04-28 00:18:32 +08:00
parent bd27b94909
commit 3373e5ade9
7 changed files with 511 additions and 83 deletions
@@ -0,0 +1,199 @@
package service
import (
"context"
"fmt"
"sort"
"strconv"
"strings"
corelimiter "github.com/go-gost/core/limiter"
connlimiter "github.com/go-gost/core/limiter/conn"
trafficlimiter "github.com/go-gost/core/limiter/traffic"
xtraffic "github.com/go-gost/x/limiter/traffic"
"github.com/go-gost/x/registry"
)
func resolveTrafficLimiter(names string) trafficlimiter.TrafficLimiter {
parts := splitLimiterNames(names)
if len(parts) == 0 {
return nil
}
if len(parts) == 1 {
return resolveSingleTrafficLimiter(parts[0])
}
limiters := make([]trafficlimiter.TrafficLimiter, 0, len(parts))
for _, part := range parts {
if lim := resolveSingleTrafficLimiter(part); lim != nil {
limiters = append(limiters, lim)
}
}
if len(limiters) == 0 {
return nil
}
if len(limiters) == 1 {
return limiters[0]
}
return &compositeTrafficLimiter{limiters: limiters}
}
func resolveSingleTrafficLimiter(name string) trafficlimiter.TrafficLimiter {
lim := registry.TrafficLimiterRegistry().Get(name)
if lim != nil {
return lim
}
if val, err := strconv.Atoi(name); err == nil && val > 0 {
return xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption(fmt.Sprintf("%s %dB %dB", xtraffic.ServiceLimitKey, val, val)),
)
}
return xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption(fmt.Sprintf("%s %s %s", xtraffic.ServiceLimitKey, name, name)),
)
}
func resolveConnLimiter(names string) connlimiter.ConnLimiter {
parts := splitLimiterNames(names)
if len(parts) == 0 {
return nil
}
if len(parts) == 1 {
return registry.ConnLimiterRegistry().Get(parts[0])
}
limiters := make([]connlimiter.ConnLimiter, 0, len(parts))
for _, part := range parts {
if lim := registry.ConnLimiterRegistry().Get(part); lim != nil {
limiters = append(limiters, lim)
}
}
if len(limiters) == 0 {
return nil
}
if len(limiters) == 1 {
return limiters[0]
}
return &compositeConnLimiter{limiters: limiters}
}
func splitLimiterNames(names string) []string {
parts := strings.Split(names, ",")
out := make([]string, 0, len(parts))
for _, part := range parts {
if part = strings.TrimSpace(part); part != "" {
out = append(out, part)
}
}
return out
}
type compositeTrafficLimiter struct {
limiters []trafficlimiter.TrafficLimiter
}
func (l *compositeTrafficLimiter) In(ctx context.Context, key string, opts ...corelimiter.Option) trafficlimiter.Limiter {
limiters := make([]trafficlimiter.Limiter, 0, len(l.limiters))
for _, child := range l.limiters {
if lim := child.In(ctx, key, opts...); lim != nil {
limiters = append(limiters, lim)
}
}
return newCompositeTrafficChildLimiter(limiters)
}
func (l *compositeTrafficLimiter) Out(ctx context.Context, key string, opts ...corelimiter.Option) trafficlimiter.Limiter {
limiters := make([]trafficlimiter.Limiter, 0, len(l.limiters))
for _, child := range l.limiters {
if lim := child.Out(ctx, key, opts...); lim != nil {
limiters = append(limiters, lim)
}
}
return newCompositeTrafficChildLimiter(limiters)
}
type compositeTrafficChildLimiter struct {
limiters []trafficlimiter.Limiter
}
func newCompositeTrafficChildLimiter(limiters []trafficlimiter.Limiter) trafficlimiter.Limiter {
if len(limiters) == 0 {
return nil
}
if len(limiters) == 1 {
return limiters[0]
}
sort.Slice(limiters, func(i, j int) bool {
return limiters[i].Limit() < limiters[j].Limit()
})
return &compositeTrafficChildLimiter{limiters: limiters}
}
func (l *compositeTrafficChildLimiter) Wait(ctx context.Context, n int) int {
for _, lim := range l.limiters {
if v := lim.Wait(ctx, n); v < n {
n = v
}
}
return n
}
func (l *compositeTrafficChildLimiter) Limit() int {
if len(l.limiters) == 0 {
return 0
}
return l.limiters[0].Limit()
}
func (l *compositeTrafficChildLimiter) Set(n int) {}
type compositeConnLimiter struct {
limiters []connlimiter.ConnLimiter
}
func (l *compositeConnLimiter) Limiter(key string) connlimiter.Limiter {
limiters := make([]connlimiter.Limiter, 0, len(l.limiters))
for _, child := range l.limiters {
if lim := child.Limiter(key); lim != nil {
limiters = append(limiters, lim)
}
}
return newCompositeConnChildLimiter(limiters)
}
type compositeConnChildLimiter struct {
limiters []connlimiter.Limiter
}
func newCompositeConnChildLimiter(limiters []connlimiter.Limiter) connlimiter.Limiter {
if len(limiters) == 0 {
return nil
}
if len(limiters) == 1 {
return limiters[0]
}
sort.Slice(limiters, func(i, j int) bool {
return limiters[i].Limit() < limiters[j].Limit()
})
return &compositeConnChildLimiter{limiters: limiters}
}
func (l *compositeConnChildLimiter) Allow(n int) (allowed bool) {
var i int
for i = range l.limiters {
if allowed = l.limiters[i].Allow(n); !allowed {
break
}
}
if !allowed && i > 0 && n > 0 {
for _, lim := range l.limiters[:i] {
lim.Allow(-n)
}
}
return allowed
}
func (l *compositeConnChildLimiter) Limit() int {
if len(l.limiters) == 0 {
return 0
}
return l.limiters[0].Limit()
}
@@ -0,0 +1,79 @@
package service
import (
"context"
"io"
"testing"
corelimiter "github.com/go-gost/core/limiter"
corelogger "github.com/go-gost/core/logger"
xconn "github.com/go-gost/x/limiter/conn"
xtraffic "github.com/go-gost/x/limiter/traffic"
xlogger "github.com/go-gost/x/logger"
"github.com/go-gost/x/registry"
)
func TestResolveTrafficLimiterComposesCommaSeparatedNames(t *testing.T) {
const totalName = "test_total_speed_composite"
const ruleName = "test_rule_speed_composite"
registry.TrafficLimiterRegistry().Unregister(totalName)
registry.TrafficLimiterRegistry().Unregister(ruleName)
defer registry.TrafficLimiterRegistry().Unregister(totalName)
defer registry.TrafficLimiterRegistry().Unregister(ruleName)
logger := xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))
if err := registry.TrafficLimiterRegistry().Register(totalName, xtraffic.NewTrafficLimiter(xtraffic.LimitsOption("$ 10B 10B"), xtraffic.LoggerOption(logger))); err != nil {
t.Fatalf("register total limiter: %v", err)
}
if err := registry.TrafficLimiterRegistry().Register(ruleName, xtraffic.NewTrafficLimiter(xtraffic.LimitsOption("0.0.0.0/0 3B 3B"), xtraffic.LoggerOption(logger))); err != nil {
t.Fatalf("register rule limiter: %v", err)
}
lim := resolveTrafficLimiter(totalName + "," + ruleName)
if lim == nil {
t.Fatalf("expected composite traffic limiter")
}
serviceLimiter := lim.In(context.Background(), "192.0.2.1:1000", corelimiter.ScopeOption(corelimiter.ScopeService))
if serviceLimiter == nil || serviceLimiter.Limit() != 10 {
t.Fatalf("expected service-scope total limiter 10, got %#v", serviceLimiter)
}
connLimiter := lim.In(context.Background(), "192.0.2.1:1000", corelimiter.ScopeOption(corelimiter.ScopeConn))
if connLimiter == nil || connLimiter.Limit() != 3 {
t.Fatalf("expected conn-scope per-IP limiter 3, got %#v", connLimiter)
}
}
func TestResolveConnLimiterComposesCommaSeparatedNames(t *testing.T) {
const totalName = "test_total_conn_composite"
const ruleName = "test_rule_conn_composite"
registry.ConnLimiterRegistry().Unregister(totalName)
registry.ConnLimiterRegistry().Unregister(ruleName)
defer registry.ConnLimiterRegistry().Unregister(totalName)
defer registry.ConnLimiterRegistry().Unregister(ruleName)
logger := xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))
if err := registry.ConnLimiterRegistry().Register(totalName, xconn.NewConnLimiter(xconn.LimitsOption("$ 2"), xconn.LoggerOption(logger))); err != nil {
t.Fatalf("register total conn limiter: %v", err)
}
if err := registry.ConnLimiterRegistry().Register(ruleName, xconn.NewConnLimiter(xconn.LimitsOption("$$ 1"), xconn.LoggerOption(logger))); err != nil {
t.Fatalf("register rule conn limiter: %v", err)
}
lim := resolveConnLimiter(totalName + "," + ruleName)
if lim == nil {
t.Fatalf("expected composite conn limiter")
}
clientLimiter := lim.Limiter("192.0.2.1")
if clientLimiter == nil || clientLimiter.Limit() != 1 {
t.Fatalf("expected composite client limiter with strictest limit 1, got %#v", clientLimiter)
}
if !clientLimiter.Allow(1) {
t.Fatalf("expected first connection to be allowed")
}
if clientLimiter.Allow(1) {
t.Fatalf("expected per-IP rule limiter to reject second connection")
}
if !lim.Limiter("192.0.2.2").Allow(1) {
t.Fatalf("expected another client to share total limiter but have independent per-IP capacity")
}
}
+2 -17
View File
@@ -3,7 +3,6 @@ package service
import (
"fmt"
"runtime"
"strconv"
"strings"
"time"
@@ -31,7 +30,6 @@ import (
logger_parser "github.com/go-gost/x/config/parsing/logger"
selector_parser "github.com/go-gost/x/config/parsing/selector"
tls_util "github.com/go-gost/x/internal/util/tls"
xtraffic "github.com/go-gost/x/limiter/traffic"
cache_limiter "github.com/go-gost/x/limiter/traffic/cache"
"github.com/go-gost/x/metadata"
mdutil "github.com/go-gost/x/metadata/util"
@@ -185,20 +183,7 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
var trafficLimiter listener.Option
if cfg.Limiter != "" {
lim := registry.TrafficLimiterRegistry().Get(cfg.Limiter)
if lim == nil {
// Try to parse as simple number (bandwidth in bytes/sec)
if val, err := strconv.Atoi(cfg.Limiter); err == nil && val > 0 {
lim = xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption(fmt.Sprintf("%s %dB %dB", xtraffic.ServiceLimitKey, val, val)),
)
}
if lim == nil {
lim = xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption(fmt.Sprintf("%s %s %s", xtraffic.ServiceLimitKey, cfg.Limiter, cfg.Limiter)),
)
}
}
lim := resolveTrafficLimiter(cfg.Limiter)
trafficLimiter = listener.TrafficLimiterOption(
cache_limiter.NewCachedTrafficLimiter(
lim,
@@ -216,7 +201,7 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
listener.AuthOption(auth_parser.Info(cfg.Listener.Auth)),
listener.TLSConfigOption(tlsConfig),
listener.AdmissionOption(xadmission.AdmissionGroup(admissions...)),
listener.ConnLimiterOption(registry.ConnLimiterRegistry().Get(cfg.CLimiter)),
listener.ConnLimiterOption(resolveConnLimiter(cfg.CLimiter)),
listener.ServiceOption(cfg.Name),
listener.ProxyProtocolOption(ppv),
listener.StatsOption(pStats),