diff --git a/dns/transport/multiplexer.go b/dns/transport/multiplexer.go index a26caab88..6b275ae9d 100644 --- a/dns/transport/multiplexer.go +++ b/dns/transport/multiplexer.go @@ -6,17 +6,35 @@ import ( "net" "sync" "sync/atomic" + "time" E "github.com/sagernet/sing/common/exceptions" mDNS "github.com/miekg/dns" ) +const ( + reuseStateUnknown int32 = iota + reuseStateProbing + reuseStateSupported + reuseStateUnsupported +) + +const ( + reuseProbeTimeout = 5 * time.Second + reuseProbeRetryInterval = time.Minute + reuseDemoteFailureLimit = 3 + + reuseProbeQueryIdA uint16 = 1 + reuseProbeQueryIdB uint16 = 2 +) + 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) retryReadError bool + probeReuse bool } type queryMultiplexer struct { @@ -26,6 +44,13 @@ type queryMultiplexer struct { queryAccess sync.Mutex queryId uint16 queries map[uint16]*pendingQuery + + reuseState atomic.Int32 + demoteFailures atomic.Int32 + + probeAccess sync.Mutex + probeEpoch uint32 + lastProbeTime time.Time } type multiplexConn struct { @@ -76,6 +101,14 @@ func (m *queryMultiplexer) Close() error { } func (m *queryMultiplexer) Reset() { + if m.options.probeReuse { + m.probeAccess.Lock() + m.probeEpoch++ + m.reuseState.Store(reuseStateUnknown) + m.lastProbeTime = time.Time{} + m.probeAccess.Unlock() + m.demoteFailures.Store(0) + } m.connection.Reset() } @@ -95,7 +128,170 @@ 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) + m.dispatch(ctx, message, callback, true) +} + +func (m *queryMultiplexer) dispatch(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error), retryReadError bool) { + if m.options.probeReuse && m.reuseState.Load() != reuseStateSupported { + m.maybeStartProbe(ctx, message) + go m.exchangeSingle(ctx, message, callback) + return + } + m.exchangeAsync(ctx, message, callback, retryReadError) +} + +func (m *queryMultiplexer) exchangeSingle(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) { + conn, err := m.options.dial(ctx) + if err != nil { + callback(nil, err) + return + } + defer conn.Close() + stop := context.AfterFunc(ctx, func() { + conn.Close() + }) + defer stop() + err = m.options.write(conn, message, message.Id) + if err != nil { + ctxErr := ctx.Err() + if ctxErr != nil { + callback(nil, ctxErr) + return + } + callback(nil, E.Cause(err, "write request")) + return + } + for { + var response *mDNS.Msg + response, err = m.options.readNext(conn) + if err != nil { + ctxErr := ctx.Err() + if ctxErr != nil { + callback(nil, ctxErr) + return + } + callback(nil, E.Cause(err, "read response")) + return + } + if response == nil { + continue + } + response.Id = message.Id + callback(response, nil) + return + } +} + +func (m *queryMultiplexer) maybeStartProbe(ctx context.Context, message *mDNS.Msg) { + if len(message.Question) == 0 { + return + } + m.probeAccess.Lock() + if m.reuseState.Load() == reuseStateProbing { + m.probeAccess.Unlock() + return + } + if !m.lastProbeTime.IsZero() && time.Since(m.lastProbeTime) < reuseProbeRetryInterval { + m.probeAccess.Unlock() + return + } + m.reuseState.Store(reuseStateProbing) + m.lastProbeTime = time.Now() + epoch := m.probeEpoch + m.probeAccess.Unlock() + go m.runReuseProbe(context.WithoutCancel(ctx), message.Question[0].Name, epoch) +} + +func (m *queryMultiplexer) runReuseProbe(ctx context.Context, questionName string, epoch uint32) { + supported, dialFailed := m.executeReuseProbe(ctx, questionName) + m.probeAccess.Lock() + defer m.probeAccess.Unlock() + if m.probeEpoch != epoch { + return + } + switch { + case supported: + m.reuseState.Store(reuseStateSupported) + m.demoteFailures.Store(0) + case dialFailed: + m.reuseState.Store(reuseStateUnknown) + default: + m.reuseState.Store(reuseStateUnsupported) + } +} + +func (m *queryMultiplexer) executeReuseProbe(ctx context.Context, questionName string) (supported bool, dialFailed bool) { + probeCtx, cancel := context.WithTimeout(ctx, reuseProbeTimeout) + defer cancel() + conn, err := m.options.dial(probeCtx) + if err != nil { + return false, true + } + defer conn.Close() + stop := context.AfterFunc(probeCtx, func() { + conn.Close() + }) + defer stop() + queryA := new(mDNS.Msg) + queryA.SetQuestion(questionName, mDNS.TypeA) + queryAAAA := new(mDNS.Msg) + queryAAAA.SetQuestion(questionName, mDNS.TypeAAAA) + err = m.options.write(conn, queryA, reuseProbeQueryIdA) + if err == nil { + err = m.options.write(conn, queryAAAA, reuseProbeQueryIdB) + } + if err != nil { + return false, false + } + var seenA, seenAAAA bool + for !seenA || !seenAAAA { + var response *mDNS.Msg + response, err = m.options.readNext(conn) + if err != nil { + return false, false + } + if response == nil { + continue + } + switch response.Id { + case reuseProbeQueryIdA: + seenA = true + case reuseProbeQueryIdB: + seenAAAA = true + } + } + return true, false +} + +func (m *queryMultiplexer) recordConnDeath(conn *multiplexConn) { + if !m.options.probeReuse || m.reuseState.Load() != reuseStateSupported { + return + } + if conn.readEpoch.Load() == 0 { + return + } + m.queryAccess.Lock() + var pendingOnConn int + for _, pending := range m.queries { + if pending.conn == conn { + pendingOnConn++ + } + } + m.queryAccess.Unlock() + if pendingOnConn == 0 { + m.demoteFailures.Store(0) + return + } + if m.demoteFailures.Add(1) < reuseDemoteFailureLimit { + return + } + m.probeAccess.Lock() + if m.reuseState.Load() == reuseStateSupported { + m.reuseState.Store(reuseStateUnsupported) + m.lastProbeTime = time.Now() + } + m.probeAccess.Unlock() + m.demoteFailures.Store(0) } func (m *queryMultiplexer) exchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error), retryReadError bool) { @@ -180,7 +376,7 @@ func (m *queryMultiplexer) completeConnDone(queryId uint16, connCtx context.Cont connErr := context.Cause(connCtx) _, readFailed := connErr.(*queryMultiplexerReadError) if pending.retryCtx != nil && readFailed { - m.exchangeAsync(pending.retryCtx, pending.message, pending.callback, false) + m.dispatch(pending.retryCtx, pending.message, pending.callback, false) return } pending.callback(nil, connErr) @@ -232,6 +428,7 @@ func (m *queryMultiplexer) recvLoop(conn *multiplexConn) { for { message, err := m.options.readNext(conn) if err != nil { + m.recordConnDeath(conn) m.connection.Invalidate(conn, &queryMultiplexerReadError{cause: err}) return } diff --git a/dns/transport/multiplexer_test.go b/dns/transport/multiplexer_test.go index 8f59ce03d..b5bd02b62 100644 --- a/dns/transport/multiplexer_test.go +++ b/dns/transport/multiplexer_test.go @@ -5,6 +5,7 @@ import ( "errors" "io" "net" + "sync/atomic" "testing" "time" @@ -67,17 +68,24 @@ func TestTCPTransportRetriesReadErrorOnReusedConn(t *testing.T) { 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() + multiplexer := newQueryMultiplexer(queryMultiplexerOptions{ + dial: func(ctx context.Context) (net.Conn, error) { + return net.Dial("tcp", listener.Addr().String()) + }, + write: func(conn net.Conn, message *mDNS.Msg, queryId uint16) error { + return WriteMessage(conn, queryId, message) + }, + readNext: func(conn net.Conn) (*mDNS.Msg, error) { + return ReadMessage(conn) + }, + retryReadError: true, + }) + defer multiplexer.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) + _, err = multiplexer.Exchange(ctx, firstMessage) cancel() if err != nil { t.Fatal("first query failed: ", err) @@ -86,7 +94,7 @@ func TestTCPTransportRetriesReadErrorOnReusedConn(t *testing.T) { secondMessage := new(mDNS.Msg) secondMessage.SetQuestion("second.example.com.", mDNS.TypeAAAA) ctx, cancel = context.WithTimeout(context.Background(), time.Second) - _, err = transport.Exchange(ctx, secondMessage) + _, err = multiplexer.Exchange(ctx, secondMessage) cancel() if err != nil { t.Fatal("second query failed: ", err) @@ -101,6 +109,290 @@ func TestTCPTransportRetriesReadErrorOnReusedConn(t *testing.T) { } } +func newTestTCPTransport(t *testing.T, listener net.Listener) *TCPTransport { + transportDialer, err := dialer.NewDefault(context.Background(), option.DialerOptions{}) + if err != nil { + t.Fatal(err) + } + return NewTCPRaw(boxDNS.NewTransportAdapter(C.DNSTypeTCP, "test", nil), transportDialer, M.SocksaddrFromNet(listener.Addr())) +} + +func testExchange(transport *TCPTransport, questionName string) error { + message := new(mDNS.Msg) + message.SetQuestion(questionName, mDNS.TypeA) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + _, err := transport.Exchange(ctx, message) + return err +} + +func TestTCPTransportSingleQueryServer(t *testing.T) { + t.Parallel() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + var accepted atomic.Int32 + go func() { + for { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + accepted.Add(1) + go func() { + defer conn.Close() + request, readErr := ReadMessage(conn) + if readErr != nil { + return + } + response := new(mDNS.Msg) + response.SetReply(request) + WriteMessage(conn, request.Id, response) + }() + } + }() + + transport := newTestTCPTransport(t, listener) + defer transport.Close() + + const queryCount = 8 + results := make(chan error, queryCount) + for range queryCount { + go func() { + results <- testExchange(transport, "example.com.") + }() + } + for range queryCount { + err = <-results + if err != nil { + t.Fatal("query failed: ", err) + } + } + deadline := time.Now().Add(time.Second) + for accepted.Load() < queryCount+1 { + if time.Now().After(deadline) { + t.Fatal("expected a probe connection, accepted ", accepted.Load()) + } + time.Sleep(10 * time.Millisecond) + } + time.Sleep(100 * time.Millisecond) + if count := accepted.Load(); count != queryCount+1 { + t.Fatal("expected one connection per query plus probe, accepted ", count) + } +} + +func TestTCPTransportProbeEnablesReuse(t *testing.T) { + t.Parallel() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + var maxServedOnConn atomic.Int32 + go func() { + for { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + go func() { + defer conn.Close() + var served int32 + for { + request, readErr := ReadMessage(conn) + if readErr != nil { + return + } + served++ + for { + current := maxServedOnConn.Load() + if served <= current || maxServedOnConn.CompareAndSwap(current, served) { + break + } + } + response := new(mDNS.Msg) + response.SetReply(request) + WriteMessage(conn, request.Id, response) + } + }() + } + }() + + transport := newTestTCPTransport(t, listener) + defer transport.Close() + + deadline := time.Now().Add(3 * time.Second) + for maxServedOnConn.Load() < 3 { + if time.Now().After(deadline) { + t.Fatal("reuse was not enabled after successful probe") + } + err = testExchange(transport, "example.com.") + if err != nil { + t.Fatal("query failed: ", err) + } + time.Sleep(10 * time.Millisecond) + } + + const burstCount = 5 + results := make(chan error, burstCount) + for range burstCount { + go func() { + results <- testExchange(transport, "example.com.") + }() + } + for range burstCount { + err = <-results + if err != nil { + t.Fatal("burst query failed: ", err) + } + } +} + +func TestTCPTransportDemotesBrokenReuse(t *testing.T) { + t.Parallel() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + var accepted atomic.Int32 + go func() { + for { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + accepted.Add(1) + go func() { + defer conn.Close() + for served := 0; ; served++ { + request, readErr := ReadMessage(conn) + if readErr != nil { + return + } + if served >= 2 { + return + } + response := new(mDNS.Msg) + response.SetReply(request) + WriteMessage(conn, request.Id, response) + } + }() + } + }() + + transport := newTestTCPTransport(t, listener) + defer transport.Close() + + deadline := time.Now().Add(3 * time.Second) + for { + before := accepted.Load() + err = testExchange(transport, "example.com.") + if err != nil { + t.Fatal("query failed: ", err) + } + if accepted.Load() == before { + break + } + if time.Now().After(deadline) { + t.Fatal("reuse was not enabled after successful probe") + } + } + + for range 15 { + err = testExchange(transport, "example.com.") + if err != nil { + t.Fatal("query failed during demotion: ", err) + } + } + if transport.multiplexer.reuseState.Load() != reuseStateUnsupported { + t.Fatal("expected demotion to single connection mode") + } + + time.Sleep(100 * time.Millisecond) + before := accepted.Load() + const singleCount = 4 + for range singleCount { + err = testExchange(transport, "example.com.") + if err != nil { + t.Fatal("query failed after demotion: ", err) + } + } + if count := accepted.Load() - before; count != singleCount { + t.Fatal("expected one connection per query after demotion, got ", count) + } +} + +func TestTCPTransportSilentPipelineServer(t *testing.T) { + t.Parallel() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + go func() { + for { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + return + } + go func() { + defer conn.Close() + request, readErr := ReadMessage(conn) + if readErr != nil { + return + } + conn.SetReadDeadline(time.Now().Add(300 * time.Millisecond)) + _, secondErr := ReadMessage(conn) + if secondErr == nil { + conn.SetReadDeadline(time.Time{}) + io.Copy(io.Discard, conn) + return + } + var netErr net.Error + if !errors.As(secondErr, &netErr) || !netErr.Timeout() { + return + } + conn.SetReadDeadline(time.Time{}) + response := new(mDNS.Msg) + response.SetReply(request) + WriteMessage(conn, request.Id, response) + }() + } + }() + + transport := newTestTCPTransport(t, listener) + defer transport.Close() + + const queryCount = 5 + results := make(chan error, queryCount) + for range queryCount { + go func() { + results <- testExchange(transport, "example.com.") + }() + } + for range queryCount { + err = <-results + if err != nil { + t.Fatal("query failed: ", err) + } + } + + deadline := time.Now().Add(8 * time.Second) + for transport.multiplexer.reuseState.Load() != reuseStateUnsupported { + if time.Now().After(deadline) { + t.Fatal("expected probe timeout to disable reuse") + } + time.Sleep(100 * time.Millisecond) + } + err = testExchange(transport, "example.com.") + if err != nil { + t.Fatal("query failed after probe timeout: ", err) + } +} + 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 cd2eb9975..d406aeb26 100644 --- a/dns/transport/tcp.go +++ b/dns/transport/tcp.go @@ -71,6 +71,7 @@ func NewTCPRaw(adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socks return ReadMessage(conn) }, retryReadError: true, + probeReuse: true, }) return t } diff --git a/dns/transport/tls.go b/dns/transport/tls.go index f05edd33e..da1db319a 100644 --- a/dns/transport/tls.go +++ b/dns/transport/tls.go @@ -77,6 +77,7 @@ func NewTLSRaw(logger logger.ContextLogger, adapter dns.TransportAdapter, dialer return ReadMessage(conn) }, retryReadError: true, + probeReuse: true, }) return t }