diff --git a/proxy/proxy.go b/proxy/proxy.go index ab7d7abc5..5cb9d7a11 100644 --- a/proxy/proxy.go +++ b/proxy/proxy.go @@ -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 diff --git a/transport/internet/reality/reality.go b/transport/internet/reality/reality.go index 1407151b5..9ccae6686 100644 --- a/transport/internet/reality/reality.go +++ b/transport/internet/reality/reality.go @@ -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 { diff --git a/transport/internet/tls/tls.go b/transport/internet/tls/tls.go index df5d1cbd7..f27fa17b0 100644 --- a/transport/internet/tls/tls.go +++ b/transport/internet/tls/tls.go @@ -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() })