diff --git a/dns/transport/multiplexer.go b/dns/transport/multiplexer.go index 51f3c325c..a26caab88 100644 --- a/dns/transport/multiplexer.go +++ b/dns/transport/multiplexer.go @@ -13,9 +13,10 @@ import ( ) type queryMultiplexerOptions struct { - dial func(ctx context.Context) (net.Conn, error) - write func(conn net.Conn, message *mDNS.Msg, queryId uint16) error - readNext func(conn net.Conn) (*mDNS.Msg, error) + dial func(ctx context.Context) (net.Conn, error) + write func(conn net.Conn, message *mDNS.Msg, queryId uint16) error + readNext func(conn net.Conn) (*mDNS.Msg, error) + retryReadError bool } type queryMultiplexer struct { @@ -32,13 +33,26 @@ type multiplexConn struct { readEpoch atomic.Uint64 } +type queryMultiplexerReadError struct { + cause error +} + +func (e *queryMultiplexerReadError) Error() string { + return e.cause.Error() +} + +func (e *queryMultiplexerReadError) Unwrap() error { + return e.cause +} + type pendingQuery struct { conn *multiplexConn - originalId uint16 + message *mDNS.Msg readEpoch uint64 callback func(response *mDNS.Msg, err error) stopContext func() bool stopConn func() bool + retryCtx context.Context } func newQueryMultiplexer(options queryMultiplexerOptions) *queryMultiplexer { @@ -81,6 +95,10 @@ func (m *queryMultiplexer) Exchange(ctx context.Context, message *mDNS.Msg) (*mD } func (m *queryMultiplexer) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) { + m.exchangeAsync(ctx, message, callback, true) +} + +func (m *queryMultiplexer) exchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error), retryReadError bool) { for firstAttempt := true; ; firstAttempt = false { conn, connCtx, created, err := m.connection.AcquireShared(ctx, m.dialConn) if err != nil { @@ -90,7 +108,7 @@ func (m *queryMultiplexer) ExchangeAsync(ctx context.Context, message *mDNS.Msg, if created { go m.recvLoop(conn) } - queryId, err := m.register(ctx, connCtx, conn, message.Id, callback) + queryId, err := m.register(ctx, connCtx, conn, message, callback, retryReadError && m.options.retryReadError && !created) if err != nil { m.connection.Release(conn, true) callback(nil, err) @@ -121,7 +139,7 @@ func (m *queryMultiplexer) dialConn(ctx context.Context) (*multiplexConn, error) return &multiplexConn{Conn: conn}, nil } -func (m *queryMultiplexer) register(ctx context.Context, connCtx context.Context, conn *multiplexConn, originalId uint16, callback func(response *mDNS.Msg, err error)) (uint16, error) { +func (m *queryMultiplexer) register(ctx context.Context, connCtx context.Context, conn *multiplexConn, message *mDNS.Msg, callback func(response *mDNS.Msg, err error), retryReadError bool) (uint16, error) { m.queryAccess.Lock() defer m.queryAccess.Unlock() start := m.queryId @@ -136,21 +154,38 @@ func (m *queryMultiplexer) register(ctx context.Context, connCtx context.Context } queryId := m.queryId pending := &pendingQuery{ - conn: conn, - originalId: originalId, - readEpoch: conn.readEpoch.Load(), - callback: callback, + conn: conn, + message: message, + readEpoch: conn.readEpoch.Load(), + callback: callback, + } + if retryReadError { + pending.retryCtx = ctx } m.queries[queryId] = pending pending.stopContext = context.AfterFunc(ctx, func() { m.completeContextDone(queryId, ctx) }) pending.stopConn = context.AfterFunc(connCtx, func() { - m.complete(queryId, nil, context.Cause(connCtx), false) + m.completeConnDone(queryId, connCtx) }) return queryId, nil } +func (m *queryMultiplexer) completeConnDone(queryId uint16, connCtx context.Context) { + pending := m.take(queryId) + if pending == nil { + return + } + connErr := context.Cause(connCtx) + _, readFailed := connErr.(*queryMultiplexerReadError) + if pending.retryCtx != nil && readFailed { + m.exchangeAsync(pending.retryCtx, pending.message, pending.callback, false) + return + } + pending.callback(nil, connErr) +} + func (m *queryMultiplexer) take(queryId uint16) *pendingQuery { m.queryAccess.Lock() pending, loaded := m.queries[queryId] @@ -174,7 +209,7 @@ func (m *queryMultiplexer) complete(queryId uint16, response *mDNS.Msg, err erro m.connection.Release(pending.conn, true) } if response != nil { - response.Id = pending.originalId + response.Id = pending.message.Id } pending.callback(response, err) } @@ -197,7 +232,7 @@ func (m *queryMultiplexer) recvLoop(conn *multiplexConn) { for { message, err := m.options.readNext(conn) if err != nil { - m.connection.Invalidate(conn, err) + m.connection.Invalidate(conn, &queryMultiplexerReadError{cause: err}) return } conn.readEpoch.Add(1) diff --git a/dns/transport/multiplexer_test.go b/dns/transport/multiplexer_test.go index 413a82e7b..8f59ce03d 100644 --- a/dns/transport/multiplexer_test.go +++ b/dns/transport/multiplexer_test.go @@ -8,9 +8,99 @@ import ( "testing" "time" + "github.com/sagernet/sing-box/common/dialer" + C "github.com/sagernet/sing-box/constant" + boxDNS "github.com/sagernet/sing-box/dns" + "github.com/sagernet/sing-box/option" + M "github.com/sagernet/sing/common/metadata" + mDNS "github.com/miekg/dns" ) +func TestTCPTransportRetriesReadErrorOnReusedConn(t *testing.T) { + t.Parallel() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + serverDone := make(chan error, 1) + go func() { + firstConn, acceptErr := listener.Accept() + if acceptErr != nil { + serverDone <- acceptErr + return + } + firstRequest, readErr := ReadMessage(firstConn) + if readErr != nil { + firstConn.Close() + serverDone <- readErr + return + } + firstResponse := new(mDNS.Msg) + firstResponse.SetReply(firstRequest) + writeErr := WriteMessage(firstConn, firstRequest.Id, firstResponse) + if writeErr != nil { + firstConn.Close() + serverDone <- writeErr + return + } + _, readErr = ReadMessage(firstConn) + firstConn.Close() + if readErr != nil { + serverDone <- readErr + return + } + secondConn, acceptErr := listener.Accept() + if acceptErr != nil { + serverDone <- acceptErr + return + } + defer secondConn.Close() + secondRequest, readErr := ReadMessage(secondConn) + if readErr != nil { + serverDone <- readErr + return + } + secondResponse := new(mDNS.Msg) + secondResponse.SetReply(secondRequest) + serverDone <- WriteMessage(secondConn, secondRequest.Id, secondResponse) + }() + + transportDialer, err := dialer.NewDefault(context.Background(), option.DialerOptions{}) + if err != nil { + t.Fatal(err) + } + transport := NewTCPRaw(boxDNS.NewTransportAdapter(C.DNSTypeTCP, "test", nil), transportDialer, M.SocksaddrFromNet(listener.Addr())) + defer transport.Close() + + firstMessage := new(mDNS.Msg) + firstMessage.SetQuestion("first.example.com.", mDNS.TypeA) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + _, err = transport.Exchange(ctx, firstMessage) + cancel() + if err != nil { + t.Fatal("first query failed: ", err) + } + + secondMessage := new(mDNS.Msg) + secondMessage.SetQuestion("second.example.com.", mDNS.TypeAAAA) + ctx, cancel = context.WithTimeout(context.Background(), time.Second) + _, err = transport.Exchange(ctx, secondMessage) + cancel() + if err != nil { + t.Fatal("second query failed: ", err) + } + select { + case err = <-serverDone: + if err != nil { + t.Fatal("DNS server failed: ", err) + } + case <-time.After(time.Second): + t.Fatal("DNS server did not finish") + } +} + func TestMultiplexerTimeoutInvalidatesConn(t *testing.T) { t.Parallel() listener, err := net.Listen("tcp", "127.0.0.1:0") diff --git a/dns/transport/tcp.go b/dns/transport/tcp.go index 45f3cda74..cd2eb9975 100644 --- a/dns/transport/tcp.go +++ b/dns/transport/tcp.go @@ -70,6 +70,7 @@ func NewTCPRaw(adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socks readNext: func(conn net.Conn) (*mDNS.Msg, error) { return ReadMessage(conn) }, + retryReadError: true, }) return t } diff --git a/dns/transport/tls.go b/dns/transport/tls.go index 9ec6bf406..f05edd33e 100644 --- a/dns/transport/tls.go +++ b/dns/transport/tls.go @@ -76,6 +76,7 @@ func NewTLSRaw(logger logger.ContextLogger, adapter dns.TransportAdapter, dialer readNext: func(conn net.Conn) (*mDNS.Msg, error) { return ReadMessage(conn) }, + retryReadError: true, }) return t }