From 3310474137726c3839ace91ba7206b8e02b617a4 Mon Sep 17 00:00:00 2001 From: Shtorm <108103062+shtorm-7@users.noreply.github.com> Date: Sat, 8 Aug 2026 10:11:10 +0300 Subject: [PATCH] Fix limiters --- protocol/limiter/bandwidth/outbound.go | 18 ++++++++++--- protocol/limiter/bandwidth/strategy.go | 23 ++++++++++++++-- protocol/limiter/traffic/outbound.go | 18 ++++++++++--- protocol/limiter/traffic/strategy.go | 37 +++++++++++++++++++------- 4 files changed, 77 insertions(+), 19 deletions(-) diff --git a/protocol/limiter/bandwidth/outbound.go b/protocol/limiter/bandwidth/outbound.go index efa2f083..efc6204c 100644 --- a/protocol/limiter/bandwidth/outbound.go +++ b/protocol/limiter/bandwidth/outbound.go @@ -106,7 +106,12 @@ func (h *Outbound) DialContext(ctx context.Context, network string, destination if err != nil { return nil, err } - return h.strategy.wrapConn(ctx, conn, adapter.ContextFrom(ctx), true) + wrappedConn, err := h.strategy.wrapConn(ctx, conn, adapter.ContextFrom(ctx), true) + if err != nil { + conn.Close() + return nil, err + } + return wrappedConn, nil } func (h *Outbound) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) { @@ -114,12 +119,17 @@ func (h *Outbound) ListenPacket(ctx context.Context, destination M.Socksaddr) (n if err != nil { return nil, err } - return h.strategy.wrapPacketConn(ctx, conn, adapter.ContextFrom(ctx), true) + wrappedConn, err := h.strategy.wrapPacketConn(ctx, conn, adapter.ContextFrom(ctx), true) + if err != nil { + conn.Close() + return nil, err + } + return wrappedConn, nil } func (h *Outbound) NewConnectionEx(ctx context.Context, conn net.Conn, metadata adapter.InboundContext, onClose N.CloseHandlerFunc) { ctx = adapter.WithContext(ctx, &metadata) - conn, err := h.strategy.wrapConn(ctx, conn, &metadata, false) + wrappedConn, err := h.strategy.wrapConn(ctx, conn, &metadata, false) if err != nil { h.logger.ErrorContext(ctx, err) N.CloseOnHandshakeFailure(conn, onClose, err) @@ -127,7 +137,7 @@ func (h *Outbound) NewConnectionEx(ctx context.Context, conn net.Conn, metadata } metadata.Inbound = h.Tag() metadata.InboundType = h.Type() - h.router.RouteConnectionEx(ctx, conn, metadata, onClose) + h.router.RouteConnectionEx(ctx, wrappedConn, metadata, onClose) return } diff --git a/protocol/limiter/bandwidth/strategy.go b/protocol/limiter/bandwidth/strategy.go index 87821910..14604e90 100644 --- a/protocol/limiter/bandwidth/strategy.go +++ b/protocol/limiter/bandwidth/strategy.go @@ -2,6 +2,7 @@ package bandwidth import ( "context" + "io" "net" "strconv" "sync" @@ -191,7 +192,7 @@ func (s *UsersBandwidthStrategy) getStrategy(ctx context.Context, metadata *adap } type bwConnEntry struct { - conn net.Conn + conn io.Closer } type ManagerBandwidthStrategy struct { @@ -251,7 +252,25 @@ func (s *ManagerBandwidthStrategy) wrapPacketConn(ctx context.Context, conn net. if !ok { return nil, E.New("user strategy not found: ", user) } - return strategy.wrapPacketConn(ctx, conn, metadata, reverse) + wrapped, err := strategy.wrapPacketConn(ctx, conn, metadata, reverse) + if err != nil { + return nil, err + } + entry := &bwConnEntry{conn: conn} + s.mtx.Lock() + s.conns[user] = append(s.conns[user], entry) + s.mtx.Unlock() + return onclose.NewPacketConn(wrapped, func() { + s.mtx.Lock() + entries := s.conns[user] + for i, e := range entries { + if e == entry { + s.conns[user] = append(entries[:i], entries[i+1:]...) + break + } + } + s.mtx.Unlock() + }), nil } func (s *ManagerBandwidthStrategy) UpdateStrategies(strategies map[string]BandwidthStrategy) { diff --git a/protocol/limiter/traffic/outbound.go b/protocol/limiter/traffic/outbound.go index 0d7a329c..088cc334 100644 --- a/protocol/limiter/traffic/outbound.go +++ b/protocol/limiter/traffic/outbound.go @@ -82,7 +82,12 @@ func (h *Outbound) DialContext(ctx context.Context, network string, destination if err != nil { return nil, err } - return h.strategy.wrapConn(ctx, conn, adapter.ContextFrom(ctx), true) + wrappedConn, err := h.strategy.wrapConn(ctx, conn, adapter.ContextFrom(ctx), true) + if err != nil { + conn.Close() + return nil, err + } + return wrappedConn, nil } func (h *Outbound) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) { @@ -90,11 +95,16 @@ func (h *Outbound) ListenPacket(ctx context.Context, destination M.Socksaddr) (n if err != nil { return nil, err } - return h.strategy.wrapPacketConn(ctx, conn, adapter.ContextFrom(ctx), true) + wrappedConn, err := h.strategy.wrapPacketConn(ctx, conn, adapter.ContextFrom(ctx), true) + if err != nil { + conn.Close() + return nil, err + } + return wrappedConn, nil } func (h *Outbound) NewConnectionEx(ctx context.Context, conn net.Conn, metadata adapter.InboundContext, onClose N.CloseHandlerFunc) { - conn, err := h.strategy.wrapConn(ctx, conn, &metadata, false) + wrappedConn, err := h.strategy.wrapConn(ctx, conn, &metadata, false) if err != nil { if err.Error() != "traffic limit exceeded" { h.logger.ErrorContext(ctx, err) @@ -104,7 +114,7 @@ func (h *Outbound) NewConnectionEx(ctx context.Context, conn net.Conn, metadata } metadata.Inbound = h.Tag() metadata.InboundType = h.Type() - h.router.RouteConnectionEx(ctx, conn, metadata, onClose) + h.router.RouteConnectionEx(ctx, wrappedConn, metadata, onClose) } func (h *Outbound) NewPacketConnectionEx(ctx context.Context, conn N.PacketConn, metadata adapter.InboundContext, onClose N.CloseHandlerFunc) { diff --git a/protocol/limiter/traffic/strategy.go b/protocol/limiter/traffic/strategy.go index 71d4be1b..68a58b5d 100644 --- a/protocol/limiter/traffic/strategy.go +++ b/protocol/limiter/traffic/strategy.go @@ -2,6 +2,7 @@ package traffic import ( "context" + "io" "net" "sync" @@ -37,11 +38,11 @@ func NewDefaultWrapStrategy(limiterStrategy TrafficLimiterStrategy, connWrapper func (s *DefaultWrapStrategy) wrapConn(ctx context.Context, conn net.Conn, metadata *adapter.InboundContext, reverse bool) (net.Conn, error) { limiter, err := s.limiterStrategy.getLimiter(ctx, metadata) if err != nil { - return conn, err + return nil, err } _, err = limiter.Reserve(0) if err != nil { - return conn, err + return nil, err } return s.connWrapper(ctx, conn, limiter, reverse), nil } @@ -49,11 +50,11 @@ func (s *DefaultWrapStrategy) wrapConn(ctx context.Context, conn net.Conn, metad func (s *DefaultWrapStrategy) wrapPacketConn(ctx context.Context, conn net.PacketConn, metadata *adapter.InboundContext, reverse bool) (net.PacketConn, error) { limiter, err := s.limiterStrategy.getLimiter(ctx, metadata) if err != nil { - return conn, err + return nil, err } _, err = limiter.Reserve(0) if err != nil { - return conn, err + return nil, err } return s.packetConnWrapper(ctx, conn, limiter, reverse), nil } @@ -73,7 +74,7 @@ func (s *GlobalTrafficStrategy) getLimiter(ctx context.Context, metadata *adapte } type connEntry struct { - conn net.Conn + conn io.Closer } type ManagerTrafficStrategy struct { @@ -94,7 +95,7 @@ func (s *ManagerTrafficStrategy) wrapConn(ctx context.Context, conn net.Conn, me if err != nil { return nil, err } - wrapped, err := strategy.wrapConn(ctx, conn, metadata, reverse) + wrappedConn, err := strategy.wrapConn(ctx, conn, metadata, reverse) if err != nil { return nil, err } @@ -102,7 +103,7 @@ func (s *ManagerTrafficStrategy) wrapConn(ctx context.Context, conn net.Conn, me s.mtx.Lock() s.conns[user] = append(s.conns[user], entry) s.mtx.Unlock() - return onclose.NewConn(wrapped, func() { + return onclose.NewConn(wrappedConn, func() { s.mtx.Lock() entries := s.conns[user] for i, e := range entries { @@ -116,11 +117,29 @@ func (s *ManagerTrafficStrategy) wrapConn(ctx context.Context, conn net.Conn, me } func (s *ManagerTrafficStrategy) wrapPacketConn(ctx context.Context, conn net.PacketConn, metadata *adapter.InboundContext, reverse bool) (net.PacketConn, error) { - strategy, _, err := s.getStrategy(ctx, metadata) + strategy, user, err := s.getStrategy(ctx, metadata) if err != nil { return nil, err } - return strategy.wrapPacketConn(ctx, conn, metadata, reverse) + wrappedConn, err := strategy.wrapPacketConn(ctx, conn, metadata, reverse) + if err != nil { + return nil, err + } + entry := &connEntry{conn: conn} + s.mtx.Lock() + s.conns[user] = append(s.conns[user], entry) + s.mtx.Unlock() + return onclose.NewPacketConn(wrappedConn, func() { + s.mtx.Lock() + entries := s.conns[user] + for i, e := range entries { + if e == entry { + s.conns[user] = append(entries[:i], entries[i+1:]...) + break + } + } + s.mtx.Unlock() + }), nil } func (s *ManagerTrafficStrategy) getStrategy(ctx context.Context, metadata *adapter.InboundContext) (TrafficStrategy, string, error) {