diff --git a/.gitignore b/.gitignore index 3b00d91..e353435 100644 --- a/.gitignore +++ b/.gitignore @@ -188,9 +188,6 @@ go.work.sum # 编译器输出 vendor/ -# 流量统计或其他运行时生成的文件 -traffic/ - # ========================================== # Docker # ========================================== diff --git a/go-gost/x/limiter/traffic/cache/limiter.go b/go-gost/x/limiter/traffic/cache/limiter.go new file mode 100644 index 0000000..c52d685 --- /dev/null +++ b/go-gost/x/limiter/traffic/cache/limiter.go @@ -0,0 +1,161 @@ +package limiter + +import ( + "context" + "time" + + "github.com/go-gost/core/limiter" + "github.com/go-gost/core/limiter/traffic" + "github.com/go-gost/x/internal/util/cache" +) + +const ( + defaultRefreshInterval = 30 * time.Second + defaultCleanupInterval = 60 * time.Second +) + +type options struct { + refreshInterval time.Duration + cleanupInterval time.Duration + scope string +} + +type Option func(*options) + +func RefreshIntervalOption(interval time.Duration) Option { + return func(o *options) { + o.refreshInterval = interval + } +} + +func CleanupIntervalOption(interval time.Duration) Option { + return func(o *options) { + o.cleanupInterval = interval + } +} + +func ScopeOption(scope string) Option { + return func(o *options) { + o.scope = scope + } +} + +type cachedTrafficLimiter struct { + inLimits *cache.Cache + outLimits *cache.Cache + limiter traffic.TrafficLimiter + options options +} + +func NewCachedTrafficLimiter(limiter traffic.TrafficLimiter, opts ...Option) traffic.TrafficLimiter { + if limiter == nil { + return nil + } + + var options options + for _, opt := range opts { + opt(&options) + } + + if options.refreshInterval == 0 { + options.refreshInterval = defaultRefreshInterval + } + if options.refreshInterval < time.Second { + options.refreshInterval = time.Second + } + + if options.cleanupInterval == 0 { + options.cleanupInterval = defaultCleanupInterval + } + if options.cleanupInterval < time.Second { + options.cleanupInterval = time.Second + } + + lim := &cachedTrafficLimiter{ + inLimits: cache.NewCache(options.cleanupInterval), + outLimits: cache.NewCache(options.cleanupInterval), + limiter: limiter, + options: options, + } + return lim +} + +func (p *cachedTrafficLimiter) In(ctx context.Context, key string, opts ...limiter.Option) traffic.Limiter { + if p.limiter == nil { + return nil + } + + var options limiter.Options + for _, opt := range opts { + opt(&options) + } + + if p.options.scope != "" && p.options.scope != options.Scope { + return nil + } + + item := p.inLimits.Get(key) + lim, _ := item.Value().(traffic.Limiter) + if !item.Expired() { + return lim + } + + limNew := p.limiter.In(ctx, key, opts...) + if limNew == nil { + limNew = lim + } + if item == nil || !p.equal(lim, limNew) { + p.inLimits.Set(key, cache.NewItem(limNew, p.options.refreshInterval)) + return limNew + } + + p.inLimits.Set(key, cache.NewItem(lim, p.options.refreshInterval)) + + return lim +} + +func (p *cachedTrafficLimiter) Out(ctx context.Context, key string, opts ...limiter.Option) traffic.Limiter { + if p.limiter == nil { + return nil + } + + var options limiter.Options + for _, opt := range opts { + opt(&options) + } + + if p.options.scope != "" && p.options.scope != options.Scope { + return nil + } + + item := p.outLimits.Get(key) + lim, _ := item.Value().(traffic.Limiter) + if !item.Expired() { + return lim + } + + limNew := p.limiter.Out(ctx, key, opts...) + if limNew == nil { + limNew = lim + } + if item == nil || !p.equal(lim, limNew) { + p.outLimits.Set(key, cache.NewItem(limNew, p.options.refreshInterval)) + return limNew + } + + p.outLimits.Set(key, cache.NewItem(lim, p.options.refreshInterval)) + + return lim +} + +func (p *cachedTrafficLimiter) equal(lim1, lim2 traffic.Limiter) bool { + if lim1 == lim2 { + return true + } + + if lim1 == nil || lim2 == nil { + return false + } + + return lim1.Limit() == lim2.Limit() +} diff --git a/go-gost/x/limiter/traffic/generator.go b/go-gost/x/limiter/traffic/generator.go new file mode 100644 index 0000000..2e17835 --- /dev/null +++ b/go-gost/x/limiter/traffic/generator.go @@ -0,0 +1,31 @@ +package traffic + +import ( + limiter "github.com/go-gost/core/limiter/traffic" +) + +type limitGenerator struct { + in int + out int +} + +func newLimitGenerator(in, out int) *limitGenerator { + return &limitGenerator{ + in: in, + out: out, + } +} + +func (p *limitGenerator) In() limiter.Limiter { + if p == nil || p.in <= 0 { + return nil + } + return NewLimiter(p.in) +} + +func (p *limitGenerator) Out() limiter.Limiter { + if p == nil || p.out <= 0 { + return nil + } + return NewLimiter(p.out) +} diff --git a/go-gost/x/limiter/traffic/limiter.go b/go-gost/x/limiter/traffic/limiter.go new file mode 100644 index 0000000..0a100d8 --- /dev/null +++ b/go-gost/x/limiter/traffic/limiter.go @@ -0,0 +1,76 @@ +package traffic + +import ( + "context" + "fmt" + "sort" + "strconv" + + limiter "github.com/go-gost/core/limiter/traffic" + "golang.org/x/time/rate" +) + +type llimiter struct { + limiter *rate.Limiter +} + +func NewLimiter(r int) limiter.Limiter { + return &llimiter{ + limiter: rate.NewLimiter(rate.Limit(r), r), + } +} + +func (l *llimiter) Wait(ctx context.Context, n int) int { + if l.limiter.Burst() < n { + n = l.limiter.Burst() + } + l.limiter.WaitN(ctx, n) + return n +} + +func (l *llimiter) Limit() int { + return int(l.limiter.Limit()) +} + +func (l *llimiter) Set(n int) { + l.limiter.SetLimit(rate.Limit(n)) + l.limiter.SetBurst(n) +} + +func (l *llimiter) String() string { + return strconv.Itoa(int(l.limiter.Limit())) +} + +type limiterGroup struct { + limiters []limiter.Limiter +} + +func newLimiterGroup(limiters ...limiter.Limiter) *limiterGroup { + sort.Slice(limiters, func(i, j int) bool { + return limiters[i].Limit() < limiters[j].Limit() + }) + return &limiterGroup{limiters: limiters} +} + +func (l *limiterGroup) Wait(ctx context.Context, n int) int { + for i := range l.limiters { + if v := l.limiters[i].Wait(ctx, n); v < n { + n = v + } + } + return n +} + +func (l *limiterGroup) Limit() int { + if len(l.limiters) == 0 { + return 0 + } + + return l.limiters[0].Limit() +} + +func (l *limiterGroup) Set(n int) {} + +func (l *limiterGroup) String() string { + return fmt.Sprintf("%v", l.limiters) +} diff --git a/go-gost/x/limiter/traffic/plugin/grpc.go b/go-gost/x/limiter/traffic/plugin/grpc.go new file mode 100644 index 0000000..4a888ac --- /dev/null +++ b/go-gost/x/limiter/traffic/plugin/grpc.go @@ -0,0 +1,111 @@ +package traffic + +import ( + "context" + "io" + + "github.com/go-gost/core/limiter" + "github.com/go-gost/core/limiter/traffic" + "github.com/go-gost/core/logger" + "github.com/go-gost/plugin/limiter/traffic/proto" + "github.com/go-gost/x/internal/plugin" + xtraffic "github.com/go-gost/x/limiter/traffic" + "google.golang.org/grpc" +) + +type grpcPlugin struct { + conn grpc.ClientConnInterface + client proto.LimiterClient + log logger.Logger +} + +// NewGRPCPlugin creates a traffic limiter plugin based on gRPC. +func NewGRPCPlugin(name string, addr string, opts ...plugin.Option) traffic.TrafficLimiter { + var options plugin.Options + for _, opt := range opts { + opt(&options) + } + + log := logger.Default().WithFields(map[string]any{ + "kind": "limiter", + "limiter": name, + }) + conn, err := plugin.NewGRPCConn(addr, &options) + if err != nil { + log.Error(err) + } + + p := &grpcPlugin{ + conn: conn, + log: log, + } + if conn != nil { + p.client = proto.NewLimiterClient(conn) + } + return p +} + +func (p *grpcPlugin) In(ctx context.Context, key string, opts ...limiter.Option) traffic.Limiter { + if p.client == nil { + return nil + } + + var options limiter.Options + for _, opt := range opts { + opt(&options) + } + + r, err := p.client.Limit(ctx, + &proto.LimitRequest{ + Service: options.Service, + Scope: options.Scope, + Network: options.Network, + Addr: options.Addr, + Client: options.Client, + Src: options.Src, + }) + if err != nil { + p.log.Error(err) + return nil + } + + return xtraffic.NewLimiter(int(r.In)) +} + +func (p *grpcPlugin) Out(ctx context.Context, key string, opts ...limiter.Option) traffic.Limiter { + if p.client == nil { + return nil + } + + var options limiter.Options + for _, opt := range opts { + opt(&options) + } + + r, err := p.client.Limit(ctx, + &proto.LimitRequest{ + Service: options.Service, + Scope: options.Scope, + Network: options.Network, + Addr: options.Addr, + Client: options.Client, + Src: options.Src, + }) + if err != nil { + p.log.Error(err) + return nil + } + + return xtraffic.NewLimiter(int(r.Out)) +} + +func (p *grpcPlugin) Close() error { + if p.conn == nil { + return nil + } + + if closer, ok := p.conn.(io.Closer); ok { + return closer.Close() + } + return nil +} diff --git a/go-gost/x/limiter/traffic/plugin/http.go b/go-gost/x/limiter/traffic/plugin/http.go new file mode 100644 index 0000000..2842940 --- /dev/null +++ b/go-gost/x/limiter/traffic/plugin/http.go @@ -0,0 +1,161 @@ +package traffic + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + + "github.com/go-gost/core/limiter" + "github.com/go-gost/core/limiter/traffic" + "github.com/go-gost/core/logger" + "github.com/go-gost/x/internal/plugin" + xtraffic "github.com/go-gost/x/limiter/traffic" +) + +type httpPluginRequest struct { + Service string `json:"service"` + Scope string `json:"scope"` + Network string `json:"network"` + Addr string `json:"addr"` + Client string `json:"client"` + Src string `json:"src"` +} + +type httpPluginResponse struct { + In int64 `json:"in"` + Out int64 `json:"out"` +} + +type httpPlugin struct { + url string + client *http.Client + header http.Header + log logger.Logger +} + +// NewHTTPPlugin creates a traffic limiter plugin based on HTTP. +func NewHTTPPlugin(name string, url string, opts ...plugin.Option) traffic.TrafficLimiter { + var options plugin.Options + for _, opt := range opts { + opt(&options) + } + + return &httpPlugin{ + url: url, + client: plugin.NewHTTPClient(&options), + header: options.Header, + log: logger.Default().WithFields(map[string]any{ + "kind": "limiter", + "limiter": name, + }), + } +} + +func (p *httpPlugin) In(ctx context.Context, key string, opts ...limiter.Option) traffic.Limiter { + if p.client == nil { + return nil + } + + var options limiter.Options + for _, opt := range opts { + opt(&options) + } + + rb := httpPluginRequest{ + Service: options.Service, + Scope: options.Scope, + Network: options.Network, + Addr: options.Addr, + Client: options.Client, + Src: options.Src, + } + v, err := json.Marshal(&rb) + if err != nil { + p.log.Error(err) + return nil + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.url, bytes.NewReader(v)) + if err != nil { + p.log.Error(err) + return nil + } + + if p.header != nil { + req.Header = p.header.Clone() + } + req.Header.Set("Content-Type", "application/json") + resp, err := p.client.Do(req) + if err != nil { + p.log.Error(err) + return nil + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + p.log.Errorf("server return non-200 code: %s", resp.Status) + return nil + } + + res := httpPluginResponse{} + if err := json.NewDecoder(resp.Body).Decode(&res); err != nil { + p.log.Error(err) + return nil + } + return xtraffic.NewLimiter(int(res.In)) +} + +func (p *httpPlugin) Out(ctx context.Context, key string, opts ...limiter.Option) traffic.Limiter { + if p.client == nil { + return nil + } + + var options limiter.Options + for _, opt := range opts { + opt(&options) + } + + rb := httpPluginRequest{ + Service: options.Service, + Scope: options.Scope, + Network: options.Network, + Addr: options.Addr, + Client: options.Client, + Src: options.Src, + } + v, err := json.Marshal(&rb) + if err != nil { + p.log.Error(err) + return nil + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.url, bytes.NewReader(v)) + if err != nil { + p.log.Error(err) + return nil + } + + if p.header != nil { + req.Header = p.header.Clone() + } + req.Header.Set("Content-Type", "application/json") + resp, err := p.client.Do(req) + if err != nil { + p.log.Error(err) + return nil + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + p.log.Errorf("server return non-200 code: %s", resp.Status) + return nil + } + + res := httpPluginResponse{} + if err := json.NewDecoder(resp.Body).Decode(&res); err != nil { + p.log.Error(err) + return nil + } + return xtraffic.NewLimiter(int(res.Out)) +} diff --git a/go-gost/x/limiter/traffic/traffic.go b/go-gost/x/limiter/traffic/traffic.go new file mode 100644 index 0000000..6f9c2b2 --- /dev/null +++ b/go-gost/x/limiter/traffic/traffic.go @@ -0,0 +1,626 @@ +package traffic + +import ( + "bufio" + "context" + "io" + "net" + "strings" + "sync" + "time" + + "github.com/alecthomas/units" + "github.com/go-gost/core/limiter" + "github.com/go-gost/core/limiter/traffic" + "github.com/go-gost/core/logger" + "github.com/go-gost/x/internal/loader" + "github.com/patrickmn/go-cache" + "github.com/yl2chen/cidranger" +) + +const ( + ServiceLimitKey = "$" + ConnLimitKey = "$$" +) + +const ( + defaultExpiration = 15 * time.Second + cleanupInterval = 30 * time.Second +) + +type options struct { + limits []string + fileLoader loader.Loader + redisLoader loader.Loader + httpLoader loader.Loader + period time.Duration + logger logger.Logger +} + +type Option func(opts *options) + +func LimitsOption(limits ...string) Option { + return func(opts *options) { + opts.limits = limits + } +} + +func ReloadPeriodOption(period time.Duration) Option { + return func(opts *options) { + opts.period = period + } +} + +func FileLoaderOption(fileLoader loader.Loader) Option { + return func(opts *options) { + opts.fileLoader = fileLoader + } +} + +func RedisLoaderOption(redisLoader loader.Loader) Option { + return func(opts *options) { + opts.redisLoader = redisLoader + } +} + +func HTTPLoaderOption(httpLoader loader.Loader) Option { + return func(opts *options) { + opts.httpLoader = httpLoader + } +} + +func LoggerOption(logger logger.Logger) Option { + return func(opts *options) { + opts.logger = logger + } +} + +type limitValue struct { + in int + out int +} + +type trafficLimiter struct { + generators sync.Map + cidrGenerators cidranger.Ranger + // connection level in/out limits + connInLimits *cache.Cache + connOutLimits *cache.Cache + // service level in/out limits + inLimits *cache.Cache + outLimits *cache.Cache + mu sync.RWMutex + cancelFunc context.CancelFunc + options options +} + +func NewTrafficLimiter(opts ...Option) traffic.TrafficLimiter { + var options options + for _, opt := range opts { + opt(&options) + } + + ctx, cancel := context.WithCancel(context.TODO()) + lim := &trafficLimiter{ + cidrGenerators: cidranger.NewPCTrieRanger(), + connInLimits: cache.New(defaultExpiration, cleanupInterval), + connOutLimits: cache.New(defaultExpiration, cleanupInterval), + inLimits: cache.New(defaultExpiration, cleanupInterval), + outLimits: cache.New(defaultExpiration, cleanupInterval), + options: options, + cancelFunc: cancel, + } + + if err := lim.reload(ctx); err != nil { + options.logger.Warnf("reload: %v", err) + } + if lim.options.period > 0 { + go lim.periodReload(ctx) + } + return lim +} + +// In obtains a traffic input limiter based on key. +// For connection scope, the key should be client connection address. +func (l *trafficLimiter) In(ctx context.Context, key string, opts ...limiter.Option) traffic.Limiter { + var options limiter.Options + for _, opt := range opts { + opt(&options) + } + + switch options.Scope { + case limiter.ScopeService: + if lim, ok := l.inLimits.Get(ServiceLimitKey); ok && lim != nil { + return lim.(traffic.Limiter) + } + return nil + + case limiter.ScopeClient: + return nil + + case limiter.ScopeConn: + fallthrough + default: + } + + var lims []traffic.Limiter + + // connection level limiter + if lim, ok := l.connInLimits.Get(key); ok { + if lim != nil { + // cached connection level limiter + lims = append(lims, lim.(traffic.Limiter)) + // reset expiration + l.connInLimits.Set(key, lim, defaultExpiration) + } + } else { + // generate a new connection level limiter and cache it + if v, ok := l.generators.Load(ConnLimitKey); ok && v != nil { + lim := v.(*limitGenerator).In() + if lim != nil { + lims = append(lims, lim) + l.connInLimits.Set(key, lim, defaultExpiration) + } + } + } + + host, _, _ := net.SplitHostPort(key) + // IP level limiter + if lim, ok := l.inLimits.Get(host); ok { + // cached IP limiter + if lim != nil { + lims = append(lims, lim.(traffic.Limiter)) + } + } else { + l.mu.RLock() + ranger := l.cidrGenerators + l.mu.RUnlock() + + // CIDR level limiter + if p, _ := ranger.ContainingNetworks(net.ParseIP(host)); len(p) > 0 { + if v, _ := p[0].(*cidrLimitEntry); v != nil { + if lim := v.generator.In(); lim != nil { + lims = append(lims, lim) + l.inLimits.Set(host, lim, cache.NoExpiration) + } + } + } + } + + var lim traffic.Limiter + if len(lims) > 0 { + lim = newLimiterGroup(lims...) + } + + if lim != nil && l.options.logger != nil { + l.options.logger.Debugf("input limit for %s: %s", key, lim) + } + + return lim +} + +// Out obtains a traffic output limiter based on key. +// For connection scope, the key should be client connection address. +func (l *trafficLimiter) Out(ctx context.Context, key string, opts ...limiter.Option) traffic.Limiter { + var options limiter.Options + for _, opt := range opts { + opt(&options) + } + + switch options.Scope { + case limiter.ScopeService: + if lim, ok := l.outLimits.Get(ServiceLimitKey); ok && lim != nil { + return lim.(traffic.Limiter) + } + return nil + + case limiter.ScopeClient: + return nil + + case limiter.ScopeConn: + fallthrough + default: + } + + var lims []traffic.Limiter + + // connection level limiter + if lim, ok := l.connOutLimits.Get(key); ok { + if lim != nil { + // cached connection level limiter + lims = append(lims, lim.(traffic.Limiter)) + // reset expiration + l.connOutLimits.Set(key, lim, defaultExpiration) + } + } else { + // generate a new connection level limiter + if v, ok := l.generators.Load(ConnLimitKey); ok && v != nil { + lim := v.(*limitGenerator).Out() + if lim != nil { + lims = append(lims, lim) + l.connOutLimits.Set(key, lim, defaultExpiration) + } + } + } + + host, _, _ := net.SplitHostPort(key) + // IP level limiter + if lim, ok := l.outLimits.Get(host); ok { + if lim != nil { + // cached IP level limiter + lims = append(lims, lim.(traffic.Limiter)) + } + } else { + l.mu.RLock() + ranger := l.cidrGenerators + l.mu.RUnlock() + + // CIDR level limiter + if p, _ := ranger.ContainingNetworks(net.ParseIP(host)); len(p) > 0 { + if v, _ := p[0].(*cidrLimitEntry); v != nil { + if lim := v.generator.Out(); lim != nil { + lims = append(lims, lim) + l.outLimits.Set(host, lim, cache.NoExpiration) + } + } + } + } + + var lim traffic.Limiter + if len(lims) > 0 { + lim = newLimiterGroup(lims...) + } + + if lim != nil && l.options.logger != nil { + l.options.logger.Debugf("output limit for %s: %s", key, lim) + } + + return lim +} + +func (l *trafficLimiter) periodReload(ctx context.Context) error { + period := l.options.period + if period < time.Second { + period = time.Second + } + ticker := time.NewTicker(period) + defer ticker.Stop() + + for { + select { + case <-ticker.C: + if err := l.reload(ctx); err != nil { + l.options.logger.Warnf("reload: %v", err) + // return err + } + case <-ctx.Done(): + return ctx.Err() + } + } +} + +func (l *trafficLimiter) reload(ctx context.Context) error { + values, err := l.load(ctx) + if err != nil { + return err + } + + // service level limiter, never expired + { + value := values[ServiceLimitKey] + if v, _ := l.inLimits.Get(ServiceLimitKey); v != nil { + lim := v.(traffic.Limiter) + if value.in <= 0 { + l.inLimits.Delete(ServiceLimitKey) + } else { + lim.Set(value.in) + } + } else { + if value.in > 0 { + l.inLimits.Set(ServiceLimitKey, NewLimiter(value.in), cache.NoExpiration) + } + } + + if v, _ := l.outLimits.Get(ServiceLimitKey); v != nil { + lim := v.(traffic.Limiter) + if value.out <= 0 { + l.outLimits.Delete(ServiceLimitKey) + } else { + lim.Set(value.out) + } + } else { + if value.out > 0 { + l.outLimits.Set(ServiceLimitKey, NewLimiter(value.out), cache.NoExpiration) + } + } + delete(values, ServiceLimitKey) + } + + // connection level limiters + { + value := values[ConnLimitKey] + + var in, out int + if v, _ := l.generators.Load(ConnLimitKey); v != nil { + in, out = v.(*limitGenerator).in, v.(*limitGenerator).out + } + l.generators.Store(ConnLimitKey, newLimitGenerator(value.in, value.out)) + + if value.in <= 0 { + l.connInLimits.Flush() + } else { + if in != value.in { + for _, item := range l.connInLimits.Items() { + if v := item.Object; v != nil { + v.(traffic.Limiter).Set(in) + } + } + } + } + + if value.out <= 0 { + l.connOutLimits.Flush() + } else { + if out != value.out { + for _, item := range l.connOutLimits.Items() { + if v := item.Object; v != nil { + v.(traffic.Limiter).Set(out) + } + } + } + } + delete(values, ConnLimitKey) + } + + cidrGenerators := cidranger.NewPCTrieRanger() + // IP/CIDR level limiters + { + // snapshot of the current limiters + inLimits := l.inLimits.Items() + outLimits := l.outLimits.Items() + + delete(inLimits, ServiceLimitKey) + delete(outLimits, ServiceLimitKey) + + for key, value := range values { + if _, ipNet, _ := net.ParseCIDR(key); ipNet != nil { + cidrGenerators.Insert(&cidrLimitEntry{ + ipNet: *ipNet, + generator: newLimitGenerator(value.in, value.out), + }) + continue + } + + if v, _ := l.inLimits.Get(key); v != nil { + lim := v.(traffic.Limiter) + if value.in <= 0 { + l.inLimits.Delete(key) + } else { + lim.Set(value.in) + } + delete(inLimits, key) + } else { + if value.in > 0 { + l.inLimits.Set(key, NewLimiter(value.in), cache.NoExpiration) + } + } + + if v, _ := l.outLimits.Get(key); v != nil { + lim := v.(traffic.Limiter) + if value.out <= 0 { + l.outLimits.Delete(key) + } else { + lim.Set(value.out) + } + delete(outLimits, key) + } else { + if value.out > 0 { + l.outLimits.Set(key, NewLimiter(value.out), cache.NoExpiration) + } + } + } + + // check the CIDR for remain limiters, clean the unmatched ones. + for k, v := range inLimits { + if p, _ := cidrGenerators.ContainingNetworks(net.ParseIP(k)); len(p) > 0 { + if le, _ := p[0].(*cidrLimitEntry); le != nil { + in := le.generator.in + if in <= 0 { + l.inLimits.Delete(k) + continue + } + lim := v.Object.(traffic.Limiter) + if lim.Limit() != in { + lim.Set(in) + } + } + } else { + l.inLimits.Delete(k) + } + } + for k, v := range outLimits { + if p, _ := cidrGenerators.ContainingNetworks(net.ParseIP(k)); len(p) > 0 { + if le, _ := p[0].(*cidrLimitEntry); le != nil { + out := le.generator.out + if out <= 0 { + l.outLimits.Delete(k) + continue + } + lim := v.Object.(traffic.Limiter) + if lim.Limit() != out { + lim.Set(out) + } + delete(outLimits, k) + } + } else { + l.outLimits.Delete(k) + } + } + } + + l.mu.Lock() + defer l.mu.Unlock() + + l.cidrGenerators = cidrGenerators + + return nil +} + +func (l *trafficLimiter) load(ctx context.Context) (values map[string]limitValue, err error) { + values = make(map[string]limitValue) + + for _, v := range l.options.limits { + key, in, out := l.parseLimit(v) + if key == "" { + continue + } + values[key] = limitValue{in: in, out: out} + } + + if l.options.fileLoader != nil { + if lister, ok := l.options.fileLoader.(loader.Lister); ok { + list, er := lister.List(ctx) + if er != nil { + l.options.logger.Warnf("file loader: %v", er) + } + for _, s := range list { + key, in, out := l.parseLimit(l.parseLine(s)) + if key == "" { + continue + } + values[key] = limitValue{in: in, out: out} + } + } else { + r, er := l.options.fileLoader.Load(ctx) + if er != nil { + l.options.logger.Warnf("file loader: %v", er) + } + patterns, _ := l.parsePatterns(r) + for _, s := range patterns { + key, in, out := l.parseLimit(l.parseLine(s)) + if key == "" { + continue + } + values[key] = limitValue{in: in, out: out} + } + } + } + if l.options.redisLoader != nil { + if lister, ok := l.options.redisLoader.(loader.Lister); ok { + list, er := lister.List(ctx) + if er != nil { + l.options.logger.Warnf("redis loader: %v", er) + } + for _, s := range list { + key, in, out := l.parseLimit(l.parseLine(s)) + if key == "" { + continue + } + values[key] = limitValue{in: in, out: out} + } + } else { + r, er := l.options.redisLoader.Load(ctx) + if er != nil { + l.options.logger.Warnf("redis loader: %v", er) + } + patterns, _ := l.parsePatterns(r) + for _, s := range patterns { + key, in, out := l.parseLimit(l.parseLine(s)) + if key == "" { + continue + } + values[key] = limitValue{in: in, out: out} + } + } + } + if l.options.httpLoader != nil { + r, er := l.options.httpLoader.Load(ctx) + if er != nil { + l.options.logger.Warnf("http loader: %v", er) + } + patterns, _ := l.parsePatterns(r) + for _, s := range patterns { + key, in, out := l.parseLimit(l.parseLine(s)) + if key == "" { + continue + } + values[key] = limitValue{in: in, out: out} + } + } + + l.options.logger.Debugf("load items %d", len(values)) + return +} + +func (l *trafficLimiter) parsePatterns(r io.Reader) (patterns []string, err error) { + if r == nil { + return + } + + scanner := bufio.NewScanner(r) + for scanner.Scan() { + if line := l.parseLine(scanner.Text()); line != "" { + patterns = append(patterns, line) + } + } + + err = scanner.Err() + return +} + +func (l *trafficLimiter) parseLine(s string) string { + if n := strings.IndexByte(s, '#'); n >= 0 { + s = s[:n] + } + return strings.TrimSpace(s) +} + +func (l *trafficLimiter) parseLimit(s string) (key string, in, out int) { + s = strings.Replace(s, "\t", " ", -1) + s = strings.TrimSpace(s) + if s == "" { + return + } + + var ss []string + for _, v := range strings.Split(s, " ") { + if v != "" { + ss = append(ss, v) + } + } + if len(ss) < 2 { + return + } + + key = ss[0] + if v, _ := units.ParseBase2Bytes(ss[1]); v > 0 { + in = int(v) + } + if len(ss) > 2 { + if v, _ := units.ParseBase2Bytes(ss[2]); v > 0 { + out = int(v) + } + } + + return +} + +func (l *trafficLimiter) Close() error { + l.cancelFunc() + if l.options.fileLoader != nil { + l.options.fileLoader.Close() + } + if l.options.redisLoader != nil { + l.options.redisLoader.Close() + } + return nil +} + +type cidrLimitEntry struct { + ipNet net.IPNet + generator *limitGenerator +} + +func (p *cidrLimitEntry) Network() net.IPNet { + return p.ipNet +} diff --git a/go-gost/x/limiter/traffic/wrapper/conn.go b/go-gost/x/limiter/traffic/wrapper/conn.go new file mode 100644 index 0000000..0797966 --- /dev/null +++ b/go-gost/x/limiter/traffic/wrapper/conn.go @@ -0,0 +1,420 @@ +package wrapper + +import ( + "bytes" + "context" + "errors" + "io" + "net" + "syscall" + + "github.com/go-gost/core/limiter" + "github.com/go-gost/core/limiter/traffic" + "github.com/go-gost/core/metadata" + xnet "github.com/go-gost/x/internal/net" + "github.com/go-gost/x/internal/net/udp" +) + +var ( + errUnsupport = errors.New("unsupported operation") +) + +// limitConn is a Conn with traffic limiter supported. +type limitConn struct { + net.Conn + rbuf bytes.Buffer + limiter traffic.TrafficLimiter + opts []limiter.Option + key string +} + +func WrapConn(c net.Conn, tlimiter traffic.TrafficLimiter, key string, opts ...limiter.Option) net.Conn { + if tlimiter == nil { + return c + } + + return &limitConn{ + Conn: c, + limiter: tlimiter, + opts: opts, + key: key, + } +} + +func (c *limitConn) Read(b []byte) (n int, err error) { + limiter := c.limiter.In(context.Background(), c.key, c.opts...) + if limiter == nil || limiter.Limit() <= 0 { + return c.Conn.Read(b) + } + + if c.rbuf.Len() > 0 { + burst := len(b) + if c.rbuf.Len() < burst { + burst = c.rbuf.Len() + } + lim := limiter.Wait(context.Background(), burst) + return c.rbuf.Read(b[:lim]) + } + + nn, err := c.Conn.Read(b) + if err != nil { + return nn, err + } + + n = limiter.Wait(context.Background(), nn) + if n < nn { + if _, err = c.rbuf.Write(b[n:nn]); err != nil { + return 0, err + } + } + + return +} + +func (c *limitConn) Write(b []byte) (n int, err error) { + limiter := c.limiter.Out(context.Background(), c.key, c.opts...) + if limiter == nil || limiter.Limit() <= 0 { + return c.Conn.Write(b) + } + + nn := 0 + for len(b) > 0 { + nn, err = c.Conn.Write(b[:limiter.Wait(context.Background(), len(b))]) + n += nn + if err != nil { + return + } + b = b[nn:] + } + + return +} + +func (c *limitConn) SyscallConn() (rc syscall.RawConn, err error) { + if sc, ok := c.Conn.(syscall.Conn); ok { + rc, err = sc.SyscallConn() + return + } + err = errUnsupport + return +} + +func (c *limitConn) Metadata() metadata.Metadata { + if md, ok := c.Conn.(metadata.Metadatable); ok { + return md.Metadata() + } + return nil +} + +type packetConn struct { + net.PacketConn + limiter traffic.TrafficLimiter + opts []limiter.Option + key string +} + +func WrapPacketConn(pc net.PacketConn, lim traffic.TrafficLimiter, key string, opts ...limiter.Option) net.PacketConn { + if lim == nil { + return pc + } + return &packetConn{ + PacketConn: pc, + limiter: lim, + opts: opts, + key: key, + } +} + +func (c *packetConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { + for { + n, addr, err = c.PacketConn.ReadFrom(p) + if err != nil { + return + } + + limiter := c.limiter.In(context.Background(), c.key, c.opts...) + if limiter == nil || limiter.Limit() <= 0 { + return + } + + // discard when exceed the limit size. + if limiter.Wait(context.Background(), n) < n { + continue + } + + return + } +} + +func (c *packetConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { + // discard when exceed the limit size. + limiter := c.limiter.Out(context.Background(), c.key, c.opts...) + if limiter != nil && limiter.Limit() > 0 && + limiter.Wait(context.Background(), len(p)) < len(p) { + n = len(p) + return + } + + return c.PacketConn.WriteTo(p, addr) +} + +func (c *packetConn) Metadata() metadata.Metadata { + if md, ok := c.PacketConn.(metadata.Metadatable); ok { + return md.Metadata() + } + return nil +} + +type udpConn struct { + net.PacketConn + limiter traffic.TrafficLimiter + opts []limiter.Option + key string +} + +func WrapUDPConn(pc net.PacketConn, limiter traffic.TrafficLimiter, key string, opts ...limiter.Option) udp.Conn { + return &udpConn{ + PacketConn: pc, + limiter: limiter, + opts: opts, + key: key, + } +} + +func (c *udpConn) RemoteAddr() net.Addr { + if nc, ok := c.PacketConn.(xnet.RemoteAddr); ok { + return nc.RemoteAddr() + } + return nil +} + +func (c *udpConn) SetReadBuffer(n int) error { + if nc, ok := c.PacketConn.(xnet.SetBuffer); ok { + return nc.SetReadBuffer(n) + } + return errUnsupport +} + +func (c *udpConn) SetWriteBuffer(n int) error { + if nc, ok := c.PacketConn.(xnet.SetBuffer); ok { + return nc.SetWriteBuffer(n) + } + return errUnsupport +} + +func (c *udpConn) Read(b []byte) (n int, err error) { + nc, ok := c.PacketConn.(io.Reader) + if !ok { + err = errUnsupport + return + } + + for { + n, err = nc.Read(b) + if err != nil { + return + } + + if c.limiter == nil { + return + } + + limiter := c.limiter.In(context.Background(), c.key, c.opts...) + if limiter == nil || limiter.Limit() <= 0 { + return + } + + // discard when exceed the limit size. + if limiter.Wait(context.Background(), n) < n { + continue + } + + return + } +} + +func (c *udpConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { + for { + n, addr, err = c.PacketConn.ReadFrom(p) + if err != nil { + return + } + + if c.limiter == nil { + return + } + + limiter := c.limiter.In(context.Background(), c.key, c.opts...) + if limiter == nil || limiter.Limit() <= 0 { + return + } + + // discard when exceed the limit size. + if limiter.Wait(context.Background(), n) < n { + continue + } + + return + } +} + +func (c *udpConn) ReadFromUDP(b []byte) (n int, addr *net.UDPAddr, err error) { + nc, ok := c.PacketConn.(udp.ReadUDP) + if !ok { + err = errUnsupport + return + } + + for { + n, addr, err = nc.ReadFromUDP(b) + if err != nil { + return + } + + if c.limiter == nil { + return + } + + limiter := c.limiter.In(context.Background(), c.key, c.opts...) + if limiter == nil || limiter.Limit() <= 0 { + return + } + + // discard when exceed the limit size. + if limiter.Wait(context.Background(), n) < n { + continue + } + + return + } +} + +func (c *udpConn) ReadMsgUDP(b, oob []byte) (n, oobn, flags int, addr *net.UDPAddr, err error) { + nc, ok := c.PacketConn.(udp.ReadUDP) + if !ok { + err = errUnsupport + return + } + + for { + n, oobn, flags, addr, err = nc.ReadMsgUDP(b, oob) + if err != nil { + return + } + + if c.limiter == nil { + return + } + + limiter := c.limiter.In(context.Background(), c.key, c.opts...) + if limiter == nil || limiter.Limit() <= 0 { + return + } + + // discard when exceed the limit size. + if limiter.Wait(context.Background(), n) < n { + continue + } + return + } +} + +func (c *udpConn) Write(p []byte) (n int, err error) { + nc, ok := c.PacketConn.(io.Writer) + if !ok { + err = errUnsupport + return + } + + if c.limiter != nil { + // discard when exceed the limit size. + limiter := c.limiter.Out(context.Background(), c.key, c.opts...) + if limiter != nil && limiter.Limit() > 0 && + limiter.Wait(context.Background(), len(p)) < len(p) { + n = len(p) + return + } + } + + n, err = nc.Write(p) + return +} + +func (c *udpConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { + if c.limiter != nil { + // discard when exceed the limit size. + limiter := c.limiter.Out(context.Background(), c.key, c.opts...) + if limiter != nil && limiter.Limit() > 0 && + limiter.Wait(context.Background(), len(p)) < len(p) { + n = len(p) + return + } + } + + n, err = c.PacketConn.WriteTo(p, addr) + return +} + +func (c *udpConn) WriteToUDP(p []byte, addr *net.UDPAddr) (n int, err error) { + nc, ok := c.PacketConn.(udp.WriteUDP) + if !ok { + err = errUnsupport + return + } + + if c.limiter != nil { + // discard when exceed the limit size. + limiter := c.limiter.Out(context.Background(), c.key, c.opts...) + if limiter != nil && limiter.Limit() > 0 && + limiter.Wait(context.Background(), len(p)) < len(p) { + n = len(p) + return + } + } + + n, err = nc.WriteToUDP(p, addr) + return +} + +func (c *udpConn) WriteMsgUDP(p, oob []byte, addr *net.UDPAddr) (n, oobn int, err error) { + nc, ok := c.PacketConn.(udp.WriteUDP) + if !ok { + err = errUnsupport + return + } + + if c.limiter != nil { + // discard when exceed the limit size. + limiter := c.limiter.Out(context.Background(), c.key, c.opts...) + if limiter != nil && limiter.Limit() > 0 && + limiter.Wait(context.Background(), len(p)) < len(p) { + n = len(p) + return + } + } + + n, oobn, err = nc.WriteMsgUDP(p, oob, addr) + return +} + +func (c *udpConn) SyscallConn() (rc syscall.RawConn, err error) { + if nc, ok := c.PacketConn.(xnet.SyscallConn); ok { + return nc.SyscallConn() + } + err = errUnsupport + return +} + +func (c *udpConn) SetDSCP(n int) error { + if nc, ok := c.PacketConn.(xnet.SetDSCP); ok { + return nc.SetDSCP(n) + } + return nil +} + +func (c *udpConn) Metadata() metadata.Metadata { + if md, ok := c.PacketConn.(metadata.Metadatable); ok { + return md.Metadata() + } + return nil +} diff --git a/go-gost/x/limiter/traffic/wrapper/io.go b/go-gost/x/limiter/traffic/wrapper/io.go new file mode 100644 index 0000000..4dbb5ce --- /dev/null +++ b/go-gost/x/limiter/traffic/wrapper/io.go @@ -0,0 +1,81 @@ +package wrapper + +import ( + "bytes" + "context" + "io" + + "github.com/go-gost/core/limiter" + "github.com/go-gost/core/limiter/traffic" +) + +// readWriter is an io.ReadWriter with traffic limiter supported. +type readWriter struct { + io.ReadWriter + rbuf bytes.Buffer + limiter traffic.TrafficLimiter + opts []limiter.Option + key string +} + +func WrapReadWriter(limiter traffic.TrafficLimiter, rw io.ReadWriter, key string, opts ...limiter.Option) io.ReadWriter { + if limiter == nil { + return rw + } + + return &readWriter{ + ReadWriter: rw, + limiter: limiter, + opts: opts, + key: key, + } +} + +func (p *readWriter) Read(b []byte) (n int, err error) { + limiter := p.limiter.In(context.Background(), p.key, p.opts...) + if limiter == nil || limiter.Limit() <= 0 { + return p.ReadWriter.Read(b) + } + + if p.rbuf.Len() > 0 { + burst := len(b) + if p.rbuf.Len() < burst { + burst = p.rbuf.Len() + } + lim := limiter.Wait(context.Background(), burst) + return p.rbuf.Read(b[:lim]) + } + + nn, err := p.ReadWriter.Read(b) + if err != nil { + return nn, err + } + + n = limiter.Wait(context.Background(), nn) + if n < nn { + if _, err = p.rbuf.Write(b[n:nn]); err != nil { + return 0, err + } + } + + return +} + +func (p *readWriter) Write(b []byte) (n int, err error) { + limiter := p.limiter.Out(context.Background(), p.key, p.opts...) + if limiter == nil || limiter.Limit() <= 0 { + return p.ReadWriter.Write(b) + } + + nn := 0 + for len(b) > 0 { + nn, err = p.ReadWriter.Write(b[:limiter.Wait(context.Background(), len(b))]) + n += nn + if err != nil { + return + } + b = b[nn:] + } + + return +} diff --git a/go-gost/x/limiter/traffic/wrapper/listener.go b/go-gost/x/limiter/traffic/wrapper/listener.go new file mode 100644 index 0000000..8af59e6 --- /dev/null +++ b/go-gost/x/limiter/traffic/wrapper/listener.go @@ -0,0 +1,40 @@ +package wrapper + +import ( + "net" + + "github.com/go-gost/core/limiter" + "github.com/go-gost/core/limiter/traffic" + traffic_limiter "github.com/go-gost/x/limiter/traffic" +) + +type listener struct { + net.Listener + limiter traffic.TrafficLimiter + service string +} + +func WrapListener(service string, ln net.Listener, limiter traffic.TrafficLimiter) net.Listener { + if limiter == nil { + return ln + } + + return &listener{ + Listener: ln, + limiter: limiter, + service: service, + } +} + +func (ln *listener) Accept() (net.Conn, error) { + c, err := ln.Listener.Accept() + if err != nil { + return nil, err + } + + return WrapConn(c, ln.limiter, traffic_limiter.ServiceLimitKey, + limiter.ScopeOption(limiter.ScopeService), + limiter.ServiceOption(ln.service), + limiter.NetworkOption(ln.Addr().Network()), + ), nil +} diff --git a/go-gost/x/traffic/memory_manager.go b/go-gost/x/traffic/memory_manager.go new file mode 100644 index 0000000..0137091 --- /dev/null +++ b/go-gost/x/traffic/memory_manager.go @@ -0,0 +1,91 @@ +package traffic + +import ( + "context" + "sync" + "sync/atomic" +) + +// MemoryManager 内存流量管理器 +type MemoryManager struct { + mu sync.RWMutex + stats map[string]*trafficStats +} + +// trafficStats 流量统计 +type trafficStats struct { + upload atomic.Int64 + download atomic.Int64 +} + +// NewMemoryManager 创建内存流量管理器 +func NewMemoryManager() *MemoryManager { + return &MemoryManager{ + stats: make(map[string]*trafficStats), + } +} + +// RecordTraffic 记录流量 +func (m *MemoryManager) RecordTraffic(ctx context.Context, service string, upload, download int64) error { + m.mu.RLock() + stats, exists := m.stats[service] + m.mu.RUnlock() + + if !exists { + m.mu.Lock() + stats, exists = m.stats[service] + if !exists { + stats = &trafficStats{} + m.stats[service] = stats + } + m.mu.Unlock() + } + + if upload > 0 { + stats.upload.Add(upload) + } + if download > 0 { + stats.download.Add(download) + } + + return nil +} + +// GetAllServicesStats 获取所有服务的流量统计 +func (m *MemoryManager) GetAllServicesStats(ctx context.Context) (map[string]map[string]int64, error) { + m.mu.RLock() + defer m.mu.RUnlock() + + result := make(map[string]map[string]int64) + for service, stats := range m.stats { + result[service] = map[string]int64{ + "upload": stats.upload.Load(), + "download": stats.download.Load(), + } + } + + return result, nil +} + +// ClearAllTrafficStats 清零所有流量统计 +func (m *MemoryManager) ClearAllTrafficStats(ctx context.Context) error { + m.mu.RLock() + defer m.mu.RUnlock() + + for _, stats := range m.stats { + stats.upload.Store(0) + stats.download.Store(0) + } + + return nil +} + +// Close 关闭管理器(内存管理器无需特殊清理) +func (m *MemoryManager) Close() error { + return nil +} + +// TestConnection 测试连接(内存管理器总是返回成功) +func (m *MemoryManager) TestConnection(ctx context.Context) error { + return nil +} diff --git a/go-gost/x/traffic/traffic.go b/go-gost/x/traffic/traffic.go new file mode 100644 index 0000000..2265571 --- /dev/null +++ b/go-gost/x/traffic/traffic.go @@ -0,0 +1,33 @@ +package traffic + +import ( + "context" + "sync" +) + +// Manager 流量管理器接口 +type Manager interface { + RecordTraffic(ctx context.Context, service string, upload, download int64) error + GetAllServicesStats(ctx context.Context) (map[string]map[string]int64, error) + ClearAllTrafficStats(ctx context.Context) error + Close() error + TestConnection(ctx context.Context) error +} + +var ( + globalManager Manager + once sync.Once +) + +// GetGlobalManager 获取全局流量管理器实例 +func GetGlobalManager() Manager { + once.Do(func() { + globalManager = NewMemoryManager() + }) + return globalManager +} + +// SetGlobalManager 设置全局流量管理器(用于测试) +func SetGlobalManager(m Manager) { + globalManager = m +}