fix: apply proxy protocol and max connection settings

This commit is contained in:
sagit
2026-04-27 10:54:25 +08:00
committed by GitHub
parent 58d2e89147
commit 2ca3849917
11 changed files with 478 additions and 60 deletions
@@ -16,6 +16,7 @@ import (
"github.com/go-gost/core/recorder"
ctxvalue "github.com/go-gost/x/ctx"
xnet "github.com/go-gost/x/internal/net"
"github.com/go-gost/x/internal/net/proxyproto"
"github.com/go-gost/x/internal/util/forwarder"
"github.com/go-gost/x/internal/util/sniffing"
tls_util "github.com/go-gost/x/internal/util/tls"
@@ -252,6 +253,8 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand
}
defer cc.Close()
cc = proxyproto.WrapClientConn(h.md.proxyProtocol, conn.RemoteAddr(), conn.LocalAddr(), cc)
if err := xnet.Transport(conn, cc); err != nil {
if marker := target.Marker(); marker != nil {
marker.Mark()
@@ -14,6 +14,7 @@ import (
type metadata struct {
readTimeout time.Duration
proxyProtocol int
httpKeepalive bool
sniffing bool
@@ -38,6 +39,7 @@ func (h *forwardHandler) parseMetadata(md mdata.Metadata) (err error) {
if h.md.readTimeout <= 0 {
h.md.readTimeout = 15 * time.Second
}
h.md.proxyProtocol = mdutil.GetInt(md, "proxyProtocol")
h.md.httpKeepalive = mdutil.GetBool(md, "http.keepalive")
@@ -0,0 +1,111 @@
package local
import (
"bufio"
"context"
"net"
"testing"
"time"
"github.com/go-gost/core/chain"
"github.com/go-gost/core/handler"
"github.com/go-gost/core/hop"
xlogger "github.com/go-gost/x/logger"
xmd "github.com/go-gost/x/metadata"
proxyproto "github.com/pires/go-proxyproto"
)
type proxyProtocolTestHop struct {
node *chain.Node
}
func (h proxyProtocolTestHop) Select(context.Context, ...hop.SelectOption) *chain.Node {
return h.node
}
func (h proxyProtocolTestHop) Nodes() []*chain.Node {
return []*chain.Node{h.node}
}
type proxyProtocolTestRouter struct{}
func (r proxyProtocolTestRouter) Options() *chain.RouterOptions {
return &chain.RouterOptions{}
}
func (r proxyProtocolTestRouter) Dial(ctx context.Context, network, address string) (net.Conn, error) {
var d net.Dialer
return d.DialContext(ctx, network, address)
}
func (r proxyProtocolTestRouter) Bind(context.Context, string, string, ...chain.BindOption) (net.Listener, error) {
return nil, net.ErrClosed
}
func TestLocalForwardHandlerSendsProxyProtocolToTarget(t *testing.T) {
targetListener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen target: %v", err)
}
defer targetListener.Close()
entryListener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen entry: %v", err)
}
defer entryListener.Close()
h := NewHandler(
handler.RouterOption(proxyProtocolTestRouter{}),
handler.LoggerOption(xlogger.Nop()),
)
forwarder := h.(handler.Forwarder)
forwarder.Forward(proxyProtocolTestHop{node: chain.NewNode("target", targetListener.Addr().String())})
if err := h.Init(xmd.NewMetadata(map[string]any{"proxyProtocol": 2})); err != nil {
t.Fatalf("init handler: %v", err)
}
handleErr := make(chan error, 1)
acceptErr := make(chan error, 1)
go func() {
serverConn, err := entryListener.Accept()
if err != nil {
acceptErr <- err
return
}
handleErr <- h.Handle(context.Background(), serverConn)
}()
clientConn, err := net.Dial("tcp", entryListener.Addr().String())
if err != nil {
t.Fatalf("dial entry: %v", err)
}
defer clientConn.Close()
targetConn, err := targetListener.Accept()
if err != nil {
t.Fatalf("accept target: %v", err)
}
defer targetConn.Close()
if err := targetConn.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil {
t.Fatalf("set target deadline: %v", err)
}
header, err := proxyproto.Read(bufio.NewReader(targetConn))
if err != nil {
t.Fatalf("read proxy protocol header: %v", err)
}
if header.Version != 2 {
t.Fatalf("expected proxy protocol v2, got v%d", header.Version)
}
_ = clientConn.Close()
_ = targetConn.Close()
select {
case err := <-acceptErr:
t.Fatalf("accept entry: %v", err)
case <-handleErr:
case <-time.After(2 * time.Second):
t.Fatal("handler did not return after closing connections")
}
}