From 7a9ba8bd817560b840c0aeb8d2c992abc6c814ee Mon Sep 17 00:00:00 2001 From: sagitchu Date: Mon, 27 Apr 2026 23:05:31 +0800 Subject: [PATCH] fix: apply per-client limits to udp listener --- go-gost/x/limiter/conn/conn_test.go | 30 +++++++ go-gost/x/limiter/traffic/traffic_test.go | 28 +++++++ go-gost/x/listener/udp/listener.go | 43 +++++++++- go-gost/x/listener/udp/listener_test.go | 95 +++++++++++++++++++++++ 4 files changed, 194 insertions(+), 2 deletions(-) create mode 100644 go-gost/x/limiter/conn/conn_test.go create mode 100644 go-gost/x/limiter/traffic/traffic_test.go create mode 100644 go-gost/x/listener/udp/listener_test.go diff --git a/go-gost/x/limiter/conn/conn_test.go b/go-gost/x/limiter/conn/conn_test.go new file mode 100644 index 0000000..063a4e9 --- /dev/null +++ b/go-gost/x/limiter/conn/conn_test.go @@ -0,0 +1,30 @@ +package conn + +import ( + "io" + "testing" + + corelogger "github.com/go-gost/core/logger" + xlogger "github.com/go-gost/x/logger" +) + +func TestIPLimitKeyCreatesIndependentLimiters(t *testing.T) { + limiter := NewConnLimiter( + LimitsOption("$$ 1"), + LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))), + ) + first := limiter.Limiter("192.0.2.1") + second := limiter.Limiter("192.0.2.2") + if first == nil || second == nil { + t.Fatalf("expected non-nil per-IP limiters") + } + if !first.Allow(1) { + t.Fatalf("expected first IP first connection to be allowed") + } + if first.Allow(1) { + t.Fatalf("expected first IP second connection to be rejected") + } + if !second.Allow(1) { + t.Fatalf("expected second IP first connection to be allowed independently") + } +} diff --git a/go-gost/x/limiter/traffic/traffic_test.go b/go-gost/x/limiter/traffic/traffic_test.go new file mode 100644 index 0000000..702cc76 --- /dev/null +++ b/go-gost/x/limiter/traffic/traffic_test.go @@ -0,0 +1,28 @@ +package traffic + +import ( + "context" + "io" + "testing" + + corelogger "github.com/go-gost/core/logger" + xlogger "github.com/go-gost/x/logger" +) + +func TestCIDRLimitCreatesIndependentClientLimiters(t *testing.T) { + limiter := NewTrafficLimiter( + LimitsOption("0.0.0.0/0 2B 2B"), + LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))), + ) + first := limiter.In(context.Background(), "192.0.2.1:1000") + second := limiter.In(context.Background(), "192.0.2.2:1000") + if first == nil || second == nil { + t.Fatalf("expected non-nil CIDR client limiters") + } + if first == second { + t.Fatalf("expected different clients to receive independent limiter instances") + } + if first.Limit() != 2 || second.Limit() != 2 { + t.Fatalf("expected both limits to be 2, got %d and %d", first.Limit(), second.Limit()) + } +} diff --git a/go-gost/x/listener/udp/listener.go b/go-gost/x/listener/udp/listener.go index 6f21f7c..23aff31 100644 --- a/go-gost/x/listener/udp/listener.go +++ b/go-gost/x/listener/udp/listener.go @@ -10,6 +10,7 @@ import ( admission "github.com/go-gost/x/admission/wrapper" xnet "github.com/go-gost/x/internal/net" "github.com/go-gost/x/internal/net/udp" + climiter "github.com/go-gost/x/limiter/conn/wrapper" traffic_limiter "github.com/go-gost/x/limiter/traffic" limiter_wrapper "github.com/go-gost/x/limiter/traffic/wrapper" metrics "github.com/go-gost/x/metrics/wrapper" @@ -70,7 +71,7 @@ func (l *udpListener) Init(md md.Metadata) (err error) { limiter.NetworkOption(conn.LocalAddr().Network()), ) - l.ln = udp.NewListener(conn, &udp.ListenConfig{ + ln := udp.NewListener(conn, &udp.ListenConfig{ Backlog: l.md.backlog, ReadQueueSize: l.md.readQueueSize, ReadBufferSize: l.md.readBufferSize, @@ -78,11 +79,49 @@ func (l *udpListener) Init(md md.Metadata) (err error) { TTL: l.md.ttl, Logger: l.logger, }) + l.ln = ln return } func (l *udpListener) Accept() (conn net.Conn, err error) { - return l.ln.Accept() + conn, err = l.ln.Accept() + if err != nil { + return + } + + if l.options.ConnLimiter != nil { + host, _, _ := net.SplitHostPort(conn.RemoteAddr().String()) + if lim := l.options.ConnLimiter.Limiter(host); lim != nil { + if !lim.Allow(1) { + _ = conn.Close() + return closedConn{Conn: conn}, nil + } + conn = climiter.WrapConn(lim, conn) + } + } + + conn = limiter_wrapper.WrapConn( + conn, + l.options.TrafficLimiter, + conn.RemoteAddr().String(), + limiter.ScopeOption(limiter.ScopeConn), + limiter.ServiceOption(l.options.Service), + limiter.NetworkOption(conn.LocalAddr().Network()), + limiter.SrcOption(conn.RemoteAddr().String()), + ) + return +} + +type closedConn struct { + net.Conn +} + +func (c closedConn) Read([]byte) (int, error) { + return 0, net.ErrClosed +} + +func (c closedConn) Write([]byte) (int, error) { + return 0, net.ErrClosed } func (l *udpListener) Addr() net.Addr { diff --git a/go-gost/x/listener/udp/listener_test.go b/go-gost/x/listener/udp/listener_test.go new file mode 100644 index 0000000..9b56847 --- /dev/null +++ b/go-gost/x/listener/udp/listener_test.go @@ -0,0 +1,95 @@ +package udp + +import ( + "io" + "net" + "testing" + "time" + + corelistener "github.com/go-gost/core/listener" + corelogger "github.com/go-gost/core/logger" + xconn "github.com/go-gost/x/limiter/conn" + xlogger "github.com/go-gost/x/logger" +) + +func TestAcceptAppliesConnLimiterAndReleasesOnClose(t *testing.T) { + ln := NewListener( + corelistener.AddrOption("127.0.0.1:0"), + corelistener.ConnLimiterOption(xconn.NewConnLimiter( + xconn.LimitsOption("$$ 1"), + xconn.LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))), + )), + corelistener.LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))), + ) + if err := ln.Init(nil); err != nil { + t.Fatalf("init listener: %v", err) + } + defer ln.Close() + + addr := ln.Addr().String() + client, err := net.Dial("udp", addr) + if err != nil { + t.Fatalf("dial udp listener: %v", err) + } + defer client.Close() + if _, err := client.Write([]byte("first")); err != nil { + t.Fatalf("write first packet: %v", err) + } + first, err := acceptWithTimeout(t, ln, time.Second) + if err != nil { + t.Fatalf("accept first conn: %v", err) + } + + blockedClient, err := net.Dial("udp", addr) + if err != nil { + t.Fatalf("dial blocked udp client: %v", err) + } + defer blockedClient.Close() + if _, err := blockedClient.Write([]byte("blocked")); err != nil { + t.Fatalf("write blocked packet: %v", err) + } + blocked, err := acceptWithTimeout(t, ln, time.Second) + if err != nil { + t.Fatalf("expected blocked same-IP pseudo-connection to be returned closed: %v", err) + } + buf := make([]byte, 16) + if _, err := blocked.Read(buf); err == nil { + _ = blocked.Close() + t.Fatalf("expected blocked same-IP pseudo-connection to be closed") + } + _ = blocked.Close() + _ = first.Close() + + reopenedClient, err := net.Dial("udp", addr) + if err != nil { + t.Fatalf("dial reopened udp client: %v", err) + } + defer reopenedClient.Close() + if _, err := reopenedClient.Write([]byte("after-close")); err != nil { + t.Fatalf("write after close packet: %v", err) + } + reopened, err := acceptWithTimeout(t, ln, time.Second) + if err != nil { + t.Fatalf("expected same client to be accepted after close: %v", err) + } + _ = reopened.Close() +} + +func acceptWithTimeout(t *testing.T, ln corelistener.Listener, timeout time.Duration) (net.Conn, error) { + t.Helper() + type result struct { + conn net.Conn + err error + } + ch := make(chan result, 1) + go func() { + conn, err := ln.Accept() + ch <- result{conn: conn, err: err} + }() + select { + case res := <-ch: + return res.conn, res.err + case <-time.After(timeout): + return nil, net.ErrClosed + } +}