mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
fix: apply per-client limits to udp listener
This commit is contained in:
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user