diff --git a/infra/conf/transport_finalmask.go b/infra/conf/transport_finalmask.go index 9dde6ca99..c46016584 100644 --- a/infra/conf/transport_finalmask.go +++ b/infra/conf/transport_finalmask.go @@ -14,7 +14,6 @@ import ( googleuuid "github.com/google/uuid" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/net" - "github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet/finalmask/fragment" "github.com/xtls/xray-core/transport/internet/finalmask/header/custom" "github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm" @@ -909,22 +908,13 @@ func (c *Realm) Build() (proto.Message, error) { } type UDPHop struct { - Sockopt *SocketConfig `json:"sockopt"` - Mode string `json:"mode"` - Interval Int32Range `json:"interval"` - RemotePorts PortList `json:"remotePorts"` - RemoteIPs []string `json:"remoteIPs"` + Mode string `json:"mode"` + Interval Int32Range `json:"interval"` + RemoteIPs []string `json:"remoteIPs"` + RemotePorts PortList `json:"remotePorts"` } func (c *UDPHop) Build() (proto.Message, error) { - var sockopt *internet.SocketConfig - if c.Sockopt != nil { - var err error - sockopt, err = c.Sockopt.Build() - if err != nil { - return nil, err - } - } var local, remote, remoteOnce bool for _, mode := range strings.Split(c.Mode, ",") { switch strings.ToLower(mode) { @@ -953,14 +943,13 @@ func (c *UDPHop) Build() (proto.Message, error) { return nil, errors.New("invalid ip ", ip) } return &udphop.Config{ - Sockopt: sockopt, Local: local, Remote: remote, RemoteOnce: remoteOnce, IntervalMin: int64(c.Interval.From), IntervalMax: int64(c.Interval.To), - RemotePorts: c.RemotePorts.Build().Ports(), RemoteIPs: remoteIPs, + RemotePorts: c.RemotePorts.Build().Ports(), }, nil } diff --git a/proxy/wireguard/client.go b/proxy/wireguard/client.go index ef0cc1e2a..c43758eb8 100644 --- a/proxy/wireguard/client.go +++ b/proxy/wireguard/client.go @@ -28,6 +28,7 @@ import ( "github.com/xtls/xray-core/features/stats" "github.com/xtls/xray-core/transport" "github.com/xtls/xray-core/transport/internet" + "github.com/xtls/xray-core/transport/internet/finalmask" "golang.zx2c4.com/wireguard/device" ) @@ -293,26 +294,26 @@ func (h *Handler) init(ctx context.Context) error { if err != nil { return nil, err } - conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings) - if err != nil { - return nil, err - } var pktConn net.PacketConn - switch c := conn.(type) { - case *internet.PacketConnWrapper: - pktConn = c.PacketConn - case *cnc.Connection: - pktConn = &internet.FakePacketConn{Conn: c} - default: - panic(reflect.TypeOf(c)) - } - if h.streamSettings.UdpmaskManager != nil { - newConn, err := h.streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn) + if h.streamSettings.FinalMask != nil { + conn, err := h.streamSettings.FinalMask.DialUDP(ctx, dest) if err != nil { - pktConn.Close() - return nil, errors.New("mask err").Base(err) + return nil, errors.New("failed to dial to dest").Base(err) + } + pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn + } else { + conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings) + if err != nil { + return nil, errors.New("failed to dial to dest").Base(err) + } + switch c := conn.(type) { + case *internet.PacketConnWrapper: + pktConn = c.PacketConn + case *cnc.Connection: + pktConn = &internet.FakePacketConn{Conn: c} + default: + panic(reflect.TypeOf(c)) } - pktConn = newConn } if h.uplinkCounter != nil || h.downlinkCounter != nil { pktConn = &PacketCounterConnection{ diff --git a/proxy/wireguard/server.go b/proxy/wireguard/server.go index 51416c435..a9a0a73ca 100644 --- a/proxy/wireguard/server.go +++ b/proxy/wireguard/server.go @@ -258,18 +258,16 @@ func (s *Server) Start() error { return errors.New("address is domain") } listenFunc := func() (net.PacketConn, error) { - pktConn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}, s.streamSettings.SocketSettings) + var pktConn net.PacketConn + var err error + if s.streamSettings.FinalMask != nil { + pktConn, err = s.streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}) + } else { + pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}, s.streamSettings.SocketSettings) + } if err != nil { return nil, err } - if s.streamSettings.UdpmaskManager != nil { - newConn, err := s.streamSettings.UdpmaskManager.WrapPacketConnServer(pktConn) - if err != nil { - pktConn.Close() - return nil, errors.New("mask err").Base(err) - } - pktConn = newConn - } if s.uplinkCounter != nil || s.downlinkCounter != nil { pktConn = &PacketCounterConnection{ PacketConn: pktConn, diff --git a/testing/scenarios/wireguard_test.go b/testing/scenarios/wireguard_test.go index 18b0ef987..ffb717cec 100644 --- a/testing/scenarios/wireguard_test.go +++ b/testing/scenarios/wireguard_test.go @@ -65,6 +65,7 @@ func TestWireguard(t *testing.T) { ProxySettings: serial.ToTypedMessage(&freedom.Config{ FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}}, }), + SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{}), }, }, } @@ -104,6 +105,7 @@ func TestWireguard(t *testing.T) { AllowedIps: []string{"0.0.0.0/0", "::0/0"}, }}, }), + SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{}), }, }, } diff --git a/transport/internet/finalmask/finalmask.go b/transport/internet/finalmask/finalmask.go index db723de56..179f958f0 100644 --- a/transport/internet/finalmask/finalmask.go +++ b/transport/internet/finalmask/finalmask.go @@ -2,103 +2,291 @@ package finalmask import ( "context" - "net" + "fmt" "slices" "github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" ) -type Udpmask interface { - WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) - WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) +type Dialer struct { + DialTCP func(net.Destination) (net.Conn, error) + DialUDP func(net.Destination) (net.Conn, error) } -type UdpmaskManager struct { - udpmasks []Udpmask +type ListenConfig struct { + Listen func(net.Addr) (net.Listener, error) + ListenPacket func(net.Addr) (net.PacketConn, error) } -func NewUdpmaskManager(udpmasks []Udpmask) *UdpmaskManager { - slices.Reverse(udpmasks) - return &UdpmaskManager{udpmasks: udpmasks} +type TCPMask interface { + WrapConnClient(net.Conn, *net.Destination, *Dialer) (net.Conn, error) + WrapConnServer(net.Conn) (net.Conn, error) + // Listen(net.Listener) (net.Listener, error) } -func (m *UdpmaskManager) WrapPacketConnClient(raw net.PacketConn) (net.PacketConn, error) { - var sizes []int - var conns []net.PacketConn - for i, mask := range m.udpmasks { - if _, ok := mask.(headerConn); ok { - conn, err := mask.WrapPacketConnClient(nil, i, len(m.udpmasks)-1) - if err != nil { - return nil, err - } - sizes = append(sizes, conn.(headerSize).Size()) - conns = append(conns, conn) - } else { - if len(conns) > 0 { - raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw} - sizes = nil - conns = nil - } - var err error - raw, err = mask.WrapPacketConnClient(raw, i, len(m.udpmasks)-1) - if err != nil { - return nil, err +type UDPMask interface { + WrapPacketConnClient(net.PacketConn, *net.Destination, *Dialer) (net.PacketConn, error) + WrapPacketConnServer(net.PacketConn, net.Addr, *ListenConfig) (net.PacketConn, error) +} + +type FinalMask struct { + tcpMasks []TCPMask + udpMasks []UDPMask + dialTCP func(context.Context, net.Destination) (net.Conn, error) + listen func(context.Context, net.Addr) (net.Listener, error) + dialUDP func(context.Context, net.Destination) (net.PacketConn, net.Addr, error) + listenPacket func(context.Context, net.Addr) (net.PacketConn, error) +} + +func NewFinalMask(tcpMasks []TCPMask, udpMasks []UDPMask, dialTCP func(context.Context, net.Destination) (net.Conn, error), listen func(context.Context, net.Addr) (net.Listener, error), dialUDP func(context.Context, net.Destination) (net.PacketConn, net.Addr, error), listenPacket func(context.Context, net.Addr) (net.PacketConn, error)) *FinalMask { + slices.Reverse(tcpMasks) + slices.Reverse(udpMasks) + return &FinalMask{ + tcpMasks: tcpMasks, + udpMasks: udpMasks, + dialTCP: dialTCP, + dialUDP: dialUDP, + listen: listen, + listenPacket: listenPacket, + } +} + +func (fm *FinalMask) DialTCP(ctx context.Context, dest net.Destination) (net.Conn, error) { + if len(fm.tcpMasks) == 0 { + return fm.dialTCP(ctx, dest) + } + for i := range fm.tcpMasks { + if i > 0 { + if _, ok := fm.tcpMasks[i].(interface{ HandleDial() }); ok { + return nil, fmt.Errorf("incorrect index: %d %T", i, fm.tcpMasks[i]) } } } - - if len(conns) > 0 { - raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw} - sizes = nil - conns = nil + var conn net.Conn + var err error + if _, ok := fm.tcpMasks[0].(interface{ HandleDial() }); !ok { + conn, err = fm.dialTCP(ctx, dest) + if err != nil { + return nil, err + } } - return raw, nil + dialer := &Dialer{ + DialTCP: func(dest net.Destination) (net.Conn, error) { + return fm.dialTCP(ctx, dest) + }, + DialUDP: func(dest net.Destination) (net.Conn, error) { + conn, addr, err := fm.dialUDP(ctx, dest) + if err != nil { + return nil, err + } + return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err + }, + } + for i := range fm.tcpMasks { + var newConn net.Conn + newConn, err = fm.tcpMasks[i].WrapConnClient(conn, &dest, dialer) + if err != nil { + _ = conn.Close() + return nil, err + } + conn = newConn + } + return conn, nil } -func (m *UdpmaskManager) WrapPacketConnServer(raw net.PacketConn) (net.PacketConn, error) { - var sizes []int - var conns []net.PacketConn - for i, mask := range m.udpmasks { - if _, ok := mask.(headerConn); ok { - conn, err := mask.WrapPacketConnServer(nil, i, len(m.udpmasks)-1) - if err != nil { - return nil, err +func (fm *FinalMask) Listen(ctx context.Context, addr net.Addr) (net.Listener, error) { + if len(fm.tcpMasks) == 0 { + return fm.listen(ctx, addr) + } + off := 0 + listener, err := fm.listen(ctx, addr) + if err != nil { + return nil, err + } + for i := range fm.tcpMasks { + if _, ok := fm.tcpMasks[i].(interface { + Listen(net.Listener) (net.Listener, error) + }); ok { + if i-off == 0 { + l, err := fm.tcpMasks[i].(interface { + Listen(net.Listener) (net.Listener, error) + }).Listen(listener) + if err != nil { + listener.Close() + return nil, err + } + listener = l + } else { + l, err := fm.tcpMasks[i].(interface { + Listen(net.Listener) (net.Listener, error) + }).Listen(&TCPListener{Listener: listener, tcpMasks: fm.tcpMasks[off:i]}) + if err != nil { + listener.Close() + return nil, err + } + listener = l } - sizes = append(sizes, conn.(headerSize).Size()) - conns = append(conns, conn) - } else { - if len(conns) > 0 { - raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw} - sizes = nil - conns = nil - } - var err error - raw, err = mask.WrapPacketConnServer(raw, i, len(m.udpmasks)-1) - if err != nil { - return nil, err + off = i + 1 + } + } + if off < len(fm.tcpMasks) { + return &TCPListener{Listener: listener, tcpMasks: fm.tcpMasks[off:]}, nil + } + return listener, nil +} + +func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Conn, error) { + if len(fm.udpMasks) == 0 { + conn, addr, err := fm.dialUDP(ctx, dest) + if err != nil { + return nil, err + } + return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil + } + for i := range fm.udpMasks { + if i > 0 { + if _, ok := fm.udpMasks[i].(interface{ HandleDial() }); ok { + return nil, fmt.Errorf("incorrect index: %d %T", i, fm.udpMasks[i]) } } } - + var conn net.PacketConn + var addr net.Addr + var err error + if _, ok := fm.udpMasks[0].(interface{ HandleDial() }); !ok { + conn, addr, err = fm.dialUDP(ctx, dest) + if err != nil { + return nil, err + } + } + dialer := &Dialer{ + DialTCP: func(dest net.Destination) (net.Conn, error) { + return fm.dialTCP(ctx, dest) + }, + DialUDP: func(dest net.Destination) (net.Conn, error) { + conn, addr, err := fm.dialUDP(ctx, dest) + if err != nil { + return nil, err + } + return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err + }, + } + var sizes []int + var conns []net.PacketConn + for i := range fm.udpMasks { + var newConn net.PacketConn + if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok { + newConn, err = fm.udpMasks[i].WrapPacketConnClient(nil, nil, nil) + if err != nil { + _ = conn.Close() + return nil, err + } + sizes = append(sizes, newConn.(interface{ Size() int }).Size()) + conns = append(conns, newConn) + } else { + if len(conns) > 0 { + conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns} + sizes = nil + conns = nil + } + newConn, err = fm.udpMasks[i].WrapPacketConnClient(conn, &dest, dialer) + if err != nil { + _ = conn.Close() + return nil, err + } + conn = newConn + } + } if len(conns) > 0 { - raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw} + conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns} sizes = nil conns = nil } - return raw, nil + if addr == nil { + addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}} + } + return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil +} + +func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) { + if len(fm.udpMasks) == 0 { + return fm.listenPacket(ctx, addr) + } + for i := range fm.udpMasks { + if i > 0 { + if _, ok := fm.udpMasks[i].(interface{ HandleListen() }); ok { + return nil, fmt.Errorf("incorrect index: %d %T", i, fm.udpMasks[i]) + } + } + } + var conn net.PacketConn + var err error + if _, ok := fm.udpMasks[0].(interface{ HandleListen() }); !ok { + conn, err = fm.listenPacket(ctx, addr) + if err != nil { + return nil, err + } + } + lc := &ListenConfig{ + Listen: func(addr net.Addr) (net.Listener, error) { return fm.listen(ctx, addr) }, + ListenPacket: func(addr net.Addr) (net.PacketConn, error) { return fm.listenPacket(ctx, addr) }, + } + var sizes []int + var conns []net.PacketConn + for i := range fm.udpMasks { + var newConn net.PacketConn + if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok { + newConn, err = fm.udpMasks[i].WrapPacketConnServer(nil, nil, nil) + if err != nil { + _ = conn.Close() + return nil, err + } + sizes = append(sizes, newConn.(interface{ Size() int }).Size()) + conns = append(conns, newConn) + } else { + if len(conns) > 0 { + conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns} + sizes = nil + conns = nil + } + newConn, err = fm.udpMasks[i].WrapPacketConnServer(conn, addr, lc) + if err != nil { + _ = conn.Close() + return nil, err + } + conn = newConn + } + } + if len(conns) > 0 { + conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns} + sizes = nil + conns = nil + } + return conn, nil } const ( UDPSize = 4096 ) -type headerConn interface { - HeaderConn() +type PacketConnWrapper struct { + net.PacketConn + udpAddr net.Addr } -type headerSize interface { - Size() int +func (c *PacketConnWrapper) RemoteAddr() net.Addr { + return c.udpAddr +} + +func (c *PacketConnWrapper) Read(b []byte) (n int, err error) { + n, _, err = c.PacketConn.ReadFrom(b) + return +} + +func (c *PacketConnWrapper) Write(b []byte) (n int, err error) { + return c.PacketConn.WriteTo(b, c.udpAddr) } type headerManagerConn struct { @@ -191,72 +379,27 @@ func (c *headerManagerConn) WriteTo(p []byte, addr net.Addr) (n int, err error) return len(p), nil } -type Tcpmask interface { - WrapConnClient(net.Conn) (net.Conn, error) - WrapConnServer(net.Conn) (net.Conn, error) -} - -type TcpmaskManager struct { - tcpmasks []Tcpmask -} - -func NewTcpmaskManager(tcpmasks []Tcpmask) *TcpmaskManager { - slices.Reverse(tcpmasks) - return &TcpmaskManager{tcpmasks: tcpmasks} -} - -func (m *TcpmaskManager) WrapConnClient(raw net.Conn) (net.Conn, error) { - var err error - for _, mask := range m.tcpmasks { - raw, err = mask.WrapConnClient(raw) - if err != nil { - return nil, err - } - } - return raw, nil -} - -func (m *TcpmaskManager) WrapConnServer(raw net.Conn) (net.Conn, error) { - var err error - for _, mask := range m.tcpmasks { - raw, err = mask.WrapConnServer(raw) - if err != nil { - return nil, err - } - } - return raw, nil -} - -func (m *TcpmaskManager) WrapListener(l net.Listener) (net.Listener, error) { - return NewTcpListener(m, l) -} - -type tcpListener struct { - m *TcpmaskManager +type TCPListener struct { net.Listener + tcpMasks []TCPMask } -func NewTcpListener(m *TcpmaskManager, l net.Listener) (net.Listener, error) { - return &tcpListener{ - m: m, - Listener: l, - }, nil -} - -func (l *tcpListener) Accept() (net.Conn, error) { +func (l *TCPListener) Accept() (net.Conn, error) { conn, err := l.Listener.Accept() if err != nil { return conn, err } - newConn, err := l.m.WrapConnServer(conn) - if err != nil { - errors.LogDebugInner(context.Background(), err, "mask err") - _ = conn.Close() - return nil, err + for i := range l.tcpMasks { + var newConn net.Conn + newConn, err = l.tcpMasks[i].WrapConnServer(conn) + if err != nil { + _ = conn.Close() + return nil, err + } + conn = newConn } - - return newConn, nil + return conn, nil } type TcpMaskConn interface { diff --git a/transport/internet/finalmask/fragment/config.go b/transport/internet/finalmask/fragment/config.go index 9610c0596..905b5232f 100644 --- a/transport/internet/finalmask/fragment/config.go +++ b/transport/internet/finalmask/fragment/config.go @@ -1,11 +1,14 @@ package fragment -import "net" +import ( + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" +) -func (c *Config) WrapConnClient(raw net.Conn) (net.Conn, error) { - return NewConnClient(c, raw, false) +func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) { + return NewConnClient(c, conn, false) } -func (c *Config) WrapConnServer(raw net.Conn) (net.Conn, error) { - return NewConnServer(c, raw, true) +func (c *Config) WrapConnServer(conn net.Conn) (net.Conn, error) { + return NewConnServer(c, conn, true) } diff --git a/transport/internet/finalmask/header/custom/config.go b/transport/internet/finalmask/header/custom/config.go index 7e2eaab74..58861bc7e 100644 --- a/transport/internet/finalmask/header/custom/config.go +++ b/transport/internet/finalmask/header/custom/config.go @@ -1,29 +1,30 @@ package custom import ( - "net" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) -func (c *TCPConfig) WrapConnClient(raw net.Conn) (net.Conn, error) { - return NewConnClientTCP(c, raw) +func (c *TCPConfig) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) { + return NewConnClientTCP(c, conn) } -func (c *TCPConfig) WrapConnServer(raw net.Conn) (net.Conn, error) { - return NewConnServerTCP(c, raw) +func (c *TCPConfig) WrapConnServer(conn net.Conn) (net.Conn, error) { + return NewConnServerTCP(c, conn) } -func (c *UDPConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnClientUDP(c, raw) +func (c *UDPConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewConnClientUDP(c, conn) } -func (c *UDPConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnServerUDP(c, raw) +func (c *UDPConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewConnServerUDP(c, conn) } -func (c *UDPStandaloneConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnClientUDPStandalone(c, raw) +func (c *UDPStandaloneConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewConnClientUDPStandalone(c, conn) } -func (c *UDPStandaloneConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnServerUDPStandalone(c, raw) +func (c *UDPStandaloneConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewConnServerUDPStandalone(c, conn) } diff --git a/transport/internet/finalmask/header/custom/metadata_test.go b/transport/internet/finalmask/header/custom/metadata_test.go index 0c426cea8..9c4f14da4 100644 --- a/transport/internet/finalmask/header/custom/metadata_test.go +++ b/transport/internet/finalmask/header/custom/metadata_test.go @@ -9,8 +9,6 @@ import ( "strings" "testing" "time" - - "github.com/xtls/xray-core/transport/internet/finalmask" ) func TestMetadataEvaluatorRejectsUnknownName(t *testing.T) { @@ -156,7 +154,7 @@ func TestMetadataUDPStandaloneWriteUsesRemotePort(t *testing.T) { } defer serverRaw.Close() - client, err := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg}).WrapPacketConnClient(clientRaw) + client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -301,7 +299,7 @@ func TestMetadataTCPHandshakeUsesEndpointPorts(t *testing.T) { } defer serverRaw.Close() - client, err := clientCfg.WrapConnClient(clientRaw) + client, err := clientCfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } diff --git a/transport/internet/finalmask/header/custom/state_test.go b/transport/internet/finalmask/header/custom/state_test.go index 4ef0a0477..d315408e3 100644 --- a/transport/internet/finalmask/header/custom/state_test.go +++ b/transport/internet/finalmask/header/custom/state_test.go @@ -5,8 +5,6 @@ import ( "net" "testing" "time" - - "github.com/xtls/xray-core/transport/internet/finalmask" ) func mustSendRecvUDP(t *testing.T, from net.PacketConn, to net.PacketConn, msg []byte) { @@ -48,7 +46,6 @@ func TestStateUDPResponseReusesPriorCapturedValues(t *testing.T) { }, }, } - maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg}) clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0") if err != nil { @@ -62,11 +59,11 @@ func TestStateUDPResponseReusesPriorCapturedValues(t *testing.T) { } defer serverRaw.Close() - client, err := maskManager.WrapPacketConnClient(clientRaw) + client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } - server, err := maskManager.WrapPacketConnServer(serverRaw) + server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil) if err != nil { t.Fatal(err) } diff --git a/transport/internet/finalmask/header/custom/tcp_runtime_test.go b/transport/internet/finalmask/header/custom/tcp_runtime_test.go index 9f10e833c..ed30791e5 100644 --- a/transport/internet/finalmask/header/custom/tcp_runtime_test.go +++ b/transport/internet/finalmask/header/custom/tcp_runtime_test.go @@ -37,7 +37,7 @@ func TestDSLTCPHandshakeReusesCapturedValue(t *testing.T) { defer clientRaw.Close() defer serverRaw.Close() - client, err := cfg.WrapConnClient(clientRaw) + client, err := cfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -117,7 +117,7 @@ func TestDSLTCPClientRejectsMismatchedResponseSequence(t *testing.T) { defer clientRaw.Close() defer serverRaw.Close() - client, err := clientCfg.WrapConnClient(clientRaw) + client, err := clientCfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } diff --git a/transport/internet/finalmask/mkcp/aes128gcm/config.go b/transport/internet/finalmask/mkcp/aes128gcm/config.go index d7cd0d41b..4e0b63f73 100644 --- a/transport/internet/finalmask/mkcp/aes128gcm/config.go +++ b/transport/internet/finalmask/mkcp/aes128gcm/config.go @@ -1,15 +1,16 @@ package aes128gcm import ( - "net" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) func (c *Config) HeaderConn() {} -func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnClient(c, raw) +func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewConnClient(c, conn) } -func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnServer(c, raw) +func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewConnServer(c, conn) } diff --git a/transport/internet/finalmask/mkcp/header/config.go b/transport/internet/finalmask/mkcp/header/config.go index 0b5c67be6..403381e88 100644 --- a/transport/internet/finalmask/mkcp/header/config.go +++ b/transport/internet/finalmask/mkcp/header/config.go @@ -1,15 +1,16 @@ package header import ( - "net" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) func (c *Config) HeaderConn() {} -func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnClient(c, raw) +func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewConnClient(c, conn) } -func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnServer(c, raw) +func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewConnServer(c, conn) } diff --git a/transport/internet/finalmask/mkcp/original/config.go b/transport/internet/finalmask/mkcp/original/config.go index 98bf964a9..9602ad0cb 100644 --- a/transport/internet/finalmask/mkcp/original/config.go +++ b/transport/internet/finalmask/mkcp/original/config.go @@ -1,15 +1,16 @@ package original import ( - "net" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) func (c *Config) HeaderConn() {} -func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnClient(c, raw) +func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewConnClient(c, conn) } -func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnServer(c, raw) +func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewConnServer(c, conn) } diff --git a/transport/internet/finalmask/noise/config.go b/transport/internet/finalmask/noise/config.go index ecb5db292..f6b02afcc 100644 --- a/transport/internet/finalmask/noise/config.go +++ b/transport/internet/finalmask/noise/config.go @@ -1,11 +1,14 @@ package noise -import "net" +import ( + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" +) -func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnClient(c, raw) +func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewConnClient(c, conn) } -func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewConnServer(c, raw) +func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewConnServer(c, conn) } diff --git a/transport/internet/finalmask/realm/config.go b/transport/internet/finalmask/realm/config.go index 4b8232f78..4ad7dc904 100644 --- a/transport/internet/finalmask/realm/config.go +++ b/transport/internet/finalmask/realm/config.go @@ -1,23 +1,14 @@ package realm import ( - "net" - - "github.com/xtls/xray-core/common/errors" - "github.com/xtls/xray-core/transport/internet" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) -func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - _, ok1 := raw.(*internet.FakePacketConn) - if level != 0 || ok1 { - return nil, errors.New("realm requires being at the outermost level") - } - return NewConnClient(c, raw) +func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewConnClient(c, conn) } -func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - if level != 0 { - return nil, errors.New("realm requires being at the outermost level") - } - return NewConnServer(c, raw) +func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewConnServer(c, conn) } diff --git a/transport/internet/finalmask/salamander/config.go b/transport/internet/finalmask/salamander/config.go index 7d019141e..99a0f501d 100644 --- a/transport/internet/finalmask/salamander/config.go +++ b/transport/internet/finalmask/salamander/config.go @@ -1,23 +1,24 @@ package salamander import ( - "net" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) func (c *Config) HeaderConn() {} -func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewSalamanderConnClient(c, raw) +func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewSalamanderConnClient(c, conn) } -func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewSalamanderConnServer(c, raw) +func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewSalamanderConnServer(c, conn) } -func (c *GeckoConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewGeckoConnClient(c, raw) +func (c *GeckoConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewGeckoConnClient(c, conn) } -func (c *GeckoConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - return NewGeckoConnServer(c, raw) +func (c *GeckoConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewGeckoConnServer(c, conn) } diff --git a/transport/internet/finalmask/sudoku/config.go b/transport/internet/finalmask/sudoku/config.go index f160a8b52..c4dcac9b6 100644 --- a/transport/internet/finalmask/sudoku/config.go +++ b/transport/internet/finalmask/sudoku/config.go @@ -1,19 +1,18 @@ package sudoku import ( - "net" - - "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) // Sudoku in finalmask mode is a pure appearance transform with no standalone handshake. // TCP always keeps classic sudoku on uplink and uses packed downlink optimization on server writes. -func (c *Config) WrapConnClient(raw net.Conn) (net.Conn, error) { - return newPackedDirectionalConn(raw, c, true) +func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) { + return newPackedDirectionalConn(conn, c, true) } -func (c *Config) WrapConnServer(raw net.Conn) (net.Conn, error) { - return newPackedDirectionalConn(raw, c, false) +func (c *Config) WrapConnServer(conn net.Conn) (net.Conn, error) { + return newPackedDirectionalConn(conn, c, false) } func newPackedDirectionalConn(raw net.Conn, config *Config, readPacked bool) (net.Conn, error) { @@ -36,16 +35,10 @@ func newPackedDirectionalConn(raw net.Conn, config *Config, readPacked bool) (ne return newWrappedConn(raw, reader, writer), nil } -func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - if level != levelCount { - return nil, errors.New("sudoku udp mask must be the innermost mask in chain") - } - return NewUDPConn(raw, c) +func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewUDPConn(conn, c) } -func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - if level != levelCount { - return nil, errors.New("sudoku udp mask must be the innermost mask in chain") - } - return NewUDPConn(raw, c) +func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewUDPConn(conn, c) } diff --git a/transport/internet/finalmask/tcp_test.go b/transport/internet/finalmask/tcp_test.go index 2ae888d0b..c8eca4c94 100644 --- a/transport/internet/finalmask/tcp_test.go +++ b/transport/internet/finalmask/tcp_test.go @@ -2,12 +2,14 @@ package finalmask_test import ( "bytes" + "context" "io" - "net" + gonet "net" "strings" "testing" "time" + "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/finalmask/header/custom" ) @@ -20,11 +22,14 @@ func mustSendRecvTcp( ) { t.Helper() + waitCh := make(chan error) + go func() { _, err := from.Write(msg) if err != nil { - t.Error(err) + t.Fatal(err) } + close(waitCh) }() buf := make([]byte, 1024) @@ -40,18 +45,23 @@ func mustSendRecvTcp( if !bytes.Equal(buf[:n], msg) { t.Fatalf("unexpected data %q", buf[:n]) } + + <-waitCh } type layerMaskTcp struct { name string - mask finalmask.Tcpmask + mask finalmask.TCPMask } type failingWrapMask struct{} -func (failingWrapMask) TCP() {} -func (f failingWrapMask) WrapConnClient(raw net.Conn) (net.Conn, error) { return raw, nil } -func (f failingWrapMask) WrapConnServer(raw net.Conn) (net.Conn, error) { +func (failingWrapMask) TCP() {} +func (f failingWrapMask) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) { + return conn, nil +} + +func (f failingWrapMask) WrapConnServer(conn net.Conn) (net.Conn, error) { return nil, io.ErrClosedPipe } @@ -92,32 +102,31 @@ func TestConnReadWrite(t *testing.T) { t.Run(c.name, func(t *testing.T) { mask := c.mask - maskManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{mask}) + dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) { + return net.Dial("tcp", dest.NetAddr()) + } + listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) { + return net.Listen("tcp", addr.String()) + } + finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{mask}, nil, dialTCP, listen, nil, nil) - ln, err := net.Listen("tcp", "127.0.0.1:0") + listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()}) if err != nil { t.Fatal(err) } + t.Cleanup(func() { listener.Close() }) - client, err := net.Dial("tcp", ln.Addr().String()) + client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port))) if err != nil { t.Fatal(err) } + t.Cleanup(func() { client.Close() }) - client, err = maskManager.WrapConnClient(client) - if err != nil { - t.Fatal(err) - } - - server, err := ln.Accept() - if err != nil { - t.Fatal(err) - } - - server, err = maskManager.WrapConnServer(server) + server, err := listener.Accept() if err != nil { t.Fatal(err) } + t.Cleanup(func() { server.Close() }) _ = client.SetDeadline(time.Now().Add(time.Second)) _ = server.SetDeadline(time.Now().Add(time.Second)) @@ -150,34 +159,32 @@ func TestTCPcustomStaticHandshakeRoundTrip(t *testing.T) { }, }, } - maskManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{cfg}) - ln, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatal(err) + dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) { + return net.Dial("tcp", dest.NetAddr()) } - defer ln.Close() + listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) { + return net.Listen("tcp", addr.String()) + } + finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{cfg}, nil, dialTCP, listen, nil, nil) - clientRaw, err := net.Dial("tcp", ln.Addr().String()) + listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()}) if err != nil { t.Fatal(err) } - defer clientRaw.Close() + defer listener.Close() - serverRaw, err := ln.Accept() + client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port))) if err != nil { t.Fatal(err) } - defer serverRaw.Close() + defer client.Close() - client, err := maskManager.WrapConnClient(clientRaw) - if err != nil { - t.Fatal(err) - } - server, err := maskManager.WrapConnServer(serverRaw) + server, err := listener.Accept() if err != nil { t.Fatal(err) } + defer server.Close() _ = client.SetDeadline(time.Now().Add(time.Second)) _ = server.SetDeadline(time.Now().Add(time.Second)) @@ -220,11 +227,11 @@ func TestTCPcustomClientRejectsMismatchedServerSequence(t *testing.T) { }, } - clientRaw, serverRaw := net.Pipe() + clientRaw, serverRaw := gonet.Pipe() defer clientRaw.Close() defer serverRaw.Close() - client, err := clientCfg.WrapConnClient(clientRaw) + client, err := clientCfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -257,42 +264,37 @@ func TestTCPcustomClientRejectsMismatchedServerSequence(t *testing.T) { } func TestTCPWrapListenerRejectsImmediateWrapErrors(t *testing.T) { - clientManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{failingWrapMask{}}) - serverManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{failingWrapMask{}}) + dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) { + return net.Dial("tcp", dest.NetAddr()) + } + listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) { + return net.Listen("tcp", addr.String()) + } + finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{failingWrapMask{}}, nil, dialTCP, listen, nil, nil) - rawLn, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatal(err) - } - defer rawLn.Close() - - ln, err := serverManager.WrapListener(rawLn) + listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()}) if err != nil { t.Fatal(err) } + defer listener.Close() accepted := make(chan struct { conn net.Conn err error }, 1) go func() { - conn, err := ln.Accept() + conn, err := listener.Accept() accepted <- struct { conn net.Conn err error }{conn: conn, err: err} }() - clientRaw, err := net.Dial("tcp", rawLn.Addr().String()) - if err != nil { - t.Fatal(err) - } - defer clientRaw.Close() - - client, err := clientManager.WrapConnClient(clientRaw) + client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port))) if err != nil { t.Fatal(err) } + defer client.Close() _ = client.SetDeadline(time.Now().Add(time.Second)) diff --git a/transport/internet/finalmask/udp_test.go b/transport/internet/finalmask/udp_test.go index 95233e02e..7d5ae5e39 100644 --- a/transport/internet/finalmask/udp_test.go +++ b/transport/internet/finalmask/udp_test.go @@ -2,13 +2,15 @@ package finalmask_test import ( "bytes" + "context" "encoding/binary" "io" - "net" + gonet "net" "sync/atomic" "testing" "time" + "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/proxy" "github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/finalmask/header/custom" @@ -51,7 +53,7 @@ func mustSendRecv( type layerMask struct { name string - mask finalmask.Udpmask + mask finalmask.UDPMask layers int } @@ -213,25 +215,23 @@ func newStandaloneStunLikeUDPServerConfig() *custom.UDPStandaloneConfig { func newUDPClientServerPair(t *testing.T, cfg *custom.UDPStandaloneConfig) (net.PacketConn, net.PacketConn, net.PacketConn, net.PacketConn) { t.Helper() - clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0") + clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = clientRaw.Close() }) - serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0") + serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = serverRaw.Close() }) - maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg}) - - client, err := maskManager.WrapPacketConnClient(clientRaw) + client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } - server, err := maskManager.WrapPacketConnServer(serverRaw) + server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -348,31 +348,39 @@ func TestPacketConnReadWrite(t *testing.T) { if layers <= 0 { layers = 1 } - masks := make([]finalmask.Udpmask, 0, layers) + masks := make([]finalmask.UDPMask, 0, layers) for i := 0; i < layers; i++ { masks = append(masks, mask) } - maskManager := finalmask.NewUdpmaskManager(masks) - client, err := net.ListenPacket("udp", "127.0.0.1:0") + dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) { + udpAddr, err := net.ResolveUDPAddr("udp", dest.NetAddr()) + if err != nil { + return nil, nil, err + } + conn, err := gonet.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + return nil, nil, err + } + return conn, udpAddr, nil + } + listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) { + return gonet.ListenPacket(addr.Network(), addr.String()) + } + finalMask := finalmask.NewFinalMask(nil, masks, nil, nil, dialUDP, listenPacket) + + server, err := finalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: net.LocalHostIP.IP()}) if err != nil { t.Fatal(err) } + t.Cleanup(func() { server.Close() }) - client, err = maskManager.WrapPacketConnClient(client) - if err != nil { - t.Fatal(err) - } - - server, err := net.ListenPacket("udp", "127.0.0.1:0") - if err != nil { - t.Fatal(err) - } - - server, err = maskManager.WrapPacketConnServer(server) + clientConn, err := finalMask.DialUDP(context.Background(), net.UDPDestination(net.IPAddress(server.LocalAddr().(*net.UDPAddr).IP), net.Port(server.LocalAddr().(*net.UDPAddr).Port))) if err != nil { t.Fatal(err) } + t.Cleanup(func() { clientConn.Close() }) + client := clientConn.(*finalmask.PacketConnWrapper).PacketConn _ = client.SetDeadline(time.Now().Add(time.Second)) _ = server.SetDeadline(time.Now().Add(time.Second)) @@ -397,21 +405,20 @@ func TestUDPcustomStaticHeaderWireShape(t *testing.T) { {Rand: 1, RandMin: 0x30, RandMax: 0x40}, }, } - maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg}) - clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0") + clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") if err != nil { t.Fatal(err) } defer clientRaw.Close() - serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0") + serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") if err != nil { t.Fatal(err) } defer serverRaw.Close() - client, err := maskManager.WrapPacketConnClient(clientRaw) + client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -642,11 +649,11 @@ func TestSudokuBDD(t *testing.T) { Ascii: "prefer_ascii", } - clientRaw, serverRaw := net.Pipe() + clientRaw, serverRaw := gonet.Pipe() defer clientRaw.Close() defer serverRaw.Close() - clientConn, err := cfg.WrapConnClient(clientRaw) + clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -683,11 +690,11 @@ func TestSudokuBDD(t *testing.T) { PaddingMax: 0, } - clientRaw, serverRaw := net.Pipe() + clientRaw, serverRaw := gonet.Pipe() defer clientRaw.Close() defer serverRaw.Close() - clientConn, err := cfg.WrapConnClient(clientRaw) + clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -738,10 +745,10 @@ func TestSudokuBDD(t *testing.T) { countWireBytes := func(wrapServer func(net.Conn, *sudoku.Config) (net.Conn, error), cfg *sudoku.Config) int64 { t.Helper() - clientRaw, serverRaw := net.Pipe() + clientRaw, serverRaw := gonet.Pipe() watchedServerRaw := &countingConn{Conn: serverRaw} - clientConn, err := cfg.WrapConnClient(clientRaw) + clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -793,11 +800,11 @@ func TestSudokuBDD(t *testing.T) { CustomTables: []string{"xpxvvpvv", "vxpvxvvp"}, } - clientRaw, serverRaw := net.Pipe() + clientRaw, serverRaw := gonet.Pipe() defer clientRaw.Close() defer serverRaw.Close() - clientConn, err := cfg.WrapConnClient(clientRaw) + clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -835,11 +842,11 @@ func TestSudokuBDD(t *testing.T) { PaddingMax: 0, } - clientRaw, serverRaw := net.Pipe() + clientRaw, serverRaw := gonet.Pipe() defer clientRaw.Close() defer serverRaw.Close() - clientConn, err := cfg.WrapConnClient(clientRaw) + clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -868,19 +875,6 @@ func TestSudokuBDD(t *testing.T) { } }) - t.Run("GivenSudokuUDPMask_WhenNotInnermost_ThenWrapFails", func(t *testing.T) { - cfg := &sudoku.Config{Password: "sudoku-udp"} - raw, err := net.ListenPacket("udp", "127.0.0.1:0") - if err != nil { - t.Fatal(err) - } - defer raw.Close() - - if _, err := cfg.WrapPacketConnClient(raw, 0, 1); err == nil { - t.Fatal("expected innermost check failure") - } - }) - t.Run("GivenSudokuMultiTableUDPMask_WhenClientSendsMultipleDatagrams_ThenPayloadMatches", func(t *testing.T) { cfg := &sudoku.Config{ Password: "sudoku-udp-multi", @@ -889,25 +883,24 @@ func TestSudokuBDD(t *testing.T) { PaddingMin: 0, PaddingMax: 0, } - maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg}) - clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0") + clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") if err != nil { t.Fatal(err) } defer clientRaw.Close() - serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0") + serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") if err != nil { t.Fatal(err) } defer serverRaw.Close() - client, err := maskManager.WrapPacketConnClient(clientRaw) + client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } - server, err := maskManager.WrapPacketConnServer(serverRaw) + server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -961,7 +954,7 @@ func TestSudokuBDD(t *testing.T) { } defer serverRaw.Close() - clientConn, err := cfg.WrapConnClient(clientRaw) + clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -1008,11 +1001,11 @@ func TestSudokuBDD(t *testing.T) { Ascii: "prefer_entropy", } - clientRaw, serverRaw := net.Pipe() + clientRaw, serverRaw := gonet.Pipe() defer clientRaw.Close() defer serverRaw.Close() - clientConn, err := cfg.WrapConnClient(clientRaw) + clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } @@ -1032,11 +1025,11 @@ func TestSudokuBDD(t *testing.T) { Ascii: "prefer_entropy", } - clientRaw, serverRaw := net.Pipe() + clientRaw, serverRaw := gonet.Pipe() defer clientRaw.Close() defer serverRaw.Close() - clientConn, err := cfg.WrapConnClient(clientRaw) + clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) if err != nil { t.Fatal(err) } diff --git a/transport/internet/finalmask/udphop/config.go b/transport/internet/finalmask/udphop/config.go index 5e8fda5fe..6fd612e47 100644 --- a/transport/internet/finalmask/udphop/config.go +++ b/transport/internet/finalmask/udphop/config.go @@ -1,20 +1,17 @@ package udphop import ( - "net" - "github.com/xtls/xray-core/common/errors" - "github.com/xtls/xray-core/transport/internet" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) -func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - _, ok1 := raw.(*internet.FakePacketConn) - if level != 0 || ok1 { - return nil, errors.New("udphop requires being at the outermost level") - } - return NewUDPHopConn(c, raw) +func (c *Config) HandleDial() {} + +func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewUDPHopConn(c, dest, dialer) } -func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { +func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { return nil, errors.New("udphop: client only") } diff --git a/transport/internet/finalmask/udphop/config.pb.go b/transport/internet/finalmask/udphop/config.pb.go index 6b4291c92..798446b2b 100644 --- a/transport/internet/finalmask/udphop/config.pb.go +++ b/transport/internet/finalmask/udphop/config.pb.go @@ -7,7 +7,6 @@ package udphop import ( - internet "github.com/xtls/xray-core/transport/internet" protoreflect "google.golang.org/protobuf/reflect/protoreflect" protoimpl "google.golang.org/protobuf/runtime/protoimpl" reflect "reflect" @@ -24,14 +23,13 @@ const ( type Config struct { state protoimpl.MessageState `protogen:"open.v1"` - Sockopt *internet.SocketConfig `protobuf:"bytes,1,opt,name=sockopt,proto3" json:"sockopt,omitempty"` Local bool `protobuf:"varint,2,opt,name=local,proto3" json:"local,omitempty"` Remote bool `protobuf:"varint,3,opt,name=remote,proto3" json:"remote,omitempty"` RemoteOnce bool `protobuf:"varint,4,opt,name=remote_once,json=remoteOnce,proto3" json:"remote_once,omitempty"` IntervalMin int64 `protobuf:"varint,5,opt,name=interval_min,json=intervalMin,proto3" json:"interval_min,omitempty"` IntervalMax int64 `protobuf:"varint,6,opt,name=interval_max,json=intervalMax,proto3" json:"interval_max,omitempty"` - RemotePorts []uint32 `protobuf:"varint,7,rep,packed,name=remote_ports,json=remotePorts,proto3" json:"remote_ports,omitempty"` - RemoteIPs []string `protobuf:"bytes,8,rep,name=remoteIPs,proto3" json:"remoteIPs,omitempty"` + RemoteIPs []string `protobuf:"bytes,7,rep,name=remoteIPs,proto3" json:"remoteIPs,omitempty"` + RemotePorts []uint32 `protobuf:"varint,8,rep,packed,name=remote_ports,json=remotePorts,proto3" json:"remote_ports,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } @@ -66,13 +64,6 @@ func (*Config) Descriptor() ([]byte, []int) { return file_transport_internet_finalmask_udphop_config_proto_rawDescGZIP(), []int{0} } -func (x *Config) GetSockopt() *internet.SocketConfig { - if x != nil { - return x.Sockopt - } - return nil -} - func (x *Config) GetLocal() bool { if x != nil { return x.Local @@ -108,16 +99,16 @@ func (x *Config) GetIntervalMax() int64 { return 0 } -func (x *Config) GetRemotePorts() []uint32 { +func (x *Config) GetRemoteIPs() []string { if x != nil { - return x.RemotePorts + return x.RemoteIPs } return nil } -func (x *Config) GetRemoteIPs() []string { +func (x *Config) GetRemotePorts() []uint32 { if x != nil { - return x.RemoteIPs + return x.RemotePorts } return nil } @@ -126,17 +117,16 @@ var File_transport_internet_finalmask_udphop_config_proto protoreflect.FileDescr const file_transport_internet_finalmask_udphop_config_proto_rawDesc = "" + "\n" + - "0transport/internet/finalmask/udphop/config.proto\x12(xray.transport.internet.finalmask.udphop\x1a\x1ftransport/internet/config.proto\"\x9f\x02\n" + - "\x06Config\x12?\n" + - "\asockopt\x18\x01 \x01(\v2%.xray.transport.internet.SocketConfigR\asockopt\x12\x14\n" + + "0transport/internet/finalmask/udphop/config.proto\x12(xray.transport.internet.finalmask.udphop\"\xe4\x01\n" + + "\x06Config\x12\x14\n" + "\x05local\x18\x02 \x01(\bR\x05local\x12\x16\n" + "\x06remote\x18\x03 \x01(\bR\x06remote\x12\x1f\n" + "\vremote_once\x18\x04 \x01(\bR\n" + "remoteOnce\x12!\n" + "\finterval_min\x18\x05 \x01(\x03R\vintervalMin\x12!\n" + - "\finterval_max\x18\x06 \x01(\x03R\vintervalMax\x12!\n" + - "\fremote_ports\x18\a \x03(\rR\vremotePorts\x12\x1c\n" + - "\tremoteIPs\x18\b \x03(\tR\tremoteIPsB\x9a\x01\n" + + "\finterval_max\x18\x06 \x01(\x03R\vintervalMax\x12\x1c\n" + + "\tremoteIPs\x18\a \x03(\tR\tremoteIPs\x12!\n" + + "\fremote_ports\x18\b \x03(\rR\vremotePortsJ\x04\b\x01\x10\x02B\x9a\x01\n" + ",com.xray.transport.internet.finalmask.udphopP\x01Z=github.com/xtls/xray-core/transport/internet/finalmask/udphop\xaa\x02(Xray.Transport.Internet.Finalmask.Udphopb\x06proto3" var ( @@ -153,16 +143,14 @@ func file_transport_internet_finalmask_udphop_config_proto_rawDescGZIP() []byte var file_transport_internet_finalmask_udphop_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1) var file_transport_internet_finalmask_udphop_config_proto_goTypes = []any{ - (*Config)(nil), // 0: xray.transport.internet.finalmask.udphop.Config - (*internet.SocketConfig)(nil), // 1: xray.transport.internet.SocketConfig + (*Config)(nil), // 0: xray.transport.internet.finalmask.udphop.Config } var file_transport_internet_finalmask_udphop_config_proto_depIdxs = []int32{ - 1, // 0: xray.transport.internet.finalmask.udphop.Config.sockopt:type_name -> xray.transport.internet.SocketConfig - 1, // [1:1] is the sub-list for method output_type - 1, // [1:1] is the sub-list for method input_type - 1, // [1:1] is the sub-list for extension type_name - 1, // [1:1] is the sub-list for extension extendee - 0, // [0:1] is the sub-list for field type_name + 0, // [0:0] is the sub-list for method output_type + 0, // [0:0] is the sub-list for method input_type + 0, // [0:0] is the sub-list for extension type_name + 0, // [0:0] is the sub-list for extension extendee + 0, // [0:0] is the sub-list for field type_name } func init() { file_transport_internet_finalmask_udphop_config_proto_init() } diff --git a/transport/internet/finalmask/udphop/config.proto b/transport/internet/finalmask/udphop/config.proto index 48b8b1da6..2556c6a45 100644 --- a/transport/internet/finalmask/udphop/config.proto +++ b/transport/internet/finalmask/udphop/config.proto @@ -6,16 +6,14 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/udph option java_package = "com.xray.transport.internet.finalmask.udphop"; option java_multiple_files = true; -import "transport/internet/config.proto"; - message Config { - xray.transport.internet.SocketConfig sockopt = 1; + reserved 1; bool local = 2; bool remote = 3; bool remote_once = 4; int64 interval_min = 5; int64 interval_max = 6; - repeated uint32 remote_ports = 7; - repeated string remoteIPs = 8; + repeated string remoteIPs = 7; + repeated uint32 remote_ports = 8; } diff --git a/transport/internet/finalmask/udphop/conn.go b/transport/internet/finalmask/udphop/conn.go index 13c3288df..a71a40439 100644 --- a/transport/internet/finalmask/udphop/conn.go +++ b/transport/internet/finalmask/udphop/conn.go @@ -6,9 +6,7 @@ import ( goerrors "errors" "io" mrand "math/rand" - gonet "net" "net/netip" - "reflect" "sync" "time" @@ -16,8 +14,6 @@ import ( "github.com/xtls/xray-core/common/crypto" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/net" - "github.com/xtls/xray-core/common/net/cnc" - "github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet/finalmask" ) @@ -34,16 +30,14 @@ type packet struct { } type udpHopConn struct { - conn net.PacketConn - sockopt *internet.SocketConfig - local bool - remote bool - remoteOnce bool + dialer *finalmask.Dialer + local bool + remote bool intervalMin int64 intervalMax int64 - remotePorts []uint32 remoteIPs []netip.Prefix + remotePorts []uint32 deadline time.Time readDeadline time.Time @@ -55,10 +49,10 @@ type udpHopConn struct { readCh chan packet closeCh chan struct{} wg sync.WaitGroup - mu sync.Mutex + mu sync.RWMutex } -func NewUDPHopConn(c *Config, raw net.PacketConn) (net.PacketConn, error) { +func NewUDPHopConn(c *Config, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { if c.IntervalMin < 5 || c.IntervalMax < 5 { return nil, errors.New("invalid interval") } @@ -66,22 +60,40 @@ func NewUDPHopConn(c *Config, raw net.PacketConn) (net.PacketConn, error) { for _, ip := range c.RemoteIPs { remoteIPs = append(remoteIPs, netip.MustParsePrefix(ip)) } - conn := &udpHopConn{ - conn: raw, - sockopt: c.Sockopt, - local: c.Local, - remote: c.Remote, - remoteOnce: c.RemoteOnce, + remotePorts := c.RemotePorts + if c.Remote || c.RemoteOnce { + if len(remoteIPs) > 0 { + dest.Address = net.IPAddress(randPrefix(remoteIPs[mrand.Intn(len(remoteIPs))])) + } + if len(remotePorts) > 0 { + dest.Port = net.Port(remotePorts[mrand.Intn(len(remotePorts))]) + } + } + conn, err := dialer.DialUDP(*dest) + if err != nil { + return nil, err + } + cur := conn.(*finalmask.PacketConnWrapper).PacketConn + addr := conn.RemoteAddr().(*net.UDPAddr) + client := &udpHopConn{ + dialer: dialer, + local: c.Local, + remote: c.Remote, intervalMin: c.IntervalMin, intervalMax: c.IntervalMax, - remotePorts: c.RemotePorts, remoteIPs: remoteIPs, + remotePorts: remotePorts, + cur: cur, + addr: addr, readCh: make(chan packet), closeCh: make(chan struct{}), } - return conn, nil + go client.run() + client.wg.Add(1) + go client.recv(client.cur) + return client, nil } func (c *udpHopConn) closed() bool { @@ -93,61 +105,67 @@ func (c *udpHopConn) closed() bool { } } -func (c *udpHopConn) hop(addr *net.UDPAddr) { +func (c *udpHopConn) run() { + ticker := time.NewTicker(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1))) + defer ticker.Stop() + for { + select { + case <-c.closeCh: + return + case <-ticker.C: + ticker.Reset(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1))) + c.hop() + } + } +} + +func (c *udpHopConn) hop() { + c.mu.Lock() + defer c.mu.Unlock() if c.closed() { return } - newAddr := &net.UDPAddr{IP: addr.IP, Port: addr.Port} - newConn := c.conn - if c.remote || c.remoteOnce && c.addr == nil { - if len(c.remotePorts) > 0 { - newAddr.Port = int(c.remotePorts[mrand.Intn(len(c.remotePorts))]) - } + oldIP := c.addr.IP + oldPort := c.addr.Port + if c.remote { if len(c.remoteIPs) > 0 { - newAddr.IP = randPrefix(c.remoteIPs[mrand.Intn(len(c.remoteIPs))]) + c.addr.IP = randPrefix(c.remoteIPs[mrand.Intn(len(c.remoteIPs))]) + } + if len(c.remotePorts) > 0 { + c.addr.Port = int(c.remotePorts[mrand.Intn(len(c.remotePorts))]) } } if c.local { - raw, err := internet.DialSystem(context.Background(), net.UDPDestination(net.IPAddress(newAddr.IP), net.Port(newAddr.Port)), c.sockopt) + conn, err := c.dialer.DialUDP(net.UDPDestination(net.IPAddress(c.addr.IP), net.Port(c.addr.Port))) if err != nil { + c.addr.IP = oldIP + c.addr.Port = oldPort errors.LogErrorInner(context.Background(), err, "hop err") return } - switch c := raw.(type) { - case *internet.PacketConnWrapper: - newConn = c.PacketConn - case *cnc.Connection: - newConn = &internet.FakePacketConn{Conn: c} - default: - panic(reflect.TypeOf(c)) - } - newConn.SetDeadline(c.deadline) - newConn.SetReadDeadline(c.readDeadline) - newConn.SetWriteDeadline(c.writeDeadline) + conn.SetDeadline(c.deadline) + conn.SetReadDeadline(c.readDeadline) + conn.SetWriteDeadline(c.writeDeadline) if c.pre != nil { _ = c.pre.Close() } c.pre = c.cur + c.cur = conn.(*finalmask.PacketConnWrapper).PacketConn c.wg.Add(1) - go c.recv(newConn) + go c.recv(c.cur) } - c.addr = newAddr - c.cur = newConn } func (c *udpHopConn) recv(conn net.PacketConn) { defer c.wg.Done() for { - if c.closed() { - return - } p := pool.Get().([]byte) n, addr, err := conn.ReadFrom(p) if err != nil { pool.Put(p[:cap(p)]) - if goerrors.Is(err, io.EOF) || goerrors.Is(err, io.ErrClosedPipe) || goerrors.Is(err, gonet.ErrClosed) { - break + if c.closed() { + return } var netErr net.Error if goerrors.As(err, &netErr) && netErr.Timeout() { @@ -156,9 +174,10 @@ func (c *udpHopConn) recv(conn net.PacketConn) { case <-c.closeCh: return } + continue } errors.LogErrorInner(context.Background(), err, "recv err") - continue + return } select { case c.readCh <- packet{p: p[:n], addr: addr}: @@ -169,22 +188,6 @@ func (c *udpHopConn) recv(conn net.PacketConn) { } } -func (c *udpHopConn) hopLoop() { - ticker := time.NewTicker(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1))) - defer ticker.Stop() - for { - select { - case <-ticker.C: - ticker.Reset(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1))) - c.mu.Lock() - c.hop(c.addr) - c.mu.Unlock() - case <-c.closeCh: - return - } - } -} - func (c *udpHopConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { packet, ok := <-c.readCh if ok { @@ -194,21 +197,12 @@ func (c *udpHopConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { } return n, packet.addr, packet.err } - return 0, nil, io.EOF + return 0, nil, io.ErrClosedPipe } func (c *udpHopConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { - c.mu.Lock() - defer c.mu.Unlock() - - if c.cur == nil { - c.hop(addr.(*net.UDPAddr)) - if c.cur == nil { - return 0, nil - } - go c.hopLoop() - } - + c.mu.RLock() + defer c.mu.RUnlock() _, err = c.cur.WriteTo(p, c.addr) if err != nil { errors.LogErrorInner(context.Background(), err, "send err") @@ -227,15 +221,12 @@ func (c *udpHopConn) Close() error { if c.pre != nil { _ = c.pre.Close() } - if c.cur != nil { - _ = c.cur.Close() - } - _ = c.conn.Close() + _ = c.cur.Close() c.wg.Wait() select { - case p := <-c.readCh: - if p.p != nil { - pool.Put(p.p[:cap(p.p)]) + case packet := <-c.readCh: + if packet.p != nil { + pool.Put(packet.p[:cap(packet.p)]) } default: } @@ -244,7 +235,9 @@ func (c *udpHopConn) Close() error { } func (c *udpHopConn) LocalAddr() net.Addr { - return c.conn.LocalAddr() + c.mu.RLock() + defer c.mu.RUnlock() + return c.cur.LocalAddr() } func (c *udpHopConn) SetDeadline(t time.Time) error { @@ -254,10 +247,7 @@ func (c *udpHopConn) SetDeadline(t time.Time) error { if c.pre != nil { _ = c.pre.SetDeadline(t) } - if c.cur != nil { - _ = c.cur.SetDeadline(t) - } - return nil + return c.cur.SetDeadline(t) } func (c *udpHopConn) SetReadDeadline(t time.Time) error { @@ -267,10 +257,7 @@ func (c *udpHopConn) SetReadDeadline(t time.Time) error { if c.pre != nil { _ = c.pre.SetReadDeadline(t) } - if c.cur != nil { - _ = c.cur.SetReadDeadline(t) - } - return nil + return c.cur.SetReadDeadline(t) } func (c *udpHopConn) SetWriteDeadline(t time.Time) error { @@ -280,10 +267,7 @@ func (c *udpHopConn) SetWriteDeadline(t time.Time) error { if c.pre != nil { _ = c.pre.SetWriteDeadline(t) } - if c.cur != nil { - _ = c.cur.SetWriteDeadline(t) - } - return nil + return c.cur.SetWriteDeadline(t) } func randPrefix(p netip.Prefix) []byte { diff --git a/transport/internet/finalmask/xdns/config.go b/transport/internet/finalmask/xdns/config.go index 46241476e..7bae597ab 100644 --- a/transport/internet/finalmask/xdns/config.go +++ b/transport/internet/finalmask/xdns/config.go @@ -1,21 +1,14 @@ package xdns import ( - "net" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) -func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - // _, ok1 := raw.(*internet.FakePacketConn) - // _, ok2 := raw.(*udphop.UdpHopPacketConn) - // if level != 0 || ok1 || ok2 { - // return nil, errors.New("xdns requires being at the outermost level") - // } - return NewConnClient(c, raw) +func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + return NewConnClient(c, conn) } -func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - // if level != 0 { - // return nil, errors.New("xdns requires being at the outermost level") - // } - return NewConnServer(c, raw) +func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewConnServer(c, conn) } diff --git a/transport/internet/finalmask/xicmp/client.go b/transport/internet/finalmask/xicmp/client.go index 93db4e307..cabdb1942 100644 --- a/transport/internet/finalmask/xicmp/client.go +++ b/transport/internet/finalmask/xicmp/client.go @@ -8,8 +8,7 @@ import ( goerrors "errors" "fmt" "io" - mathrand "math/rand" - "net" + mrand "math/rand" "net/netip" "sync" "time" @@ -17,6 +16,7 @@ import ( "github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/transport/internet/finalmask" "golang.org/x/net/icmp" "golang.org/x/net/ipv4" @@ -36,11 +36,11 @@ type packet struct { } type xicmpConnClient struct { - conn net.PacketConn icmp4 *icmp.PacketConn icmp6 *icmp.PacketConn udp bool ips []netip.Addr + ip net.IP clientID [8]byte id int seq int @@ -50,7 +50,7 @@ type xicmpConnClient struct { mu sync.Mutex } -func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) { +func NewConnClient(c *Config, dest *net.Destination) (net.PacketConn, error) { var icmp4, icmp6 *icmp.PacketConn var err4, err6 error if c.DGRAM { @@ -69,17 +69,24 @@ func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) { ips = append(ips, netip.MustParseAddr(ip)) } + var ip net.IP + if len(ips) > 0 { + ip = ips[mrand.Intn(len(ips))].AsSlice() + } else { + ip = dest.Address.IP() + } + var clientID [8]byte common.Must2(rand.Read(clientID[:])) conn := &xicmpConnClient{ - conn: raw, icmp4: icmp4, icmp6: icmp6, udp: c.DGRAM, ips: ips, + ip: ip, clientID: clientID, - id: mathrand.Intn(65536), + id: mrand.Intn(65536), seq: 1, readCh: make(chan packet), closeCh: make(chan struct{}), @@ -92,10 +99,6 @@ func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) { return conn, nil } -func (c *xicmpConnClient) ring(a, b uint16) uint16 { - return min(a-b, b-a) -} - func (c *xicmpConnClient) closed() bool { select { case <-c.closeCh: @@ -110,12 +113,11 @@ func (c *xicmpConnClient) recv4() { var b [finalmask.UDPSize]byte for { - if c.closed() { - return - } - n, addr, err := c.icmp4.ReadFrom(b[:]) if err != nil { + if c.closed() { + return + } var netErr net.Error if goerrors.As(err, &netErr) && netErr.Timeout() { select { @@ -125,9 +127,10 @@ func (c *xicmpConnClient) recv4() { case <-c.closeCh: return } + continue } - errors.LogErrorInner(context.Background(), err, "recv4 err") - continue + errors.LogErrorInner(context.Background(), err, "recv err 4") + return } msg, err := icmp.ParseMessage(1, b[:n]) @@ -150,10 +153,6 @@ func (c *xicmpConnClient) recv4() { continue } - if c.ring(uint16(echo.Seq), uint16(c.seq)) > 1000 { - continue - } - if len(echo.Data) > 8 && bytes.Equal(echo.Data[:8], c.clientID[:]) { continue } @@ -182,12 +181,11 @@ func (c *xicmpConnClient) recv6() { var b [finalmask.UDPSize]byte for { - if c.closed() { - return - } - n, addr, err := c.icmp6.ReadFrom(b[:]) if err != nil { + if c.closed() { + return + } var netErr net.Error if goerrors.As(err, &netErr) && netErr.Timeout() { select { @@ -197,9 +195,10 @@ func (c *xicmpConnClient) recv6() { case <-c.closeCh: return } + continue } - errors.LogErrorInner(context.Background(), err, "recv6 err") - continue + errors.LogErrorInner(context.Background(), err, "recv err 6") + return } msg, err := icmp.ParseMessage(58, b[:n]) @@ -222,10 +221,6 @@ func (c *xicmpConnClient) recv6() { continue } - if c.ring(uint16(echo.Seq), uint16(c.seq)) > 1000 { - continue - } - if len(echo.Data) > 8 && bytes.Equal(echo.Data[:8], c.clientID[:]) { continue } @@ -273,9 +268,9 @@ func (c *xicmpConnClient) WriteTo(p []byte, addr net.Addr) (n int, err error) { c.seq %= 65536 c.mu.Unlock() - ip := addr.(*net.UDPAddr).IP + ip := c.ip if len(c.ips) > 0 { - ip = c.ips[mathrand.Intn(len(c.ips))].AsSlice() + ip = c.ips[mrand.Intn(len(c.ips))].AsSlice() } if c.udp { @@ -314,7 +309,6 @@ func (c *xicmpConnClient) Close() error { close(c.closeCh) _ = c.icmp4.Close() _ = c.icmp6.Close() - _ = c.conn.Close() c.wg.Wait() select { case p := <-c.readCh: @@ -328,7 +322,7 @@ func (c *xicmpConnClient) Close() error { } func (c *xicmpConnClient) LocalAddr() net.Addr { - return c.conn.LocalAddr() + return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} } func (c *xicmpConnClient) SetDeadline(t time.Time) error { diff --git a/transport/internet/finalmask/xicmp/config.go b/transport/internet/finalmask/xicmp/config.go index f99784c1a..cd4d0d00b 100644 --- a/transport/internet/finalmask/xicmp/config.go +++ b/transport/internet/finalmask/xicmp/config.go @@ -1,23 +1,23 @@ package xicmp import ( - "net" + "errors" - "github.com/xtls/xray-core/common/errors" - "github.com/xtls/xray-core/transport/internet" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) -func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - _, ok1 := raw.(*internet.FakePacketConn) - if level != 0 || ok1 { - return nil, errors.New("xicmp requires being at the outermost level") +func (c *Config) HandleDial() {} + +func (c *Config) HandleListen() {} + +func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { + if dest.Address.Family().IsDomain() && len(c.IPs) == 0 { + return nil, errors.New("empty ip addresses") } - return NewConnClient(c, raw) + return NewConnClient(c, dest) } -func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) { - if level != 0 { - return nil, errors.New("xicmp requires being at the outermost level") - } - return NewConnServer(c, raw) +func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { + return NewConnServer(c) } diff --git a/transport/internet/finalmask/xicmp/server.go b/transport/internet/finalmask/xicmp/server.go index f3fb429ba..117638923 100644 --- a/transport/internet/finalmask/xicmp/server.go +++ b/transport/internet/finalmask/xicmp/server.go @@ -37,7 +37,6 @@ type record struct { } type xicmpConnServer struct { - conn net.PacketConn icmp4 *icmp.PacketConn icmp6 *icmp.PacketConn ips map[netip.Addr]struct{} @@ -48,7 +47,7 @@ type xicmpConnServer struct { mu sync.Mutex } -func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) { +func NewConnServer(c *Config) (net.PacketConn, error) { icmp4, err := icmp.ListenPacket("ip4:icmp", "0.0.0.0") if err != nil { return nil, err @@ -64,7 +63,6 @@ func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) { } conn := &xicmpConnServer{ - conn: raw, icmp4: icmp4, icmp6: icmp6, ips: ips, @@ -115,12 +113,11 @@ func (c *xicmpConnServer) recv4() { var b [finalmask.UDPSize]byte for { - if c.closed() { - return - } - n, addr, err := c.icmp4.ReadFrom(b[:]) if err != nil { + if c.closed() { + return + } var netErr net.Error if goerrors.As(err, &netErr) && netErr.Timeout() { select { @@ -130,9 +127,10 @@ func (c *xicmpConnServer) recv4() { case <-c.closeCh: return } + continue } - errors.LogErrorInner(context.Background(), err, "recv4 err") - continue + errors.LogErrorInner(context.Background(), err, "recv err 4") + return } msg, err := icmp.ParseMessage(1, b[:n]) @@ -195,12 +193,11 @@ func (c *xicmpConnServer) recv6() { var b [finalmask.UDPSize]byte for { - if c.closed() { - return - } - n, addr, err := c.icmp6.ReadFrom(b[:]) if err != nil { + if c.closed() { + return + } var netErr net.Error if goerrors.As(err, &netErr) && netErr.Timeout() { select { @@ -210,9 +207,10 @@ func (c *xicmpConnServer) recv6() { case <-c.closeCh: return } + continue } - errors.LogErrorInner(context.Background(), err, "recv6 err") - continue + errors.LogErrorInner(context.Background(), err, "recv err 6") + return } msg, err := icmp.ParseMessage(58, b[:n]) @@ -330,7 +328,6 @@ func (c *xicmpConnServer) Close() error { close(c.closeCh) _ = c.icmp4.Close() _ = c.icmp6.Close() - _ = c.conn.Close() c.wg.Wait() select { case p := <-c.readCh: @@ -344,7 +341,7 @@ func (c *xicmpConnServer) Close() error { } func (c *xicmpConnServer) LocalAddr() net.Addr { - return c.conn.LocalAddr() + return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} } func (c *xicmpConnServer) SetDeadline(t time.Time) error { diff --git a/transport/internet/finalmask/xicmp/server_oob.go b/transport/internet/finalmask/xicmp/server_oob.go index d8bdb93c3..8d1c26db8 100644 --- a/transport/internet/finalmask/xicmp/server_oob.go +++ b/transport/internet/finalmask/xicmp/server_oob.go @@ -39,7 +39,6 @@ type record struct { } type xicmpConnServer struct { - conn net.PacketConn icmp4 *icmp.PacketConn icmp6 *icmp.PacketConn ipv4PC *ipv4.PacketConn @@ -52,7 +51,7 @@ type xicmpConnServer struct { mu sync.Mutex } -func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) { +func NewConnServer(c *Config) (net.PacketConn, error) { icmp4, err := icmp.ListenPacket("ip4:icmp", "0.0.0.0") if err != nil { return nil, err @@ -68,7 +67,6 @@ func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) { } conn := &xicmpConnServer{ - conn: raw, icmp4: icmp4, icmp6: icmp6, ipv4PC: icmp4.IPv4PacketConn(), @@ -124,12 +122,11 @@ func (c *xicmpConnServer) recv4() { var b [finalmask.UDPSize]byte for { - if c.closed() { - return - } - n, cm, addr, err := c.ipv4PC.ReadFrom(b[:]) if err != nil { + if c.closed() { + return + } var netErr net.Error if goerrors.As(err, &netErr) && netErr.Timeout() { select { @@ -139,9 +136,10 @@ func (c *xicmpConnServer) recv4() { case <-c.closeCh: return } + continue } - errors.LogErrorInner(context.Background(), err, "recv4 err") - continue + errors.LogErrorInner(context.Background(), err, "recv err 4") + return } msg, err := icmp.ParseMessage(1, b[:n]) @@ -205,12 +203,11 @@ func (c *xicmpConnServer) recv6() { var b [finalmask.UDPSize]byte for { - if c.closed() { - return - } - n, cm, addr, err := c.ipv6PC.ReadFrom(b[:]) if err != nil { + if c.closed() { + return + } var netErr net.Error if goerrors.As(err, &netErr) && netErr.Timeout() { select { @@ -220,9 +217,10 @@ func (c *xicmpConnServer) recv6() { case <-c.closeCh: return } + continue } - errors.LogErrorInner(context.Background(), err, "recv6 err") - continue + errors.LogErrorInner(context.Background(), err, "recv err 6") + return } msg, err := icmp.ParseMessage(58, b[:n]) @@ -341,7 +339,6 @@ func (c *xicmpConnServer) Close() error { close(c.closeCh) _ = c.icmp4.Close() _ = c.icmp6.Close() - _ = c.conn.Close() c.wg.Wait() select { case p := <-c.readCh: @@ -355,7 +352,7 @@ func (c *xicmpConnServer) Close() error { } func (c *xicmpConnServer) LocalAddr() net.Addr { - return c.conn.LocalAddr() + return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} } func (c *xicmpConnServer) SetDeadline(t time.Time) error { diff --git a/transport/internet/finalmask/xmc/config.go b/transport/internet/finalmask/xmc/config.go index 863eb97ba..8ee16b2e4 100644 --- a/transport/internet/finalmask/xmc/config.go +++ b/transport/internet/finalmask/xmc/config.go @@ -2,10 +2,12 @@ package xmc import ( "fmt" - "net" + + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" ) -func (c *Config) WrapConnClient(conn net.Conn) (net.Conn, error) { +func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) { profiles, err := profilesFromConfig(c.Profiles) if err != nil { return nil, fmt.Errorf("minecraft finalmask: %w", err) diff --git a/transport/internet/grpc/dial.go b/transport/internet/grpc/dial.go index a54d7c02f..a7d57ed77 100644 --- a/transport/internet/grpc/dial.go +++ b/transport/internet/grpc/dial.go @@ -83,7 +83,6 @@ func getGrpcClient(ctx context.Context, dest net.Destination, streamSettings *in } tlsConfig := tls.ConfigFromStreamSettings(streamSettings) realityConfig := reality.ConfigFromStreamSettings(streamSettings) - sockopt := streamSettings.SocketSettings grpcSettings := streamSettings.ProtocolSettings.(*Config) if client, found := globalDialerMap[dialerConf{dest, streamSettings}]; found && client.GetState() != connectivity.Shutdown { @@ -124,17 +123,13 @@ func getGrpcClient(ctx context.Context, dest net.Destination, streamSettings *in gctx = session.ContextWithOutbounds(gctx, session.OutboundsFromContext(ctx)) gctx = session.ContextWithTimeoutOnly(gctx, true) - c, err := internet.DialSystem(gctx, net.TCPDestination(address, port), sockopt) + var c net.Conn + if streamSettings.FinalMask != nil { + c, err = streamSettings.FinalMask.DialTCP(gctx, net.TCPDestination(address, port)) + } else { + c, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings) + } if err == nil { - if streamSettings.TcpmaskManager != nil { - newConn, err := streamSettings.TcpmaskManager.WrapConnClient(c) - if err != nil { - c.Close() - return nil, errors.New("mask err").Base(err) - } - c = newConn - } - if tlsConfig != nil { config := tlsConfig.GetTLSConfig(tls.WithDestination(dest)) if fingerprint := tls.GetFingerprint(tlsConfig.Fingerprint); fingerprint != nil { diff --git a/transport/internet/grpc/hub.go b/transport/internet/grpc/hub.go index 91bd2ab60..d8c35bc43 100644 --- a/transport/internet/grpc/hub.go +++ b/transport/internet/grpc/hub.go @@ -104,28 +104,20 @@ func Listen(ctx context.Context, address net.Address, port net.Port, settings *i go func() { var streamListener net.Listener var err error + var addr net.Addr if port == net.Port(0) { // unix - streamListener, err = internet.ListenSystem(ctx, &net.UnixAddr{ - Name: address.Domain(), - Net: "unix", - }, settings.SocketSettings) - if err != nil { - errors.LogErrorInner(ctx, err, "failed to listen on ", address) - return - } + addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"} } else { // tcp - streamListener, err = internet.ListenSystem(ctx, &net.TCPAddr{ - IP: address.IP(), - Port: int(port), - }, settings.SocketSettings) - if err != nil { - errors.LogErrorInner(ctx, err, "failed to listen on ", address, ":", port) - return - } + addr = &net.TCPAddr{IP: address.IP(), Port: int(port)} } - - if settings.TcpmaskManager != nil { - streamListener, _ = settings.TcpmaskManager.WrapListener(streamListener) + if settings.FinalMask != nil { + streamListener, err = settings.FinalMask.Listen(ctx, addr) + } else { + streamListener, err = internet.ListenSystem(ctx, addr, settings.SocketSettings) + } + if err != nil { + errors.LogErrorInner(ctx, err, "failed to listen on ", address, ":", port) + return } errors.LogDebug(ctx, "gRPC listen for service name `"+grpcSettings.getServiceName()+"` tun `"+grpcSettings.getTunStreamName()+"` multi tun `"+grpcSettings.getTunMultiStreamName()+"`") diff --git a/transport/internet/httpupgrade/dialer.go b/transport/internet/httpupgrade/dialer.go index 571797f61..d05ae8f47 100644 --- a/transport/internet/httpupgrade/dialer.go +++ b/transport/internet/httpupgrade/dialer.go @@ -46,21 +46,18 @@ func (c *ConnRF) Read(b []byte) (int, error) { func dialhttpUpgrade(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (net.Conn, error) { transportConfiguration := streamSettings.ProtocolSettings.(*Config) - pconn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) + var pconn net.Conn + var err error + if streamSettings.FinalMask != nil { + pconn, err = streamSettings.FinalMask.DialTCP(ctx, dest) + } else { + pconn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings) + } if err != nil { errors.LogErrorInner(ctx, err, "failed to dial to ", dest) return nil, err } - if streamSettings.TcpmaskManager != nil { - newConn, err := streamSettings.TcpmaskManager.WrapConnClient(pconn) - if err != nil { - pconn.Close() - return nil, errors.New("mask err").Base(err) - } - pconn = newConn - } - var conn net.Conn var requestURL url.URL tConfig := tls.ConfigFromStreamSettings(streamSettings) diff --git a/transport/internet/httpupgrade/hub.go b/transport/internet/httpupgrade/hub.go index cbf6a0d47..9a6429447 100644 --- a/transport/internet/httpupgrade/hub.go +++ b/transport/internet/httpupgrade/hub.go @@ -124,29 +124,21 @@ func ListenHTTPUpgrade(ctx context.Context, address net.Address, port net.Port, } var listener net.Listener var err error + var addr net.Addr if port == net.Port(0) { // unix - listener, err = internet.ListenSystem(ctx, &net.UnixAddr{ - Name: address.Domain(), - Net: "unix", - }, streamSettings.SocketSettings) - if err != nil { - return nil, errors.New("failed to listen unix domain socket(for HttpUpgrade) on ", address).Base(err) - } - errors.LogInfo(ctx, "listening unix domain socket(for HttpUpgrade) on ", address) + addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"} } else { // tcp - listener, err = internet.ListenSystem(ctx, &net.TCPAddr{ - IP: address.IP(), - Port: int(port), - }, streamSettings.SocketSettings) - if err != nil { - return nil, errors.New("failed to listen TCP(for HttpUpgrade) on ", address, ":", port).Base(err) - } - errors.LogInfo(ctx, "listening TCP(for HttpUpgrade) on ", address, ":", port) + addr = &net.TCPAddr{IP: address.IP(), Port: int(port)} } - - if streamSettings.TcpmaskManager != nil { - listener, _ = streamSettings.TcpmaskManager.WrapListener(listener) + if streamSettings.FinalMask != nil { + listener, err = streamSettings.FinalMask.Listen(ctx, addr) + } else { + listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings) } + if err != nil { + return nil, errors.New("failed to listen ", addr.Network(), "(for HttpUpgrade) on ", address, ":", port).Base(err) + } + errors.LogInfo(ctx, "listening ", addr.Network(), "(for HttpUpgrade) on ", address, ":", port) if streamSettings.SocketSettings != nil && streamSettings.SocketSettings.AcceptProxyProtocol { errors.LogWarning(ctx, "accepting PROXY protocol") diff --git a/transport/internet/hysteria/dialer.go b/transport/internet/hysteria/dialer.go index 5a32f8160..f37cdd11c 100644 --- a/transport/internet/hysteria/dialer.go +++ b/transport/internet/hysteria/dialer.go @@ -2,7 +2,7 @@ package hysteria import ( "context" - go_tls "crypto/tls" + gotls "crypto/tls" "net/http" "net/url" "reflect" @@ -28,12 +28,12 @@ import ( type client struct { sync.Mutex - dest net.Destination - config *Config - tlsConfig *go_tls.Config - socketConfig *internet.SocketConfig - udpmaskManager *finalmask.UdpmaskManager - quicParams *internet.QuicParams + dest net.Destination + config *Config + tlsConfig *gotls.Config + socketConfig *internet.SocketConfig + finalMask *finalmask.FinalMask + quicParams *internet.QuicParams conn *quic.Conn tr *quic.Transport @@ -113,30 +113,29 @@ func (c *client) dial(ctx context.Context) error { // } var pktConn net.PacketConn - var udpAddr *net.UDPAddr - - raw, err := internet.DialSystem(ctx, c.dest, c.socketConfig) - if err != nil { - return errors.New("failed to dial to dest").Base(err) - } - switch c := raw.(type) { - case *internet.PacketConnWrapper: - pktConn = c.PacketConn - udpAddr = raw.RemoteAddr().(*net.UDPAddr) - case *cnc.Connection: - pktConn = &internet.FakePacketConn{Conn: c} - udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port} - default: - panic(reflect.TypeOf(c)) - } - - if c.udpmaskManager != nil { - newConn, err := c.udpmaskManager.WrapPacketConnClient(pktConn) + var udpAddr net.Addr + if c.finalMask != nil { + conn, err := c.finalMask.DialUDP(ctx, c.dest) if err != nil { - pktConn.Close() - return errors.New("mask err").Base(err) + return errors.New("failed to dial to dest").Base(err) + } + pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn + udpAddr = conn.RemoteAddr() + } else { + conn, err := internet.DialSystem(ctx, c.dest, c.socketConfig) + if err != nil { + return errors.New("failed to dial to dest").Base(err) + } + switch c := conn.(type) { + case *internet.PacketConnWrapper: + pktConn = c.PacketConn + udpAddr = c.RemoteAddr() + case *cnc.Connection: + pktConn = &internet.FakePacketConn{Conn: c} + udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}} + default: + panic(reflect.TypeOf(c)) } - pktConn = newConn } tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO} @@ -150,7 +149,7 @@ func (c *client) dial(ctx context.Context) error { rt := &http3.Transport{ TLSClientConfig: c.tlsConfig, QUICConfig: quicConfig, - Dial: func(ctx context.Context, _ string, tlsCfg *go_tls.Config, cfg *quic.Config) (*quic.Conn, error) { + Dial: func(ctx context.Context, _ string, tlsCfg *gotls.Config, cfg *quic.Config) (*quic.Conn, error) { qc, err := tr.DialEarly(ctx, udpAddr, tlsCfg, cfg) if err != nil { return nil, err @@ -316,12 +315,12 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me c = manager.m[dialerConf{dest, streamSettings}] if c == nil { c = &client{ - dest: dest, - config: streamSettings.ProtocolSettings.(*Config), - tlsConfig: tlsConfig.GetTLSConfig(tls.WithDestination(dest)), - socketConfig: streamSettings.SocketSettings, - udpmaskManager: streamSettings.UdpmaskManager, - quicParams: streamSettings.QuicParams, + dest: dest, + config: streamSettings.ProtocolSettings.(*Config), + tlsConfig: tlsConfig.GetTLSConfig(tls.WithDestination(dest)), + socketConfig: streamSettings.SocketSettings, + finalMask: streamSettings.FinalMask, + quicParams: streamSettings.QuicParams, } manager.m[dialerConf{dest, streamSettings}] = c } diff --git a/transport/internet/hysteria/hub.go b/transport/internet/hysteria/hub.go index b95fc621b..0b3f449cd 100644 --- a/transport/internet/hysteria/hub.go +++ b/transport/internet/hysteria/hub.go @@ -316,20 +316,17 @@ func Listen(ctx context.Context, address net.Address, port net.Port, streamSetti quicConfig.MaxIncomingStreams = 1024 } - pktConn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings) + var pktConn net.PacketConn + var err error + if streamSettings.FinalMask != nil { + pktConn, err = streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}) + } else { + pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings) + } if err != nil { return nil, err } - if streamSettings.UdpmaskManager != nil { - newConn, err := streamSettings.UdpmaskManager.WrapPacketConnServer(pktConn) - if err != nil { - pktConn.Close() - return nil, errors.New("mask err").Base(err) - } - pktConn = newConn - } - var k *quic.StatelessResetKey if !quicParams.DisableStatelessReset { k = &quic.StatelessResetKey{} diff --git a/transport/internet/kcp/dialer.go b/transport/internet/kcp/dialer.go index 175998ec7..8ec4d1973 100644 --- a/transport/internet/kcp/dialer.go +++ b/transport/internet/kcp/dialer.go @@ -3,7 +3,6 @@ package kcp import ( "context" "io" - reflect "reflect" "sync/atomic" "github.com/xtls/xray-core/common" @@ -11,7 +10,6 @@ import ( "github.com/xtls/xray-core/common/dice" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/net" - "github.com/xtls/xray-core/common/net/cnc" "github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet/stat" "github.com/xtls/xray-core/transport/internet/tls" @@ -51,36 +49,17 @@ func DialKCP(ctx context.Context, dest net.Destination, streamSettings *internet dest.Network = net.Network_UDP errors.LogInfo(ctx, "dialing mKCP to ", dest) - conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) + var conn net.Conn + var err error + if streamSettings.FinalMask != nil { + conn, err = streamSettings.FinalMask.DialUDP(ctx, dest) + } else { + conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings) + } if err != nil { return nil, errors.New("failed to dial to dest: ", err).AtWarning().Base(err) } - if streamSettings.UdpmaskManager != nil { - var pktConn net.PacketConn - var udpAddr *net.UDPAddr - switch c := conn.(type) { - case *internet.PacketConnWrapper: - pktConn = c.PacketConn - udpAddr = c.RemoteAddr().(*net.UDPAddr) - case *cnc.Connection: - pktConn = &internet.FakePacketConn{Conn: c} - udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port} - default: - panic(reflect.TypeOf(c)) - } - newConn, err := streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn) - if err != nil { - pktConn.Close() - return nil, errors.New("mask err").Base(err) - } - pktConn = newConn - conn = &internet.PacketConnWrapper{ - PacketConn: pktConn, - Dest: udpAddr, - } - } - kcpSettings := streamSettings.ProtocolSettings.(*Config) reader := &KCPPacketReader{} diff --git a/transport/internet/memory_settings.go b/transport/internet/memory_settings.go index 02fb247fb..770cf82cc 100644 --- a/transport/internet/memory_settings.go +++ b/transport/internet/memory_settings.go @@ -1,7 +1,12 @@ package internet import ( + "context" + "reflect" + + "github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/net/cnc" "github.com/xtls/xray-core/transport/internet/finalmask" ) @@ -12,8 +17,7 @@ type MemoryStreamConfig struct { ProtocolSettings interface{} SecurityType string SecuritySettings interface{} - TcpmaskManager *finalmask.TcpmaskManager - UdpmaskManager *finalmask.UdpmaskManager + FinalMask *finalmask.FinalMask QuicParams *QuicParams SocketSettings *SocketConfig DownloadSettings *MemoryStreamConfig @@ -51,33 +55,53 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) { mss.SecuritySettings = ess } - if s != nil && len(s.Tcpmasks) > 0 { - var masks []finalmask.Tcpmask - for _, msg := range s.Tcpmasks { - instance, err := msg.GetInstance() - if err != nil { - return nil, err - } - masks = append(masks, instance.(finalmask.Tcpmask)) + var tcpMasks []finalmask.TCPMask + var udpMasks []finalmask.UDPMask + + if s != nil { + for i := range s.Tcpmasks { + instance := common.Must2(s.Tcpmasks[i].GetInstance()) + tcpMasks = append(tcpMasks, instance.(finalmask.TCPMask)) + } + for i := range s.Udpmasks { + instance := common.Must2(s.Udpmasks[i].GetInstance()) + udpMasks = append(udpMasks, instance.(finalmask.UDPMask)) } - mss.TcpmaskManager = finalmask.NewTcpmaskManager(masks) } + dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) { + return DialSystem(ctx, dest, mss.SocketSettings) + } + listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) { + return ListenSystem(ctx, addr, mss.SocketSettings) + } + dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) { + conn, err := DialSystem(ctx, dest, mss.SocketSettings) + if err != nil { + return nil, nil, err + } + var newConn net.PacketConn + var udpAddr net.Addr + switch c := conn.(type) { + case *PacketConnWrapper: + newConn = c.PacketConn + udpAddr = conn.RemoteAddr() + case *cnc.Connection: + newConn = &FakePacketConn{Conn: c} + udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}, Port: 0} + default: + panic(reflect.TypeOf(c)) + } + return newConn, udpAddr, nil + } + listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) { + return ListenSystemPacket(ctx, addr, mss.SocketSettings) + } + mss.FinalMask = finalmask.NewFinalMask(tcpMasks, udpMasks, dialTCP, listen, dialUDP, listenPacket) + if s != nil && s.QuicParams != nil { mss.QuicParams = s.QuicParams } - if s != nil && len(s.Udpmasks) > 0 { - var masks []finalmask.Udpmask - for _, msg := range s.Udpmasks { - instance, err := msg.GetInstance() - if err != nil { - return nil, err - } - masks = append(masks, instance.(finalmask.Udpmask)) - } - mss.UdpmaskManager = finalmask.NewUdpmaskManager(masks) - } - return mss, nil } diff --git a/transport/internet/splithttp/dialer.go b/transport/internet/splithttp/dialer.go index 9bac57a18..e896516ed 100644 --- a/transport/internet/splithttp/dialer.go +++ b/transport/internet/splithttp/dialer.go @@ -8,7 +8,7 @@ import ( "net/http" "net/http/httptrace" "net/url" - reflect "reflect" + "reflect" "runtime" "strconv" "sync" @@ -25,6 +25,7 @@ import ( "github.com/xtls/xray-core/common/signal/done" "github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet/browser_dialer" + "github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/hysteria/congestion" "github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr" "github.com/xtls/xray-core/transport/internet/reality" @@ -116,20 +117,17 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea transportConfig := streamSettings.ProtocolSettings.(*Config) dialContext := func(ctxInner context.Context) (net.Conn, error) { - conn, err := internet.DialSystem(ctxInner, dest, streamSettings.SocketSettings) + var conn net.Conn + var err error + if streamSettings.FinalMask != nil { + conn, err = streamSettings.FinalMask.DialTCP(ctxInner, dest) + } else { + conn, err = internet.DialSystem(ctxInner, dest, streamSettings.SocketSettings) + } if err != nil { return nil, err } - if streamSettings.TcpmaskManager != nil { - newConn, err := streamSettings.TcpmaskManager.WrapConnClient(conn) - if err != nil { - conn.Close() - return nil, errors.New("mask err").Base(err) - } - conn = newConn - } - if realityConfig != nil { return reality.UClient(conn, realityConfig, ctxInner, dest) } @@ -196,30 +194,29 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea TLSClientConfig: gotlsConfig, Dial: func(ctx context.Context, addr string, tlsCfg *gotls.Config, cfg *quic.Config) (*quic.Conn, error) { var pktConn net.PacketConn - var udpAddr *net.UDPAddr - - raw, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) - if err != nil { - return nil, errors.New("failed to dial to dest").Base(err) - } - switch c := raw.(type) { - case *internet.PacketConnWrapper: - pktConn = c.PacketConn - udpAddr = raw.RemoteAddr().(*net.UDPAddr) - case *cnc.Connection: - pktConn = &internet.FakePacketConn{Conn: c} - udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port} - default: - panic(reflect.TypeOf(c)) - } - - if streamSettings.UdpmaskManager != nil { - newConn, err := streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn) + var udpAddr net.Addr + if streamSettings.FinalMask != nil { + conn, err := streamSettings.FinalMask.DialUDP(ctx, dest) if err != nil { - pktConn.Close() - return nil, errors.New("mask err").Base(err) + return nil, errors.New("failed to dial to dest").Base(err) + } + pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn + udpAddr = conn.RemoteAddr() + } else { + conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) + if err != nil { + return nil, errors.New("failed to dial to dest").Base(err) + } + switch c := conn.(type) { + case *internet.PacketConnWrapper: + pktConn = c.PacketConn + udpAddr = c.RemoteAddr() + case *cnc.Connection: + pktConn = &internet.FakePacketConn{Conn: c} + udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}} + default: + panic(reflect.TypeOf(c)) } - pktConn = newConn } tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO} diff --git a/transport/internet/splithttp/hub.go b/transport/internet/splithttp/hub.go index 431d8f186..c4573b961 100644 --- a/transport/internet/splithttp/hub.go +++ b/transport/internet/splithttp/hub.go @@ -463,31 +463,17 @@ func ListenXH(ctx context.Context, address net.Address, port net.Port, streamSet l.isH3 = len(tlsConfig.NextProtos) == 1 && tlsConfig.NextProtos[0] == "h3" var err error - if port == net.Port(0) { // unix - l.listener, err = internet.ListenSystem(ctx, &net.UnixAddr{ - Name: address.Domain(), - Net: "unix", - }, streamSettings.SocketSettings) - if err != nil { - return nil, errors.New("failed to listen UNIX domain socket for XHTTP on ", address).Base(err) + if l.isH3 { + var pktConn net.PacketConn + var err error + if streamSettings.FinalMask != nil { + pktConn, err = streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}) + } else { + pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings) } - errors.LogInfo(ctx, "listening UNIX domain socket for XHTTP on ", address) - } else if l.isH3 { // quic - Conn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{ - IP: address.IP(), - Port: int(port), - }, streamSettings.SocketSettings) if err != nil { return nil, errors.New("failed to listen UDP for XHTTP/3 on ", address, ":", port).Base(err) } - if streamSettings.UdpmaskManager != nil { - newConn, err := streamSettings.UdpmaskManager.WrapPacketConnServer(Conn) - if err != nil { - Conn.Close() - return nil, errors.New("mask err").Base(err) - } - Conn = newConn - } quicParams := streamSettings.QuicParams if quicParams == nil { @@ -512,7 +498,7 @@ func ListenXH(ctx context.Context, address net.Address, port net.Port, streamSet common.Must2(rand.Read((*k)[:])) } - tr := &quic.Transport{Conn: Conn, DisableGSO: quicParams.DisableGSO, StatelessResetKey: k} + tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO, StatelessResetKey: k} l.h3listener, err = tr.ListenEarly(tlsConfig, quicConfig) if err != nil { @@ -534,21 +520,24 @@ func ListenXH(ctx context.Context, address net.Address, port net.Port, streamSet errors.LogErrorInner(ctx, err, "failed to serve HTTP/3 for XHTTP/3") } _ = tr.Close() - _ = Conn.Close() + _ = pktConn.Close() }() - } else { // tcp - l.listener, err = internet.ListenSystem(ctx, &net.TCPAddr{ - IP: address.IP(), - Port: int(port), - }, streamSettings.SocketSettings) - if err != nil { - return nil, errors.New("failed to listen TCP for XHTTP on ", address, ":", port).Base(err) + } else { + var addr net.Addr + if port == net.Port(0) { // unix + addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"} + } else { // tcp + addr = &net.TCPAddr{IP: address.IP(), Port: int(port)} } - errors.LogInfo(ctx, "listening TCP for XHTTP on ", address, ":", port) - } - - if !l.isH3 && streamSettings.TcpmaskManager != nil { - l.listener, _ = streamSettings.TcpmaskManager.WrapListener(l.listener) + if streamSettings.FinalMask != nil { + l.listener, err = streamSettings.FinalMask.Listen(ctx, addr) + } else { + l.listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings) + } + if err != nil { + return nil, errors.New("failed to listen ", addr.Network(), " for XHTTP on ", address, ":", port).Base(err) + } + errors.LogInfo(ctx, "listening ", addr.Network(), " for XHTTP on ", address, ":", port) } // tcp/unix (h1/h2) diff --git a/transport/internet/system_dialer.go b/transport/internet/system_dialer.go index b80e9a295..2ff7693de 100644 --- a/transport/internet/system_dialer.go +++ b/transport/internet/system_dialer.go @@ -235,5 +235,5 @@ func (c *FakePacketConn) WriteTo(p []byte, _ net.Addr) (n int, err error) { } func (c *FakePacketConn) LocalAddr() net.Addr { - return &net.UDPAddr{IP: c.Conn.LocalAddr().(*net.TCPAddr).IP, Port: c.Conn.LocalAddr().(*net.TCPAddr).Port} + return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} } diff --git a/transport/internet/tcp/dialer.go b/transport/internet/tcp/dialer.go index 92fa7557f..680ee83b2 100644 --- a/transport/internet/tcp/dialer.go +++ b/transport/internet/tcp/dialer.go @@ -19,18 +19,15 @@ import ( // Dial dials a new TCP connection to the given destination. func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) { errors.LogInfo(ctx, "dialing TCP to ", dest) - conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) - if err != nil { - return nil, err + var conn net.Conn + var err error + if streamSettings.FinalMask != nil { + conn, err = streamSettings.FinalMask.DialTCP(ctx, dest) + } else { + conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings) } - - if streamSettings.TcpmaskManager != nil { - newConn, err := streamSettings.TcpmaskManager.WrapConnClient(conn) - if err != nil { - conn.Close() - return nil, errors.New("mask err").Base(err) - } - conn = newConn + if err != nil { + return nil, errors.New("failed to dial to dest").Base(err) } if config := tls.ConfigFromStreamSettings(streamSettings); config != nil { diff --git a/transport/internet/tcp/hub.go b/transport/internet/tcp/hub.go index c68d55dd9..bb8099ed3 100644 --- a/transport/internet/tcp/hub.go +++ b/transport/internet/tcp/hub.go @@ -41,29 +41,21 @@ func ListenTCP(ctx context.Context, address net.Address, port net.Port, streamSe } var listener net.Listener var err error + var addr net.Addr if port == net.Port(0) { // unix - listener, err = internet.ListenSystem(ctx, &net.UnixAddr{ - Name: address.Domain(), - Net: "unix", - }, streamSettings.SocketSettings) - if err != nil { - return nil, errors.New("failed to listen Unix Domain Socket on ", address).Base(err) - } - errors.LogInfo(ctx, "listening Unix Domain Socket on ", address) + addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"} + } else { // tcp + addr = &net.TCPAddr{IP: address.IP(), Port: int(port)} + } + if streamSettings.FinalMask != nil { + listener, err = streamSettings.FinalMask.Listen(ctx, addr) } else { - listener, err = internet.ListenSystem(ctx, &net.TCPAddr{ - IP: address.IP(), - Port: int(port), - }, streamSettings.SocketSettings) - if err != nil { - return nil, errors.New("failed to listen TCP on ", address, ":", port).Base(err) - } - errors.LogInfo(ctx, "listening TCP on ", address, ":", port) + listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings) } - - if streamSettings.TcpmaskManager != nil { - listener, _ = streamSettings.TcpmaskManager.WrapListener(listener) + if err != nil { + return nil, errors.New("failed to listen ", addr.Network(), " on ", address, ":", port).Base(err) } + errors.LogInfo(ctx, "listening ", addr.Network(), " on ", address, ":", port) if streamSettings.SocketSettings != nil && streamSettings.SocketSettings.AcceptProxyProtocol { errors.LogWarning(ctx, "accepting PROXY protocol") diff --git a/transport/internet/udp/dialer.go b/transport/internet/udp/dialer.go index a81d22230..538c7888d 100644 --- a/transport/internet/udp/dialer.go +++ b/transport/internet/udp/dialer.go @@ -2,12 +2,9 @@ package udp import ( "context" - "reflect" "github.com/xtls/xray-core/common" - "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/net" - "github.com/xtls/xray-core/common/net/cnc" "github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet/stat" ) @@ -15,40 +12,14 @@ import ( func init() { common.Must(internet.RegisterTransportDialer(protocolName, func(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) { - var sockopt *internet.SocketConfig - if streamSettings != nil { - sockopt = streamSettings.SocketSettings - } - conn, err := internet.DialSystem(ctx, dest, sockopt) - if err != nil { - return nil, err - } - - if streamSettings != nil && streamSettings.UdpmaskManager != nil { - var pktConn net.PacketConn - var udpAddr *net.UDPAddr - switch c := conn.(type) { - case *internet.PacketConnWrapper: - pktConn = c.PacketConn - udpAddr = c.RemoteAddr().(*net.UDPAddr) - case *cnc.Connection: - pktConn = &internet.FakePacketConn{Conn: c} - udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port} - default: - panic(reflect.TypeOf(c)) - } - newConn, err := streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn) - if err != nil { - pktConn.Close() - return nil, errors.New("mask err").Base(err) - } - pktConn = newConn - conn = &internet.PacketConnWrapper{ - PacketConn: pktConn, - Dest: udpAddr, + if streamSettings != nil && streamSettings.FinalMask != nil { + return streamSettings.FinalMask.DialUDP(ctx, dest) + } else { + var sockopt *internet.SocketConfig + if streamSettings != nil && streamSettings.SocketSettings != nil { + sockopt = streamSettings.SocketSettings } + return internet.DialSystem(ctx, dest, sockopt) } - - return conn, nil })) } diff --git a/transport/internet/udp/hub.go b/transport/internet/udp/hub.go index 5d29d2030..9ea553197 100644 --- a/transport/internet/udp/hub.go +++ b/transport/internet/udp/hub.go @@ -58,24 +58,15 @@ func ListenUDP(ctx context.Context, address net.Address, port net.Port, streamSe } var err error - hub.conn, err = internet.ListenSystemPacket(ctx, &net.UDPAddr{ - IP: address.IP(), - Port: int(port), - }, sockopt) + if streamSettings.FinalMask != nil { + hub.conn, err = streamSettings.FinalMask.ListenPacket(ctx, &net.UDPAddr{IP: address.IP(), Port: int(port)}) + } else { + hub.conn, err = internet.ListenSystemPacket(ctx, &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings) + } if err != nil { return nil, err } - raw := hub.conn - - if streamSettings.UdpmaskManager != nil { - hub.conn, err = streamSettings.UdpmaskManager.WrapPacketConnServer(raw) - if err != nil { - raw.Close() - return nil, errors.New("mask err").Base(err) - } - } - errors.LogInfo(ctx, "listening UDP on ", address, ":", port) hub.udpConn, _ = hub.conn.(*net.UDPConn) hub.cache = make(chan *udp.Packet, hub.capacity) diff --git a/transport/internet/websocket/dialer.go b/transport/internet/websocket/dialer.go index 199514a5c..0e04505b5 100644 --- a/transport/internet/websocket/dialer.go +++ b/transport/internet/websocket/dialer.go @@ -48,20 +48,16 @@ func dialWebSocket(ctx context.Context, dest net.Destination, streamSettings *in dialer := &websocket.Dialer{ NetDial: func(network, addr string) (net.Conn, error) { - conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) + var conn net.Conn + var err error + if streamSettings.FinalMask != nil { + conn, err = streamSettings.FinalMask.DialTCP(ctx, dest) + } else { + conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings) + } if err != nil { - return nil, err + return nil, errors.New("failed to dial to dest").Base(err) } - - if streamSettings.TcpmaskManager != nil { - newConn, err := streamSettings.TcpmaskManager.WrapConnClient(conn) - if err != nil { - conn.Close() - return nil, errors.New("mask err").Base(err) - } - conn = newConn - } - return conn, err }, ReadBufferSize: 4 * 1024, @@ -79,19 +75,15 @@ func dialWebSocket(ctx context.Context, dest net.Destination, streamSettings *in if fingerprint := tls.GetFingerprint(tConfig.Fingerprint); fingerprint != nil { dialer.NetDialTLSContext = func(_ context.Context, _, addr string) (net.Conn, error) { // Like the NetDial in the dialer - pconn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) - if err != nil { - errors.LogErrorInner(ctx, err, "failed to dial to "+addr) - return nil, err + var pconn net.Conn + var err error + if streamSettings.FinalMask != nil { + pconn, err = streamSettings.FinalMask.DialTCP(ctx, dest) + } else { + pconn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings) } - - if streamSettings.TcpmaskManager != nil { - newConn, err := streamSettings.TcpmaskManager.WrapConnClient(pconn) - if err != nil { - pconn.Close() - return nil, errors.New("mask err").Base(err) - } - pconn = newConn + if err != nil { + return nil, errors.New("failed to dial to dest").Base(err) } // TLS and apply the handshake diff --git a/transport/internet/websocket/hub.go b/transport/internet/websocket/hub.go index 42dba880f..db1cc7c36 100644 --- a/transport/internet/websocket/hub.go +++ b/transport/internet/websocket/hub.go @@ -97,29 +97,21 @@ func ListenWS(ctx context.Context, address net.Address, port net.Port, streamSet } var listener net.Listener var err error + var addr net.Addr if port == net.Port(0) { // unix - listener, err = internet.ListenSystem(ctx, &net.UnixAddr{ - Name: address.Domain(), - Net: "unix", - }, streamSettings.SocketSettings) - if err != nil { - return nil, errors.New("failed to listen unix domain socket(for WS) on ", address).Base(err) - } - errors.LogInfo(ctx, "listening unix domain socket(for WS) on ", address) + addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"} } else { // tcp - listener, err = internet.ListenSystem(ctx, &net.TCPAddr{ - IP: address.IP(), - Port: int(port), - }, streamSettings.SocketSettings) - if err != nil { - return nil, errors.New("failed to listen TCP(for WS) on ", address, ":", port).Base(err) - } - errors.LogInfo(ctx, "listening TCP(for WS) on ", address, ":", port) + addr = &net.TCPAddr{IP: address.IP(), Port: int(port)} } - - if streamSettings.TcpmaskManager != nil { - listener, _ = streamSettings.TcpmaskManager.WrapListener(listener) + if streamSettings.FinalMask != nil { + listener, err = streamSettings.FinalMask.Listen(ctx, addr) + } else { + listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings) } + if err != nil { + return nil, errors.New("failed to listen ", addr.Network(), "(for WS) on ", address, ":", port).Base(err) + } + errors.LogInfo(ctx, "listening ", addr.Network(), "(for WS) on ", address, ":", port) if streamSettings.SocketSettings != nil && streamSettings.SocketSettings.AcceptProxyProtocol { errors.LogWarning(ctx, "accepting PROXY protocol") diff --git a/transport/internet/xdrive/client.go b/transport/internet/xdrive/client.go index 3b14d009b..9cbd14607 100644 --- a/transport/internet/xdrive/client.go +++ b/transport/internet/xdrive/client.go @@ -55,17 +55,14 @@ func newServiceClient(streamSettings *internet.MemoryStreamConfig, timeout time. } } - conn, err := internet.DialSystem(ctx, target, sockopt) - if err != nil { - return nil, host, err + var conn net.Conn + if streamSettings.FinalMask != nil { + conn, err = streamSettings.FinalMask.DialTCP(ctx, target) + } else { + conn, err = internet.DialSystem(ctx, target, sockopt) } - if streamSettings != nil && streamSettings.TcpmaskManager != nil { - masked, err := streamSettings.TcpmaskManager.WrapConnClient(conn) - if err != nil { - conn.Close() - return nil, host, errors.New("mask err").Base(err) - } - conn = masked + if err != nil { + return nil, host, errors.New("failed to dial to dest").Base(err) } return conn, host, nil }