fix: preserve udp packet semantics with limiters

This commit is contained in:
sagitchu
2026-04-27 23:13:16 +08:00
parent 7a9ba8bd81
commit edfe2a2372
2 changed files with 158 additions and 12 deletions
+93 -12
View File
@@ -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()
}
+65
View File
@@ -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()