From edfe2a2372d70fc14b0fc06a601bcb19522f2725 Mon Sep 17 00:00:00 2001 From: sagitchu Date: Mon, 27 Apr 2026 23:13:16 +0800 Subject: [PATCH] fix: preserve udp packet semantics with limiters --- go-gost/x/listener/udp/listener.go | 105 +++++++++++++++++++++--- go-gost/x/listener/udp/listener_test.go | 65 +++++++++++++++ 2 files changed, 158 insertions(+), 12 deletions(-) diff --git a/go-gost/x/listener/udp/listener.go b/go-gost/x/listener/udp/listener.go index 23aff31..3ac3119 100644 --- a/go-gost/x/listener/udp/listener.go +++ b/go-gost/x/listener/udp/listener.go @@ -2,15 +2,17 @@ package udp import ( "net" + "sync" + "time" "github.com/go-gost/core/limiter" + conn_limiter "github.com/go-gost/core/limiter/conn" "github.com/go-gost/core/listener" "github.com/go-gost/core/logger" md "github.com/go-gost/core/metadata" 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" @@ -94,26 +96,77 @@ func (l *udpListener) Accept() (conn net.Conn, err error) { if lim := l.options.ConnLimiter.Limiter(host); lim != nil { if !lim.Allow(1) { _ = conn.Close() - return closedConn{Conn: conn}, nil + return newClosedConn(conn), nil } - conn = climiter.WrapConn(lim, conn) + conn = wrapConnLimiter(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()), - ) + if pc, ok := conn.(net.PacketConn); ok { + conn = limiter_wrapper.WrapUDPConn( + pc, + 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 connLimiterConn struct { + net.Conn + net.PacketConn + limiter conn_limiter.Limiter + once sync.Once +} + +func wrapConnLimiter(limiter conn_limiter.Limiter, conn net.Conn) net.Conn { + pc, ok := conn.(net.PacketConn) + if !ok { + return conn + } + return &connLimiterConn{ + Conn: conn, + PacketConn: pc, + limiter: limiter, + } +} + +func (c *connLimiterConn) Close() (err error) { + c.once.Do(func() { + c.limiter.Allow(-1) + err = c.Conn.Close() + }) + return +} + +func (c *connLimiterConn) LocalAddr() net.Addr { + return c.Conn.LocalAddr() +} + +func (c *connLimiterConn) SetDeadline(t time.Time) error { + return c.Conn.SetDeadline(t) +} + +func (c *connLimiterConn) SetReadDeadline(t time.Time) error { + return c.Conn.SetReadDeadline(t) +} + +func (c *connLimiterConn) SetWriteDeadline(t time.Time) error { + return c.Conn.SetWriteDeadline(t) +} + type closedConn struct { net.Conn + net.PacketConn +} + +func newClosedConn(conn net.Conn) net.Conn { + pc, _ := conn.(net.PacketConn) + return closedConn{Conn: conn, PacketConn: pc} } func (c closedConn) Read([]byte) (int, error) { @@ -124,6 +177,34 @@ func (c closedConn) Write([]byte) (int, error) { return 0, net.ErrClosed } +func (c closedConn) ReadFrom([]byte) (int, net.Addr, error) { + return 0, nil, net.ErrClosed +} + +func (c closedConn) WriteTo([]byte, net.Addr) (int, error) { + return 0, net.ErrClosed +} + +func (c closedConn) Close() error { + return c.Conn.Close() +} + +func (c closedConn) LocalAddr() net.Addr { + return c.Conn.LocalAddr() +} + +func (c closedConn) SetDeadline(t time.Time) error { + return c.Conn.SetDeadline(t) +} + +func (c closedConn) SetReadDeadline(t time.Time) error { + return c.Conn.SetReadDeadline(t) +} + +func (c closedConn) SetWriteDeadline(t time.Time) error { + return c.Conn.SetWriteDeadline(t) +} + func (l *udpListener) Addr() net.Addr { return l.ln.Addr() } diff --git a/go-gost/x/listener/udp/listener_test.go b/go-gost/x/listener/udp/listener_test.go index 9b56847..504042a 100644 --- a/go-gost/x/listener/udp/listener_test.go +++ b/go-gost/x/listener/udp/listener_test.go @@ -9,9 +9,61 @@ import ( corelistener "github.com/go-gost/core/listener" 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" ) +func TestAcceptWithLimitersPreservesPacketConn(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.TrafficLimiterOption(xtraffic.NewTrafficLimiter( + xtraffic.LimitsOption("$$ 1024B 1024B"), + xtraffic.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() + + client, err := net.Dial("udp", ln.Addr().String()) + if err != nil { + t.Fatalf("dial udp listener: %v", err) + } + defer client.Close() + if _, err := client.Write([]byte("packet")); err != nil { + t.Fatalf("write packet: %v", err) + } + + conn, err := acceptWithTimeout(t, ln, time.Second) + if err != nil { + t.Fatalf("accept conn: %v", err) + } + defer conn.Close() + + packetConn, ok := conn.(net.PacketConn) + if !ok { + t.Fatalf("expected accepted UDP conn with limiters to implement net.PacketConn, got %T", conn) + } + + buf := make([]byte, 16) + n, addr, err := packetConn.ReadFrom(buf) + if err != nil { + t.Fatalf("read packet: %v", err) + } + if string(buf[:n]) != "packet" { + t.Fatalf("expected original datagram, got %q", string(buf[:n])) + } + if addr == nil || addr.String() != client.LocalAddr().String() { + t.Fatalf("expected client addr %v, got %v", client.LocalAddr(), addr) + } +} + func TestAcceptAppliesConnLimiterAndReleasesOnClose(t *testing.T) { ln := NewListener( corelistener.AddrOption("127.0.0.1:0"), @@ -57,6 +109,19 @@ func TestAcceptAppliesConnLimiterAndReleasesOnClose(t *testing.T) { _ = blocked.Close() t.Fatalf("expected blocked same-IP pseudo-connection to be closed") } + packetConn, ok := blocked.(net.PacketConn) + if !ok { + _ = blocked.Close() + t.Fatalf("expected blocked same-IP pseudo-connection to preserve net.PacketConn, got %T", blocked) + } + if _, _, err := packetConn.ReadFrom(buf); err == nil { + _ = blocked.Close() + t.Fatalf("expected blocked same-IP packet connection to be closed") + } + if _, err := packetConn.WriteTo([]byte("blocked"), client.LocalAddr()); err == nil { + _ = blocked.Close() + t.Fatalf("expected blocked same-IP packet write to be closed") + } _ = blocked.Close() _ = first.Close()