From 0b5391447ed783d5b8d26ec509f476b6091801ef Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sat, 29 Aug 2026 15:59:41 +0800 Subject: [PATCH] Fix websocket early data with smux and yamux multiplex --- transport/v2raywebsocket/client.go | 21 ++++++--------- transport/v2raywebsocket/conn.go | 42 ++++++++++++++---------------- 2 files changed, 28 insertions(+), 35 deletions(-) diff --git a/transport/v2raywebsocket/client.go b/transport/v2raywebsocket/client.go index a071e8d1..7a5bb252 100644 --- a/transport/v2raywebsocket/client.go +++ b/transport/v2raywebsocket/client.go @@ -73,11 +73,18 @@ func NewClient(ctx context.Context, dialer N.Dialer, serverAddr M.Socksaddr, opt }, nil } -func (c *Client) dialContext(ctx context.Context, requestURL *url.URL, headers http.Header) (*WebsocketConn, error) { +func (c *Client) DialContext(ctx context.Context) (net.Conn, error) { conn, err := c.dialer.DialContext(ctx, N.NetworkTCP, c.serverAddr) if err != nil { return nil, err } + if c.maxEarlyData > 0 { + return &EarlyWebsocketConn{Client: c, rawConn: conn, create: make(chan struct{})}, nil + } + return c.upgrade(conn, &c.requestURL, c.headers) +} + +func (c *Client) upgrade(conn net.Conn, requestURL *url.URL, headers http.Header) (*WebsocketConn, error) { var deadlineConn net.Conn if deadline.NeedAdditionalReadDeadline(conn) { deadlineConn = deadline.NewConn(conn) @@ -108,18 +115,6 @@ func (c *Client) dialContext(ctx context.Context, requestURL *url.URL, headers h return NewConn(conn, nil, ws.StateClientSide), nil } -func (c *Client) DialContext(ctx context.Context) (net.Conn, error) { - if c.maxEarlyData <= 0 { - conn, err := c.dialContext(ctx, &c.requestURL, c.headers) - if err != nil { - return nil, err - } - return conn, nil - } else { - return &EarlyWebsocketConn{Client: c, ctx: ctx, create: make(chan struct{})}, nil - } -} - func (c *Client) Close() error { return nil } diff --git a/transport/v2raywebsocket/conn.go b/transport/v2raywebsocket/conn.go index 4939e737..ef277233 100644 --- a/transport/v2raywebsocket/conn.go +++ b/transport/v2raywebsocket/conn.go @@ -1,7 +1,6 @@ package v2raywebsocket import ( - "context" "encoding/base64" "errors" "io" @@ -16,7 +15,6 @@ import ( "github.com/sagernet/sing/common/buf" "github.com/sagernet/sing/common/debug" E "github.com/sagernet/sing/common/exceptions" - M "github.com/sagernet/sing/common/metadata" "github.com/sagernet/ws" "github.com/sagernet/ws/wsutil" ) @@ -135,11 +133,11 @@ func (c *WebsocketConn) Upstream() any { type EarlyWebsocketConn struct { *Client - ctx context.Context - conn atomic.Pointer[WebsocketConn] - access sync.Mutex - create chan struct{} - err error + rawConn net.Conn + conn atomic.Pointer[WebsocketConn] + access sync.Mutex + create chan struct{} + err error } func (c *EarlyWebsocketConn) Read(b []byte) (n int, err error) { @@ -172,14 +170,14 @@ func (c *EarlyWebsocketConn) writeRequest(content []byte) error { if c.earlyDataHeaderName == "" { requestURL := c.requestURL requestURL.Path += earlyDataString - conn, err = c.dialContext(c.ctx, &requestURL, c.headers) + conn, err = c.upgrade(c.rawConn, &requestURL, c.headers) } else { headers := c.headers.Clone() headers.Set(c.earlyDataHeaderName, earlyDataString) - conn, err = c.dialContext(c.ctx, &c.requestURL, headers) + conn, err = c.upgrade(c.rawConn, &c.requestURL, headers) } } else { - conn, err = c.dialContext(c.ctx, &c.requestURL, c.headers) + conn, err = c.upgrade(c.rawConn, &c.requestURL, c.headers) } if err != nil { return err @@ -240,26 +238,26 @@ func (c *EarlyWebsocketConn) WriteBuffer(buffer *buf.Buffer) error { func (c *EarlyWebsocketConn) Close() error { conn := c.conn.Load() - if conn == nil { + if conn != nil { + return conn.Close() + } + c.rawConn.Close() + c.access.Lock() + defer c.access.Unlock() + if c.conn.Load() != nil || c.err != nil { return nil } - return conn.Close() + c.err = net.ErrClosed + close(c.create) + return nil } func (c *EarlyWebsocketConn) LocalAddr() net.Addr { - conn := c.conn.Load() - if conn == nil { - return M.Socksaddr{} - } - return conn.LocalAddr() + return c.rawConn.LocalAddr() } func (c *EarlyWebsocketConn) RemoteAddr() net.Addr { - conn := c.conn.Load() - if conn == nil { - return M.Socksaddr{} - } - return conn.RemoteAddr() + return c.rawConn.RemoteAddr() } func (c *EarlyWebsocketConn) SetDeadline(t time.Time) error {