fix(quic,v2ray): close the sockets quic-go was never going to close
DialEarly with a packet conn the caller made sets a flag that means quic-go does not own it: closing the transport only stops reading from the socket. Neither DNS transport closed it. On the QUIC one it was closed on a failed handshake and never on success, so every redial — idle timeout, retry error, engine reload — left a UDP socket for the life of the process. On the HTTP/3 one the library drives its own reconnects, so the leak compounds without anything in our code looking wrong. That is the same shape as v2rayquic's, where offerNew overwrote the raw conn on every reconnect without closing the previous one. Both are now owned by a watcher tied to the connection's own context, so the socket lives exactly as long as the connection does. This matters more than it did last week: the shipped resolvers are DoH, and DNS is intercepted by default now, so the whole network's query stream rides this path on a router with 512 MB. The same upstream commit fixes both halves. We had taken the v2ray half and not the DNS one — the third time this session a paired fix arrived half-applied, and the first of those cost a day of debugging. These two files are now byte-identical to upstream so a rebase cannot reopen it. Also from that family: websocket and httpupgrade leaked their conn on failed handshakes, and a QUIC stream's Close did not release a blocked write. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -126,6 +126,12 @@ func (t *HTTP3Transport) newTransport() *http3.Transport {
|
||||
conn.Close()
|
||||
return nil, dialErr
|
||||
}
|
||||
// quic-go does not take ownership of the packet conn passed to
|
||||
// DialEarly: when the connection ends it only stops reading.
|
||||
go func() {
|
||||
<-quicConn.Context().Done()
|
||||
conn.Close()
|
||||
}()
|
||||
return quicConn, nil
|
||||
},
|
||||
TLSClientConfig: t.tlsConfig,
|
||||
|
||||
@@ -0,0 +1,351 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/quic-go"
|
||||
"github.com/sagernet/quic-go/http3"
|
||||
sbTLS "github.com/sagernet/sing-box/common/tls"
|
||||
C "github.com/sagernet/sing-box/constant"
|
||||
"github.com/sagernet/sing-box/dns"
|
||||
"github.com/sagernet/sing-box/dns/transport"
|
||||
"github.com/sagernet/sing-box/option"
|
||||
"github.com/sagernet/sing/common"
|
||||
"github.com/sagernet/sing/common/logger"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
|
||||
mDNS "github.com/miekg/dns"
|
||||
)
|
||||
|
||||
var _ N.Dialer = (*trackingDialer)(nil)
|
||||
|
||||
// These tests pin down who owns the UDP socket handed to quic-go.
|
||||
//
|
||||
// quic-go's Dial/DialEarly take a net.PacketConn but do NOT take ownership of
|
||||
// it: quic.setupTransport() builds a Transport with createdConn=false, and
|
||||
// Transport.Close() then only calls conn.SetReadDeadline(time.Now()) instead of
|
||||
// conn.Close(). So every QUIC connection torn down here — idle timeout, a
|
||||
// retryable error, an engine reload calling Reset() — used to strand the UDP
|
||||
// socket that carried it for the rest of the process's life. On a router that
|
||||
// resolves through DoQ/DoH3 for months that is an unbounded fd leak.
|
||||
//
|
||||
// Both tests reconnect once and assert the socket from the FIRST connection is
|
||||
// actually closed. Without the `<-conn.Context().Done() -> rawConn.Close()`
|
||||
// watchdogs in quic.go / http3.go they fail on that assertion.
|
||||
|
||||
type trackedConn struct {
|
||||
net.Conn
|
||||
closeOnce sync.Once
|
||||
closed chan struct{}
|
||||
}
|
||||
|
||||
func (c *trackedConn) Close() error {
|
||||
c.closeOnce.Do(func() { close(c.closed) })
|
||||
return c.Conn.Close()
|
||||
}
|
||||
|
||||
// trackingDialer hands out real UDP sockets and remembers every one of them.
|
||||
type trackingDialer struct {
|
||||
access sync.Mutex
|
||||
conns []*trackedConn
|
||||
}
|
||||
|
||||
func (d *trackingDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
||||
conn, err := (&net.Dialer{}).DialContext(ctx, network, destination.String())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tracked := &trackedConn{Conn: conn, closed: make(chan struct{})}
|
||||
d.access.Lock()
|
||||
d.conns = append(d.conns, tracked)
|
||||
d.access.Unlock()
|
||||
return tracked, nil
|
||||
}
|
||||
|
||||
func (d *trackingDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
||||
return net.ListenUDP("udp", nil)
|
||||
}
|
||||
|
||||
func (d *trackingDialer) count() int {
|
||||
d.access.Lock()
|
||||
defer d.access.Unlock()
|
||||
return len(d.conns)
|
||||
}
|
||||
|
||||
func (d *trackingDialer) at(index int) *trackedConn {
|
||||
d.access.Lock()
|
||||
defer d.access.Unlock()
|
||||
return d.conns[index]
|
||||
}
|
||||
|
||||
func (d *trackingDialer) closeAll() {
|
||||
d.access.Lock()
|
||||
defer d.access.Unlock()
|
||||
for _, conn := range d.conns {
|
||||
conn.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func requireClosed(t *testing.T, conn *trackedConn, what string) {
|
||||
t.Helper()
|
||||
select {
|
||||
case <-conn.closed:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatalf("%s: the UDP socket of the retired QUIC connection was never closed — quic-go does not own it, we must", what)
|
||||
}
|
||||
}
|
||||
|
||||
func requireDialed(t *testing.T, dialer *trackingDialer, want int) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if dialer.count() >= want {
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("expected at least %d dial(s), got %d", want, dialer.count())
|
||||
}
|
||||
|
||||
func testServerTLSConfig(t *testing.T, nextProtos []string) *tls.Config {
|
||||
t.Helper()
|
||||
certificate, err := sbTLS.GenerateKeyPair(nil, nil, nil, "localhost")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return &tls.Config{
|
||||
Certificates: []tls.Certificate{*certificate},
|
||||
NextProtos: nextProtos,
|
||||
MinVersion: tls.VersionTLS13,
|
||||
}
|
||||
}
|
||||
|
||||
func testClientTLSConfig(t *testing.T, nextProtos []string) sbTLS.Config {
|
||||
t.Helper()
|
||||
config, err := sbTLS.NewClient(context.Background(), logger.NOP(), "localhost", option.OutboundTLSOptions{
|
||||
Enabled: true,
|
||||
Insecure: true,
|
||||
ServerName: "localhost",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
config.SetNextProtos(nextProtos)
|
||||
return config
|
||||
}
|
||||
|
||||
// startDoQServer serves a minimal DoQ responder and returns its address.
|
||||
func startDoQServer(t *testing.T) M.Socksaddr {
|
||||
t.Helper()
|
||||
listener, err := quic.ListenAddr("127.0.0.1:0", testServerTLSConfig(t, []string{"doq"}), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(func() {
|
||||
cancel()
|
||||
listener.Close()
|
||||
})
|
||||
go func() {
|
||||
for {
|
||||
conn, acceptErr := listener.Accept(ctx)
|
||||
if acceptErr != nil {
|
||||
return
|
||||
}
|
||||
go func(conn *quic.Conn) {
|
||||
for {
|
||||
stream, streamErr := conn.AcceptStream(ctx)
|
||||
if streamErr != nil {
|
||||
return
|
||||
}
|
||||
go func(stream *quic.Stream) {
|
||||
defer stream.Close()
|
||||
request, readErr := transport.ReadMessage(stream)
|
||||
if readErr != nil {
|
||||
return
|
||||
}
|
||||
response := new(mDNS.Msg)
|
||||
response.SetReply(request)
|
||||
transport.WriteMessage(stream, 0, response)
|
||||
}(stream)
|
||||
}
|
||||
}(conn)
|
||||
}
|
||||
}()
|
||||
return M.ParseSocksaddr(listener.Addr().String())
|
||||
}
|
||||
|
||||
func testQuery() *mDNS.Msg {
|
||||
message := new(mDNS.Msg)
|
||||
message.SetQuestion("example.com.", mDNS.TypeA)
|
||||
return message
|
||||
}
|
||||
|
||||
func TestQUICTransportClosesPacketConnOnReconnect(t *testing.T) {
|
||||
t.Parallel()
|
||||
serverAddr := startDoQServer(t)
|
||||
dialer := &trackingDialer{}
|
||||
t.Cleanup(dialer.closeAll)
|
||||
|
||||
dnsTransport := &Transport{
|
||||
TransportAdapter: dns.NewTransportAdapter(C.DNSTypeQUIC, "test-doq", nil),
|
||||
dialer: dialer,
|
||||
serverAddr: serverAddr,
|
||||
tlsConfig: testClientTLSConfig(t, []string{"doq"}),
|
||||
connection: transport.NewConnPool(transport.ConnPoolOptions[*quic.Conn]{
|
||||
Mode: transport.ConnPoolSingle,
|
||||
IsAlive: func(conn *quic.Conn) bool {
|
||||
return conn != nil && !common.Done(conn.Context())
|
||||
},
|
||||
Close: func(conn *quic.Conn, _ error) {
|
||||
conn.CloseWithError(0, "")
|
||||
},
|
||||
}),
|
||||
}
|
||||
t.Cleanup(func() { dnsTransport.Close() })
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if _, err := dnsTransport.Exchange(ctx, testQuery()); err != nil {
|
||||
t.Fatal("first exchange: ", err)
|
||||
}
|
||||
requireDialed(t, dialer, 1)
|
||||
first := dialer.at(0)
|
||||
|
||||
// Retire the connection the way a retryable error or an engine reload does.
|
||||
dnsTransport.Reset()
|
||||
requireClosed(t, first, "Reset()")
|
||||
|
||||
// The reconnect must still work, on a fresh socket.
|
||||
if _, err := dnsTransport.Exchange(ctx, testQuery()); err != nil {
|
||||
t.Fatal("second exchange: ", err)
|
||||
}
|
||||
requireDialed(t, dialer, 2)
|
||||
second := dialer.at(1)
|
||||
if second == first {
|
||||
t.Fatal("expected a new UDP socket for the reconnect")
|
||||
}
|
||||
|
||||
if err := dnsTransport.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
requireClosed(t, second, "Close()")
|
||||
}
|
||||
|
||||
func TestHTTP3TransportClosesPacketConnOnReconnect(t *testing.T) {
|
||||
t.Parallel()
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/dns-query", func(writer http.ResponseWriter, request *http.Request) {
|
||||
message, err := readRequestMessage(request)
|
||||
if err != nil {
|
||||
writer.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
response := new(mDNS.Msg)
|
||||
response.SetReply(message)
|
||||
rawResponse, err := response.Pack()
|
||||
if err != nil {
|
||||
writer.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
writer.Header().Set("Content-Type", transport.MimeType)
|
||||
writer.Write(rawResponse)
|
||||
})
|
||||
listener, err := quic.ListenAddrEarly("127.0.0.1:0", testServerTLSConfig(t, []string{http3.NextProtoH3}), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server := &http3.Server{Handler: mux}
|
||||
go server.ServeListener(listener)
|
||||
t.Cleanup(func() {
|
||||
server.Close()
|
||||
listener.Close()
|
||||
})
|
||||
serverAddr := M.ParseSocksaddr(listener.Addr().String())
|
||||
|
||||
dialer := &trackingDialer{}
|
||||
t.Cleanup(dialer.closeAll)
|
||||
|
||||
stdConfig := &tls.Config{
|
||||
InsecureSkipVerify: true,
|
||||
ServerName: "localhost",
|
||||
NextProtos: []string{http3.NextProtoH3},
|
||||
MinVersion: tls.VersionTLS13,
|
||||
}
|
||||
dnsTransport := &HTTP3Transport{
|
||||
TransportAdapter: dns.NewTransportAdapter(C.DNSTypeHTTP3, "test-doh3", nil),
|
||||
logger: logger.NOP(),
|
||||
dialer: dialer,
|
||||
destination: &url.URL{Scheme: "https", Host: "localhost", Path: "/dns-query"},
|
||||
headers: http.Header{},
|
||||
serverAddr: serverAddr,
|
||||
tlsConfig: stdConfig,
|
||||
}
|
||||
dnsTransport.transport = dnsTransport.newTransport()
|
||||
t.Cleanup(func() { dnsTransport.Close() })
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if _, err = dnsTransport.Exchange(ctx, testQuery()); err != nil {
|
||||
t.Fatal("first exchange: ", err)
|
||||
}
|
||||
requireDialed(t, dialer, 1)
|
||||
first := dialer.at(0)
|
||||
|
||||
dnsTransport.Reset()
|
||||
requireClosed(t, first, "Reset()")
|
||||
|
||||
if _, err = dnsTransport.Exchange(ctx, testQuery()); err != nil {
|
||||
t.Fatal("second exchange: ", err)
|
||||
}
|
||||
requireDialed(t, dialer, 2)
|
||||
second := dialer.at(1)
|
||||
if second == first {
|
||||
t.Fatal("expected a new UDP socket for the reconnect")
|
||||
}
|
||||
|
||||
if err = dnsTransport.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
requireClosed(t, second, "Close()")
|
||||
}
|
||||
|
||||
func readRequestMessage(request *http.Request) (*mDNS.Msg, error) {
|
||||
defer request.Body.Close()
|
||||
rawMessage := make([]byte, 4096)
|
||||
n, err := readFull(request.Body, rawMessage)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var message mDNS.Msg
|
||||
err = message.Unpack(rawMessage[:n])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &message, nil
|
||||
}
|
||||
|
||||
func readFull(reader interface{ Read([]byte) (int, error) }, buffer []byte) (int, error) {
|
||||
var total int
|
||||
for total < len(buffer) {
|
||||
n, err := reader.Read(buffer[total:])
|
||||
total += n
|
||||
if err != nil {
|
||||
if total > 0 {
|
||||
return total, nil
|
||||
}
|
||||
return total, err
|
||||
}
|
||||
}
|
||||
return total, nil
|
||||
}
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/quic-go"
|
||||
"github.com/sagernet/sing-box/adapter"
|
||||
@@ -117,6 +118,12 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
|
||||
rawConn.Close()
|
||||
return nil, E.Cause(err, "establish QUIC connection")
|
||||
}
|
||||
// quic-go does not take ownership of the packet conn passed to
|
||||
// DialEarly: when the connection ends it only stops reading.
|
||||
go func() {
|
||||
<-earlyConnection.Context().Done()
|
||||
rawConn.Close()
|
||||
}()
|
||||
return earlyConnection, nil
|
||||
})
|
||||
if err != nil {
|
||||
@@ -144,6 +151,11 @@ func (t *Transport) exchange(ctx context.Context, message *mDNS.Msg, conn *quic.
|
||||
return nil, E.Cause(err, "open stream")
|
||||
}
|
||||
defer stream.CancelRead(0)
|
||||
stopWatch := context.AfterFunc(ctx, func() {
|
||||
stream.CancelRead(0)
|
||||
_ = stream.SetWriteDeadline(time.Now())
|
||||
})
|
||||
defer stopWatch()
|
||||
err = transport.WriteMessage(stream, 0, message)
|
||||
if err != nil {
|
||||
stream.Close()
|
||||
|
||||
@@ -0,0 +1,194 @@
|
||||
package route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/sagernet/sing-box/adapter"
|
||||
C "github.com/sagernet/sing-box/constant"
|
||||
"github.com/sagernet/sing-box/log"
|
||||
"github.com/sagernet/sing-box/option"
|
||||
R "github.com/sagernet/sing-box/route/rule"
|
||||
"github.com/sagernet/sing/common/json/badoption"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// The pre-match path (adapter.JudgeFlow -> Router.PreMatch, used by the TUN/
|
||||
// WireGuard-endpoint flow dispatcher) used to prepare only fakeip and the IP
|
||||
// version. Everything a rule matches on that the inbound cannot know — the
|
||||
// connection owner and the neighbor behind the source address — was resolved in
|
||||
// matchRule only, so a `source_mac_address` / `source_hostname` rule silently
|
||||
// failed to match in pre-match and the flow fell through to the default
|
||||
// outbound. Upstream b911fb078 shares one prepareMatchMetadata between both
|
||||
// paths; these tests pin that.
|
||||
|
||||
// stubNeighborResolver answers for exactly one address.
|
||||
type stubNeighborResolver struct {
|
||||
address netip.Addr
|
||||
mac net.HardwareAddr
|
||||
hostname string
|
||||
}
|
||||
|
||||
func (r *stubNeighborResolver) LookupMAC(address netip.Addr) (net.HardwareAddr, bool) {
|
||||
if address != r.address || r.mac == nil {
|
||||
return nil, false
|
||||
}
|
||||
return r.mac, true
|
||||
}
|
||||
|
||||
func (r *stubNeighborResolver) LookupHostname(address netip.Addr) (string, bool) {
|
||||
if address != r.address || r.hostname == "" {
|
||||
return "", false
|
||||
}
|
||||
return r.hostname, true
|
||||
}
|
||||
|
||||
func (r *stubNeighborResolver) LookupAddresses(hostname string) []netip.Addr {
|
||||
if hostname != r.hostname {
|
||||
return nil
|
||||
}
|
||||
return []netip.Addr{r.address}
|
||||
}
|
||||
|
||||
func (r *stubNeighborResolver) Start() error { return nil }
|
||||
func (r *stubNeighborResolver) Close() error { return nil }
|
||||
|
||||
// stubDNSRouter / stubDNSTransportManager implement only what
|
||||
// prepareMatchMetadata reaches; every other method is left to the embedded nil
|
||||
// interface and would panic if it were ever called.
|
||||
type stubDNSRouter struct {
|
||||
adapter.DNSRouter
|
||||
}
|
||||
|
||||
func (s *stubDNSRouter) LookupReverseMapping(netip.Addr) (string, bool) { return "", false }
|
||||
|
||||
type stubDNSTransportManager struct {
|
||||
adapter.DNSTransportManager
|
||||
}
|
||||
|
||||
func (s *stubDNSTransportManager) FakeIP() adapter.FakeIPTransport { return nil }
|
||||
|
||||
// stubOutboundManager's default outbound supports no network at all, so a flow
|
||||
// that reaches preMatchFlow bails out with PreMatchContinue instead of nil-
|
||||
// dereferencing. That is exactly the pre-fix verdict we assert against.
|
||||
type stubOutboundManager struct {
|
||||
adapter.OutboundManager
|
||||
defaultOutbound adapter.Outbound
|
||||
}
|
||||
|
||||
func (s *stubOutboundManager) Default() adapter.Outbound { return s.defaultOutbound }
|
||||
|
||||
type stubNoNetworkOutbound struct {
|
||||
adapter.Outbound
|
||||
}
|
||||
|
||||
func (o *stubNoNetworkOutbound) Tag() string { return "stub" }
|
||||
func (o *stubNoNetworkOutbound) Type() string { return "direct" }
|
||||
func (o *stubNoNetworkOutbound) Network() []string { return nil }
|
||||
|
||||
func newPreMatchTestRouter(t *testing.T, resolver adapter.NeighborResolver, rules ...option.Rule) *Router {
|
||||
t.Helper()
|
||||
logger := log.NewNOPFactory().NewLogger("test")
|
||||
router := &Router{
|
||||
ctx: context.Background(),
|
||||
logger: logger,
|
||||
dns: &stubDNSRouter{},
|
||||
dnsTransport: &stubDNSTransportManager{},
|
||||
outbound: &stubOutboundManager{defaultOutbound: &stubNoNetworkOutbound{}},
|
||||
neighborResolver: resolver,
|
||||
needFindNeighbor: true,
|
||||
}
|
||||
for i, ruleOptions := range rules {
|
||||
rule, err := R.NewRule(router.ctx, logger, ruleOptions, false)
|
||||
require.NoError(t, err, "build rule[%d]", i)
|
||||
router.rules = append(router.rules, rule)
|
||||
}
|
||||
return router
|
||||
}
|
||||
|
||||
func rejectOnSourceMAC(macAddress string) option.Rule {
|
||||
return option.Rule{
|
||||
Type: C.RuleTypeDefault,
|
||||
DefaultOptions: option.DefaultRule{
|
||||
RawDefaultRule: option.RawDefaultRule{
|
||||
SourceMACAddress: badoption.Listable[string]{macAddress},
|
||||
},
|
||||
RuleAction: option.RuleAction{
|
||||
Action: C.RuleActionTypeReject,
|
||||
RejectOptions: option.RejectActionOptions{Method: C.RuleActionRejectMethodDefault},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func rejectOnSourceHostname(hostname string) option.Rule {
|
||||
return option.Rule{
|
||||
Type: C.RuleTypeDefault,
|
||||
DefaultOptions: option.DefaultRule{
|
||||
RawDefaultRule: option.RawDefaultRule{
|
||||
SourceHostname: badoption.Listable[string]{hostname},
|
||||
},
|
||||
RuleAction: option.RuleAction{
|
||||
Action: C.RuleActionTypeReject,
|
||||
RejectOptions: option.RejectActionOptions{Method: C.RuleActionRejectMethodDefault},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func preMatchMetadata() adapter.InboundContext {
|
||||
return adapter.InboundContext{
|
||||
Inbound: "tun-in",
|
||||
InboundType: C.TypeTun,
|
||||
Network: N.NetworkUDP,
|
||||
Source: M.ParseSocksaddr("192.168.1.5:41234"),
|
||||
Destination: M.ParseSocksaddr("1.1.1.1:443"),
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreMatchResolvesNeighborMAC(t *testing.T) {
|
||||
t.Parallel()
|
||||
mac, err := net.ParseMAC("de:ad:be:ef:00:01")
|
||||
require.NoError(t, err)
|
||||
resolver := &stubNeighborResolver{
|
||||
address: netip.MustParseAddr("192.168.1.5"),
|
||||
mac: mac,
|
||||
hostname: "kitchen-tv",
|
||||
}
|
||||
router := newPreMatchTestRouter(t, resolver, rejectOnSourceMAC("de:ad:be:ef:00:01"))
|
||||
result := router.PreMatch(preMatchMetadata(), nil)
|
||||
require.Equal(t, adapter.PreMatchReject, result.Action,
|
||||
"source_mac_address rule must match in pre-match; the MAC has to be resolved there too")
|
||||
}
|
||||
|
||||
func TestPreMatchResolvesNeighborHostname(t *testing.T) {
|
||||
t.Parallel()
|
||||
resolver := &stubNeighborResolver{
|
||||
address: netip.MustParseAddr("192.168.1.5"),
|
||||
hostname: "kitchen-tv",
|
||||
}
|
||||
router := newPreMatchTestRouter(t, resolver, rejectOnSourceHostname("kitchen-tv"))
|
||||
result := router.PreMatch(preMatchMetadata(), nil)
|
||||
require.Equal(t, adapter.PreMatchReject, result.Action,
|
||||
"source_hostname rule must match in pre-match; the hostname has to be resolved there too")
|
||||
}
|
||||
|
||||
// A source the neighbor resolver does not know must still fall through, not
|
||||
// match on a half-filled metadata.
|
||||
func TestPreMatchNeighborMissDoesNotMatch(t *testing.T) {
|
||||
t.Parallel()
|
||||
mac, err := net.ParseMAC("de:ad:be:ef:00:01")
|
||||
require.NoError(t, err)
|
||||
resolver := &stubNeighborResolver{
|
||||
address: netip.MustParseAddr("192.168.1.9"),
|
||||
mac: mac,
|
||||
}
|
||||
router := newPreMatchTestRouter(t, resolver, rejectOnSourceMAC("de:ad:be:ef:00:01"))
|
||||
result := router.PreMatch(preMatchMetadata(), nil)
|
||||
require.Equal(t, adapter.PreMatchContinue, result.Action)
|
||||
}
|
||||
+28
-25
@@ -319,22 +319,14 @@ func (r *Router) PreMatch(metadata adapter.InboundContext, firstPacket []byte) a
|
||||
metadata.PreMatch = true
|
||||
continueResult := adapter.PreMatchResult{Action: adapter.PreMatchContinue}
|
||||
packetDestination := metadata.Destination
|
||||
if metadata.Destination.Addr.IsValid() && r.dnsTransport.FakeIP() != nil && r.dnsTransport.FakeIP().Store().Contains(metadata.Destination.Addr) {
|
||||
domain, loaded := r.dnsTransport.FakeIP().Store().Lookup(metadata.Destination.Addr)
|
||||
if !loaded || domain == "" {
|
||||
return continueResult
|
||||
}
|
||||
metadata.OriginDestination = metadata.Destination
|
||||
metadata.Destination = M.Socksaddr{
|
||||
Fqdn: domain,
|
||||
Port: metadata.Destination.Port,
|
||||
}
|
||||
metadata.FakeIP = true
|
||||
}
|
||||
if metadata.Destination.IsIPv4() {
|
||||
metadata.IPVersion = 4
|
||||
} else if metadata.Destination.IsIPv6() {
|
||||
metadata.IPVersion = 6
|
||||
// lx: pre-match used to prepare only fakeip + IP version, so process/neighbor
|
||||
// rule items (process_name, source_mac_address, source_hostname, …) never had
|
||||
// their metadata filled here and silently failed to match — they were resolved
|
||||
// in matchRule only. Both paths now share prepareMatchMetadata (upstream
|
||||
// b911fb078).
|
||||
err := r.prepareMatchMetadata(ctx, &metadata)
|
||||
if err != nil {
|
||||
return continueResult
|
||||
}
|
||||
for currentRuleIndex, currentRule := range r.rules {
|
||||
metadata.ResetRuleCache()
|
||||
@@ -540,13 +532,11 @@ func (r *Router) preMatchFlow(ctx context.Context, metadata *adapter.InboundCont
|
||||
return result
|
||||
}
|
||||
|
||||
func (r *Router) matchRule(
|
||||
ctx context.Context, metadata *adapter.InboundContext,
|
||||
inputConn net.Conn, inputPacketConn N.PacketConn,
|
||||
) (
|
||||
selectedRule adapter.Rule, selectedRuleIndex int,
|
||||
buffers []*buf.Buffer, packetBuffers []*N.PacketBuffer, fatalErr error,
|
||||
) {
|
||||
// prepareMatchMetadata fills in everything a rule may match on but the inbound
|
||||
// cannot know: the connection owner, the neighbor (MAC/hostname) behind the
|
||||
// source address, the fakeip / reverse-mapped domain and the IP version. Shared
|
||||
// by matchRule and PreMatch — see the note at the PreMatch call site.
|
||||
func (r *Router) prepareMatchMetadata(ctx context.Context, metadata *adapter.InboundContext) error {
|
||||
r.searchProcessInfo(ctx, metadata)
|
||||
if r.neighborResolver != nil && metadata.SourceMACAddress == nil && metadata.Source.Addr.IsValid() {
|
||||
mac, macFound := r.neighborResolver.LookupMAC(metadata.Source.Addr)
|
||||
@@ -568,8 +558,7 @@ func (r *Router) matchRule(
|
||||
if metadata.Destination.Addr.IsValid() && r.dnsTransport.FakeIP() != nil && r.dnsTransport.FakeIP().Store().Contains(metadata.Destination.Addr) {
|
||||
domain, loaded := r.dnsTransport.FakeIP().Store().Lookup(metadata.Destination.Addr)
|
||||
if !loaded {
|
||||
fatalErr = E.New("missing fakeip record, try enable `experimental.cache_file`")
|
||||
return
|
||||
return E.New("missing fakeip record, try enable `experimental.cache_file`")
|
||||
}
|
||||
if domain != "" {
|
||||
metadata.OriginDestination = metadata.Destination
|
||||
@@ -592,6 +581,20 @@ func (r *Router) matchRule(
|
||||
} else if metadata.Destination.IsIPv6() {
|
||||
metadata.IPVersion = 6
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Router) matchRule(
|
||||
ctx context.Context, metadata *adapter.InboundContext,
|
||||
inputConn net.Conn, inputPacketConn N.PacketConn,
|
||||
) (
|
||||
selectedRule adapter.Rule, selectedRuleIndex int,
|
||||
buffers []*buf.Buffer, packetBuffers []*N.PacketBuffer, fatalErr error,
|
||||
) {
|
||||
fatalErr = r.prepareMatchMetadata(ctx, metadata)
|
||||
if fatalErr != nil {
|
||||
return
|
||||
}
|
||||
|
||||
match:
|
||||
for currentRuleIndex, currentRule := range r.rules {
|
||||
|
||||
@@ -87,22 +87,27 @@ func (c *Client) DialContext(ctx context.Context) (net.Conn, error) {
|
||||
request.Header.Set("Upgrade", "websocket")
|
||||
err = request.Write(conn)
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
bufReader := std_bufio.NewReader(conn)
|
||||
response, err := http.ReadResponse(bufReader, request)
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
if response.StatusCode != 101 ||
|
||||
!strings.EqualFold(response.Header.Get("Connection"), "upgrade") ||
|
||||
!strings.EqualFold(response.Header.Get("Upgrade"), "websocket") {
|
||||
conn.Close()
|
||||
response.Body.Close()
|
||||
return nil, E.New("v2ray-http-upgrade: unexpected status: ", response.Status)
|
||||
}
|
||||
if bufReader.Buffered() > 0 {
|
||||
buffer := buf.NewSize(bufReader.Buffered())
|
||||
_, err = buffer.ReadFullFrom(bufReader, buffer.Len())
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
conn = bufio.NewCachedConn(conn, buffer)
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
package v2rayhttpupgrade
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing-box/option"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// A failed http-upgrade handshake used to return the error while leaving the
|
||||
// dialed TCP connection open: nothing in the caller chain owns a conn that was
|
||||
// never returned. On a router that keeps ~380 nodes under a health checker, one
|
||||
// leaked descriptor per failed handshake is a slow death. Upstream 0f1763877.
|
||||
|
||||
type trackedConn struct {
|
||||
net.Conn
|
||||
closed atomic.Bool
|
||||
}
|
||||
|
||||
func (c *trackedConn) Close() error {
|
||||
c.closed.Store(true)
|
||||
return c.Conn.Close()
|
||||
}
|
||||
|
||||
// dialRecorder dials the real (test) listener and remembers every conn handed
|
||||
// out, so the test can assert on the exact conn the client was given.
|
||||
type dialRecorder struct {
|
||||
access sync.Mutex
|
||||
conns []*trackedConn
|
||||
}
|
||||
|
||||
func (d *dialRecorder) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
||||
conn, err := new(net.Dialer).DialContext(ctx, network, destination.String())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tracked := &trackedConn{Conn: conn}
|
||||
d.access.Lock()
|
||||
d.conns = append(d.conns, tracked)
|
||||
d.access.Unlock()
|
||||
return tracked, nil
|
||||
}
|
||||
|
||||
func (d *dialRecorder) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
func (d *dialRecorder) only(t *testing.T) *trackedConn {
|
||||
t.Helper()
|
||||
d.access.Lock()
|
||||
defer d.access.Unlock()
|
||||
require.Len(t, d.conns, 1, "client must have dialed exactly once")
|
||||
return d.conns[0]
|
||||
}
|
||||
|
||||
// serveOnce accepts one connection, drains the request and writes raw response
|
||||
// bytes back.
|
||||
func serveOnce(t *testing.T, response string) net.Listener {
|
||||
t.Helper()
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { listener.Close() })
|
||||
go func() {
|
||||
conn, acceptErr := listener.Accept()
|
||||
if acceptErr != nil {
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
buffer := make([]byte, 4096)
|
||||
conn.SetReadDeadline(time.Now().Add(5 * time.Second))
|
||||
conn.Read(buffer)
|
||||
if response != "" {
|
||||
conn.Write([]byte(response))
|
||||
}
|
||||
}()
|
||||
return listener
|
||||
}
|
||||
|
||||
func newTestClient(t *testing.T, dialer *dialRecorder, listener net.Listener) *Client {
|
||||
t.Helper()
|
||||
client, err := NewClient(
|
||||
context.Background(),
|
||||
dialer,
|
||||
M.ParseSocksaddr(listener.Addr().String()),
|
||||
option.V2RayHTTPUpgradeOptions{Path: "/"},
|
||||
nil,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
return client
|
||||
}
|
||||
|
||||
func TestClientClosesConnOnUnexpectedStatus(t *testing.T) {
|
||||
t.Parallel()
|
||||
listener := serveOnce(t, "HTTP/1.1 403 Forbidden\r\nContent-Length: 0\r\n\r\n")
|
||||
dialer := &dialRecorder{}
|
||||
client := newTestClient(t, dialer, listener)
|
||||
|
||||
conn, err := client.DialContext(context.Background())
|
||||
require.Error(t, err)
|
||||
require.Nil(t, conn)
|
||||
require.True(t, dialer.only(t).closed.Load(),
|
||||
"the dialed conn must be closed when the upgrade is refused")
|
||||
}
|
||||
|
||||
func TestClientClosesConnOnResponseFailure(t *testing.T) {
|
||||
t.Parallel()
|
||||
// The server hangs up without answering: http.ReadResponse fails.
|
||||
listener := serveOnce(t, "")
|
||||
dialer := &dialRecorder{}
|
||||
client := newTestClient(t, dialer, listener)
|
||||
|
||||
conn, err := client.DialContext(context.Background())
|
||||
require.Error(t, err)
|
||||
require.Nil(t, conn)
|
||||
require.True(t, dialer.only(t).closed.Load(),
|
||||
"the dialed conn must be closed when the response cannot be read")
|
||||
}
|
||||
@@ -78,6 +78,12 @@ func (c *Client) offerNew() (*quic.Conn, error) {
|
||||
packetConn.Close()
|
||||
return nil, err
|
||||
}
|
||||
// quic-go does not take ownership of the packet conn passed to Dial:
|
||||
// when the connection ends it only stops reading.
|
||||
go func() {
|
||||
<-quicConn.Context().Done()
|
||||
packetConn.Close()
|
||||
}()
|
||||
c.conn.Store(quicConn)
|
||||
c.rawConn = udpConn
|
||||
return quicConn, nil
|
||||
|
||||
@@ -0,0 +1,203 @@
|
||||
//go:build with_quic
|
||||
|
||||
package v2rayquic
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/quic-go"
|
||||
"github.com/sagernet/sing-box/common/tls"
|
||||
"github.com/sagernet/sing-box/log"
|
||||
"github.com/sagernet/sing-box/option"
|
||||
qtls "github.com/sagernet/sing-quic"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// quic-go does not take ownership of the packet conn handed to Dial: when the
|
||||
// QUIC connection ends it only stops reading from it. offerNew() then dials a
|
||||
// fresh one and overwrites c.rawConn, so the previous UDP socket was leaked for
|
||||
// the lifetime of the process — one per reconnect, on a box with 512 MB and a
|
||||
// health checker that reconnects constantly. Upstream 7067276170.
|
||||
|
||||
const testALPN = "shater-test"
|
||||
|
||||
type trackedConn struct {
|
||||
net.Conn
|
||||
closed atomic.Bool
|
||||
}
|
||||
|
||||
func (c *trackedConn) Close() error {
|
||||
c.closed.Store(true)
|
||||
return c.Conn.Close()
|
||||
}
|
||||
|
||||
type dialRecorder struct {
|
||||
access sync.Mutex
|
||||
conns []*trackedConn
|
||||
}
|
||||
|
||||
func (d *dialRecorder) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
||||
conn, err := new(net.Dialer).DialContext(ctx, network, destination.String())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tracked := &trackedConn{Conn: conn}
|
||||
d.access.Lock()
|
||||
d.conns = append(d.conns, tracked)
|
||||
d.access.Unlock()
|
||||
return tracked, nil
|
||||
}
|
||||
|
||||
func (d *dialRecorder) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
func (d *dialRecorder) only(t *testing.T) *trackedConn {
|
||||
t.Helper()
|
||||
d.access.Lock()
|
||||
defer d.access.Unlock()
|
||||
require.Len(t, d.conns, 1, "client must have dialed exactly once")
|
||||
return d.conns[0]
|
||||
}
|
||||
|
||||
// serveQUIC brings up a real QUIC listener on localhost with a self-signed
|
||||
// certificate, runs handler for every accepted connection, and returns its
|
||||
// address.
|
||||
func serveQUIC(t *testing.T, handler func(conn *quic.Conn)) M.Socksaddr {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
logger := log.NewNOPFactory().NewLogger("test")
|
||||
|
||||
keyPem, certificatePem, err := tls.GenerateCertificate(nil, nil, time.Now, "localhost", time.Now().Add(time.Hour))
|
||||
require.NoError(t, err)
|
||||
serverTLSConfig, err := tls.NewSTDServer(ctx, logger, option.InboundTLSOptions{
|
||||
Enabled: true,
|
||||
ServerName: "localhost",
|
||||
ALPN: []string{testALPN},
|
||||
Certificate: []string{string(certificatePem)},
|
||||
Key: []string{string(keyPem)},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, serverTLSConfig.Start())
|
||||
t.Cleanup(func() { serverTLSConfig.Close() })
|
||||
|
||||
packetConn, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { packetConn.Close() })
|
||||
|
||||
listener, err := qtls.Listen(packetConn, serverTLSConfig, &quic.Config{})
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { listener.Close() })
|
||||
|
||||
go func() {
|
||||
for {
|
||||
conn, acceptErr := listener.Accept(ctx)
|
||||
if acceptErr != nil {
|
||||
return
|
||||
}
|
||||
go handler(conn)
|
||||
}
|
||||
}()
|
||||
return M.ParseSocksaddr(packetConn.LocalAddr().String())
|
||||
}
|
||||
|
||||
func newTestClient(t *testing.T, dialer *dialRecorder, serverAddr M.Socksaddr) *Client {
|
||||
t.Helper()
|
||||
clientTLSConfig, err := tls.NewSTDClient(context.Background(), log.NewNOPFactory().NewLogger("test"), "localhost", option.OutboundTLSOptions{
|
||||
Enabled: true,
|
||||
Insecure: true,
|
||||
ServerName: "localhost",
|
||||
ALPN: []string{testALPN},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
transport, err := NewClient(context.Background(), dialer, serverAddr, option.V2RayQUICOptions{}, clientTLSConfig)
|
||||
require.NoError(t, err)
|
||||
client, isClient := transport.(*Client)
|
||||
require.True(t, isClient)
|
||||
return client
|
||||
}
|
||||
|
||||
func TestClientClosesPacketConnWhenConnectionEnds(t *testing.T) {
|
||||
// Drop the connection right after the handshake: this is the server-side
|
||||
// reset / idle timeout the client must survive without leaking its socket.
|
||||
serverAddr := serveQUIC(t, func(conn *quic.Conn) {
|
||||
conn.CloseWithError(0, "bye")
|
||||
})
|
||||
dialer := &dialRecorder{}
|
||||
client := newTestClient(t, dialer, serverAddr)
|
||||
t.Cleanup(func() { client.Close() })
|
||||
|
||||
quicConn, err := client.offer()
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, quicConn)
|
||||
|
||||
// The server hangs up; the client keeps its Client alive (a health checker
|
||||
// would simply dial again later).
|
||||
select {
|
||||
case <-quicConn.Context().Done():
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("server never closed the QUIC connection")
|
||||
}
|
||||
|
||||
tracked := dialer.only(t)
|
||||
require.Eventually(t, tracked.closed.Load, 5*time.Second, 10*time.Millisecond,
|
||||
"the UDP socket behind a dead QUIC connection must be closed, not leaked until Client.Close()")
|
||||
}
|
||||
|
||||
// quic-go's Stream.Close() does not unblock a Write parked on flow control. The
|
||||
// writer goroutine (for us: the copy loop of a proxied connection) then survives
|
||||
// its own connection forever. Closing has to push the write deadline into the
|
||||
// past as well.
|
||||
func TestStreamCloseUnblocksBlockedWrite(t *testing.T) {
|
||||
serverIsDone := make(chan struct{})
|
||||
t.Cleanup(func() { close(serverIsDone) })
|
||||
// Accept the stream but never read from it, so the client's writes fill the
|
||||
// receive window and block.
|
||||
serverAddr := serveQUIC(t, func(conn *quic.Conn) {
|
||||
_, err := conn.AcceptStream(context.Background())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
<-serverIsDone
|
||||
})
|
||||
dialer := &dialRecorder{}
|
||||
client := newTestClient(t, dialer, serverAddr)
|
||||
t.Cleanup(func() { client.Close() })
|
||||
|
||||
stream, err := client.DialContext(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
writeDone := make(chan error, 1)
|
||||
go func() {
|
||||
payload := make([]byte, 64*1024)
|
||||
// 32 MiB is far past any quic-go receive window, so this must park.
|
||||
for range 512 {
|
||||
_, writeErr := stream.Write(payload)
|
||||
if writeErr != nil {
|
||||
writeDone <- writeErr
|
||||
return
|
||||
}
|
||||
}
|
||||
writeDone <- nil
|
||||
}()
|
||||
|
||||
select {
|
||||
case err = <-writeDone:
|
||||
t.Fatal("the write never blocked, the test proves nothing: ", err)
|
||||
case <-time.After(time.Second):
|
||||
}
|
||||
|
||||
require.NoError(t, stream.Close())
|
||||
select {
|
||||
case <-writeDone:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("Write stayed blocked after Close: the writer goroutine is leaked")
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package v2rayquic
|
||||
|
||||
import (
|
||||
"net"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/quic-go"
|
||||
qtls "github.com/sagernet/sing-quic"
|
||||
@@ -37,5 +38,8 @@ func (s *StreamWrapper) Upstream() any {
|
||||
func (s *StreamWrapper) Close() error {
|
||||
s.CancelRead(0)
|
||||
s.Stream.Close()
|
||||
// quic-go's Stream.Close does not unblock a Write blocked on flow control,
|
||||
// but a past write deadline does; buffered data and the FIN are unaffected.
|
||||
s.Stream.SetWriteDeadline(time.Now())
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -93,12 +93,14 @@ func (c *Client) dialContext(ctx context.Context, requestURL *url.URL, headers h
|
||||
reader, _, err := ws.Dialer{Header: ws.HandshakeHeaderHTTP(headers), Protocols: protocols}.Upgrade(deadlineConn, requestURL)
|
||||
deadlineConn.SetDeadline(time.Time{})
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
if reader != nil {
|
||||
buffer := buf.NewSize(reader.Buffered())
|
||||
_, err = buffer.ReadFullFrom(reader, buffer.Len())
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
conn = bufio.NewCachedConn(conn, buffer)
|
||||
|
||||
@@ -0,0 +1,122 @@
|
||||
package v2raywebsocket
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing-box/option"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// A refused websocket upgrade used to return the error while leaving the dialed
|
||||
// TCP connection open — the conn was never returned, so nobody else could close
|
||||
// it. With ~380 nodes under a health checker every failing node leaks one
|
||||
// descriptor per probe. Upstream 0f1763877.
|
||||
|
||||
type trackedConn struct {
|
||||
net.Conn
|
||||
closed atomic.Bool
|
||||
}
|
||||
|
||||
func (c *trackedConn) Close() error {
|
||||
c.closed.Store(true)
|
||||
return c.Conn.Close()
|
||||
}
|
||||
|
||||
type dialRecorder struct {
|
||||
access sync.Mutex
|
||||
conns []*trackedConn
|
||||
}
|
||||
|
||||
func (d *dialRecorder) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
||||
conn, err := new(net.Dialer).DialContext(ctx, network, destination.String())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tracked := &trackedConn{Conn: conn}
|
||||
d.access.Lock()
|
||||
d.conns = append(d.conns, tracked)
|
||||
d.access.Unlock()
|
||||
return tracked, nil
|
||||
}
|
||||
|
||||
func (d *dialRecorder) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
func (d *dialRecorder) only(t *testing.T) *trackedConn {
|
||||
t.Helper()
|
||||
d.access.Lock()
|
||||
defer d.access.Unlock()
|
||||
require.Len(t, d.conns, 1, "client must have dialed exactly once")
|
||||
return d.conns[0]
|
||||
}
|
||||
|
||||
func serveOnce(t *testing.T, response string) net.Listener {
|
||||
t.Helper()
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { listener.Close() })
|
||||
go func() {
|
||||
conn, acceptErr := listener.Accept()
|
||||
if acceptErr != nil {
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
buffer := make([]byte, 4096)
|
||||
conn.SetReadDeadline(time.Now().Add(5 * time.Second))
|
||||
conn.Read(buffer)
|
||||
if response != "" {
|
||||
conn.Write([]byte(response))
|
||||
}
|
||||
}()
|
||||
return listener
|
||||
}
|
||||
|
||||
func newTestClient(t *testing.T, dialer *dialRecorder, listener net.Listener) *Client {
|
||||
t.Helper()
|
||||
transport, err := NewClient(
|
||||
context.Background(),
|
||||
dialer,
|
||||
M.ParseSocksaddr(listener.Addr().String()),
|
||||
option.V2RayWebsocketOptions{Path: "/"},
|
||||
nil,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
client, isClient := transport.(*Client)
|
||||
require.True(t, isClient)
|
||||
return client
|
||||
}
|
||||
|
||||
func TestClientClosesConnOnRefusedUpgrade(t *testing.T) {
|
||||
t.Parallel()
|
||||
listener := serveOnce(t, "HTTP/1.1 403 Forbidden\r\nContent-Length: 0\r\n\r\n")
|
||||
dialer := &dialRecorder{}
|
||||
client := newTestClient(t, dialer, listener)
|
||||
|
||||
conn, err := client.DialContext(context.Background())
|
||||
require.Error(t, err)
|
||||
require.Nil(t, conn)
|
||||
require.True(t, dialer.only(t).closed.Load(),
|
||||
"the dialed conn must be closed when the websocket upgrade is refused")
|
||||
}
|
||||
|
||||
func TestClientClosesConnOnHandshakeEOF(t *testing.T) {
|
||||
t.Parallel()
|
||||
// The server hangs up mid-handshake: ws.Dialer.Upgrade fails on read.
|
||||
listener := serveOnce(t, "")
|
||||
dialer := &dialRecorder{}
|
||||
client := newTestClient(t, dialer, listener)
|
||||
|
||||
conn, err := client.DialContext(context.Background())
|
||||
require.Error(t, err)
|
||||
require.Nil(t, conn)
|
||||
require.True(t, dialer.only(t).closed.Load(),
|
||||
"the dialed conn must be closed when the handshake cannot complete")
|
||||
}
|
||||
@@ -187,6 +187,7 @@ func (c *EarlyWebsocketConn) writeRequest(content []byte) error {
|
||||
if len(lateData) > 0 {
|
||||
_, err = conn.Write(lateData)
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user