Compare commits

...
17 changed files with 137 additions and 100 deletions
+20
View File
@@ -0,0 +1,20 @@
package net
// PacketConnWrapper wraps a PacketConn into a Conn with a fixed destination address.
type PacketConnWrapper struct {
PacketConn
Dest Addr
}
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
n, _, err := c.PacketConn.ReadFrom(p)
return n, err
}
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
return c.PacketConn.WriteTo(p, c.Dest)
}
func (c *PacketConnWrapper) RemoteAddr() Addr {
return c.Dest
}
+4 -4
View File
@@ -467,7 +467,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
if statConn != nil {
counter = statConn.ReadCounter
}
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
if c, ok := iConn.(*net.PacketConnWrapper); ok {
isOverridden := false
if UDPOverride.Address != nil || UDPOverride.Port != 0 {
isOverridden = true
@@ -487,7 +487,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
}
type PacketReader struct {
*internet.PacketConnWrapper
*net.PacketConnWrapper
stats.Counter
Handler *Handler
DefaultRule *FinalRule
@@ -542,7 +542,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
if statConn != nil {
counter = statConn.WriteCounter
}
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
if c, ok := iConn.(*net.PacketConnWrapper); ok {
// If DialDest is a domain, it will be resolved in dialer
// check this behavior and add it to map
resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]()
@@ -563,7 +563,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
}
type PacketWriter struct {
*internet.PacketConnWrapper
*net.PacketConnWrapper
stats.Counter
*Handler
DefaultRule *FinalRule
+1 -1
View File
@@ -151,7 +151,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
}
defer conn.Close()
uc := &wireguard.UDPConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = uc
+15
View File
@@ -277,6 +277,7 @@ func (w *VisionReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
w.ob.CanSpliceCopy = 1
}
}
SuppressOuterCloseNotify(w.conn)
readerConn, readCounter, _ := UnwrapRawConn(w.conn)
w.directReadCounter = readCounter
w.Reader = buf.NewReader(readerConn)
@@ -340,6 +341,7 @@ func (w *VisionWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
// w.ob.CanSpliceCopy = 1
// }
}
SuppressOuterCloseNotify(w.conn)
rawConn, _, writerCounter := UnwrapRawConn(w.conn)
w.Writer = buf.NewWriter(rawConn)
w.directWriteCounter = writerCounter
@@ -669,6 +671,19 @@ func XtlsFilterTls(buffer buf.MultiBuffer, trafficState *TrafficState, ctx conte
}
}
type CloseNotifySuppressor interface {
SuppressCloseNotify()
}
// Close our local TLS conn instance might send a incorrect close_notify alert
// if the XTLS direct copy mode is enabled and cause TLS BAD_RECORD_MAC on users' browser
// Close the underlying connection directly to avoid this issue.
func SuppressOuterCloseNotify(conn net.Conn) {
if suppressor, ok := stat.TryUnwrapStatsConn(conn).(CloseNotifySuppressor); ok {
suppressor.SuppressCloseNotify()
}
}
// UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it
func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) {
var readCounter, writerCounter stats.Counter
+3 -4
View File
@@ -27,7 +27,6 @@ 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"
)
@@ -200,7 +199,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
}
defer conn.Close()
c := &UDPConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = c
@@ -264,14 +263,14 @@ func (h *Handler) init(ctx context.Context) error {
if err != nil {
return nil, errors.New("failed to dial to dest").Base(err)
}
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
pktConn = conn.(*net.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:
case *net.PacketConnWrapper:
pktConn = c.PacketConn
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
+2 -2
View File
@@ -21,7 +21,7 @@ import (
"syscall"
"time"
"github.com/xtls/xray-core/transport/internet"
xnet "github.com/xtls/xray-core/common/net"
"golang.zx2c4.com/wireguard/tun"
"golang.org/x/net/dns/dnsmessage"
@@ -220,7 +220,7 @@ func (tun *netTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, error
if err != nil {
return nil, err
}
return &internet.PacketConnWrapper{
return &xnet.PacketConnWrapper{
PacketConn: conn,
Dest: net.UDPAddrFromAddrPort(raddr),
}, nil
+2 -2
View File
@@ -16,7 +16,7 @@ import (
"github.com/vishvananda/netlink"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet"
xnet "github.com/xtls/xray-core/common/net"
"golang.zx2c4.com/wireguard/tun"
)
@@ -263,7 +263,7 @@ func (tun *kernelTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, er
if err != nil {
return nil, err
}
return &internet.PacketConnWrapper{
return &xnet.PacketConnWrapper{
PacketConn: conn,
Dest: net.UDPAddrFromAddrPort(raddr),
}, nil
+4 -22
View File
@@ -82,7 +82,7 @@ func (fm *FinalMask) DialTCP(ctx context.Context, dest net.Destination) (net.Con
if err != nil {
return nil, err
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err
},
}
for i := range fm.tcpMasks {
@@ -144,7 +144,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
if err != nil {
return nil, err
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil
}
for i := range fm.udpMasks {
if i > 0 {
@@ -171,7 +171,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
if err != nil {
return nil, err
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err
},
}
var sizes []int
@@ -208,7 +208,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
if addr == nil {
addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil
}
func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
@@ -272,24 +272,6 @@ const (
UDPSize = 4096
)
type PacketConnWrapper struct {
net.PacketConn
udpAddr net.Addr
}
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 {
net.PacketConn
+1 -1
View File
@@ -380,7 +380,7 @@ func TestPacketConnReadWrite(t *testing.T) {
t.Fatal(err)
}
t.Cleanup(func() { clientConn.Close() })
client := clientConn.(*finalmask.PacketConnWrapper).PacketConn
client := clientConn.(*net.PacketConnWrapper).PacketConn
_ = client.SetDeadline(time.Now().Add(time.Second))
_ = server.SetDeadline(time.Now().Add(time.Second))
+2 -2
View File
@@ -73,7 +73,7 @@ func NewUDPHopConn(c *Config, dest *net.Destination, dialer *finalmask.Dialer) (
if err != nil {
return nil, err
}
cur := conn.(*finalmask.PacketConnWrapper).PacketConn
cur := conn.(*net.PacketConnWrapper).PacketConn
addr := conn.RemoteAddr().(*net.UDPAddr)
client := &udpHopConn{
dialer: dialer,
@@ -150,7 +150,7 @@ func (c *udpHopConn) hop() {
_ = c.pre.Close()
}
c.pre = c.cur
c.cur = conn.(*finalmask.PacketConnWrapper).PacketConn
c.cur = conn.(*net.PacketConnWrapper).PacketConn
c.wg.Add(1)
go c.recv(c.cur)
}
+2 -2
View File
@@ -119,7 +119,7 @@ func (c *client) dial(ctx context.Context) error {
if err != nil {
return errors.New("failed to dial to dest").Base(err)
}
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
pktConn = conn.(*net.PacketConnWrapper).PacketConn
udpAddr = conn.RemoteAddr()
} else {
conn, err := internet.DialSystem(ctx, c.dest, c.socketConfig)
@@ -127,7 +127,7 @@ func (c *client) dial(ctx context.Context) error {
return errors.New("failed to dial to dest").Base(err)
}
switch c := conn.(type) {
case *internet.PacketConnWrapper:
case *net.PacketConnWrapper:
pktConn = c.PacketConn
udpAddr = c.RemoteAddr()
case *cnc.Connection:
+2 -3
View File
@@ -19,7 +19,6 @@ import (
"github.com/xtls/xray-core/common/net/cnc"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/transport/internet"
"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/masque/connectip"
@@ -80,7 +79,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
if err != nil {
return nil, errors.New("failed to dial to dest").Base(err)
}
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
pktConn = conn.(*net.PacketConnWrapper).PacketConn
udpAddr = conn.RemoteAddr()
} else {
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
@@ -88,7 +87,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
return nil, errors.New("failed to dial to dest").Base(err)
}
switch c := conn.(type) {
case *internet.PacketConnWrapper:
case *net.PacketConnWrapper:
pktConn = c.PacketConn
udpAddr = c.RemoteAddr()
case *cnc.Connection:
+30 -31
View File
@@ -54,11 +54,10 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) {
mss.SecurityType = s.SecurityType
mss.SecuritySettings = ess
}
if s != nil && (len(s.Tcpmasks) != 0 || len(s.Udpmasks) != 0) {
var tcpMasks []finalmask.TCPMask
var udpMasks []finalmask.UDPMask
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))
@@ -67,37 +66,37 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) {
instance := common.Must2(s.Udpmasks[i].GetInstance())
udpMasks = append(udpMasks, instance.(finalmask.UDPMask))
}
}
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
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
return DialSystem(ctx, dest, mss.SocketSettings)
}
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))
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
return ListenSystem(ctx, addr, mss.SocketSettings)
}
return newConn, udpAddr, nil
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 *net.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)
}
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
+29 -4
View File
@@ -18,6 +18,7 @@ import (
"regexp"
"strings"
"sync"
"sync/atomic"
"time"
"unsafe"
@@ -36,6 +37,18 @@ import (
type Conn struct {
*reality.Conn
suppressCloseNotify atomic.Bool
}
func (c *Conn) SuppressCloseNotify() {
c.suppressCloseNotify.Store(true)
}
func (c *Conn) Close() error {
if c.suppressCloseNotify.Load() {
return c.Conn.NetConn().Close()
}
return c.Conn.Close()
}
func (c *Conn) HandshakeAddress() net.Address {
@@ -56,10 +69,22 @@ func Server(c net.Conn, config *reality.Config) (net.Conn, error) {
type UConn struct {
*utls.UConn
Config *Config
ServerName string
AuthKey []byte
Verified bool
Config *Config
ServerName string
AuthKey []byte
Verified bool
suppressCloseNotify atomic.Bool
}
func (c *UConn) SuppressCloseNotify() {
c.suppressCloseNotify.Store(true)
}
func (c *UConn) Close() error {
if c.suppressCloseNotify.Load() {
return c.NetConn().Close()
}
return c.UConn.Close()
}
func (c *UConn) HandshakeAddress() net.Address {
+2 -3
View File
@@ -25,7 +25,6 @@ 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"
@@ -200,7 +199,7 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea
if err != nil {
return nil, errors.New("failed to dial to dest").Base(err)
}
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
pktConn = conn.(*net.PacketConnWrapper).PacketConn
udpAddr = conn.RemoteAddr()
} else {
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
@@ -208,7 +207,7 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea
return nil, errors.New("failed to dial to dest").Base(err)
}
switch c := conn.(type) {
case *internet.PacketConnWrapper:
case *net.PacketConnWrapper:
pktConn = c.PacketConn
udpAddr = c.RemoteAddr()
case *cnc.Connection:
+1 -19
View File
@@ -86,7 +86,7 @@ func (d *DefaultSystemDialer) Dial(ctx context.Context, src net.Address, dest ne
if err != nil {
return nil, err
}
return &PacketConnWrapper{
return &net.PacketConnWrapper{
PacketConn: packetConn,
Dest: destAddr,
}, nil
@@ -148,24 +148,6 @@ func (d *DefaultSystemDialer) DestIpAddress() net.IP {
return nil
}
type PacketConnWrapper struct {
net.PacketConn
Dest net.Addr
}
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
n, _, err := c.PacketConn.ReadFrom(p)
return n, err
}
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
return c.PacketConn.WriteTo(p, c.Dest)
}
func (c *PacketConnWrapper) RemoteAddr() net.Addr {
return c.Dest
}
type SystemDialerAdapter interface {
Dial(network string, address string) (net.Conn, error)
}
+17
View File
@@ -6,6 +6,7 @@ import (
"crypto/tls"
"math/big"
"slices"
"sync/atomic"
"time"
utls "github.com/refraction-networking/utls"
@@ -29,11 +30,19 @@ var (
type Conn struct {
*tls.Conn
suppressCloseNotify atomic.Bool
}
const tlsCloseTimeout = 250 * time.Millisecond
func (c *Conn) SuppressCloseNotify() {
c.suppressCloseNotify.Store(true)
}
func (c *Conn) Close() error {
if c.suppressCloseNotify.Load() {
return c.Conn.NetConn().Close()
}
timer := time.AfterFunc(tlsCloseTimeout, func() {
c.Conn.NetConn().Close()
})
@@ -74,11 +83,19 @@ func Server(c net.Conn, config *tls.Config) net.Conn {
type UConn struct {
*utls.UConn
suppressCloseNotify atomic.Bool
}
var _ Interface = (*UConn)(nil)
func (c *UConn) SuppressCloseNotify() {
c.suppressCloseNotify.Store(true)
}
func (c *UConn) Close() error {
if c.suppressCloseNotify.Load() {
return c.Conn.NetConn().Close()
}
timer := time.AfterFunc(tlsCloseTimeout, func() {
c.Conn.NetConn().Close()
})