Fix TCP DNS retry

This commit is contained in:
世界
2026-07-21 15:00:57 +08:00
parent b352013db5
commit a9d89ab2a3
4 changed files with 140 additions and 13 deletions
+48 -13
View File
@@ -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)
+90
View File
@@ -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")
+1
View File
@@ -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
}
+1
View File
@@ -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
}