fix: apply per-client limits to udp listener

This commit is contained in:
sagitchu
2026-04-27 23:05:31 +08:00
parent 46394388b1
commit 7a9ba8bd81
4 changed files with 194 additions and 2 deletions
+30
View File
@@ -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")
}
}
+28
View File
@@ -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())
}
}
+41 -2
View File
@@ -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 {
+95
View File
@@ -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
}
}