Fix TCP DNS retry
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user