mirror of
https://github.com/MetaCubeX/mihomo.git
synced 2026-10-10 04:03:11 +08:00
fix: trusttunnel client not waiting for TCP connection established
This commit is contained in:
@@ -11,6 +11,7 @@ import (
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
atomic2 "github.com/metacubex/mihomo/common/atomic"
|
||||
"github.com/metacubex/mihomo/common/httputils"
|
||||
"github.com/metacubex/mihomo/common/once"
|
||||
"github.com/metacubex/mihomo/component/dialer"
|
||||
@@ -18,6 +19,7 @@ import (
|
||||
"github.com/metacubex/mihomo/transport/vmess"
|
||||
|
||||
"github.com/metacubex/http"
|
||||
"github.com/metacubex/http/httptrace"
|
||||
"golang.org/x/exp/slices"
|
||||
)
|
||||
|
||||
@@ -139,7 +141,7 @@ func (c *Client) resetHealthCheckTimer() {
|
||||
c.healthCheckTimer.Reset(DefaultHealthCheckTimeout)
|
||||
}
|
||||
|
||||
func (c *Client) roundTrip(request *http.Request, conn *httpConn) {
|
||||
func (c *Client) roundTrip(ctx context.Context, request *http.Request, conn *httpConn) error {
|
||||
c.startOnce.Do(c.start)
|
||||
pipeReader, pipeWriter := io.Pipe()
|
||||
request.Body = pipeReader
|
||||
@@ -151,20 +153,36 @@ func (c *Client) roundTrip(request *http.Request, conn *httpConn) {
|
||||
conn.closeFn = once.OnceFunc(func() {
|
||||
c.count.Add(-1)
|
||||
})
|
||||
ctx, cancel := context.WithCancel(c.ctx) // requestCtx must alive during conn not closed
|
||||
conn.cancelFn = cancel // cancel ctx when conn closed
|
||||
requestCtx, cancel := context.WithCancel(c.ctx) // requestCtx must alive during conn not closed
|
||||
conn.cancelFn = cancel // cancel ctx when conn closed
|
||||
|
||||
// Use gotConn to detect when TCP connection is established, so we can
|
||||
// return the conn immediately without waiting for the HTTP response.
|
||||
gotConn := make(chan bool, 1)
|
||||
addrCtx := httputils.NewAddrContext(&conn.NetAddr, requestCtx)
|
||||
streamCtx := httptrace.WithClientTrace(addrCtx, &httptrace.ClientTrace{
|
||||
GotConn: func(info httptrace.GotConnInfo) {
|
||||
select {
|
||||
case gotConn <- true:
|
||||
default: // GotConn maybe called multiple times, ignore the second and later calls
|
||||
}
|
||||
},
|
||||
})
|
||||
|
||||
var requestErr atomic2.TypedValue[error]
|
||||
go func() {
|
||||
timeout := time.AfterFunc(C.DefaultTCPTimeout, cancel) // only cancel when RoundTrip timeout
|
||||
defer timeout.Stop() // RoundTrip already returned, stop the timer
|
||||
request = request.WithContext(httputils.NewAddrContext(&conn.NetAddr, ctx))
|
||||
request = request.WithContext(streamCtx)
|
||||
response, err := c.roundTripper.RoundTrip(request)
|
||||
if err != nil {
|
||||
requestErr.Store(err)
|
||||
close(gotConn)
|
||||
_ = pipeWriter.CloseWithError(err)
|
||||
_ = pipeReader.CloseWithError(err)
|
||||
conn.setup(nil, err)
|
||||
} else if response.StatusCode != http.StatusOK {
|
||||
_ = response.Body.Close()
|
||||
err = fmt.Errorf("unexpected status code: %d", response.StatusCode)
|
||||
requestErr.Store(err)
|
||||
_ = pipeWriter.CloseWithError(err)
|
||||
_ = pipeReader.CloseWithError(err)
|
||||
conn.setup(nil, err)
|
||||
@@ -173,6 +191,13 @@ func (c *Client) roundTrip(request *http.Request, conn *httpConn) {
|
||||
conn.setup(response.Body, nil)
|
||||
}
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-gotConn:
|
||||
return requestErr.Load()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) newConnectRequest(host, userAgent string) *http.Request {
|
||||
@@ -193,21 +218,33 @@ func (c *Client) newConnectRequest(host, userAgent string) *http.Request {
|
||||
func (c *Client) Dial(ctx context.Context, host string) (net.Conn, error) {
|
||||
request := c.newConnectRequest(host, TCPUserAgent)
|
||||
conn := &tcpConn{}
|
||||
c.roundTrip(request, &conn.httpConn)
|
||||
err := c.roundTrip(ctx, request, &conn.httpConn)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (c *Client) ListenPacket(ctx context.Context) (net.PacketConn, error) {
|
||||
request := c.newConnectRequest(UDPMagicAddress, UDPUserAgent)
|
||||
conn := &clientPacketConn{}
|
||||
c.roundTrip(request, &conn.httpConn)
|
||||
err := c.roundTrip(ctx, request, &conn.httpConn)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (c *Client) ListenICMP(ctx context.Context) (*IcmpConn, error) {
|
||||
request := c.newConnectRequest(ICMPMagicAddress, ICMPUserAgent)
|
||||
conn := &IcmpConn{}
|
||||
c.roundTrip(request, &conn.httpConn)
|
||||
err := c.roundTrip(ctx, request, &conn.httpConn)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user