dns: Probe connection reuse support for TCP transports

This commit is contained in:
世界 2026-07-22 16:29:39 +08:00
parent 59f4dd0830
commit 962bb4d192
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
4 changed files with 501 additions and 10 deletions

View file

@ -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
}

View file

@ -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")

View file

@ -71,6 +71,7 @@ func NewTCPRaw(adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socks
return ReadMessage(conn)
},
retryReadError: true,
probeReuse: true,
})
return t
}

View file

@ -77,6 +77,7 @@ func NewTLSRaw(logger logger.ContextLogger, adapter dns.TransportAdapter, dialer
return ReadMessage(conn)
},
retryReadError: true,
probeReuse: true,
})
return t
}