Merge tag 'v1.14.0'

This commit is contained in:
Shtorm
2026-09-04 01:08:30 +03:00
987 changed files with 141481 additions and 11302 deletions
+509 -320
View File
File diff suppressed because it is too large Load Diff
+26
View File
@@ -22,6 +22,19 @@ func logCachedResponse(logger logger.ContextLogger, ctx context.Context, respons
}
}
func logOptimisticResponse(logger logger.ContextLogger, ctx context.Context, response *dns.Msg) {
if logger == nil || len(response.Question) == 0 {
return
}
domain := FqdnToDomain(response.Question[0].Name)
logger.DebugContext(ctx, "optimistic ", domain, " ", dns.RcodeToString[response.Rcode])
for _, recordList := range [][]dns.RR{response.Answer, response.Ns, response.Extra} {
for _, record := range recordList {
logger.InfoContext(ctx, "optimistic ", dns.Type(record.Header().Rrtype).String(), " ", FormatQuestion(record.String()))
}
}
}
func logExchangedResponse(logger logger.ContextLogger, ctx context.Context, response *dns.Msg, ttl uint32) {
if logger == nil || len(response.Question) == 0 {
return
@@ -35,6 +48,19 @@ func logExchangedResponse(logger logger.ContextLogger, ctx context.Context, resp
}
}
func logRefreshedResponse(logger logger.ContextLogger, ctx context.Context, response *dns.Msg, ttl uint32) {
if logger == nil || len(response.Question) == 0 {
return
}
domain := FqdnToDomain(response.Question[0].Name)
logger.DebugContext(ctx, "refreshed ", domain, " ", dns.RcodeToString[response.Rcode], " ", ttl)
for _, recordList := range [][]dns.RR{response.Answer, response.Ns, response.Extra} {
for _, record := range recordList {
logger.InfoContext(ctx, "refreshed ", dns.Type(record.Header().Rrtype).String(), " ", FormatQuestion(record.String()))
}
}
}
func logRejectedResponse(logger logger.ContextLogger, ctx context.Context, response *dns.Msg) {
if logger == nil || len(response.Question) == 0 {
return
+3 -3
View File
@@ -6,7 +6,7 @@ import (
"github.com/miekg/dns"
)
func TruncateDNSMessage(request *dns.Msg, response *dns.Msg, headroom int) (*buf.Buffer, error) {
func TruncateDNSMessage(request *dns.Msg, response *dns.Msg, frontHeadroom int, rearHeadroom int) (*buf.Buffer, error) {
maxLen := 512
if edns0Option := request.IsEdns0(); edns0Option != nil {
if udpSize := int(edns0Option.UDPSize()); udpSize > 512 {
@@ -18,8 +18,8 @@ func TruncateDNSMessage(request *dns.Msg, response *dns.Msg, headroom int) (*buf
response = response.Copy()
response.Truncate(maxLen)
}
buffer := buf.NewSize(headroom*2 + 1 + responseLen)
buffer.Resize(headroom, 0)
buffer := buf.NewSize(frontHeadroom + responseLen + 1 + rearHeadroom)
buffer.Resize(frontHeadroom, 0)
rawMessage, err := response.PackBuffer(buffer.FreeBytes())
if err != nil {
buffer.Release()
+49
View File
@@ -2,6 +2,9 @@ package dns
import (
"net/netip"
"slices"
"github.com/sagernet/sing/common"
"github.com/miekg/dns"
)
@@ -10,6 +13,52 @@ func SetClientSubnet(message *dns.Msg, clientSubnet netip.Prefix) *dns.Msg {
return setClientSubnet(message, clientSubnet, true)
}
func clientSubnetFromMessage(message *dns.Msg) netip.Prefix {
for _, record := range message.Extra {
optRecord, isOPTRecord := record.(*dns.OPT)
if !isOPTRecord {
continue
}
for _, option := range optRecord.Option {
subnetOption, isEDNS0Subnet := option.(*dns.EDNS0_SUBNET)
if !isEDNS0Subnet {
continue
}
address, addressLoaded := netip.AddrFromSlice(subnetOption.Address)
if !addressLoaded {
return netip.Prefix{}
}
return netip.PrefixFrom(address.Unmap(), int(subnetOption.SourceNetmask))
}
}
return netip.Prefix{}
}
func removeClientSubnet(message *dns.Msg) *dns.Msg {
if !slices.ContainsFunc(message.Extra, func(record dns.RR) bool {
optRecord, isOPTRecord := record.(*dns.OPT)
if !isOPTRecord {
return false
}
return slices.ContainsFunc(optRecord.Option, func(option dns.EDNS0) bool {
return option.Option() == dns.EDNS0SUBNET
})
}) {
return message
}
message = message.Copy()
for _, record := range message.Extra {
optRecord, isOPTRecord := record.(*dns.OPT)
if !isOPTRecord {
continue
}
optRecord.Option = common.Filter(optRecord.Option, func(option dns.EDNS0) bool {
return option.Option() != dns.EDNS0SUBNET
})
}
return message
}
func setClientSubnet(message *dns.Msg, clientSubnet netip.Prefix, clone bool) *dns.Msg {
var (
optRecord *dns.OPT
+5 -4
View File
@@ -5,10 +5,11 @@ import (
)
const (
RcodeSuccess RcodeError = mDNS.RcodeSuccess
RcodeFormatError RcodeError = mDNS.RcodeFormatError
RcodeNameError RcodeError = mDNS.RcodeNameError
RcodeRefused RcodeError = mDNS.RcodeRefused
RcodeSuccess RcodeError = mDNS.RcodeSuccess
RcodeServerFailure RcodeError = mDNS.RcodeServerFailure
RcodeFormatError RcodeError = mDNS.RcodeFormatError
RcodeNameError RcodeError = mDNS.RcodeNameError
RcodeRefused RcodeError = mDNS.RcodeRefused
)
type RcodeError int
+1454 -156
View File
File diff suppressed because it is too large Load Diff
+612
View File
@@ -0,0 +1,612 @@
package dns
import (
"context"
"net/netip"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
R "github.com/sagernet/sing-box/route/rule"
mDNS "github.com/miekg/dns"
"github.com/stretchr/testify/require"
)
type fakeDNSTransport struct {
tag string
delay time.Duration
immediate bool
rcode int
address netip.Addr
exchangeErr error
access sync.Mutex
queryCount atomic.Int32
firstQueried time.Time
}
func (t *fakeDNSTransport) Start(stage adapter.StartStage) error {
return nil
}
func (t *fakeDNSTransport) Close() error {
return nil
}
func (t *fakeDNSTransport) Type() string {
return "fake"
}
func (t *fakeDNSTransport) Tag() string {
return t.tag
}
func (t *fakeDNSTransport) Dependencies() []string {
return nil
}
func (t *fakeDNSTransport) Reset() {
}
func (t *fakeDNSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
t.access.Lock()
if t.firstQueried.IsZero() {
t.firstQueried = time.Now()
}
t.access.Unlock()
t.queryCount.Add(1)
select {
case <-time.After(t.delay):
case <-ctx.Done():
return nil, ctx.Err()
}
if t.exchangeErr != nil {
return nil, t.exchangeErr
}
if t.rcode != mDNS.RcodeSuccess {
return FixedResponseStatus(message, t.rcode), nil
}
return FixedResponse(message.Id, message.Question[0], []netip.Addr{t.address}, 300), nil
}
func (t *fakeDNSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
if t.immediate {
callback(t.Exchange(ctx, message))
return
}
go func() {
callback(t.Exchange(ctx, message))
}()
}
type fakeDNSTransportManager struct {
transports map[string]adapter.DNSTransport
defaultTransport adapter.DNSTransport
}
func (m *fakeDNSTransportManager) Start(stage adapter.StartStage) error {
return nil
}
func (m *fakeDNSTransportManager) Close() error {
return nil
}
func (m *fakeDNSTransportManager) Transports() []adapter.DNSTransport {
return nil
}
func (m *fakeDNSTransportManager) Transport(tag string) (adapter.DNSTransport, bool) {
transport, loaded := m.transports[tag]
return transport, loaded
}
func (m *fakeDNSTransportManager) Default() adapter.DNSTransport {
return m.defaultTransport
}
func (m *fakeDNSTransportManager) FakeIP() adapter.FakeIPTransport {
return nil
}
func (m *fakeDNSTransportManager) Remove(tag string) error {
return nil
}
func (m *fakeDNSTransportManager) Create(ctx context.Context, logger log.ContextLogger, tag string, outboundType string, options any) error {
return nil
}
func raceTestRouter(t *testing.T, transports ...*fakeDNSTransport) *Router {
transportMap := make(map[string]adapter.DNSTransport)
for _, transport := range transports {
transportMap[transport.tag] = transport
}
return &Router{
ctx: context.Background(),
logger: log.NewNOPFactory().Logger(),
transport: &fakeDNSTransportManager{
transports: transportMap,
defaultTransport: transportMap["final"],
},
client: NewClient(ClientOptions{
Context: context.Background(),
DisableCache: true,
Logger: log.NewNOPFactory().Logger(),
}),
}
}
func raceTestRules(t *testing.T, rawRules []option.DNSRule) []adapter.DNSRule {
rules := make([]adapter.DNSRule, 0, len(rawRules))
for _, rawRule := range rawRules {
rule, err := R.NewDNSRule(context.Background(), log.NewNOPFactory().Logger(), rawRule, true, false)
require.NoError(t, err)
rules = append(rules, rule)
}
return rules
}
func raceTestExchange(router *Router, rules []adapter.DNSRule) exchangeWithRulesResult {
message := &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: 1,
RecursionDesired: true,
},
Question: []mDNS.Question{{
Name: "race.example.org.",
Qtype: mDNS.TypeA,
Qclass: mDNS.ClassINET,
}},
}
metadata := &adapter.InboundContext{
Domain: "race.example.org",
QueryType: mDNS.TypeA,
}
ctx := adapter.WithContext(context.Background(), metadata)
return router.exchangeWithRules(ctx, rules, message, adapter.DNSQueryOptions{}, false)
}
func evaluateRule(server string, tag string, speculative bool) option.DNSRule {
return option.DNSRule{
Type: "",
DefaultOptions: option.DefaultDNSRule{
DNSRuleAction: option.DNSRuleAction{
Action: C.RuleActionTypeEvaluate,
EvaluateOptions: option.DNSEvaluateActionOptions{
Server: server,
Tag: tag,
Speculative: speculative,
},
},
},
}
}
func respondRule(responseTag string, race bool, requireSuccess bool) option.DNSRule {
rule := option.DNSRule{
Type: "",
DefaultOptions: option.DefaultDNSRule{
RawDefaultDNSRule: option.RawDefaultDNSRule{
MatchResponse: &option.DNSRuleMatchResponse{Enabled: true, Tag: responseTag},
},
DNSRuleAction: option.DNSRuleAction{
Action: C.RuleActionTypeRespond,
Race: race,
},
},
}
if requireSuccess {
successRcode := option.DNSRCode(mDNS.RcodeSuccess)
rule.DefaultOptions.ResponseRcode = &successRcode
}
return rule
}
func routeRule(server string, speculative bool) option.DNSRule {
return option.DNSRule{
Type: "",
DefaultOptions: option.DefaultDNSRule{
DNSRuleAction: option.DNSRuleAction{
Action: C.RuleActionTypeRoute,
RouteOptions: option.DNSRouteActionOptions{
Server: server,
Speculative: speculative,
},
},
},
}
}
func responseAddress(t *testing.T, response *mDNS.Msg) netip.Addr {
require.NotNil(t, response)
require.Len(t, response.Answer, 1)
record, isA := response.Answer[0].(*mDNS.A)
require.True(t, isA)
address, _ := netip.AddrFromSlice(record.A)
return address.Unmap()
}
// Both evaluate queries must launch in parallel, and a failed primary must
// fall through to the secondary instead of failing the request.
func TestDNSEvaluateParallelFallback(t *testing.T) {
t.Parallel()
transportX := &fakeDNSTransport{tag: "x", delay: 200 * time.Millisecond, exchangeErr: context.DeadlineExceeded}
transportY := &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
router := raceTestRouter(t, transportX, transportY)
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
evaluateRule("y", "y", false),
respondRule("x", false, true),
respondRule("y", false, true),
})
startTime := time.Now()
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.2"), responseAddress(t, result.response))
require.Equal(t, int32(1), transportX.queryCount.Load())
require.Equal(t, int32(1), transportY.queryCount.Load())
require.Less(t, transportY.firstQueried.Sub(transportX.firstQueried), 100*time.Millisecond)
require.Less(t, time.Since(startTime), 350*time.Millisecond)
}
// The first race rule whose response arrives and matches must commit
// immediately, without waiting for the slower rule written before it.
func TestDNSRaceFastestWins(t *testing.T) {
t.Parallel()
transportX := &fakeDNSTransport{tag: "x", delay: 500 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")}
transportY := &fakeDNSTransport{tag: "y", delay: 20 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
router := raceTestRouter(t, transportX, transportY)
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
evaluateRule("y", "y", false),
respondRule("x", true, true),
respondRule("y", true, true),
})
startTime := time.Now()
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.2"), responseAddress(t, result.response))
require.Less(t, time.Since(startTime), 400*time.Millisecond)
}
// A race rule whose response completed synchronously (cache hit) before its
// rule is scanned must commit immediately instead of being blocked by an
// earlier armed race rule.
func TestDNSRaceImmediateLaterResponseWins(t *testing.T) {
t.Parallel()
transportX := &fakeDNSTransport{tag: "x", delay: 200 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")}
transportY := &fakeDNSTransport{tag: "y", immediate: true, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
router := raceTestRouter(t, transportX, transportY)
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
evaluateRule("y", "y", false),
respondRule("x", true, true),
respondRule("y", true, true),
})
startTime := time.Now()
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.2"), responseAddress(t, result.response))
require.Less(t, time.Since(startTime), 100*time.Millisecond)
}
// A synchronously completed race rule that misses disarms in place and rule
// scanning continues; the remaining race rule wins once its response arrives.
func TestDNSRaceImmediateMissContinues(t *testing.T) {
t.Parallel()
transportX := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")}
transportY := &fakeDNSTransport{tag: "y", immediate: true, rcode: mDNS.RcodeNameError}
router := raceTestRouter(t, transportX, transportY)
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
evaluateRule("y", "y", false),
respondRule("y", true, true),
respondRule("x", true, true),
})
startTime := time.Now()
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response))
require.GreaterOrEqual(t, time.Since(startTime), 90*time.Millisecond)
}
// A race route rule whose binding completed synchronously must commit its
// route immediately instead of being blocked by an earlier armed race rule.
func TestDNSRaceImmediateRouteCommits(t *testing.T) {
t.Parallel()
transportX := &fakeDNSTransport{tag: "x", delay: 200 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")}
transportY := &fakeDNSTransport{tag: "y", immediate: true, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
transportFinal := &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")}
router := raceTestRouter(t, transportX, transportY, transportFinal)
successRcode := option.DNSRCode(mDNS.RcodeSuccess)
raceRouteRule := option.DNSRule{
Type: "",
DefaultOptions: option.DefaultDNSRule{
RawDefaultDNSRule: option.RawDefaultDNSRule{
MatchResponse: &option.DNSRuleMatchResponse{Enabled: true, Tag: "y"},
ResponseRcode: &successRcode,
},
DNSRuleAction: option.DNSRuleAction{
Action: C.RuleActionTypeRoute,
RouteOptions: option.DNSRouteActionOptions{
Server: "final",
},
Race: true,
},
},
}
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
evaluateRule("y", "y", false),
respondRule("x", true, true),
raceRouteRule,
})
startTime := time.Now()
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.9"), responseAddress(t, result.response))
require.Equal(t, int32(1), transportFinal.queryCount.Load())
require.Less(t, time.Since(startTime), 100*time.Millisecond)
}
// Without race, rule order decides even when a later response arrives first.
func TestDNSOrderedReadsPreferEarlierRule(t *testing.T) {
t.Parallel()
transportX := &fakeDNSTransport{tag: "x", delay: 200 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")}
transportY := &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
router := raceTestRouter(t, transportX, transportY)
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
evaluateRule("y", "y", false),
respondRule("x", false, true),
respondRule("y", false, true),
})
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response))
}
// A pending race rule must hold back the default route: the default server
// is never queried when the race rule hits, and is queried only after the
// race decision resolved when it misses.
func TestDNSRaceBarrierProtectsDefaultRoute(t *testing.T) {
t.Parallel()
transportHit := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")}
transportFinal := &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")}
router := raceTestRouter(t, transportHit, transportFinal)
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
respondRule("x", true, true),
})
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response))
require.Equal(t, int32(0), transportFinal.queryCount.Load())
transportMiss := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeNameError}
transportFinal = &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")}
router = raceTestRouter(t, transportMiss, transportFinal)
rules = raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
respondRule("x", true, true),
})
startTime := time.Now()
result = raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.9"), responseAddress(t, result.response))
require.Equal(t, int32(1), transportFinal.queryCount.Load())
require.GreaterOrEqual(t, transportFinal.firstQueried.Sub(startTime), 90*time.Millisecond)
}
// A speculative route launches while the race decision is pending, but its
// response is only used after the race rule missed.
func TestDNSSpeculativeRoute(t *testing.T) {
t.Parallel()
transportMiss := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeNameError}
transportFinal := &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")}
router := raceTestRouter(t, transportMiss, transportFinal)
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
respondRule("x", true, true),
routeRule("final", true),
})
startTime := time.Now()
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.9"), responseAddress(t, result.response))
require.Equal(t, int32(1), transportFinal.queryCount.Load())
require.Less(t, transportFinal.firstQueried.Sub(startTime), 90*time.Millisecond)
require.GreaterOrEqual(t, time.Since(startTime), 90*time.Millisecond)
transportHit := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")}
transportFinal = &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")}
router = raceTestRouter(t, transportHit, transportFinal)
rules = raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
respondRule("x", true, true),
routeRule("final", true),
})
result = raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response))
require.Equal(t, int32(1), transportFinal.queryCount.Load())
}
// A matched rule without race must not take effect while a race rule is
// still pending: a race hit wins even when the other rule matched earlier,
// and on a race miss the other rule takes effect only after that decision.
func TestDNSNonRaceCommitWaitsForPendingRace(t *testing.T) {
t.Parallel()
transportHit := &fakeDNSTransport{tag: "x", delay: 150 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")}
transportY := &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
router := raceTestRouter(t, transportHit, transportY)
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
evaluateRule("y", "y", false),
respondRule("x", true, true),
respondRule("y", false, true),
})
startTime := time.Now()
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response))
require.GreaterOrEqual(t, time.Since(startTime), 140*time.Millisecond)
transportMiss := &fakeDNSTransport{tag: "x", delay: 150 * time.Millisecond, rcode: mDNS.RcodeNameError}
transportY = &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
router = raceTestRouter(t, transportMiss, transportY)
rules = raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
evaluateRule("y", "y", false),
respondRule("x", true, true),
respondRule("y", false, true),
})
startTime = time.Now()
result = raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.2"), responseAddress(t, result.response))
require.GreaterOrEqual(t, time.Since(startTime), 140*time.Millisecond)
}
// speculative on a route rule with match_response launches the route query as
// soon as the rule matched, while its response is only used after the pending
// race rule missed.
func TestDNSSpeculativeRouteOnBindingRule(t *testing.T) {
t.Parallel()
transportMiss := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeNameError}
transportY := &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
transportFinal := &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")}
router := raceTestRouter(t, transportMiss, transportY, transportFinal)
successRcode := option.DNSRCode(mDNS.RcodeSuccess)
boundRouteRule := option.DNSRule{
Type: "",
DefaultOptions: option.DefaultDNSRule{
RawDefaultDNSRule: option.RawDefaultDNSRule{
MatchResponse: &option.DNSRuleMatchResponse{Enabled: true, Tag: "y"},
ResponseRcode: &successRcode,
},
DNSRuleAction: option.DNSRuleAction{
Action: C.RuleActionTypeRoute,
RouteOptions: option.DNSRouteActionOptions{
Server: "final",
Speculative: true,
},
},
},
}
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
evaluateRule("y", "y", false),
respondRule("x", true, true),
boundRouteRule,
})
startTime := time.Now()
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.9"), responseAddress(t, result.response))
require.Equal(t, int32(1), transportFinal.queryCount.Load())
require.Less(t, transportFinal.firstQueried.Sub(startTime), 90*time.Millisecond)
require.GreaterOrEqual(t, time.Since(startTime), 90*time.Millisecond)
}
// A race rule that rejects its response (NXDOMAIN vs required success)
// disarms and lets the other race rule win.
func TestDNSRaceSkipsRejectedResponse(t *testing.T) {
t.Parallel()
transportX := &fakeDNSTransport{tag: "x", delay: 10 * time.Millisecond, rcode: mDNS.RcodeNameError}
transportY := &fakeDNSTransport{tag: "y", delay: 100 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
router := raceTestRouter(t, transportX, transportY)
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "x", false),
evaluateRule("y", "y", false),
respondRule("x", true, true),
respondRule("y", true, true),
})
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.2"), responseAddress(t, result.response))
}
// A logical race rule is judged once all of its referenced responses arrived:
// it wins over a slower race rule when its sub-rules match, and on a miss the
// slower race rule takes over.
func TestDNSLogicalRace(t *testing.T) {
t.Parallel()
successRcode := option.DNSRCode(mDNS.RcodeSuccess)
logicalRule := func() option.DNSRule {
return option.DNSRule{
Type: C.RuleTypeLogical,
LogicalOptions: option.LogicalDNSRule{
RawLogicalDNSRule: option.RawLogicalDNSRule{
Mode: C.LogicalTypeAnd,
Rules: []option.DNSRule{
{
Type: C.RuleTypeDefault,
DefaultOptions: option.DefaultDNSRule{
RawDefaultDNSRule: option.RawDefaultDNSRule{
MatchResponse: &option.DNSRuleMatchResponse{Enabled: true},
ResponseRcode: &successRcode,
},
},
},
{
Type: C.RuleTypeDefault,
DefaultOptions: option.DefaultDNSRule{
RawDefaultDNSRule: option.RawDefaultDNSRule{
MatchResponse: &option.DNSRuleMatchResponse{Enabled: true, Tag: "y"},
ResponseRcode: &successRcode,
},
},
},
},
},
DNSRuleAction: option.DNSRuleAction{
Action: C.RuleActionTypeRespond,
Race: true,
},
},
}
}
transportX := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")}
transportY := &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
transportZ := &fakeDNSTransport{tag: "z", delay: 250 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.3")}
router := raceTestRouter(t, transportX, transportY, transportZ)
rules := raceTestRules(t, []option.DNSRule{
evaluateRule("x", "", false),
evaluateRule("y", "y", false),
evaluateRule("z", "z", false),
logicalRule(),
respondRule("z", true, true),
})
startTime := time.Now()
result := raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response))
require.GreaterOrEqual(t, time.Since(startTime), 90*time.Millisecond)
require.Less(t, time.Since(startTime), 240*time.Millisecond)
transportX = &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeNameError}
transportY = &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")}
transportZ = &fakeDNSTransport{tag: "z", delay: 250 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.3")}
router = raceTestRouter(t, transportX, transportY, transportZ)
rules = raceTestRules(t, []option.DNSRule{
evaluateRule("x", "", false),
evaluateRule("y", "y", false),
evaluateRule("z", "z", false),
logicalRule(),
respondRule("z", true, true),
})
startTime = time.Now()
result = raceTestExchange(router, rules)
require.NoError(t, result.err)
require.Equal(t, netip.MustParseAddr("192.0.2.3"), responseAddress(t, result.response))
require.GreaterOrEqual(t, time.Since(startTime), 240*time.Millisecond)
}
+280 -92
View File
@@ -8,12 +8,14 @@ import (
"runtime"
"strings"
"sync"
"sync/atomic"
"syscall"
"time"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing-tun"
@@ -37,7 +39,12 @@ func RegisterTransport(registry *dns.TransportRegistry) {
dns.RegisterTransport[option.DHCPDNSServerOptions](registry, C.DNSTypeDHCP, NewTransport)
}
var _ adapter.DNSTransport = (*Transport)(nil)
var (
_ adapter.DNSTransport = (*Transport)(nil)
_ adapter.DNSTransportWithEnvironment = (*Transport)(nil)
)
var errInterfaceIsCellular = E.New("interface is cellular")
type Transport struct {
dns.TransportAdapter
@@ -45,15 +52,24 @@ type Transport struct {
dialer N.Dialer
logger logger.ContextLogger
networkManager adapter.NetworkManager
platformInterface adapter.PlatformInterface
interfaceName string
interfaceCallback *list.Element[tun.DefaultInterfaceUpdateCallback]
transportLock sync.RWMutex
updatedAt time.Time
lastError error
servers []M.Socksaddr
search []string
updateAccess sync.Mutex
updateCancel context.CancelFunc
refreshAccess sync.Mutex
savedState atomic.Pointer[transportState]
ndots int
attempts int
optional bool
}
type transportState struct {
updatedAt time.Time
lastError error
search []string
servers []M.Socksaddr
serverTransports []adapter.DNSTransport
}
func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, options option.DHCPDNSServerOptions) (adapter.DNSTransport, error) {
@@ -62,26 +78,29 @@ func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, opt
return nil, err
}
return &Transport{
TransportAdapter: dns.NewTransportAdapterWithLocalOptions(C.DNSTypeDHCP, tag, options.LocalDNSServerOptions),
ctx: ctx,
dialer: transportDialer,
logger: logger,
networkManager: service.FromContext[adapter.NetworkManager](ctx),
interfaceName: options.Interface,
ndots: 1,
attempts: 2,
TransportAdapter: dns.NewTransportAdapterWithLocalOptions(C.DNSTypeDHCP, tag, options.LocalDNSServerOptions),
ctx: ctx,
dialer: transportDialer,
logger: logger,
networkManager: service.FromContext[adapter.NetworkManager](ctx),
platformInterface: service.FromContext[adapter.PlatformInterface](ctx),
interfaceName: options.Interface,
ndots: 1,
attempts: 2,
}, nil
}
func NewRawTransport(transportAdapter dns.TransportAdapter, ctx context.Context, dialer N.Dialer, logger log.ContextLogger) *Transport {
return &Transport{
TransportAdapter: transportAdapter,
ctx: ctx,
dialer: dialer,
logger: logger,
networkManager: service.FromContext[adapter.NetworkManager](ctx),
ndots: 1,
attempts: 2,
TransportAdapter: transportAdapter,
ctx: ctx,
dialer: dialer,
logger: logger,
networkManager: service.FromContext[adapter.NetworkManager](ctx),
platformInterface: service.FromContext[adapter.PlatformInterface](ctx),
ndots: 1,
attempts: 2,
optional: true,
}
}
@@ -93,9 +112,13 @@ func (t *Transport) Start(stage adapter.StartStage) error {
t.interfaceCallback = t.networkManager.InterfaceMonitor().RegisterCallback(t.interfaceUpdated)
}
go func() {
_, err := t.fetch()
err := t.fetch()
if err != nil {
t.logger.Error(E.Cause(err, "fetch DNS servers"))
if errors.Is(err, errInterfaceIsCellular) && t.optional {
t.logger.Debug(E.Cause(errInterfaceIsCellular, "dhcp: fetch DNS servers"))
} else {
t.logger.Error(E.Cause(err, "dhcp: fetch DNS servers"))
}
}
}()
return nil
@@ -105,65 +128,152 @@ func (t *Transport) Close() error {
if t.interfaceCallback != nil {
t.networkManager.InterfaceMonitor().UnregisterCallback(t.interfaceCallback)
}
t.updateAccess.Lock()
updateCancel := t.updateCancel
t.updateCancel = nil
t.updateAccess.Unlock()
if updateCancel != nil {
updateCancel()
}
t.refreshAccess.Lock()
defer t.refreshAccess.Unlock()
state := t.savedState.Swap(nil)
if state != nil {
closeServerTransports(state.serverTransports)
}
return nil
}
func (t *Transport) Reset() {
t.transportLock.Lock()
t.updatedAt = time.Time{}
t.lastError = nil
t.servers = nil
t.transportLock.Unlock()
t.refreshAccess.Lock()
defer t.refreshAccess.Unlock()
state := t.savedState.Swap(nil)
if state != nil {
closeServerTransports(state.serverTransports)
}
}
func (t *Transport) Environment() []string {
state := t.savedState.Load()
if state == nil {
return nil
}
environment := make([]string, 0, len(state.servers)+len(state.search))
for _, server := range state.servers {
environment = append(environment, server.String())
}
return append(environment, state.search...)
}
func closeServerTransports(serverTransports []adapter.DNSTransport) {
for _, serverTransport := range serverTransports {
serverTransport.Close()
}
}
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
servers, err := t.fetch()
if err != nil {
return nil, err
}
if len(servers) == 0 {
return nil, E.New("dhcp: empty DNS servers from response")
}
return t.Exchange0(ctx, message, servers)
done := make(chan struct{})
var (
response *mDNS.Msg
err error
)
t.ExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
response = callbackResponse
err = callbackErr
close(done)
})
<-done
return response, err
}
func (t *Transport) Exchange0(ctx context.Context, message *mDNS.Msg, servers []M.Socksaddr) (*mDNS.Msg, error) {
question := message.Question[0]
domain := dns.FqdnToDomain(question.Name)
if len(servers) == 1 || !(message.Question[0].Qtype == mDNS.TypeA || message.Question[0].Qtype == mDNS.TypeAAAA) {
return t.exchangeSingleRequest(ctx, servers, message, domain)
} else {
return t.exchangeParallel(ctx, servers, message, domain)
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
state := t.savedState.Load()
if state == nil {
go t.exchangeCold(ctx, message, callback)
return
}
if state.lastError != nil {
callback(nil, E.Cause(state.lastError, "dhcp: fetch DNS servers"))
return
}
if len(state.serverTransports) == 0 {
go t.exchangeCold(ctx, message, callback)
return
}
if time.Since(state.updatedAt) >= C.DHCPTTL {
t.startRefresh()
}
t.exchangeWithTransports(ctx, message, state, callback)
}
func (t *Transport) exchangeCold(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
err := t.fetch()
if err != nil {
callback(nil, E.Cause(err, "dhcp: fetch DNS servers"))
return
}
state := t.savedState.Load()
if state == nil || len(state.serverTransports) == 0 {
callback(nil, E.New("dhcp: empty DNS servers from response"))
return
}
t.exchangeWithTransports(ctx, message, state, callback)
}
func (t *Transport) Fetch() []M.Socksaddr {
servers, _ := t.fetch()
return servers
state := t.savedState.Load()
if state == nil || state.lastError != nil {
return nil
}
if len(state.servers) > 0 && time.Since(state.updatedAt) >= C.DHCPTTL {
t.startRefresh()
}
return state.servers
}
func (t *Transport) fetch() ([]M.Socksaddr, error) {
t.transportLock.RLock()
updatedAt := t.updatedAt
lastError := t.lastError
servers := t.servers
t.transportLock.RUnlock()
if lastError != nil {
return nil, lastError
func (t *Transport) fetch() error {
state := t.savedState.Load()
if state != nil {
if state.lastError != nil {
return state.lastError
}
if time.Since(state.updatedAt) < C.DHCPTTL {
return nil
}
}
if time.Since(updatedAt) < C.DHCPTTL {
return servers, nil
t.refreshAccess.Lock()
defer t.refreshAccess.Unlock()
state = t.savedState.Load()
if state != nil {
if state.lastError != nil {
return state.lastError
}
if time.Since(state.updatedAt) < C.DHCPTTL {
return nil
}
}
t.transportLock.Lock()
defer t.transportLock.Unlock()
if time.Since(t.updatedAt) < C.DHCPTTL {
return t.servers, nil
return t.updateServersLocked(t.ctx)
}
func (t *Transport) startRefresh() {
if !t.refreshAccess.TryLock() {
return
}
err := t.updateServers()
if err != nil {
return servers, err
}
return t.servers, nil
go func() {
defer t.refreshAccess.Unlock()
state := t.savedState.Load()
if state != nil && time.Since(state.updatedAt) < C.DHCPTTL {
return
}
err := t.updateServersLocked(t.ctx)
if err != nil {
if errors.Is(err, errInterfaceIsCellular) && t.optional {
t.logger.Debug(E.Cause(err, "dhcp: refresh DNS servers"))
} else {
t.logger.Error(E.Cause(err, "dhcp: refresh DNS servers"))
}
}
}()
}
func (t *Transport) fetchInterface() (*control.Interface, error) {
@@ -171,44 +281,91 @@ func (t *Transport) fetchInterface() (*control.Interface, error) {
if t.networkManager.InterfaceMonitor() == nil {
return nil, E.New("missing monitor for auto DHCP, set route.auto_detect_interface")
}
defaultInterface := t.networkManager.InterfaceMonitor().DefaultInterface()
if defaultInterface == nil {
return nil, E.New("missing default interface")
if t.platformInterface != nil && t.platformInterface.UsePlatformNetworkInterfaces() {
defaultInterface := t.networkManager.DefaultNetworkInterface()
if defaultInterface == nil {
return nil, E.New("missing default interface")
}
if defaultInterface.Type == C.InterfaceTypeCellular {
return nil, errInterfaceIsCellular
}
return &defaultInterface.Interface, nil
} else {
defaultInterface := t.networkManager.InterfaceMonitor().DefaultInterface()
if defaultInterface == nil {
return nil, E.New("missing default interface")
}
return defaultInterface, nil
}
return defaultInterface, nil
} else {
return t.networkManager.InterfaceFinder().ByName(t.interfaceName)
}
}
func (t *Transport) updateServers() error {
func (t *Transport) updateServersLocked(ctx context.Context) error {
iface, err := t.fetchInterface()
if err != nil {
return E.Cause(err, "dhcp: prepare interface")
t.storeFailureLocked(err)
return E.Cause(err, "prepare interface")
}
t.logger.Notice("dhcp: query DNS servers on ", iface.Name)
fetchCtx, cancel := context.WithTimeout(t.ctx, C.DHCPTimeout)
fetchCtx, cancel := context.WithTimeout(ctx, C.DHCPTimeout)
err = t.fetchServers0(fetchCtx, iface)
cancel()
t.updatedAt = time.Now()
if err != nil {
t.lastError = err
if ctx.Err() != nil {
return err
}
t.storeFailureLocked(err)
return err
} else if len(t.servers) == 0 {
t.lastError = E.New("dhcp: empty DNS servers response")
return t.lastError
} else {
t.lastError = nil
return nil
}
state := t.savedState.Load()
if state == nil || len(state.servers) == 0 {
err = E.New("dhcp: empty DNS servers response")
t.storeFailureLocked(err)
return err
}
return nil
}
func (t *Transport) storeFailureLocked(err error) {
newState := &transportState{
updatedAt: time.Now(),
lastError: err,
}
previousState := t.savedState.Load()
if previousState != nil {
newState.search = previousState.search
newState.servers = previousState.servers
newState.serverTransports = previousState.serverTransports
}
t.savedState.Store(newState)
}
func (t *Transport) interfaceUpdated(defaultInterface *control.Interface, flags int) {
err := t.updateServers()
if err != nil {
t.logger.Error("update servers: ", err)
updateContext, updateCancel := context.WithCancel(t.ctx)
t.updateAccess.Lock()
previousCancel := t.updateCancel
t.updateCancel = updateCancel
t.updateAccess.Unlock()
if previousCancel != nil {
previousCancel()
}
go func() {
defer updateCancel()
t.refreshAccess.Lock()
err := t.updateServersLocked(updateContext)
t.refreshAccess.Unlock()
if err == nil || updateContext.Err() != nil {
return
}
if errors.Is(err, errInterfaceIsCellular) && t.optional {
t.logger.Debug(E.Cause(errInterfaceIsCellular, "dhcp: update DNS servers"))
} else {
t.logger.Error("dhcp: update DNS servers: ", err)
}
}()
}
func (t *Transport) fetchServers0(ctx context.Context, iface *control.Interface) error {
@@ -224,11 +381,15 @@ func (t *Transport) fetchServers0(ctx context.Context, iface *control.Interface)
err error
)
for range 5 {
packetConn, err = listener.ListenPacket(t.ctx, "udp4", listenAddr)
packetConn, err = listener.ListenPacket(ctx, "udp4", listenAddr)
if err == nil || !errors.Is(err, syscall.EADDRINUSE) {
break
}
time.Sleep(time.Second)
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(time.Second):
}
}
if err != nil {
return err
@@ -289,28 +450,55 @@ func (t *Transport) fetchServersResponse(iface *control.Interface, packetConn ne
continue
}
return t.recreateServers(iface, dhcpPacket)
return t.recreateServersLocked(iface, dhcpPacket)
}
}
func (t *Transport) recreateServers(iface *control.Interface, dhcpPacket *dhcpv4.DHCPv4) error {
func (t *Transport) recreateServersLocked(iface *control.Interface, dhcpPacket *dhcpv4.DHCPv4) error {
previousState := t.savedState.Load()
newState := &transportState{updatedAt: time.Now()}
if previousState != nil {
newState.search = previousState.search
}
searchList := dhcpPacket.DomainSearch()
if searchList != nil && len(searchList.Labels) > 0 {
t.search = common.Filter(common.Map(searchList.Labels, mDNS.Fqdn), func(it string) bool {
newState.search = common.Filter(common.Map(searchList.Labels, mDNS.Fqdn), func(it string) bool {
return it != "."
})
} else if dhcpPacket.DomainName() != "" {
domainName := mDNS.Fqdn(dhcpPacket.DomainName())
if domainName != "." {
t.search = []string{domainName}
newState.search = []string{domainName}
}
}
serverAddrs := common.Map(dhcpPacket.DNS(), func(it net.IP) M.Socksaddr {
newState.servers = common.Map(dhcpPacket.DNS(), func(it net.IP) M.Socksaddr {
return M.SocksaddrFrom(M.AddrFromIP(it), 53)
})
if len(serverAddrs) > 0 && !slices.Equal(t.servers, serverAddrs) {
t.logger.Notice("dhcp: updated DNS servers from ", iface.Name, ": [", strings.Join(common.Map(serverAddrs, M.Socksaddr.String), ","), "], search: [", strings.Join(t.search, ","), "]")
serversUnchanged := previousState != nil && slices.Equal(previousState.servers, newState.servers)
if len(newState.servers) > 0 && !serversUnchanged {
t.logger.Notice("dhcp: updated DNS servers from ", iface.Name, ": [", strings.Join(common.Map(newState.servers, M.Socksaddr.String), ","), "], search: [", strings.Join(newState.search, ","), "]")
}
if serversUnchanged && previousState.serverTransports != nil {
newState.serverTransports = previousState.serverTransports
t.savedState.Store(newState)
return nil
}
serverTransports := make([]adapter.DNSTransport, 0, len(newState.servers))
for _, serverAddr := range newState.servers {
serverTransport := transport.NewUDPRaw(t.logger, dns.NewTransportAdapter(C.DNSTypeUDP, "", nil), t.dialer, serverAddr)
err := serverTransport.Start(adapter.StartStateStart)
if err != nil {
for _, startedTransport := range serverTransports {
startedTransport.Close()
}
return E.Cause(err, "initialize transport for ", serverAddr)
}
serverTransports = append(serverTransports, serverTransport)
}
newState.serverTransports = serverTransports
t.savedState.Store(newState)
if previousState != nil {
closeServerTransports(previousState.serverTransports)
}
t.servers = serverAddrs
return nil
}
+27 -143
View File
@@ -2,165 +2,49 @@ package dhcp
import (
"context"
"errors"
"math/rand"
"strings"
"syscall"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
mDNS "github.com/miekg/dns"
)
func (t *Transport) exchangeSingleRequest(ctx context.Context, servers []M.Socksaddr, message *mDNS.Msg, domain string) (*mDNS.Msg, error) {
var lastErr error
for _, fqdn := range t.nameList(domain) {
response, err := t.tryOneName(ctx, servers, fqdn, message)
if err != nil {
lastErr = err
continue
}
return response, nil
func (t *Transport) exchangeWithTransports(ctx context.Context, message *mDNS.Msg, state *transportState, callback func(response *mDNS.Msg, err error)) {
question := message.Question[0]
domain := dns.FqdnToDomain(question.Name)
names := t.nameList(state.search, domain)
if len(names) == 0 {
callback(nil, E.New("invalid domain: ", domain))
return
}
return nil, lastErr
transport.ExchangeNames(ctx, names, question, func(fqdn string) transport.AsyncExchanger {
return t.newNameExchanger(message, fqdn, state.serverTransports)
}, callback)
}
func (t *Transport) exchangeParallel(ctx context.Context, servers []M.Socksaddr, message *mDNS.Msg, domain string) (*mDNS.Msg, error) {
returned := make(chan struct{})
defer close(returned)
type queryResult struct {
response *mDNS.Msg
err error
}
results := make(chan queryResult)
startRacer := func(ctx context.Context, fqdn string) {
response, err := t.tryOneName(ctx, servers, fqdn, message)
select {
case results <- queryResult{response, err}:
case <-returned:
func (t *Transport) newNameExchanger(message *mDNS.Msg, fqdn string, serverTransports []adapter.DNSTransport) transport.AsyncExchanger {
attemptExchangers := make([]transport.AsyncExchanger, 0, t.attempts*len(serverTransports))
for range t.attempts {
for _, serverTransport := range serverTransports {
attemptExchangers = append(attemptExchangers, func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
serverTransport.ExchangeAsync(ctx, transport.NewFanOutRequest(message, fqdn, true), callback)
})
}
}
queryCtx, queryCancel := context.WithCancel(ctx)
defer queryCancel()
var nameCount int
for _, fqdn := range t.nameList(domain) {
nameCount++
go startRacer(queryCtx, fqdn)
}
var errors []error
for {
select {
case <-ctx.Done():
return nil, ctx.Err()
case result := <-results:
if result.err == nil {
return result.response, nil
}
errors = append(errors, result.err)
if len(errors) == nameCount {
return nil, E.Errors(errors...)
}
}
}
}
func (t *Transport) tryOneName(ctx context.Context, servers []M.Socksaddr, fqdn string, message *mDNS.Msg) (*mDNS.Msg, error) {
sLen := len(servers)
var lastErr error
for i := 0; i < t.attempts; i++ {
for j := range sLen {
server := servers[j]
question := message.Question[0]
question.Name = fqdn
response, err := t.exchangeOne(ctx, server, question)
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
transport.ExchangeSequential(ctx, attemptExchangers, nil, func(response *mDNS.Msg, err error) {
if err != nil {
lastErr = err
continue
err = E.Cause(err, fqdn)
}
return response, nil
}
callback(response, err)
})
}
return nil, E.Cause(lastErr, fqdn)
}
func (t *Transport) exchangeOne(ctx context.Context, server M.Socksaddr, question mDNS.Question) (*mDNS.Msg, error) {
if server.Port == 0 {
server.Port = 53
}
request := &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: uint16(rand.Uint32()),
RecursionDesired: true,
AuthenticatedData: true,
},
Question: []mDNS.Question{question},
Compress: true,
}
request.SetEdns0(buf.UDPBufferSize, false)
return t.exchangeUDP(ctx, server, request)
}
func (t *Transport) exchangeUDP(ctx context.Context, server M.Socksaddr, request *mDNS.Msg) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkUDP, server)
if err != nil {
return nil, err
}
defer conn.Close()
if deadline, loaded := ctx.Deadline(); loaded && !deadline.IsZero() {
conn.SetDeadline(deadline)
}
buffer := buf.Get(buf.UDPBufferSize)
defer buf.Put(buffer)
rawMessage, err := request.PackBuffer(buffer)
if err != nil {
return nil, E.Cause(err, "pack request")
}
_, err = conn.Write(rawMessage)
if err != nil {
if errors.Is(err, syscall.EMSGSIZE) {
return t.exchangeTCP(ctx, server, request)
}
return nil, E.Cause(err, "write request")
}
n, err := conn.Read(buffer)
if err != nil {
if errors.Is(err, syscall.EMSGSIZE) {
return t.exchangeTCP(ctx, server, request)
}
return nil, E.Cause(err, "read response")
}
var response mDNS.Msg
err = response.Unpack(buffer[:n])
if err != nil {
return nil, E.Cause(err, "unpack response")
}
if response.Truncated {
return t.exchangeTCP(ctx, server, request)
}
return &response, nil
}
func (t *Transport) exchangeTCP(ctx context.Context, server M.Socksaddr, request *mDNS.Msg) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, server)
if err != nil {
return nil, err
}
defer conn.Close()
if deadline, loaded := ctx.Deadline(); loaded && !deadline.IsZero() {
conn.SetDeadline(deadline)
}
err = transport.WriteMessage(conn, 0, request)
if err != nil {
return nil, err
}
return transport.ReadMessage(conn)
}
func (t *Transport) nameList(name string) []string {
func (t *Transport) nameList(search []string, name string) []string {
l := len(name)
rooted := l > 0 && name[l-1] == '.'
if l > 254 || l == 254 && !rooted {
@@ -178,11 +62,11 @@ func (t *Transport) nameList(name string) []string {
name += "."
// l++
names := make([]string, 0, 1+len(t.search))
names := make([]string, 0, 1+len(search))
if hasNdots && !avoidDNS(name) {
names = append(names, name)
}
for _, suffix := range t.search {
for _, suffix := range search {
fqdn := name + suffix
if !avoidDNS(fqdn) && len(fqdn) <= 254 {
names = append(names, fqdn)
+157
View File
@@ -0,0 +1,157 @@
package transport
import (
"context"
"strings"
"sync"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
mDNS "github.com/miekg/dns"
)
type AsyncExchanger = func(ctx context.Context, callback func(response *mDNS.Msg, err error))
// ExchangeSequential tries exchangers in order until accept returns true
// (nil accept means err == nil); the last result is delivered as-is.
func ExchangeSequential(ctx context.Context, exchangers []AsyncExchanger, accept func(response *mDNS.Msg, err error) bool, callback func(response *mDNS.Msg, err error)) {
if len(exchangers) == 0 {
callback(nil, E.New("missing exchangers"))
return
}
if accept == nil {
accept = func(response *mDNS.Msg, err error) bool {
return err == nil
}
}
sequential := &sequentialExchange{
ctx: ctx,
exchangers: exchangers,
accept: accept,
callback: callback,
}
sequential.run(0)
}
type sequentialExchange struct {
ctx context.Context
exchangers []AsyncExchanger
accept func(response *mDNS.Msg, err error) bool
callback func(response *mDNS.Msg, err error)
}
func (s *sequentialExchange) run(index int) {
for index < len(s.exchangers) {
ctxErr := s.ctx.Err()
if ctxErr != nil {
s.callback(nil, ctxErr)
return
}
currentIndex := index
state := &sequentialCallState{}
s.exchangers[currentIndex](s.ctx, func(response *mDNS.Msg, err error) {
if currentIndex == len(s.exchangers)-1 || s.accept(response, err) {
s.callback(response, err)
return
}
state.access.Lock()
if state.returned {
state.access.Unlock()
s.run(currentIndex + 1)
return
}
state.continued = true
state.access.Unlock()
})
state.access.Lock()
state.returned = true
continued := state.continued
state.access.Unlock()
if !continued {
return
}
index = currentIndex + 1
}
}
type sequentialCallState struct {
access sync.Mutex
returned bool
continued bool
}
func ExchangeNames(ctx context.Context, names []string, question mDNS.Question, exchangerFor func(fqdn string) AsyncExchanger, callback func(response *mDNS.Msg, err error)) {
if len(names) == 0 {
callback(nil, E.New("missing name candidates"))
return
}
search := &nameSearchExchange{question: question}
nameExchangers := common.Map(names, func(fqdn string) AsyncExchanger {
return search.wrap(fqdn, exchangerFor(fqdn))
})
ExchangeSequential(ctx, nameExchangers, func(response *mDNS.Msg, err error) bool {
return err == nil && response.Rcode != mDNS.RcodeNameError
}, func(response *mDNS.Msg, err error) {
if err != nil || response.Rcode == mDNS.RcodeNameError {
search.access.Lock()
nameErrorResponse := search.nameErrorResponse
search.access.Unlock()
if nameErrorResponse != nil {
response, err = nameErrorResponse, nil
}
}
callback(response, err)
})
}
type nameSearchExchange struct {
question mDNS.Question
access sync.Mutex
nameErrorResponse *mDNS.Msg
}
func (s *nameSearchExchange) wrap(fqdn string, exchanger AsyncExchanger) AsyncExchanger {
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
exchanger(ctx, func(response *mDNS.Msg, err error) {
if err == nil {
restoreOriginalQuestion(response, fqdn, s.question)
if response.Rcode == mDNS.RcodeNameError {
s.access.Lock()
if s.nameErrorResponse == nil || fqdn == s.question.Name {
s.nameErrorResponse = response
}
s.access.Unlock()
}
}
callback(response, err)
})
}
}
// Stub resolvers discard Answer RRs whose owner name does not match the question.
func restoreOriginalQuestion(response *mDNS.Msg, fqdn string, question mDNS.Question) {
response.Question = []mDNS.Question{question}
for _, record := range response.Answer {
if strings.EqualFold(record.Header().Name, fqdn) {
record.Header().Name = question.Name
}
}
}
func NewFanOutRequest(message *mDNS.Msg, fqdn string, authenticatedData bool) *mDNS.Msg {
question := message.Question[0]
question.Name = fqdn
request := &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: message.Id,
RecursionDesired: true,
AuthenticatedData: authenticatedData,
},
Question: []mDNS.Question{question},
Compress: true,
}
request.SetEdns0(buf.UDPBufferSize, false)
return request
}
+4
View File
@@ -74,6 +74,10 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
return dns.FixedResponse(message.Id, question, []netip.Addr{address}, C.DefaultDNSTTL), nil
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
callback(t.Exchange(ctx, message))
}
func (t *Transport) Store() adapter.FakeIPStore {
return t.store
}
+4
View File
@@ -70,3 +70,7 @@ func (t *Transport) Reset() {
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
return t.strategy(ctx, message)
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
callback(t.Exchange(ctx, message))
}
+26 -3
View File
@@ -19,7 +19,10 @@ func RegisterTransport(registry *dns.TransportRegistry) {
dns.RegisterTransport[option.HostsDNSServerOptions](registry, C.DNSTypeHosts, NewTransport)
}
var _ adapter.DNSTransport = (*Transport)(nil)
var (
_ adapter.DNSTransport = (*Transport)(nil)
_ adapter.DNSTransportWithPreferredDomain = (*Transport)(nil)
)
type Transport struct {
dns.TransportAdapter
@@ -33,10 +36,14 @@ func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, opt
predefined = make(map[string][]netip.Addr)
)
if len(options.Path) == 0 {
files = append(files, NewFile(DefaultPath))
defaultFile, err := NewDefault()
if err != nil {
return nil, err
}
files = append(files, defaultFile)
} else {
for _, path := range options.Path {
files = append(files, NewFile(filemanager.BasePath(ctx, os.ExpandEnv(path))))
files = append(files, NewFile(ctx, filemanager.BasePath(ctx, os.ExpandEnv(path))))
}
}
if options.Predefined != nil {
@@ -62,6 +69,18 @@ func (t *Transport) Close() error {
func (t *Transport) Reset() {
}
func (t *Transport) PreferredDomain(domain string) bool {
if _, loaded := t.predefined[domain]; loaded {
return true
}
for _, file := range t.files {
if len(file.Lookup(domain)) > 0 {
return true
}
}
return false
}
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
question := message.Question[0]
domain := mDNS.CanonicalName(question.Name)
@@ -85,3 +104,7 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
Question: []mDNS.Question{question},
}, nil
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
callback(t.Exchange(ctx, message))
}
+17 -4
View File
@@ -2,20 +2,24 @@ package hosts
import (
"bufio"
"context"
"errors"
"io"
"net/netip"
"os"
"strings"
"sync"
"time"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/service/filemanager"
"github.com/miekg/dns"
)
const cacheMaxAge = 5 * time.Second
type File struct {
ctx context.Context
path string
access sync.Mutex
byName map[string][]netip.Addr
@@ -24,12 +28,21 @@ type File struct {
size int64
}
func NewFile(path string) *File {
func NewFile(ctx context.Context, path string) *File {
return &File{
ctx: ctx,
path: path,
}
}
func NewDefault() (*File, error) {
defaultPathResolved, err := defaultPath()
if err != nil {
return nil, E.Cause(err, "resolve default hosts path")
}
return NewFile(context.Background(), defaultPathResolved), nil
}
func (f *File) Lookup(name string) []netip.Addr {
f.access.Lock()
defer f.access.Unlock()
@@ -42,7 +55,7 @@ func (f *File) update() {
if now.Before(f.expire) && len(f.byName) > 0 {
return
}
stat, err := os.Stat(f.path)
stat, err := filemanager.Stat(f.ctx, f.path)
if err != nil {
return
}
@@ -51,7 +64,7 @@ func (f *File) update() {
return
}
byName := make(map[string][]netip.Addr)
file, err := os.Open(f.path)
file, err := filemanager.Open(f.ctx, f.path)
if err != nil {
return
}
+18 -4
View File
@@ -1,16 +1,30 @@
package hosts_test
package hosts
import (
"context"
"net/netip"
"os"
"runtime"
"testing"
"github.com/sagernet/sing-box/dns/transport/hosts"
E "github.com/sagernet/sing/common/exceptions"
"github.com/stretchr/testify/require"
)
func TestHosts(t *testing.T) {
t.Parallel()
require.Equal(t, []netip.Addr{netip.AddrFrom4([4]byte{127, 0, 0, 1}), netip.IPv6Loopback()}, hosts.NewFile("testdata/hosts").Lookup("localhost"))
require.NotEmpty(t, hosts.NewFile(hosts.DefaultPath).Lookup("localhost"))
require.Equal(t, []netip.Addr{netip.AddrFrom4([4]byte{127, 0, 0, 1}), netip.IPv6Loopback()}, NewFile(context.Background(), "testdata/hosts").Lookup("localhost"))
if runtime.GOOS != "windows" {
defaultPathResolved, err := defaultPath()
if err != nil {
t.Fatal(E.Cause(err, "resolve default hosts path"))
}
content, readErr := os.ReadFile(defaultPathResolved)
require.NoError(t, readErr)
hFile := NewFile(context.Background(), defaultPathResolved)
if len(hFile.Lookup("localhost")) == 0 {
t.Fatal("failed to resolve localhost: ", defaultPathResolved, ": \n", string(content))
}
}
}
+3 -1
View File
@@ -2,4 +2,6 @@
package hosts
var DefaultPath = "/etc/hosts"
func defaultPath() (string, error) {
return "/etc/hosts", nil
}
+5 -6
View File
@@ -2,16 +2,15 @@ package hosts
import (
"path/filepath"
"sync"
"golang.org/x/sys/windows"
)
var DefaultPath string
func init() {
var defaultPath = sync.OnceValues(func() (string, error) {
systemDirectory, err := windows.GetSystemDirectory()
if err != nil {
systemDirectory = "C:\\Windows\\System32"
return "", err
}
DefaultPath = filepath.Join(systemDirectory, "Drivers/etc/hosts")
}
return filepath.Join(systemDirectory, "Drivers", "etc", "hosts"), nil
})
+27 -10
View File
@@ -10,6 +10,7 @@ import (
"net/url"
"strconv"
"sync"
"sync/atomic"
"time"
"github.com/sagernet/sing-box/adapter"
@@ -45,6 +46,8 @@ type HTTPSTransport struct {
dialer N.Dialer
destination *url.URL
headers http.Header
serverAddr M.Socksaddr
fallback *atomic.Bool
transportAccess sync.Mutex
transport *HTTPSTransportWrapper
transportResetAt time.Time
@@ -123,13 +126,20 @@ func NewHTTPSRaw(
if tlsConfig != nil {
dialer = tls.NewDialer(dialer, tlsConfig)
}
fallback := new(atomic.Bool)
if destination.Scheme == "http" {
// plain HTTP DoH used by Tailscale
fallback.Store(true)
}
return &HTTPSTransport{
TransportAdapter: adapter,
logger: logger,
dialer: dialer,
destination: destination,
headers: headers,
transport: NewHTTPSTransportWrapper(dialer, serverAddr, destination),
serverAddr: serverAddr,
fallback: fallback,
transport: NewHTTPSTransportWrapper(dialer, serverAddr, fallback),
}
}
@@ -141,18 +151,21 @@ func (t *HTTPSTransport) Start(stage adapter.StartStage) error {
}
func (t *HTTPSTransport) Close() error {
t.transportAccess.Lock()
defer t.transportAccess.Unlock()
t.transport.CloseIdleConnections()
t.transport = t.transport.Clone()
t.Reset()
return nil
}
func (t *HTTPSTransport) Reset() {
t.transportAccess.Lock()
defer t.transportAccess.Unlock()
t.transport.CloseIdleConnections()
t.transport = t.transport.Clone()
t.resetTransportLocked()
}
func (t *HTTPSTransport) resetTransportLocked() {
oldTransport := t.transport
t.transport = NewHTTPSTransportWrapper(t.dialer, t.serverAddr, t.fallback)
t.transportResetAt = time.Now()
oldTransport.Close()
}
func (t *HTTPSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
@@ -165,15 +178,19 @@ func (t *HTTPSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS
if t.transportResetAt.After(startAt) {
return nil, err
}
t.transport.CloseIdleConnections()
t.transport = t.transport.Clone()
t.transportResetAt = time.Now()
t.resetTransportLocked()
}
return nil, err
}
return response, nil
}
func (t *HTTPSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
go func() {
callback(t.Exchange(ctx, message))
}()
}
func (t *HTTPSTransport) exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
exMessage := *message
exMessage.Id = 0
+78 -40
View File
@@ -5,7 +5,7 @@ import (
"errors"
"net"
"net/http"
"net/url"
"sync"
"sync/atomic"
"github.com/sagernet/sing-box/common/tls"
@@ -22,42 +22,50 @@ type HTTPSTransportWrapper struct {
http2Transport *http2.Transport
httpTransport *http.Transport
fallback *atomic.Bool
connAccess sync.Mutex
connections map[*httpsTrackedConn]struct{}
closed bool
}
func NewHTTPSTransportWrapper(dialer N.Dialer, serverAddr M.Socksaddr, destination *url.URL) *HTTPSTransportWrapper {
var fallback atomic.Bool
if destination.Scheme == "http" {
// plain HTTP DoH used by Tailscale
fallback.Store(true)
func NewHTTPSTransportWrapper(dialer N.Dialer, serverAddr M.Socksaddr, fallback *atomic.Bool) *HTTPSTransportWrapper {
wrapper := &HTTPSTransportWrapper{
fallback: fallback,
connections: make(map[*httpsTrackedConn]struct{}),
}
return &HTTPSTransportWrapper{
http2Transport: &http2.Transport{
DialTLSContext: func(ctx context.Context, _, _ string, _ *tls.STDConfig) (net.Conn, error) {
resultConn, err := dialer.DialContext(ctx, N.NetworkTCP, serverAddr)
if err != nil {
return nil, err
wrapper.http2Transport = &http2.Transport{
DialTLSContext: func(ctx context.Context, _, _ string, _ *tls.STDConfig) (net.Conn, error) {
resultConn, err := dialer.DialContext(ctx, N.NetworkTCP, serverAddr)
if err != nil {
return nil, err
}
if tlsConn, isTLSConn := resultConn.(tls.Conn); isTLSConn {
state := tlsConn.ConnectionState()
if state.NegotiatedProtocol != http2.NextProtoTLS {
tlsConn.Close()
fallback.Store(true)
return nil, errFallback
}
if tlsConn, isTLSConn := resultConn.(tls.Conn); isTLSConn {
state := tlsConn.ConnectionState()
if state.NegotiatedProtocol != http2.NextProtoTLS {
tlsConn.Close()
fallback.Store(true)
return nil, errFallback
}
}
return resultConn, nil
},
}
return wrapper.trackConn(resultConn)
},
httpTransport: &http.Transport{
DialContext: func(ctx context.Context, _, addr string) (net.Conn, error) {
return dialer.DialContext(ctx, N.NetworkTCP, serverAddr)
},
DialTLSContext: func(ctx context.Context, _, _ string) (net.Conn, error) {
return dialer.DialContext(ctx, N.NetworkTCP, serverAddr)
},
},
fallback: &fallback,
}
wrapper.httpTransport = &http.Transport{
DialContext: func(ctx context.Context, _, addr string) (net.Conn, error) {
resultConn, err := dialer.DialContext(ctx, N.NetworkTCP, serverAddr)
if err != nil {
return nil, err
}
return wrapper.trackConn(resultConn)
},
DialTLSContext: func(ctx context.Context, _, _ string) (net.Conn, error) {
resultConn, err := dialer.DialContext(ctx, N.NetworkTCP, serverAddr)
if err != nil {
return nil, err
}
return wrapper.trackConn(resultConn)
},
}
return wrapper
}
func (h *HTTPSTransportWrapper) RoundTrip(request *http.Request) (*http.Response, error) {
@@ -74,17 +82,47 @@ func (h *HTTPSTransportWrapper) RoundTrip(request *http.Request) (*http.Response
return response, nil
}
func (h *HTTPSTransportWrapper) CloseIdleConnections() {
func (h *HTTPSTransportWrapper) trackConn(conn net.Conn) (net.Conn, error) {
trackedConn := &httpsTrackedConn{Conn: conn, wrapper: h}
h.connAccess.Lock()
if h.closed {
h.connAccess.Unlock()
conn.Close()
return nil, net.ErrClosed
}
h.connections[trackedConn] = struct{}{}
h.connAccess.Unlock()
return trackedConn, nil
}
func (h *HTTPSTransportWrapper) Close() {
h.connAccess.Lock()
if h.closed {
h.connAccess.Unlock()
return
}
h.closed = true
connections := make([]*httpsTrackedConn, 0, len(h.connections))
for trackedConn := range h.connections {
connections = append(connections, trackedConn)
}
h.connections = nil
h.connAccess.Unlock()
for _, trackedConn := range connections {
trackedConn.Conn.Close()
}
h.http2Transport.CloseIdleConnections()
h.httpTransport.CloseIdleConnections()
}
func (h *HTTPSTransportWrapper) Clone() *HTTPSTransportWrapper {
return &HTTPSTransportWrapper{
httpTransport: h.httpTransport,
http2Transport: &http2.Transport{
DialTLSContext: h.http2Transport.DialTLSContext,
},
fallback: h.fallback,
}
type httpsTrackedConn struct {
net.Conn
wrapper *HTTPSTransportWrapper
}
func (c *httpsTrackedConn) Close() error {
c.wrapper.connAccess.Lock()
delete(c.wrapper.connections, c)
c.wrapper.connAccess.Unlock()
return c.Conn.Close()
}
+116 -39
View File
@@ -1,16 +1,18 @@
//go:build !darwin
package local
import (
"context"
"sync"
"sync/atomic"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport/hosts"
"github.com/sagernet/sing-box/dns/transport/local/systemconfig"
"github.com/sagernet/sing-box/dns/transport/mdns"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
@@ -22,16 +24,25 @@ func RegisterTransport(registry *dns.TransportRegistry) {
dns.RegisterTransport[option.LocalDNSServerOptions](registry, C.DNSTypeLocal, NewTransport)
}
var _ adapter.DNSTransport = (*Transport)(nil)
var (
_ adapter.DNSTransport = (*Transport)(nil)
_ adapter.DNSTransportWithPreferredDomain = (*Transport)(nil)
_ adapter.DNSTransportWithEnvironment = (*Transport)(nil)
)
type Transport struct {
dns.TransportAdapter
ctx context.Context
logger logger.ContextLogger
hosts *hosts.File
dialer N.Dialer
preferGo bool
resolved ResolvedResolver
ctx context.Context
logger logger.ContextLogger
preferredResolver *PreferredDomainResolver
dialer N.Dialer
preferGo bool
resolved ResolvedResolver
mdnsTransport adapter.DNSTransport
configSource *systemconfig.Source
system systemResolver
serverSet atomic.Pointer[localServerSet]
serverSetAccess sync.Mutex
}
func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, options option.LocalDNSServerOptions) (adapter.DNSTransport, error) {
@@ -39,56 +50,122 @@ func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, opt
if err != nil {
return nil, err
}
preferredResolver, err := NewPreferredDomainResolver(ctx, logger, options)
if err != nil {
return nil, err
}
return &Transport{
TransportAdapter: dns.NewTransportAdapterWithLocalOptions(C.DNSTypeLocal, tag, options),
ctx: ctx,
logger: logger,
hosts: hosts.NewFile(hosts.DefaultPath),
dialer: transportDialer,
preferGo: options.PreferGo,
TransportAdapter: dns.NewTransportAdapterWithLocalOptions(C.DNSTypeLocal, tag, options),
ctx: ctx,
logger: logger,
preferredResolver: preferredResolver,
dialer: transportDialer,
preferGo: options.PreferGo,
configSource: systemconfig.NewSource(ctx),
}, nil
}
func (t *Transport) Start(stage adapter.StartStage) error {
t.preferredResolver.Start(stage)
switch stage {
case adapter.StartStateInitialize:
if !t.preferGo {
if isSystemdResolvedManaged() {
resolvedResolver, err := NewResolvedResolver(t.ctx, t.logger)
if !t.preferGo && isSystemdResolvedManaged() {
resolvedResolver, err := NewResolvedResolver(t.ctx, t.logger)
if err == nil {
err = resolvedResolver.Start()
if err == nil {
err = resolvedResolver.Start()
if err == nil {
t.resolved = resolvedResolver
} else {
t.logger.Warn(E.Cause(err, "initialize resolved resolver"))
}
t.resolved = resolvedResolver
} else {
t.logger.Warn(E.Cause(err, "initialize resolved resolver"))
}
}
}
case adapter.StartStateStart:
if !C.IsDarwin {
t.mdnsTransport = mdns.NewRawTransport(t.TransportAdapter, t.ctx, t.logger)
}
fallthrough
default:
if t.mdnsTransport != nil {
err := t.mdnsTransport.Start(stage)
if err != nil {
return err
}
}
}
return nil
}
func (t *Transport) Close() error {
if t.resolved != nil {
return t.resolved.Close()
serverSet := t.serverSet.Swap(nil)
if serverSet != nil {
serverSet.Close()
}
return nil
t.system.close()
return common.Close(t.resolved, t.mdnsTransport, t.configSource)
}
func (t *Transport) Reset() {
serverSet := t.serverSet.Load()
if serverSet != nil {
for _, serverTransport := range serverSet.transports {
serverTransport.Reset()
}
}
t.system.reset()
t.configSource.Reset()
if t.resolved != nil {
t.resolved.Reset()
}
if t.mdnsTransport != nil {
t.mdnsTransport.Reset()
}
}
func (t *Transport) PreferredDomain(domain string) bool {
return t.preferredResolver.PreferredDomain(domain)
}
func (t *Transport) Environment() []string {
if t.resolved != nil {
return t.resolved.Environment()
}
return t.configSource.Configuration().Signature()
}
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
if t.resolved != nil {
return t.resolved.Exchange(ctx, message)
}
question := message.Question[0]
if question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA {
addresses := t.hosts.Lookup(dns.FqdnToDomain(question.Name))
if len(addresses) > 0 {
return dns.FixedResponse(message.Id, question, addresses, C.DefaultDNSTTL), nil
}
}
return t.exchange(ctx, message, question.Name)
done := make(chan struct{})
var (
response *mDNS.Msg
err error
)
t.ExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
response = callbackResponse
err = callbackErr
close(done)
})
<-done
return response, err
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
question := message.Question[0]
response := t.preferredResolver.Lookup(message)
if response != nil {
callback(response, nil)
return
}
if mdns.IsLocalDomain(question.Name) {
if C.IsDarwin {
t.systemExchangeAsync(ctx, message, callback)
return
}
t.mdnsTransport.ExchangeAsync(ctx, message, callback)
return
}
if t.resolved != nil {
t.resolved.ExchangeAsync(ctx, message, callback)
return
}
t.exchangeAsync(ctx, message, question.Name, callback)
}
+494 -106
View File
@@ -3,138 +3,526 @@
package local
import (
"cmp"
"context"
"encoding/binary"
"errors"
"io"
"net"
"os"
"sync"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport/hosts"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
dnsTransport "github.com/sagernet/sing-box/dns/transport"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/service"
mDNS "github.com/miekg/dns"
)
func RegisterTransport(registry *dns.TransportRegistry) {
dns.RegisterTransport[option.LocalDNSServerOptions](registry, C.DNSTypeLocal, NewTransport)
}
var _ adapter.DNSTransport = (*Transport)(nil)
type Transport struct {
dns.TransportAdapter
ctx context.Context
logger logger.ContextLogger
hosts *hosts.File
dialer N.Dialer
preferGo bool
fallback bool
dhcpTransport dhcpTransport
resolver net.Resolver
}
type dhcpTransport interface {
adapter.DNSTransport
Fetch() []M.Socksaddr
Exchange0(ctx context.Context, message *mDNS.Msg, servers []M.Socksaddr) (*mDNS.Msg, error)
}
func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, options option.LocalDNSServerOptions) (adapter.DNSTransport, error) {
transportDialer, err := dns.NewLocalDialer(ctx, options)
if err != nil {
return nil, err
}
transportAdapter := dns.NewTransportAdapterWithLocalOptions(C.DNSTypeLocal, tag, options)
return &Transport{
TransportAdapter: transportAdapter,
ctx: ctx,
logger: logger,
hosts: hosts.NewFile(hosts.DefaultPath),
dialer: transportDialer,
preferGo: options.PreferGo,
}, nil
}
func (t *Transport) Start(stage adapter.StartStage) error {
if stage != adapter.StartStateStart {
return nil
}
inboundManager := service.FromContext[adapter.InboundManager](t.ctx)
for _, inbound := range inboundManager.Inbounds() {
if inbound.Type() == C.TypeTun {
t.fallback = true
break
}
}
if t.fallback {
t.dhcpTransport = newDHCPTransport(t.TransportAdapter, log.ContextWithOverrideLevel(t.ctx, log.LevelDebug), t.dialer, t.logger)
if t.dhcpTransport != nil {
err := t.dhcpTransport.Start(stage)
if err != nil {
return err
func (t *Transport) systemExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
question := message.Question[0]
t.system.exchangeAsync(ctx, question.Name, question.Qtype, question.Qclass, func(response *mDNS.Msg, err error) {
if err != nil {
var rcodeError dns.RcodeError
if errors.As(err, &rcodeError) {
callback(dns.FixedResponseStatus(message, int(rcodeError)), nil)
return
}
callback(nil, err)
return
}
response.Id = message.Id
response.Response = true
response.RecursionAvailable = true
callback(response, nil)
})
}
// The mDNSResponder daemon speaks an undocumented binary protocol over a
// AF_UNIX SOCK_STREAM socket. The framing below is taken from the client
// stub of Apple's open-source mDNSResponder (mDNSShared/dnssd_ipc.h,
// dnssd_clientstub.c and uds_daemon.c). All multi-byte fields are
// big-endian. A connection opened with connection_request acts as a shared
// connection (DNSServiceCreateConnection): subsequent requests on the same
// stream carry a unique client_context in header bytes 16-24, which the
// daemon echoes back in every reply, allowing concurrent queries to be
// demultiplexed. With IPC_FLAGS_NOERRSD set the daemon does not expect the
// SCM_RIGHTS error-return socket used by Apple's stub; request errors are
// instead delivered as async_error_op replies, and success produces no
// acknowledgment at all. A query is cancelled by sending cancel_request
// with the same client_context and no payload.
const (
mdnsResponderSocketPath = "/var/run/mDNSResponder"
mdnsResponderSocketEnv = "DNSSD_UDS_PATH"
mdnsResponderVersion = 1
mdnsResponderHeaderLength = 28
mdnsResponderConnectionRequest = 1 // connection_request
mdnsResponderQueryRequest = 8 // query_request
mdnsResponderCancelRequest = 63 // cancel_request
mdnsResponderQueryReply = 68 // query_reply_op
mdnsResponderAsyncErrorReply = 73 // async_error_op
mdnsResponderFlagMoreComing = 0x1
mdnsResponderFlagAdd = 0x2
mdnsResponderFlagReturnIntermediates = 0x1000
mdnsResponderFlagShareConnection = 0x4000
mdnsResponderFlagTimeout = 0x10000
mdnsResponderIPCFlagNoErrorSocket = 0x4 // IPC_FLAGS_NOERRSD
mdnsResponderErrNoError = 0
mdnsResponderErrNoSuchName = -65538
mdnsResponderErrNoSuchRecord = -65554
mdnsResponderErrTimeout = -65568
mdnsResponderMaxReplyLength = 1 << 20
)
type systemResolver struct {
initOnce sync.Once
connection *dnsTransport.ConnPool[net.Conn]
queryAccess sync.Mutex
queryId uint64
queries map[uint64]*systemPendingQuery
}
type systemPendingQuery struct {
conn net.Conn
name string
qtype uint16
qclass uint16
answers []mDNS.RR
hasFinalAnswer bool
ready bool
callback func(response *mDNS.Msg, err error)
stopContext func() bool
stopConn func() bool
}
type systemCompletion struct {
pending *systemPendingQuery
err error
}
func (r *systemResolver) init() {
r.queries = make(map[uint64]*systemPendingQuery)
r.connection = dnsTransport.NewConnPool(dnsTransport.ConnPoolOptions[net.Conn]{
Mode: dnsTransport.ConnPoolSingle,
IsAlive: func(conn net.Conn) bool {
return conn != nil
},
Close: func(conn net.Conn, cause error) {
conn.Close()
},
})
}
func (r *systemResolver) close() {
r.initOnce.Do(r.init)
_ = r.connection.Close()
}
func (r *systemResolver) reset() {
r.initOnce.Do(r.init)
r.connection.Reset()
}
func (r *systemResolver) exchangeAsync(ctx context.Context, name string, qtype uint16, qclass uint16, callback func(response *mDNS.Msg, err error)) {
r.initOnce.Do(r.init)
for firstAttempt := true; ; firstAttempt = false {
conn, connCtx, created, err := r.connection.AcquireShared(ctx, r.dial)
if err != nil {
callback(nil, err)
return
}
if created {
go r.recvLoop(conn)
}
queryId := r.register(ctx, connCtx, conn, name, qtype, qclass, callback)
_, writeErr := conn.Write(buildQueryRequest(queryId, name, qtype, qclass))
if writeErr == nil {
return
}
pending := r.take(queryId)
r.connection.Invalidate(conn, writeErr)
if pending == nil {
return
}
if !created && firstAttempt {
continue
}
callback(nil, E.Cause(writeErr, "write mDNSResponder query"))
return
}
}
func (r *systemResolver) dial(ctx context.Context) (net.Conn, error) {
socketPath := cmp.Or(os.Getenv(mdnsResponderSocketEnv), mdnsResponderSocketPath)
var dialer net.Dialer
conn, err := dialer.DialContext(ctx, "unix", socketPath)
if err != nil {
return nil, E.Cause(err, "connect mDNSResponder")
}
stopCancel := context.AfterFunc(ctx, func() {
conn.Close()
})
err = writeConnectionRequest(conn)
stopCancel()
if err != nil {
conn.Close()
return nil, contextError(ctx, err)
}
return conn, nil
}
func writeConnectionRequest(conn net.Conn) error {
_, err := conn.Write(appendResponderHeader(make([]byte, 0, mdnsResponderHeaderLength), mdnsResponderConnectionRequest, 0, 0, 0))
if err != nil {
return E.Cause(err, "write mDNSResponder connection request")
}
var status [4]byte
_, err = io.ReadFull(conn, status[:])
if err != nil {
return E.Cause(err, "read mDNSResponder connection status")
}
statusCode := int32(binary.BigEndian.Uint32(status[:]))
if statusCode != mdnsResponderErrNoError {
return E.New("mDNSResponder connection request failed: error ", statusCode)
}
return nil
}
func (t *Transport) Close() error {
return common.Close(
t.dhcpTransport,
)
func (r *systemResolver) register(ctx context.Context, connCtx context.Context, conn net.Conn, name string, qtype uint16, qclass uint16, callback func(response *mDNS.Msg, err error)) uint64 {
r.queryAccess.Lock()
defer r.queryAccess.Unlock()
r.queryId++
queryId := r.queryId
pending := &systemPendingQuery{
conn: conn,
name: name,
qtype: qtype,
qclass: qclass,
callback: callback,
}
r.queries[queryId] = pending
pending.stopContext = context.AfterFunc(ctx, func() {
r.cancelQuery(queryId, ctx)
})
pending.stopConn = context.AfterFunc(connCtx, func() {
r.completeConnClosed(queryId, connCtx)
})
return queryId
}
func (t *Transport) Reset() {
if t.dhcpTransport != nil {
t.dhcpTransport.Reset()
func (r *systemResolver) take(queryId uint64) *systemPendingQuery {
r.queryAccess.Lock()
pending, loaded := r.queries[queryId]
if !loaded {
r.queryAccess.Unlock()
return nil
}
delete(r.queries, queryId)
r.queryAccess.Unlock()
pending.stopContext()
pending.stopConn()
return pending
}
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
question := message.Question[0]
if question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA {
addresses := t.hosts.Lookup(dns.FqdnToDomain(question.Name))
if len(addresses) > 0 {
return dns.FixedResponse(message.Id, question, addresses, C.DefaultDNSTTL), nil
}
func (r *systemResolver) cancelQuery(queryId uint64, ctx context.Context) {
pending := r.take(queryId)
if pending == nil {
return
}
if !t.fallback {
return t.exchange(ctx, message, question.Name)
_, err := pending.conn.Write(appendResponderHeader(make([]byte, 0, mdnsResponderHeaderLength), mdnsResponderCancelRequest, 0, queryId, 0))
if err != nil {
r.connection.Invalidate(pending.conn, err)
} else {
r.connection.Release(pending.conn, true)
}
if t.dhcpTransport != nil {
dhcpTransports := t.dhcpTransport.Fetch()
if len(dhcpTransports) > 0 {
return t.dhcpTransport.Exchange0(ctx, message, dhcpTransports)
}
pending.callback(nil, ctx.Err())
}
func (r *systemResolver) completeConnClosed(queryId uint64, connCtx context.Context) {
pending := r.take(queryId)
if pending == nil {
return
}
if t.preferGo {
// Assuming the user knows what they are doing, we still execute the query which will fail.
return t.exchange(ctx, message, question.Name)
pending.callback(nil, context.Cause(connCtx))
}
func (r *systemResolver) finish(pending *systemPendingQuery, err error) {
pending.stopContext()
pending.stopConn()
r.connection.Release(pending.conn, true)
if err != nil {
pending.callback(nil, err)
return
}
if question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA {
var network string
if question.Qtype == mDNS.TypeA {
network = "ip4"
} else {
network = "ip6"
}
addresses, err := t.resolver.LookupNetIP(ctx, network, question.Name)
pending.callback(&mDNS.Msg{
Question: []mDNS.Question{{Name: mDNS.Fqdn(pending.name), Qtype: pending.qtype, Qclass: pending.qclass}},
Answer: pending.answers,
}, nil)
}
func (r *systemResolver) recvLoop(conn net.Conn) {
for {
operation, clientContext, data, err := readResponderReply(conn)
if err != nil {
var dnsError *net.DNSError
if errors.As(err, &dnsError) && dnsError.IsNotFound {
return nil, dns.RcodeRefused
}
return nil, err
r.connection.Invalidate(conn, err)
return
}
switch operation {
case mdnsResponderQueryReply:
reply, parseErr := parseResponderReply(data)
if parseErr != nil {
r.connection.Invalidate(conn, parseErr)
return
}
r.handleQueryReply(clientContext, reply)
case mdnsResponderAsyncErrorReply:
if len(data) >= 12 {
r.completeQueryError(clientContext, binary.BigEndian.Uint32(data[0:4]), int32(binary.BigEndian.Uint32(data[8:12])))
}
}
return dns.FixedResponse(message.Id, question, addresses, C.DefaultDNSTTL), nil
}
return nil, E.New("only A and AAAA queries are supported on Apple platforms when using TUN and DHCP unavailable.")
}
// On a shared connection MoreComing applies collectively to all operations
// (dns_sd.h "Collective kDNSServiceFlagsMoreComing flag"): the daemon sets it
// whenever another reply, for any query, is queued behind this one. A reply
// without it is therefore a connection-wide flush point, at which every query
// that already collected its final answer is completed.
func (r *systemResolver) handleQueryReply(queryId uint64, reply mdnsResponderReply) {
var completions []systemCompletion
r.queryAccess.Lock()
pending, loaded := r.queries[queryId]
if loaded {
if reply.errorCode != mdnsResponderErrNoError {
delete(r.queries, queryId)
if len(pending.answers) > 0 {
completions = append(completions, systemCompletion{pending: pending})
} else {
completions = append(completions, systemCompletion{pending: pending, err: darwinResolverError(pending.name, reply.errorCode)})
}
} else {
if reply.flags&mdnsResponderFlagAdd != 0 && len(reply.rdata) > 0 {
record, buildErr := buildResourceRecord(reply)
if buildErr == nil {
pending.answers = append(pending.answers, record)
if record.Header().Rrtype == pending.qtype {
pending.hasFinalAnswer = true
}
}
}
if pending.hasFinalAnswer && reply.rrtype == pending.qtype {
pending.ready = true
}
}
}
if reply.flags&mdnsResponderFlagMoreComing == 0 {
completions = r.collectReadyLocked(completions)
}
r.queryAccess.Unlock()
for _, completion := range completions {
r.finish(completion.pending, completion.err)
}
}
func (r *systemResolver) completeQueryError(queryId uint64, flags uint32, errorCode int32) {
var completions []systemCompletion
r.queryAccess.Lock()
pending, loaded := r.queries[queryId]
if loaded {
delete(r.queries, queryId)
completions = append(completions, systemCompletion{pending: pending, err: darwinResolverError(pending.name, errorCode)})
}
if flags&mdnsResponderFlagMoreComing == 0 {
completions = r.collectReadyLocked(completions)
}
r.queryAccess.Unlock()
for _, completion := range completions {
r.finish(completion.pending, completion.err)
}
}
func (r *systemResolver) collectReadyLocked(completions []systemCompletion) []systemCompletion {
for queryId, pending := range r.queries {
if pending.ready {
delete(r.queries, queryId)
completions = append(completions, systemCompletion{pending: pending})
}
}
return completions
}
func appendResponderHeader(buffer []byte, operation uint32, dataLength int, clientContext uint64, ipcFlags uint32) []byte {
buffer = binary.BigEndian.AppendUint32(buffer, mdnsResponderVersion)
buffer = binary.BigEndian.AppendUint32(buffer, uint32(dataLength))
buffer = binary.BigEndian.AppendUint32(buffer, ipcFlags)
buffer = binary.BigEndian.AppendUint32(buffer, operation)
buffer = binary.BigEndian.AppendUint64(buffer, clientContext)
buffer = binary.BigEndian.AppendUint32(buffer, 0) // reg_index
return buffer
}
func buildQueryRequest(queryId uint64, name string, qtype uint16, qclass uint16) []byte {
payloadLength := 4 + 4 + len(name) + 1 + 2 + 2
message := make([]byte, 0, mdnsResponderHeaderLength+payloadLength)
message = appendResponderHeader(message, mdnsResponderQueryRequest, payloadLength, queryId, mdnsResponderIPCFlagNoErrorSocket)
message = binary.BigEndian.AppendUint32(message, mdnsResponderFlagShareConnection|mdnsResponderFlagReturnIntermediates|mdnsResponderFlagTimeout)
message = binary.BigEndian.AppendUint32(message, 0) // interfaceIndex
message = append(message, name...)
message = append(message, 0)
message = binary.BigEndian.AppendUint16(message, qtype)
message = binary.BigEndian.AppendUint16(message, qclass)
return message
}
func readResponderReply(conn net.Conn) (operation uint32, clientContext uint64, data []byte, err error) {
var header [mdnsResponderHeaderLength]byte
_, err = io.ReadFull(conn, header[:])
if err != nil {
return
}
dataLength := binary.BigEndian.Uint32(header[4:8])
if dataLength > mdnsResponderMaxReplyLength {
err = E.New("oversized mDNSResponder reply: ", dataLength)
return
}
operation = binary.BigEndian.Uint32(header[12:16])
clientContext = binary.BigEndian.Uint64(header[16:24])
data = make([]byte, dataLength)
_, err = io.ReadFull(conn, data)
return
}
type mdnsResponderReply struct {
flags uint32
errorCode int32
name string
rrtype uint16
rrclass uint16
ttl uint32
rdata []byte
}
func parseResponderReply(data []byte) (mdnsResponderReply, error) {
var reply mdnsResponderReply
reader := replyReader{data: data}
reply.flags = reader.uint32()
reader.uint32() // interfaceIndex
reply.errorCode = int32(reader.uint32())
reply.name = reader.cString()
reply.rrtype = reader.uint16()
reply.rrclass = reader.uint16()
rdlen := reader.uint16()
reply.rdata = reader.bytes(int(rdlen))
reply.ttl = reader.uint32()
if reader.err != nil {
return reply, reader.err
}
return reply, nil
}
func buildResourceRecord(reply mdnsResponderReply) (mDNS.RR, error) {
name := mDNS.Fqdn(reply.name)
nameBuffer := make([]byte, 256)
offset, err := mDNS.PackDomainName(name, nameBuffer, 0, nil, false)
if err != nil {
return nil, err
}
record := make([]byte, 0, offset+10+len(reply.rdata))
record = append(record, nameBuffer[:offset]...)
record = binary.BigEndian.AppendUint16(record, reply.rrtype)
record = binary.BigEndian.AppendUint16(record, reply.rrclass)
record = binary.BigEndian.AppendUint32(record, reply.ttl)
record = binary.BigEndian.AppendUint16(record, uint16(len(reply.rdata)))
record = append(record, reply.rdata...)
resourceRecord, _, err := mDNS.UnpackRR(record, 0)
if err != nil {
return nil, err
}
return resourceRecord, nil
}
// The daemon's NoSuchRecord conflates NXDOMAIN and NODATA, so it is reported as
// an empty NOERROR to avoid a false NXDOMAIN.
func darwinResolverError(name string, code int32) error {
switch code {
case mdnsResponderErrNoSuchRecord:
return dns.RcodeSuccess
case mdnsResponderErrNoSuchName:
return dns.RcodeNameError
case mdnsResponderErrTimeout:
return E.New("mDNSResponder query timeout for ", name)
default:
return E.New("mDNSResponder query failed for ", name, ": error ", code)
}
}
func contextError(ctx context.Context, err error) error {
ctxErr := ctx.Err()
if ctxErr != nil {
return ctxErr
}
return err
}
type replyReader struct {
data []byte
offset int
err error
}
func (r *replyReader) uint32() uint32 {
if r.err != nil || r.offset+4 > len(r.data) {
r.fail()
return 0
}
value := binary.BigEndian.Uint32(r.data[r.offset:])
r.offset += 4
return value
}
func (r *replyReader) uint16() uint16 {
if r.err != nil || r.offset+2 > len(r.data) {
r.fail()
return 0
}
value := binary.BigEndian.Uint16(r.data[r.offset:])
r.offset += 2
return value
}
func (r *replyReader) cString() string {
if r.err != nil {
return ""
}
end := r.offset
for end < len(r.data) && r.data[end] != 0 {
end++
}
if end >= len(r.data) {
r.fail()
return ""
}
value := string(r.data[r.offset:end])
r.offset = end + 1
return value
}
func (r *replyReader) bytes(length int) []byte {
if r.err != nil || length < 0 || r.offset+length > len(r.data) {
r.fail()
return nil
}
value := r.data[r.offset : r.offset+length]
r.offset += length
return value
}
func (r *replyReader) fail() {
if r.err == nil {
r.err = E.New("truncated mDNSResponder reply")
}
}
-16
View File
@@ -1,16 +0,0 @@
//go:build darwin && with_dhcp
package local
import (
"context"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport/dhcp"
"github.com/sagernet/sing-box/log"
N "github.com/sagernet/sing/common/network"
)
func newDHCPTransport(transportAdapter dns.TransportAdapter, ctx context.Context, dialer N.Dialer, logger log.ContextLogger) dhcpTransport {
return dhcp.NewRawTransport(transportAdapter, ctx, dialer, logger)
}
@@ -1,15 +0,0 @@
//go:build darwin && !with_dhcp
package local
import (
"context"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/log"
N "github.com/sagernet/sing/common/network"
)
func newDHCPTransport(transportAdapter dns.TransportAdapter, ctx context.Context, dialer N.Dialer, logger log.ContextLogger) dhcpTransport {
return nil
}
+161
View File
@@ -0,0 +1,161 @@
//go:build darwin
package local
import (
"cmp"
"context"
"net"
"os"
"sync"
"testing"
"time"
mDNS "github.com/miekg/dns"
)
// "localhost" is answered by the mDNSResponder daemon itself.
func requireMDNSResponder(t *testing.T) {
t.Helper()
socketPath := cmp.Or(os.Getenv(mdnsResponderSocketEnv), mdnsResponderSocketPath)
conn, err := net.DialTimeout("unix", socketPath, time.Second)
if err != nil {
t.Skipf("mDNSResponder not reachable at %s: %v", socketPath, err)
}
conn.Close()
}
func systemExchangeForTest(ctx context.Context, transport *Transport, message *mDNS.Msg) (*mDNS.Msg, error) {
done := make(chan struct{})
var (
response *mDNS.Msg
err error
)
transport.systemExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
response = callbackResponse
err = callbackErr
close(done)
})
<-done
return response, err
}
func TestSystemExchangeLoopback(t *testing.T) {
requireMDNSResponder(t)
transport := &Transport{}
defer transport.system.close()
for _, testCase := range []struct {
qtype uint16
expected net.IP
}{
{mDNS.TypeA, net.IPv4(127, 0, 0, 1)},
{mDNS.TypeAAAA, net.IPv6loopback},
} {
message := new(mDNS.Msg)
message.SetQuestion("localhost.", testCase.qtype)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
response, err := systemExchangeForTest(ctx, transport, message)
cancel()
if err != nil {
t.Fatalf("%s localhost: %v", mDNS.TypeToString[testCase.qtype], err)
}
if response.Id != message.Id {
t.Fatalf("%s response id %d != request id %d", mDNS.TypeToString[testCase.qtype], response.Id, message.Id)
}
if !response.Response {
t.Fatalf("%s response flag not set", mDNS.TypeToString[testCase.qtype])
}
var found bool
for _, answer := range response.Answer {
switch record := answer.(type) {
case *mDNS.A:
found = found || record.A.Equal(testCase.expected)
case *mDNS.AAAA:
found = found || record.AAAA.Equal(testCase.expected)
}
}
if !found {
t.Fatalf("%s localhost: expected %s in answer, got %v", mDNS.TypeToString[testCase.qtype], testCase.expected, response.Answer)
}
}
}
func TestSystemExchangeNoData(t *testing.T) {
requireMDNSResponder(t)
transport := &Transport{}
defer transport.system.close()
message := new(mDNS.Msg)
// localhost has no MX record, so the daemon reports NoSuchRecord.
message.SetQuestion("localhost.", mDNS.TypeMX)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
response, err := systemExchangeForTest(ctx, transport, message)
if err != nil {
t.Fatalf("MX localhost: %v", err)
}
if response.Rcode != mDNS.RcodeSuccess {
t.Fatalf("MX localhost: rcode %s, want NOERROR", mDNS.RcodeToString[response.Rcode])
}
if len(response.Answer) != 0 {
t.Fatalf("MX localhost: expected no answers, got %v", response.Answer)
}
}
func TestSystemExchangeCancel(t *testing.T) {
requireMDNSResponder(t)
transport := &Transport{}
defer transport.system.close()
message := new(mDNS.Msg)
message.SetQuestion("localhost.", mDNS.TypeA)
ctx, cancel := context.WithCancel(context.Background())
cancel()
start := time.Now()
_, err := systemExchangeForTest(ctx, transport, message)
elapsed := time.Since(start)
if err == nil {
t.Fatal("expected error for cancelled context")
}
if elapsed > time.Second {
t.Fatalf("cancellation too slow: %s", elapsed)
}
}
func TestSystemExchangeConcurrent(t *testing.T) {
requireMDNSResponder(t)
transport := &Transport{}
defer transport.system.close()
var waitGroup sync.WaitGroup
errors := make(chan error, 16)
for i := range 16 {
qtype := mDNS.TypeA
if i%2 == 1 {
qtype = mDNS.TypeAAAA
}
waitGroup.Go(func() {
message := new(mDNS.Msg)
message.SetQuestion("localhost.", qtype)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
response, exchangeErr := systemExchangeForTest(ctx, transport, message)
if exchangeErr != nil {
errors <- exchangeErr
return
}
if len(response.Answer) == 0 {
errors <- context.DeadlineExceeded
}
})
}
waitGroup.Wait()
close(errors)
for exchangeErr := range errors {
t.Fatal("concurrent query failed: ", exchangeErr)
}
transport.system.queryAccess.Lock()
pendingCount := len(transport.system.queries)
transport.system.queryAccess.Unlock()
if pendingCount != 0 {
t.Fatalf("expected no pending queries after completion, got %d", pendingCount)
}
}
+68
View File
@@ -0,0 +1,68 @@
package local
import (
"strings"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/dns"
E "github.com/sagernet/sing/common/exceptions"
mDNS "github.com/miekg/dns"
)
func buildNeighborMatchers(domains []string) ([]string, error) {
if len(domains) == 0 {
return nil, nil
}
var suffixes []string
for _, domain := range domains {
if !strings.HasPrefix(domain, ".") {
return nil, E.New("neighbor_domain entry must start with '.': ", domain)
}
suffixes = append(suffixes, mDNS.CanonicalName(domain))
}
return suffixes, nil
}
func (r *PreferredDomainResolver) lookupNeighbor(message *mDNS.Msg) *mDNS.Msg {
if r.neighborResolver == nil {
return nil
}
question := message.Question[0]
if question.Qtype != mDNS.TypeA && question.Qtype != mDNS.TypeAAAA {
return nil
}
host := extractNeighborHost(mDNS.CanonicalName(question.Name), r.neighborSuffixes)
if host == "" {
return nil
}
addresses := r.neighborResolver.LookupAddresses(host)
if len(addresses) == 0 {
return nil
}
return dns.FixedResponse(message.Id, question, addresses, C.DefaultDNSTTL)
}
func (r *PreferredDomainResolver) hasNeighborHost(domain string) bool {
if r.neighborResolver == nil {
return false
}
host := extractNeighborHost(domain, r.neighborSuffixes)
if host == "" {
return false
}
return len(r.neighborResolver.LookupAddresses(host)) > 0
}
func extractNeighborHost(canonical string, suffixes []string) string {
for _, suffix := range suffixes {
if !strings.HasSuffix(canonical, suffix) || len(canonical) <= len(suffix) {
continue
}
host := canonical[:len(canonical)-len(suffix)]
if !strings.ContainsRune(host, '.') {
return host
}
}
return ""
}
+20
View File
@@ -0,0 +1,20 @@
//go:build !darwin
package local
import (
"context"
"os"
mDNS "github.com/miekg/dns"
)
type systemResolver struct{}
func (r *systemResolver) close() {}
func (r *systemResolver) reset() {}
func (t *Transport) systemExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
callback(nil, os.ErrInvalid)
}
+76
View File
@@ -0,0 +1,76 @@
package local
import (
"context"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport/hosts"
"github.com/sagernet/sing-box/dns/transport/mdns"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common/logger"
"github.com/sagernet/sing/service"
mDNS "github.com/miekg/dns"
)
type PreferredDomainResolver struct {
ctx context.Context
logger logger.ContextLogger
hosts *hosts.File
neighborResolver adapter.NeighborResolver
neighborSuffixes []string
}
func NewPreferredDomainResolver(ctx context.Context, contextLogger logger.ContextLogger, options option.LocalDNSServerOptions) (*PreferredDomainResolver, error) {
suffixes, err := buildNeighborMatchers(options.NeighborDomain)
if err != nil {
return nil, err
}
return &PreferredDomainResolver{
ctx: ctx,
logger: contextLogger,
neighborSuffixes: suffixes,
}, nil
}
func (r *PreferredDomainResolver) Start(stage adapter.StartStage) {
switch stage {
case adapter.StartStateInitialize:
defaultHosts, err := hosts.NewDefault()
if err != nil {
r.logger.Warn(err)
} else {
r.hosts = defaultHosts
}
case adapter.StartStateStart:
router := service.FromContext[adapter.Router](r.ctx)
if router != nil {
r.neighborResolver = router.NeighborResolver()
}
}
}
func (r *PreferredDomainResolver) PreferredDomain(domain string) bool {
if r.hosts != nil {
if len(r.hosts.Lookup(dns.FqdnToDomain(domain))) > 0 {
return true
}
}
return r.hasNeighborHost(domain) || mdns.IsLocalDomain(domain)
}
func (r *PreferredDomainResolver) Lookup(message *mDNS.Msg) *mDNS.Msg {
question := message.Question[0]
if question.Qtype != mDNS.TypeA && question.Qtype != mDNS.TypeAAAA {
return nil
}
if r.hosts != nil {
addresses := r.hosts.Lookup(dns.FqdnToDomain(question.Name))
if len(addresses) > 0 {
return dns.FixedResponse(message.Id, question, addresses, C.DefaultDNSTTL)
}
}
return r.lookupNeighbor(message)
}
+3
View File
@@ -9,5 +9,8 @@ import (
type ResolvedResolver interface {
Start() error
Close() error
Reset()
Environment() []string
Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error)
ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error))
}
+132 -11
View File
@@ -19,6 +19,7 @@ import (
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing-box/service/resolved"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/control"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
@@ -57,11 +58,16 @@ type DBusResolvedResolver struct {
interfaceCallback *list.Element[tun.DefaultInterfaceUpdateCallback]
systemBus *dbus.Conn
savedServerSet atomic.Pointer[resolvedServerSet]
updateAccess sync.Mutex
updateCancel context.CancelFunc
updateRunAccess sync.Mutex
closed bool
closeOnce sync.Once
}
type resolvedServerSet struct {
servers []resolvedServer
servers []resolvedServer
signature []string
}
type resolvedServer struct {
@@ -93,7 +99,7 @@ func NewResolvedResolver(ctx context.Context, logger logger.ContextLogger) (Reso
}
func (t *DBusResolvedResolver) Start() error {
t.updateStatus()
t.updateStatus(t.ctx)
t.interfaceCallback = t.interfaceMonitor.RegisterCallback(t.updateDefaultInterface)
err := t.systemBus.BusObject().AddMatchSignal(
"org.freedesktop.DBus",
@@ -120,7 +126,17 @@ func (t *DBusResolvedResolver) Start() error {
func (t *DBusResolvedResolver) Close() error {
var closeErr error
t.closeOnce.Do(func() {
t.updateAccess.Lock()
updateCancel := t.updateCancel
t.updateCancel = nil
t.updateAccess.Unlock()
if updateCancel != nil {
updateCancel()
}
t.updateRunAccess.Lock()
t.closed = true
serverSet := t.savedServerSet.Swap(nil)
t.updateRunAccess.Unlock()
if serverSet != nil {
closeErr = serverSet.Close()
}
@@ -134,6 +150,27 @@ func (t *DBusResolvedResolver) Close() error {
return closeErr
}
func (t *DBusResolvedResolver) Reset() {
serverSet := t.savedServerSet.Load()
if serverSet == nil {
return
}
for _, server := range serverSet.servers {
server.primaryTransport.Reset()
if server.fallbackTransport != nil {
server.fallbackTransport.Reset()
}
}
}
func (t *DBusResolvedResolver) Environment() []string {
serverSet := t.savedServerSet.Load()
if serverSet == nil {
return nil
}
return serverSet.signature
}
func (t *DBusResolvedResolver) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
serverSet := t.savedServerSet.Load()
if serverSet == nil {
@@ -151,7 +188,7 @@ func (t *DBusResolvedResolver) Exchange(ctx context.Context, message *mDNS.Msg)
if err == nil {
return response, nil
}
t.updateStatus()
t.updateStatus(t.ctx)
refreshedServerSet := t.savedServerSet.Load()
if refreshedServerSet == nil || refreshedServerSet == serverSet {
return nil, err
@@ -159,6 +196,51 @@ func (t *DBusResolvedResolver) Exchange(ctx context.Context, message *mDNS.Msg)
return t.exchangeServerSet(ctx, message, refreshedServerSet)
}
func (t *DBusResolvedResolver) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
serverSet := t.savedServerSet.Load()
if serverSet == nil {
go func() {
callback(t.Exchange(ctx, message))
}()
return
}
t.exchangeServerSetAsync(ctx, message, serverSet, func(response *mDNS.Msg, err error) {
if err == nil {
callback(response, nil)
return
}
go func() {
t.updateStatus(t.ctx)
refreshedServerSet := t.savedServerSet.Load()
if refreshedServerSet == nil || refreshedServerSet == serverSet {
callback(nil, err)
return
}
t.exchangeServerSetAsync(ctx, message, refreshedServerSet, callback)
}()
})
}
func (t *DBusResolvedResolver) exchangeServerSetAsync(ctx context.Context, message *mDNS.Msg, serverSet *resolvedServerSet, callback func(response *mDNS.Msg, err error)) {
if len(serverSet.servers) == 0 {
callback(nil, E.New("link has no DNS servers configured"))
return
}
serverExchangers := make([]dnsTransport.AsyncExchanger, 0, len(serverSet.servers))
for _, server := range serverSet.servers {
serverExchangers = append(serverExchangers, func(exchangeCtx context.Context, exchangeCallback func(response *mDNS.Msg, err error)) {
server.primaryTransport.ExchangeAsync(exchangeCtx, message, func(response *mDNS.Msg, exchangeErr error) {
if exchangeErr != nil && server.fallbackTransport != nil {
server.fallbackTransport.ExchangeAsync(exchangeCtx, message, exchangeCallback)
return
}
exchangeCallback(response, exchangeErr)
})
})
}
dnsTransport.ExchangeSequential(ctx, serverExchangers, nil, callback)
}
func (t *DBusResolvedResolver) loopUpdateStatus() {
signalChan := make(chan *dbus.Signal, 1)
t.systemBus.Signal(signalChan)
@@ -172,18 +254,44 @@ func (t *DBusResolvedResolver) loopUpdateStatus() {
if !loaded || newOwner == "" {
continue
}
t.updateStatus()
t.postUpdateStatus()
case "org.freedesktop.DBus.Properties.PropertiesChanged":
if !shouldUpdateResolvedServerSet(signal) {
continue
}
t.updateStatus()
t.postUpdateStatus()
}
}
}
func (t *DBusResolvedResolver) updateStatus() {
serverSet, err := t.checkResolved(context.Background())
func (t *DBusResolvedResolver) postUpdateStatus() {
updateContext, updateCancel := context.WithCancel(t.ctx)
t.updateAccess.Lock()
previousCancel := t.updateCancel
t.updateCancel = updateCancel
t.updateAccess.Unlock()
if previousCancel != nil {
previousCancel()
}
go func() {
defer updateCancel()
t.updateStatus(updateContext)
}()
}
func (t *DBusResolvedResolver) updateStatus(ctx context.Context) {
t.updateRunAccess.Lock()
defer t.updateRunAccess.Unlock()
if t.closed || ctx.Err() != nil {
return
}
serverSet, err := t.checkResolved(ctx)
if t.closed || ctx.Err() != nil {
if serverSet != nil {
_ = serverSet.Close()
}
return
}
oldServerSet := t.savedServerSet.Swap(serverSet)
if oldServerSet != nil {
_ = oldServerSet.Close()
@@ -223,7 +331,7 @@ func (t *DBusResolvedResolver) exchangeServerSet(ctx context.Context, message *m
func (t *DBusResolvedResolver) checkResolved(ctx context.Context) (*resolvedServerSet, error) {
dbusObject := t.systemBus.Object("org.freedesktop.resolve1", "/org/freedesktop/resolve1")
err := dbusObject.Call("org.freedesktop.DBus.Peer.Ping", 0).Err
err := dbusObject.(*dbus.Object).CallWithContext(ctx, "org.freedesktop.DBus.Peer.Ping", 0).Err
if err != nil {
return nil, err
}
@@ -253,10 +361,18 @@ func (t *DBusResolvedResolver) checkResolved(ctx context.Context) (*resolvedServ
if err != nil {
return nil, err
}
err = ctx.Err()
if err != nil {
return nil, err
}
linkDNSEx, err := loadResolvedLinkDNSEx(linkObject)
if err != nil {
return nil, err
}
err = ctx.Err()
if err != nil {
return nil, err
}
linkDNS, err := loadResolvedLinkDNS(linkObject)
if err != nil {
return nil, err
@@ -270,8 +386,10 @@ func (t *DBusResolvedResolver) checkResolved(ctx context.Context) (*resolvedServ
return nil, E.New("link has no DNS servers configured")
}
serverDialer, err := dialer.NewDefault(t.ctx, option.DialerOptions{
BindInterface: defaultInterface.Name,
UDPFragmentDefault: true,
AbstractDialerOptions: option.AbstractDialerOptions{
BindInterface: defaultInterface.Name,
UDPFragmentDefault: true,
},
})
if err != nil {
return nil, err
@@ -299,6 +417,9 @@ func (t *DBusResolvedResolver) checkResolved(ctx context.Context) (*resolvedServ
}
serverSet := &resolvedServerSet{
servers: make([]resolvedServer, 0, len(serverSpecifications)),
signature: common.Map(serverSpecifications, func(it resolvedServerSpecification) string {
return M.SocksaddrFrom(it.address, it.port).String()
}),
}
for _, serverSpecification := range serverSpecifications {
server, createErr := t.createResolvedServer(serverDialer, dnsOverTLSMode, serverSpecification)
@@ -497,5 +618,5 @@ func shouldUpdateResolvedServerSet(signal *dbus.Signal) bool {
}
func (t *DBusResolvedResolver) updateDefaultInterface(defaultInterface *control.Interface, flags int) {
t.updateStatus()
t.postUpdateStatus()
}
+80 -156
View File
@@ -2,182 +2,106 @@ package local
import (
"context"
"errors"
"math/rand"
"syscall"
"time"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing-box/dns/transport/local/systemconfig"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
mDNS "github.com/miekg/dns"
)
func (t *Transport) exchange(ctx context.Context, message *mDNS.Msg, domain string) (*mDNS.Msg, error) {
systemConfig := getSystemDNSConfig(t.ctx)
if systemConfig.singleRequest || !(message.Question[0].Qtype == mDNS.TypeA || message.Question[0].Qtype == mDNS.TypeAAAA) {
return t.exchangeSingleRequest(ctx, systemConfig, message, domain)
} else {
return t.exchangeParallel(ctx, systemConfig, message, domain)
type localServerSet struct {
config *systemconfig.Config
transports []adapter.DNSTransport
}
func (s *localServerSet) Close() {
for _, serverTransport := range s.transports {
serverTransport.Close()
}
}
func (t *Transport) exchangeSingleRequest(ctx context.Context, systemConfig *dnsConfig, message *mDNS.Msg, domain string) (*mDNS.Msg, error) {
var lastErr error
for _, fqdn := range systemConfig.nameList(domain) {
response, err := t.tryOneName(ctx, systemConfig, fqdn, message)
func (t *Transport) serverSetFor(systemConfig *systemconfig.Config) (*localServerSet, error) {
serverSet := t.serverSet.Load()
if serverSet != nil && serverSet.config == systemConfig {
return serverSet, nil
}
t.serverSetAccess.Lock()
defer t.serverSetAccess.Unlock()
serverSet = t.serverSet.Load()
if serverSet != nil && serverSet.config == systemConfig {
return serverSet, nil
}
transports := make([]adapter.DNSTransport, 0, len(systemConfig.Servers))
for _, serverAddr := range systemConfig.Servers {
var serverTransport adapter.DNSTransport
if systemConfig.UseTCP {
serverTransport = transport.NewTCPRaw(dns.NewTransportAdapter(C.DNSTypeTCP, "", nil), t.dialer, serverAddr)
} else {
serverTransport = transport.NewUDPRaw(t.logger, dns.NewTransportAdapter(C.DNSTypeUDP, "", nil), t.dialer, serverAddr)
}
err := serverTransport.Start(adapter.StartStateStart)
if err != nil {
lastErr = err
continue
for _, startedTransport := range transports {
startedTransport.Close()
}
return nil, E.Cause(err, "initialize transport for ", serverAddr)
}
return response, nil
transports = append(transports, serverTransport)
}
return nil, lastErr
newServerSet := &localServerSet{
config: systemConfig,
transports: transports,
}
oldServerSet := t.serverSet.Swap(newServerSet)
if oldServerSet != nil {
oldServerSet.Close()
}
return newServerSet, nil
}
func (t *Transport) exchangeParallel(ctx context.Context, systemConfig *dnsConfig, message *mDNS.Msg, domain string) (*mDNS.Msg, error) {
returned := make(chan struct{})
defer close(returned)
type queryResult struct {
response *mDNS.Msg
err error
func (t *Transport) exchangeAsync(ctx context.Context, message *mDNS.Msg, domain string, callback func(response *mDNS.Msg, err error)) {
systemConfig := t.configSource.Configuration()
serverSet, err := t.serverSetFor(systemConfig)
if err != nil {
callback(nil, err)
return
}
results := make(chan queryResult)
startRacer := func(ctx context.Context, fqdn string) {
response, err := t.tryOneName(ctx, systemConfig, fqdn, message)
select {
case results <- queryResult{response, err}:
case <-returned:
}
}
queryCtx, queryCancel := context.WithCancel(ctx)
defer queryCancel()
var nameCount int
for _, fqdn := range systemConfig.nameList(domain) {
nameCount++
go startRacer(queryCtx, fqdn)
}
var errors []error
for {
select {
case <-ctx.Done():
return nil, ctx.Err()
case result := <-results:
if result.err == nil {
return result.response, nil
}
errors = append(errors, result.err)
if len(errors) == nameCount {
return nil, E.Errors(errors...)
}
}
names := systemConfig.NameList(domain)
if len(names) == 0 {
callback(nil, E.New("invalid domain: ", domain))
return
}
transport.ExchangeNames(ctx, names, message.Question[0], func(fqdn string) transport.AsyncExchanger {
return newNameExchanger(systemConfig, serverSet, message, fqdn)
}, callback)
}
func (t *Transport) tryOneName(ctx context.Context, config *dnsConfig, fqdn string, message *mDNS.Msg) (*mDNS.Msg, error) {
serverOffset := config.serverOffset()
sLen := uint32(len(config.servers))
var lastErr error
for i := 0; i < config.attempts; i++ {
for j := range sLen {
server := config.servers[(serverOffset+j)%sLen]
question := message.Question[0]
question.Name = fqdn
response, err := t.exchangeOne(ctx, M.ParseSocksaddr(server), question, config.timeout, config.useTCP, config.trustAD)
func newNameExchanger(systemConfig *systemconfig.Config, serverSet *localServerSet, message *mDNS.Msg, fqdn string) transport.AsyncExchanger {
serverOffset := systemConfig.ServerOffset()
serverCount := uint32(len(serverSet.transports))
attemptExchangers := make([]transport.AsyncExchanger, 0, systemConfig.Attempts*int(serverCount))
for i := 0; i < systemConfig.Attempts; i++ {
for j := range serverCount {
serverTransport := serverSet.transports[(serverOffset+j)%serverCount]
attemptExchangers = append(attemptExchangers, func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
attemptCtx, cancel := context.WithTimeout(ctx, systemConfig.Timeout)
serverTransport.ExchangeAsync(attemptCtx, transport.NewFanOutRequest(message, fqdn, systemConfig.TrustAD), func(response *mDNS.Msg, err error) {
cancel()
callback(response, err)
})
})
}
}
return func(ctx context.Context, callback func(response *mDNS.Msg, err error)) {
transport.ExchangeSequential(ctx, attemptExchangers, nil, func(response *mDNS.Msg, err error) {
if err != nil {
lastErr = err
continue
err = E.Cause(err, fqdn)
}
return response, nil
}
}
return nil, E.Cause(lastErr, fqdn)
}
func (t *Transport) exchangeOne(ctx context.Context, server M.Socksaddr, question mDNS.Question, timeout time.Duration, useTCP, ad bool) (*mDNS.Msg, error) {
if server.Port == 0 {
server.Port = 53
}
request := &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Id: uint16(rand.Uint32()),
RecursionDesired: true,
AuthenticatedData: ad,
},
Question: []mDNS.Question{question},
Compress: true,
}
request.SetEdns0(buf.UDPBufferSize, false)
if !useTCP {
return t.exchangeUDP(ctx, server, request, timeout)
} else {
return t.exchangeTCP(ctx, server, request, timeout)
callback(response, err)
})
}
}
func (t *Transport) exchangeUDP(ctx context.Context, server M.Socksaddr, request *mDNS.Msg, timeout time.Duration) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkUDP, server)
if err != nil {
return nil, err
}
defer conn.Close()
if deadline, loaded := ctx.Deadline(); loaded && !deadline.IsZero() {
newDeadline := time.Now().Add(timeout)
if deadline.After(newDeadline) {
deadline = newDeadline
}
conn.SetDeadline(deadline)
}
buffer := buf.Get(buf.UDPBufferSize)
defer buf.Put(buffer)
rawMessage, err := request.PackBuffer(buffer)
if err != nil {
return nil, E.Cause(err, "pack request")
}
_, err = conn.Write(rawMessage)
if err != nil {
if errors.Is(err, syscall.EMSGSIZE) {
return t.exchangeTCP(ctx, server, request, timeout)
}
return nil, E.Cause(err, "write request")
}
n, err := conn.Read(buffer)
if err != nil {
if errors.Is(err, syscall.EMSGSIZE) {
return t.exchangeTCP(ctx, server, request, timeout)
}
return nil, E.Cause(err, "read response")
}
var response mDNS.Msg
err = response.Unpack(buffer[:n])
if err != nil {
return nil, E.Cause(err, "unpack response")
}
if response.Truncated {
return t.exchangeTCP(ctx, server, request, timeout)
}
return &response, nil
}
func (t *Transport) exchangeTCP(ctx context.Context, server M.Socksaddr, request *mDNS.Msg, timeout time.Duration) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, server)
if err != nil {
return nil, err
}
defer conn.Close()
if deadline, loaded := ctx.Deadline(); loaded && !deadline.IsZero() {
newDeadline := time.Now().Add(timeout)
if deadline.After(newDeadline) {
deadline = newDeadline
}
conn.SetDeadline(deadline)
}
err = transport.WriteMessage(conn, 0, request)
if err != nil {
return nil, err
}
return transport.ReadMessage(conn)
}
-145
View File
@@ -1,145 +0,0 @@
//nolint:unused
package local
import (
"context"
"os"
"runtime"
"strings"
"sync"
"sync/atomic"
"time"
)
type resolverConfig struct {
initOnce sync.Once
ch chan struct{}
lastChecked time.Time
dnsConfig atomic.Pointer[dnsConfig]
}
var resolvConf resolverConfig
func getSystemDNSConfig(ctx context.Context) *dnsConfig {
resolvConf.tryUpdate(ctx, "/etc/resolv.conf")
return resolvConf.dnsConfig.Load()
}
func (conf *resolverConfig) init(ctx context.Context) {
conf.dnsConfig.Store(dnsReadConfig(ctx, "/etc/resolv.conf"))
conf.lastChecked = time.Now()
conf.ch = make(chan struct{}, 1)
}
func (conf *resolverConfig) tryUpdate(ctx context.Context, name string) {
conf.initOnce.Do(func() {
conf.init(ctx)
})
if conf.dnsConfig.Load().noReload {
return
}
if !conf.tryAcquireSema() {
return
}
defer conf.releaseSema()
now := time.Now()
if conf.lastChecked.After(now.Add(-5 * time.Second)) {
return
}
conf.lastChecked = now
if runtime.GOOS != "windows" {
var mtime time.Time
if fi, err := os.Stat(name); err == nil {
mtime = fi.ModTime()
}
if mtime.Equal(conf.dnsConfig.Load().mtime) {
return
}
}
dnsConf := dnsReadConfig(ctx, name)
conf.dnsConfig.Store(dnsConf)
}
func (conf *resolverConfig) tryAcquireSema() bool {
select {
case conf.ch <- struct{}{}:
return true
default:
return false
}
}
func (conf *resolverConfig) releaseSema() {
<-conf.ch
}
type dnsConfig struct {
servers []string
search []string
ndots int
timeout time.Duration
attempts int
rotate bool
unknownOpt bool
lookup []string
err error
mtime time.Time
soffset uint32
singleRequest bool
useTCP bool
trustAD bool
noReload bool
}
func (c *dnsConfig) serverOffset() uint32 {
if c.rotate {
return atomic.AddUint32(&c.soffset, 1) - 1 // return 0 to start
}
return 0
}
func (c *dnsConfig) nameList(name string) []string {
l := len(name)
rooted := l > 0 && name[l-1] == '.'
if l > 254 || l == 254 && !rooted {
return nil
}
if rooted {
if avoidDNS(name) {
return nil
}
return []string{name}
}
hasNdots := strings.Count(name, ".") >= c.ndots
name += "."
// l++
names := make([]string, 0, 1+len(c.search))
if hasNdots && !avoidDNS(name) {
names = append(names, name)
}
for _, suffix := range c.search {
fqdn := name + suffix
if !avoidDNS(fqdn) && len(fqdn) <= 254 {
names = append(names, fqdn)
}
}
if !hasNdots && !avoidDNS(name) {
names = append(names, name)
}
return names
}
func avoidDNS(name string) bool {
if name == "" {
return true
}
if name[len(name)-1] == '.' {
name = name[:len(name)-1]
}
return strings.HasSuffix(name, ".onion")
}
-24
View File
@@ -1,24 +0,0 @@
//nolint:unused
package local
import (
"os"
"strings"
_ "unsafe"
"github.com/miekg/dns"
)
//go:linkname defaultNS net.defaultNS
var defaultNS []string
func dnsDefaultSearch() []string {
hn, err := os.Hostname()
if err != nil {
return nil
}
if i := strings.IndexRune(hn, '.'); i >= 0 && i < len(hn)-1 {
return []string{dns.Fqdn(hn[i+1:])}
}
return nil
}
-13
View File
@@ -1,13 +0,0 @@
package local
import (
"context"
"testing"
"github.com/stretchr/testify/require"
)
func TestDNSReadConfig(t *testing.T) {
t.Parallel()
require.NoError(t, dnsReadConfig(context.Background(), "/etc/resolv.conf").err)
}
-156
View File
@@ -1,156 +0,0 @@
//go:build !windows
package local
import (
"bufio"
"context"
"net"
"net/netip"
"os"
"strings"
"time"
"github.com/miekg/dns"
)
func dnsReadConfig(_ context.Context, name string) *dnsConfig {
conf := &dnsConfig{
ndots: 1,
timeout: 5 * time.Second,
attempts: 2,
}
file, err := os.Open(name)
if err != nil {
conf.servers = defaultNS
conf.search = dnsDefaultSearch()
conf.err = err
return conf
}
defer file.Close()
fi, err := file.Stat()
if err == nil {
conf.mtime = fi.ModTime()
} else {
conf.servers = defaultNS
conf.search = dnsDefaultSearch()
conf.err = err
return conf
}
reader := bufio.NewReader(file)
var (
prefix []byte
line []byte
isPrefix bool
)
for {
line, isPrefix, err = reader.ReadLine()
if err != nil {
break
}
if isPrefix {
prefix = append(prefix, line...)
continue
} else if len(prefix) > 0 {
line = append(prefix, line...)
prefix = nil
}
if len(line) > 0 && (line[0] == ';' || line[0] == '#') {
continue
}
f := strings.Fields(string(line))
if len(f) < 1 {
continue
}
switch f[0] {
case "nameserver":
if len(f) > 1 && len(conf.servers) < 3 {
if _, err := netip.ParseAddr(f[1]); err == nil {
conf.servers = append(conf.servers, net.JoinHostPort(f[1], "53"))
}
}
case "domain":
if len(f) > 1 {
conf.search = []string{dns.Fqdn(f[1])}
}
case "search":
conf.search = make([]string, 0, len(f)-1)
for i := 1; i < len(f); i++ {
name := dns.Fqdn(f[i])
if name == "." {
continue
}
conf.search = append(conf.search, name)
}
case "options":
for _, s := range f[1:] {
switch {
case strings.HasPrefix(s, "ndots:"):
n, _, _ := dtoi(s[6:])
if n < 0 {
n = 0
} else if n > 15 {
n = 15
}
conf.ndots = n
case strings.HasPrefix(s, "timeout:"):
n, _, _ := dtoi(s[8:])
if n < 1 {
n = 1
}
conf.timeout = time.Duration(n) * time.Second
case strings.HasPrefix(s, "attempts:"):
n, _, _ := dtoi(s[9:])
if n < 1 {
n = 1
}
conf.attempts = n
case s == "rotate":
conf.rotate = true
case s == "single-request" || s == "single-request-reopen":
conf.singleRequest = true
case s == "use-vc" || s == "usevc" || s == "tcp":
conf.useTCP = true
case s == "trust-ad":
conf.trustAD = true
case s == "edns0":
case s == "no-reload":
conf.noReload = true
default:
conf.unknownOpt = true
}
}
case "lookup":
conf.lookup = f[1:]
default:
conf.unknownOpt = true
}
}
if len(conf.servers) == 0 {
conf.servers = defaultNS
}
if len(conf.search) == 0 {
conf.search = dnsDefaultSearch()
}
return conf
}
const big = 0xFFFFFF
func dtoi(s string) (n int, i int, ok bool) {
n = 0
for i = 0; i < len(s) && '0' <= s[i] && s[i] <= '9'; i++ {
n = n*10 + int(s[i]-'0')
if n >= big {
return big, i, false
}
}
if i == 0 {
return 0, 0, false
}
return n, i, true
}
-119
View File
@@ -1,119 +0,0 @@
package local
import (
"context"
"net"
"net/netip"
"os"
"strconv"
"syscall"
"time"
"unsafe"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/service"
"golang.org/x/sys/windows"
)
func dnsReadConfig(ctx context.Context, _ string) *dnsConfig {
conf := &dnsConfig{
ndots: 1,
timeout: 5 * time.Second,
attempts: 2,
}
defer func() {
if len(conf.servers) == 0 {
conf.servers = defaultNS
}
}()
addresses, err := adapterAddresses()
if err != nil {
return nil
}
var dnsAddresses []struct {
ifName string
netip.Addr
}
for _, address := range addresses {
if address.OperStatus != windows.IfOperStatusUp {
continue
}
if address.IfType == windows.IF_TYPE_TUNNEL {
continue
}
if address.FirstGatewayAddress == nil {
continue
}
for dnsServerAddress := address.FirstDnsServerAddress; dnsServerAddress != nil; dnsServerAddress = dnsServerAddress.Next {
rawSockaddr, err := dnsServerAddress.Address.Sockaddr.Sockaddr()
if err != nil {
continue
}
var dnsServerAddr netip.Addr
switch sockaddr := rawSockaddr.(type) {
case *syscall.SockaddrInet4:
dnsServerAddr = netip.AddrFrom4(sockaddr.Addr)
case *syscall.SockaddrInet6:
if sockaddr.Addr[0] == 0xfe && sockaddr.Addr[1] == 0xc0 {
// fec0/10 IPv6 addresses are site local anycast DNS
// addresses Microsoft sets by default if no other
// IPv6 DNS address is set. Site local anycast is
// deprecated since 2004, see
// https://datatracker.ietf.org/doc/html/rfc3879
continue
}
dnsServerAddr = netip.AddrFrom16(sockaddr.Addr)
if sockaddr.ZoneId != 0 {
dnsServerAddr = dnsServerAddr.WithZone(strconv.FormatInt(int64(sockaddr.ZoneId), 10))
}
default:
// Unexpected type.
continue
}
dnsAddresses = append(dnsAddresses, struct {
ifName string
netip.Addr
}{ifName: windows.UTF16PtrToString(address.FriendlyName), Addr: dnsServerAddr})
}
}
var myInterfaces []string
if networkManager := service.FromContext[adapter.NetworkManager](ctx); networkManager != nil {
myInterfaces = networkManager.InterfaceMonitor().MyInterfaces()
}
for _, address := range dnsAddresses {
if common.Contains(myInterfaces, address.ifName) {
continue
}
conf.servers = append(conf.servers, net.JoinHostPort(address.String(), "53"))
}
return conf
}
func adapterAddresses() ([]*windows.IpAdapterAddresses, error) {
var b []byte
l := uint32(15000) // recommended initial size
for {
b = make([]byte, l)
const flags = windows.GAA_FLAG_INCLUDE_PREFIX | windows.GAA_FLAG_INCLUDE_GATEWAYS
err := windows.GetAdaptersAddresses(syscall.AF_UNSPEC, flags, 0, (*windows.IpAdapterAddresses)(unsafe.Pointer(&b[0])), &l)
if err == nil {
if l == 0 {
return nil, nil
}
break
}
if err.(syscall.Errno) != syscall.ERROR_BUFFER_OVERFLOW {
return nil, os.NewSyscallError("getadaptersaddresses", err)
}
if l <= uint32(len(b)) {
return nil, os.NewSyscallError("getadaptersaddresses", err)
}
}
var aas []*windows.IpAdapterAddresses
for aa := (*windows.IpAdapterAddresses)(unsafe.Pointer(&b[0])); aa != nil; aa = aa.Next {
aas = append(aas, aa)
}
return aas, nil
}
+113
View File
@@ -0,0 +1,113 @@
package systemconfig
import (
"net/netip"
"os"
"slices"
"strconv"
"strings"
"sync/atomic"
"time"
M "github.com/sagernet/sing/common/metadata"
mDNS "github.com/miekg/dns"
)
var defaultServers = []M.Socksaddr{
M.SocksaddrFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 53),
M.SocksaddrFrom(netip.IPv6Loopback(), 53),
}
type Config struct {
Servers []M.Socksaddr
Search []string
Ndots int
Timeout time.Duration
Attempts int
Rotate bool
soffset uint32
SingleRequest bool
UseTCP bool
TrustAD bool
}
func (c *Config) Equal(other *Config) bool {
return slices.Equal(c.Servers, other.Servers) &&
slices.Equal(c.Search, other.Search) &&
c.Ndots == other.Ndots &&
c.Timeout == other.Timeout &&
c.Attempts == other.Attempts &&
c.Rotate == other.Rotate &&
c.SingleRequest == other.SingleRequest &&
c.UseTCP == other.UseTCP &&
c.TrustAD == other.TrustAD
}
func (c *Config) Signature() []string {
signature := make([]string, 0, len(c.Servers)+len(c.Search)+1)
for _, server := range c.Servers {
signature = append(signature, server.String())
}
signature = append(signature, c.Search...)
return append(signature, "ndots:"+strconv.Itoa(c.Ndots))
}
func (c *Config) ServerOffset() uint32 {
if c.Rotate {
return atomic.AddUint32(&c.soffset, 1) - 1
}
return 0
}
func (c *Config) NameList(name string) []string {
l := len(name)
rooted := l > 0 && name[l-1] == '.'
if l > 254 || l == 254 && !rooted {
return nil
}
if rooted {
if avoidDNS(name) {
return nil
}
return []string{name}
}
hasNdots := strings.Count(name, ".") >= c.Ndots
name += "."
names := make([]string, 0, 1+len(c.Search))
if hasNdots && !avoidDNS(name) {
names = append(names, name)
}
for _, suffix := range c.Search {
fqdn := name + suffix
if !avoidDNS(fqdn) && len(fqdn) <= 254 {
names = append(names, fqdn)
}
}
if !hasNdots && !avoidDNS(name) {
names = append(names, name)
}
return names
}
func avoidDNS(name string) bool {
if name == "" {
return true
}
return strings.HasSuffix(strings.TrimSuffix(name, "."), ".onion")
}
func defaultSearch() []string {
hostname, err := os.Hostname()
if err != nil {
return nil
}
_, domain, found := strings.Cut(hostname, ".")
if !found || domain == "" {
return nil
}
return []string{mDNS.Fqdn(domain)}
}
@@ -0,0 +1,364 @@
//go:build cgo
package systemconfig
/*
#include <dlfcn.h>
#include <notify.h>
#include <stdint.h>
#include <string.h>
#include <netinet/in.h>
#include <sys/socket.h>
// dnsinfo.h is not shipped in any SDK. The layouts below are DNSINFO_VERSION
// 20170629 from apple-oss-distributions/configd (#pragma pack(4)), the format
// libsystem_configuration unpacks into at runtime. dns_configuration_copy,
// dns_configuration_free and dns_configuration_notify_key are private
// libSystem exports. cgo silently drops packed struct fields that fall on
// unaligned offsets.
#pragma pack(4)
typedef struct {
struct in_addr address;
struct in_addr mask;
} box_dns_sortaddr_t;
typedef struct {
char *domain;
int32_t n_nameserver;
struct sockaddr **nameserver;
uint16_t port;
int32_t n_search;
char **search;
int32_t n_sortaddr;
box_dns_sortaddr_t **sortaddr;
char *options;
uint32_t timeout;
uint32_t search_order;
uint32_t if_index;
uint32_t flags;
uint32_t reach_flags;
uint32_t service_identifier;
char *cid;
char *if_name;
} box_dns_resolver_t;
typedef struct {
int32_t n_resolver;
box_dns_resolver_t **resolver;
int32_t n_scoped_resolver;
box_dns_resolver_t **scoped_resolver;
uint64_t generation;
int32_t n_service_specific_resolver;
box_dns_resolver_t **service_specific_resolver;
uint32_t version;
} box_dns_config_t;
#pragma pack()
static box_dns_config_t *(*box_dns_configuration_copy)(void);
static void (*box_dns_configuration_free)(box_dns_config_t *);
static void box_reverse_string(char *s) {
size_t length = strlen(s);
for (size_t i = 0; i < length / 2; i++) {
char tmp = s[i];
s[i] = s[length - 1 - i];
s[length - 1 - i] = tmp;
}
}
static int box_dnsinfo_load(void) {
if (box_dns_configuration_copy != NULL && box_dns_configuration_free != NULL) {
return 1;
}
char copy_name[] = "ypoc_noitarugifnoc_snd";
char free_name[] = "eerf_noitarugifnoc_snd";
box_reverse_string(copy_name);
box_reverse_string(free_name);
box_dns_configuration_copy = (box_dns_config_t * (*)(void)) dlsym(RTLD_DEFAULT, copy_name);
box_dns_configuration_free = (void (*)(box_dns_config_t *))dlsym(RTLD_DEFAULT, free_name);
return box_dns_configuration_copy != NULL && box_dns_configuration_free != NULL;
}
static box_dns_config_t *box_dnsinfo_copy(void) {
return box_dns_configuration_copy();
}
static void box_dnsinfo_free(box_dns_config_t *config) {
box_dns_configuration_free(config);
}
static const char *box_dnsinfo_notify_key(void) {
const char *(*notify_key)(void) = (const char *(*)(void))dlsym(RTLD_DEFAULT, "dns_configuration_notify_key");
if (notify_key != NULL) {
return notify_key();
}
return "com.apple.system.SystemConfiguration.dns_configuration";
}
static box_dns_resolver_t *box_dnsinfo_default_resolver(box_dns_config_t *config, int32_t index) {
return config->resolver[index];
}
static box_dns_resolver_t *box_dnsinfo_scoped_resolver(box_dns_config_t *config, int32_t index) {
return config->scoped_resolver[index];
}
static struct sockaddr *box_dnsinfo_nameserver(box_dns_resolver_t *resolver, int32_t index) {
return resolver->nameserver[index];
}
static const char *box_dnsinfo_search_domain(box_dns_resolver_t *resolver, int32_t index) {
return resolver->search[index];
}
*/
import "C"
import (
"context"
"encoding/binary"
"net/netip"
"strconv"
"sync"
"time"
"unsafe"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing/common"
M "github.com/sagernet/sing/common/metadata"
"github.com/sagernet/sing/service"
mDNS "github.com/miekg/dns"
)
type Source struct {
interfaceMonitor tun.DefaultInterfaceMonitor
access sync.Mutex
notifyToken C.int
notifyValid bool
stale bool
interfaceIndex int
config *Config
}
func NewSource(ctx context.Context) *Source {
source := &Source{
interfaceMonitor: service.FromContext[adapter.NetworkManager](ctx).InterfaceMonitor(),
}
if C.box_dnsinfo_load() != 0 {
var token C.int
if C.notify_register_check(C.box_dnsinfo_notify_key(), &token) == 0 {
source.notifyToken = token
source.notifyValid = true
}
}
return source
}
func (s *Source) Configuration() *Config {
interfaceIndex := s.defaultInterfaceIndex()
s.access.Lock()
defer s.access.Unlock()
interfaceChanged := s.interfaceIndex != interfaceIndex
s.interfaceIndex = interfaceIndex
changed := s.changedLocked()
if s.config != nil && !s.stale && !interfaceChanged && !changed {
return s.config
}
s.stale = false
systemInfo := copyDNSInfo()
if systemInfo == nil {
if s.config == nil {
s.config = new(dnsInfoConfig).build(interfaceIndex)
}
return s.config
}
config := systemInfo.build(interfaceIndex)
if s.config != nil && config.Equal(s.config) {
return s.config
}
s.config = config
return config
}
func (s *Source) changedLocked() bool {
if !s.notifyValid {
return true
}
var changed C.int
status := C.notify_check(s.notifyToken, &changed)
if status != 0 {
return true
}
return changed != 0
}
func (s *Source) Reset() {
s.access.Lock()
s.stale = true
s.access.Unlock()
}
func (s *Source) Close() error {
s.access.Lock()
defer s.access.Unlock()
if s.notifyValid {
C.notify_cancel(s.notifyToken)
s.notifyValid = false
}
return nil
}
func (s *Source) defaultInterfaceIndex() int {
if s.interfaceMonitor == nil {
return 0
}
defaultInterface := s.interfaceMonitor.DefaultInterface()
if defaultInterface == nil {
return 0
}
return defaultInterface.Index
}
type dnsInfoResolver struct {
interfaceIndex int
domain string
servers []M.Socksaddr
search []string
timeout time.Duration
}
type dnsInfoConfig struct {
resolvers []dnsInfoResolver
scopedResolvers []dnsInfoResolver
}
func (c *dnsInfoConfig) build(interfaceIndex int) *Config {
var selected dnsInfoResolver
if interfaceIndex != 0 {
selected = common.Find(c.scopedResolvers, func(it dnsInfoResolver) bool {
return it.interfaceIndex == interfaceIndex && len(it.servers) > 0
})
}
if len(selected.servers) == 0 {
selected = common.Find(c.resolvers, func(it dnsInfoResolver) bool {
return it.domain == "" && len(it.servers) > 0
})
}
config := &Config{
Ndots: 1,
Timeout: 5 * time.Second,
Attempts: 2,
}
if len(selected.servers) == 0 {
config.Servers = defaultServers
config.Search = defaultSearch()
return config
}
config.Servers = selected.servers
if len(selected.search) > 0 {
config.Search = selected.search
} else {
config.Search = defaultSearch()
}
if selected.timeout > 0 {
config.Timeout = selected.timeout
}
return config
}
func copyDNSInfo() *dnsInfoConfig {
if C.box_dnsinfo_load() == 0 {
return nil
}
rawConfig := C.box_dnsinfo_copy()
if rawConfig == nil {
return nil
}
defer C.box_dnsinfo_free(rawConfig)
systemInfo := new(dnsInfoConfig)
for i := C.int32_t(0); i < rawConfig.n_resolver; i++ {
rawResolver := C.box_dnsinfo_default_resolver(rawConfig, i)
if rawResolver == nil {
continue
}
systemInfo.resolvers = append(systemInfo.resolvers, parseResolver(rawResolver))
}
for i := C.int32_t(0); i < rawConfig.n_scoped_resolver; i++ {
rawResolver := C.box_dnsinfo_scoped_resolver(rawConfig, i)
if rawResolver == nil {
continue
}
systemInfo.scopedResolvers = append(systemInfo.scopedResolvers, parseResolver(rawResolver))
}
return systemInfo
}
func parseResolver(rawResolver *C.box_dns_resolver_t) dnsInfoResolver {
resolver := dnsInfoResolver{
interfaceIndex: int(rawResolver.if_index),
domain: C.GoString(rawResolver.domain),
timeout: time.Duration(rawResolver.timeout) * time.Second,
}
interfaceName := C.GoString(rawResolver.if_name)
resolverPort := uint16(rawResolver.port)
if resolverPort == 0 {
resolverPort = 53
}
for i := C.int32_t(0); i < rawResolver.n_nameserver; i++ {
rawSockaddr := C.box_dnsinfo_nameserver(rawResolver, i)
if rawSockaddr == nil {
continue
}
serverAddr, loaded := parseSockaddr(rawSockaddr, resolverPort, interfaceName)
if !loaded {
continue
}
resolver.servers = append(resolver.servers, M.SocksaddrFromNetIP(serverAddr))
}
for i := C.int32_t(0); i < rawResolver.n_search; i++ {
searchDomain := C.GoString(C.box_dnsinfo_search_domain(rawResolver, i))
if searchDomain == "" {
continue
}
searchDomain = mDNS.Fqdn(searchDomain)
if searchDomain == "." {
continue
}
resolver.search = append(resolver.search, searchDomain)
}
return resolver
}
func parseSockaddr(rawSockaddr *C.struct_sockaddr, fallbackPort uint16, zone string) (netip.AddrPort, bool) {
switch rawSockaddr.sa_family {
case C.AF_INET:
sockaddrInet := (*C.struct_sockaddr_in)(unsafe.Pointer(rawSockaddr))
addr := netip.AddrFrom4(*(*[4]byte)(unsafe.Pointer(&sockaddrInet.sin_addr)))
return netip.AddrPortFrom(addr, sockaddrPort(unsafe.Pointer(&sockaddrInet.sin_port), fallbackPort)), true
case C.AF_INET6:
sockaddrInet6 := (*C.struct_sockaddr_in6)(unsafe.Pointer(rawSockaddr))
addr := netip.AddrFrom16(*(*[16]byte)(unsafe.Pointer(&sockaddrInet6.sin6_addr)))
if addr.IsLinkLocalUnicast() {
scopeId := uint32(sockaddrInet6.sin6_scope_id)
if zone == "" && scopeId != 0 {
zone = strconv.FormatUint(uint64(scopeId), 10)
}
if zone != "" {
addr = addr.WithZone(zone)
}
}
return netip.AddrPortFrom(addr, sockaddrPort(unsafe.Pointer(&sockaddrInet6.sin6_port), fallbackPort)), true
default:
return netip.AddrPort{}, false
}
}
func sockaddrPort(rawPort unsafe.Pointer, fallbackPort uint16) uint16 {
port := binary.BigEndian.Uint16((*[2]byte)(rawPort)[:])
if port == 0 {
return fallbackPort
}
return port
}
@@ -0,0 +1,176 @@
//go:build !windows && !(darwin && cgo)
package systemconfig
import (
"bufio"
"context"
"net/netip"
"os"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
M "github.com/sagernet/sing/common/metadata"
mDNS "github.com/miekg/dns"
)
const resolvConfPath = "/etc/resolv.conf"
type Source struct {
updateAccess sync.Mutex
lastChecked time.Time
current atomic.Pointer[resolvConfig]
}
type resolvConfig struct {
config *Config
mtime time.Time
noReload bool
}
func NewSource(_ context.Context) *Source {
source := &Source{lastChecked: time.Now()}
source.current.Store(readResolvConfig(resolvConfPath))
return source
}
func (s *Source) Configuration() *Config {
s.tryUpdate()
return s.current.Load().config
}
func (s *Source) tryUpdate() {
if s.current.Load().noReload {
return
}
if !s.updateAccess.TryLock() {
return
}
defer s.updateAccess.Unlock()
now := time.Now()
if s.lastChecked.After(now.Add(-5 * time.Second)) {
return
}
s.lastChecked = now
var mtime time.Time
fileInfo, err := os.Stat(resolvConfPath)
if err == nil {
mtime = fileInfo.ModTime()
}
current := s.current.Load()
if mtime.Equal(current.mtime) {
return
}
updated := readResolvConfig(resolvConfPath)
if updated.config.Equal(current.config) {
updated.config = current.config
}
s.current.Store(updated)
}
func (s *Source) Reset() {
s.updateAccess.Lock()
s.lastChecked = time.Time{}
s.updateAccess.Unlock()
}
func (s *Source) Close() error {
return nil
}
func readResolvConfig(path string) *resolvConfig {
config := &Config{
Ndots: 1,
Timeout: 5 * time.Second,
Attempts: 2,
}
result := &resolvConfig{config: config}
file, err := os.Open(path)
if err != nil {
config.Servers = defaultServers
config.Search = defaultSearch()
return result
}
defer file.Close()
fileInfo, err := file.Stat()
if err != nil {
config.Servers = defaultServers
config.Search = defaultSearch()
return result
}
result.mtime = fileInfo.ModTime()
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := scanner.Text()
if strings.HasPrefix(line, ";") || strings.HasPrefix(line, "#") {
continue
}
fields := strings.Fields(line)
if len(fields) < 1 {
continue
}
switch fields[0] {
case "nameserver":
if len(fields) > 1 && len(config.Servers) < 3 {
serverAddr, parseErr := netip.ParseAddr(fields[1])
if parseErr == nil {
config.Servers = append(config.Servers, M.SocksaddrFrom(serverAddr, 53))
}
}
case "domain":
if len(fields) > 1 {
config.Search = []string{mDNS.Fqdn(fields[1])}
}
case "search":
config.Search = make([]string, 0, len(fields)-1)
for _, searchDomain := range fields[1:] {
name := mDNS.Fqdn(searchDomain)
if name == "." {
continue
}
config.Search = append(config.Search, name)
}
case "options":
for _, option := range fields[1:] {
switch {
case strings.HasPrefix(option, "ndots:"):
value, parseErr := strconv.Atoi(option[len("ndots:"):])
if parseErr == nil {
config.Ndots = min(max(value, 0), 15)
}
case strings.HasPrefix(option, "timeout:"):
value, parseErr := strconv.Atoi(option[len("timeout:"):])
if parseErr == nil {
config.Timeout = time.Duration(max(value, 1)) * time.Second
}
case strings.HasPrefix(option, "attempts:"):
value, parseErr := strconv.Atoi(option[len("attempts:"):])
if parseErr == nil {
config.Attempts = max(value, 1)
}
case option == "rotate":
config.Rotate = true
case option == "single-request" || option == "single-request-reopen":
config.SingleRequest = true
case option == "use-vc" || option == "usevc" || option == "tcp":
config.UseTCP = true
case option == "trust-ad":
config.TrustAD = true
case option == "no-reload":
result.noReload = true
}
}
}
}
if len(config.Servers) == 0 {
config.Servers = defaultServers
}
if len(config.Search) == 0 {
config.Search = defaultSearch()
}
return result
}
@@ -0,0 +1,181 @@
package systemconfig
import (
"context"
"net/netip"
"os"
"slices"
"strconv"
"sync"
"syscall"
"time"
"unsafe"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/control"
M "github.com/sagernet/sing/common/metadata"
"github.com/sagernet/sing/common/x/list"
"github.com/sagernet/sing/service"
"golang.org/x/sys/windows"
)
type Source struct {
interfaceMonitor tun.DefaultInterfaceMonitor
access sync.Mutex
updateCallback *list.Element[tun.DefaultInterfaceUpdateCallback]
stale bool
config *Config
}
func NewSource(ctx context.Context) *Source {
source := &Source{}
interfaceMonitor := service.FromContext[adapter.NetworkManager](ctx).InterfaceMonitor()
if interfaceMonitor != nil {
source.interfaceMonitor = interfaceMonitor
source.updateCallback = interfaceMonitor.RegisterCallback(source.interfaceUpdated)
}
return source
}
func (s *Source) Configuration() *Config {
s.access.Lock()
defer s.access.Unlock()
if s.config != nil && !s.stale && s.updateCallback != nil {
return s.config
}
s.stale = false
config := s.readConfig()
if s.config != nil && config.Equal(s.config) {
return s.config
}
s.config = config
return config
}
func (s *Source) interfaceUpdated(defaultInterface *control.Interface, flags int) {
s.access.Lock()
s.stale = true
s.access.Unlock()
}
func (s *Source) Reset() {
s.access.Lock()
s.stale = true
s.access.Unlock()
}
func (s *Source) Close() error {
s.access.Lock()
updateCallback := s.updateCallback
s.updateCallback = nil
s.access.Unlock()
if updateCallback != nil {
s.interfaceMonitor.UnregisterCallback(updateCallback)
}
return nil
}
func (s *Source) readConfig() *Config {
config := &Config{
Ndots: 1,
Timeout: 5 * time.Second,
Attempts: 2,
}
defer func() {
if len(config.Servers) == 0 {
config.Servers = defaultServers
}
if len(config.Search) == 0 {
config.Search = defaultSearch()
}
}()
addresses, err := adapterAddresses()
if err != nil {
return config
}
var dnsAddresses []struct {
ifName string
netip.Addr
}
for _, address := range addresses {
if address.OperStatus != windows.IfOperStatusUp {
continue
}
if address.IfType == windows.IF_TYPE_TUNNEL {
continue
}
if address.FirstGatewayAddress == nil {
continue
}
for dnsServerAddress := address.FirstDnsServerAddress; dnsServerAddress != nil; dnsServerAddress = dnsServerAddress.Next {
rawSockaddr, sockaddrErr := dnsServerAddress.Address.Sockaddr.Sockaddr()
if sockaddrErr != nil {
continue
}
var dnsServerAddr netip.Addr
switch sockaddr := rawSockaddr.(type) {
case *syscall.SockaddrInet4:
dnsServerAddr = netip.AddrFrom4(sockaddr.Addr)
case *syscall.SockaddrInet6:
if sockaddr.Addr[0] == 0xfe && sockaddr.Addr[1] == 0xc0 {
// fec0::/10 site local anycast addresses are set by
// Windows itself when no IPv6 DNS server is configured.
continue
}
dnsServerAddr = netip.AddrFrom16(sockaddr.Addr)
if sockaddr.ZoneId != 0 {
dnsServerAddr = dnsServerAddr.WithZone(strconv.FormatInt(int64(sockaddr.ZoneId), 10))
}
default:
continue
}
dnsAddresses = append(dnsAddresses, struct {
ifName string
netip.Addr
}{ifName: windows.UTF16PtrToString(address.FriendlyName), Addr: dnsServerAddr})
}
}
var myInterfaces []string
if s.interfaceMonitor != nil {
myInterfaces = s.interfaceMonitor.MyInterfaces()
}
var servers []M.Socksaddr
for _, address := range dnsAddresses {
if slices.Contains(myInterfaces, address.ifName) {
continue
}
servers = append(servers, M.SocksaddrFrom(address.Addr, 53))
}
config.Servers = common.Uniq(servers)
return config
}
func adapterAddresses() ([]*windows.IpAdapterAddresses, error) {
var b []byte
l := uint32(15000)
for {
b = make([]byte, l)
const flags = windows.GAA_FLAG_INCLUDE_PREFIX | windows.GAA_FLAG_INCLUDE_GATEWAYS
err := windows.GetAdaptersAddresses(syscall.AF_UNSPEC, flags, 0, (*windows.IpAdapterAddresses)(unsafe.Pointer(&b[0])), &l)
if err == nil {
if l == 0 {
return nil, nil
}
break
}
if err.(syscall.Errno) != syscall.ERROR_BUFFER_OVERFLOW {
return nil, os.NewSyscallError("getadaptersaddresses", err)
}
if l <= uint32(len(b)) {
return nil, os.NewSyscallError("getadaptersaddresses", err)
}
}
var aas []*windows.IpAdapterAddresses
for aa := (*windows.IpAdapterAddresses)(unsafe.Pointer(&b[0])); aa != nil; aa = aa.Next {
aas = append(aas, aa)
}
return aas, nil
}
+479
View File
@@ -0,0 +1,479 @@
package mdns
import (
"context"
"net"
"net/netip"
"slices"
"strings"
"time"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/dns/transport/local/systemconfig"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/control"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/json/badoption"
"github.com/sagernet/sing/common/logger"
"github.com/sagernet/sing/common/task"
"github.com/sagernet/sing/service"
mDNS "github.com/miekg/dns"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
)
const (
mdnsPort = 5353
mdnsClassTopBit = 1 << 15
mdnsTimeout = time.Second
)
var (
mdnsGroupIPv4 = net.IPv4(224, 0, 0, 251)
mdnsGroupIPv6 = net.ParseIP("ff02::fb")
mdnsLocalZones = []string{
"local.",
"254.169.in-addr.arpa.",
"8.e.f.ip6.arpa.",
"9.e.f.ip6.arpa.",
"a.e.f.ip6.arpa.",
"b.e.f.ip6.arpa.",
}
)
func IsLocalDomain(name string) bool {
canonical := mDNS.CanonicalName(name)
return common.Any(mdnsLocalZones, func(zone string) bool {
return canonical == zone || strings.HasSuffix(canonical, "."+zone)
})
}
func RegisterTransport(registry *dns.TransportRegistry) {
dns.RegisterTransport[option.MDNSDNSServerOptions](registry, C.DNSTypeMDNS, NewTransport)
}
var (
_ adapter.DNSTransport = (*Transport)(nil)
_ adapter.DNSTransportWithPreferredDomain = (*Transport)(nil)
_ adapter.DNSTransportWithEnvironment = (*Transport)(nil)
)
type Transport struct {
dns.TransportAdapter
ctx context.Context
logger logger.ContextLogger
networkManager adapter.NetworkManager
interfaceNames badoption.Listable[string]
configSource *systemconfig.Source
}
func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, options option.MDNSDNSServerOptions) (adapter.DNSTransport, error) {
return &Transport{
TransportAdapter: dns.NewTransportAdapterWithLocalOptions(C.DNSTypeMDNS, tag, options.LocalDNSServerOptions),
ctx: ctx,
logger: logger,
networkManager: service.FromContext[adapter.NetworkManager](ctx),
interfaceNames: options.Interface,
configSource: systemconfig.NewSource(ctx),
}, nil
}
func NewRawTransport(transportAdapter dns.TransportAdapter, ctx context.Context, logger log.ContextLogger) *Transport {
return &Transport{
TransportAdapter: transportAdapter,
ctx: ctx,
logger: logger,
networkManager: service.FromContext[adapter.NetworkManager](ctx),
}
}
func (t *Transport) Start(stage adapter.StartStage) error {
return nil
}
func (t *Transport) Close() error {
if t.configSource != nil {
return t.configSource.Close()
}
return nil
}
func (t *Transport) Reset() {
if t.configSource != nil {
t.configSource.Reset()
}
}
func (t *Transport) PreferredDomain(domain string) bool {
return IsLocalDomain(domain)
}
func (t *Transport) Environment() []string {
if t.configSource == nil {
return nil
}
return t.configSource.Configuration().Signature()
}
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
targets, err := t.queryTargets()
if err != nil {
return nil, E.Cause(err, "mdns: prepare interfaces")
}
request := makeQueryMessage(message)
rawMessage, err := request.Pack()
if err != nil {
return nil, E.Cause(err, "mdns: pack request")
}
deadline, loaded := ctx.Deadline()
if !loaded || deadline.IsZero() {
deadline = time.Now().Add(mdnsTimeout)
}
exchangeCtx, cancel := context.WithDeadline(ctx, deadline)
defer cancel()
results := make(chan exchangeResult, len(targets))
var group task.Group
for _, target := range targets {
group.Append0(func(ctx context.Context) error {
response, err := t.exchangeTarget(ctx, target, rawMessage, message.Question[0], deadline)
if err != nil || response != nil {
results <- exchangeResult{
response: response,
err: err,
}
}
return nil
})
}
groupErr := group.Run(exchangeCtx)
close(results)
response := newResponse(message)
seenRecords := make(map[string]bool)
var lastErr error
for result := range results {
if result.err != nil {
lastErr = result.err
t.logger.TraceContext(ctx, result.err)
continue
}
mergeResponse(response, result.response, seenRecords)
}
if len(response.Answer) > 0 || len(response.Ns) > 0 || len(response.Extra) > 0 {
return response, nil
}
if lastErr != nil {
return nil, lastErr
}
if groupErr != nil && ctx.Err() != nil {
return nil, groupErr
}
return nil, E.New("mdns: query timeout")
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
go func() {
callback(t.Exchange(ctx, message))
}()
}
type exchangeResult struct {
response *mDNS.Msg
err error
}
type queryTarget struct {
iface control.Interface
family string
}
func (t *Transport) exchangeTarget(ctx context.Context, target queryTarget, rawMessage []byte, question mDNS.Question, deadline time.Time) (*mDNS.Msg, error) {
packetConn, destination, err := t.listenPacket(ctx, target)
if err != nil {
return nil, err
}
defer packetConn.Close()
_, err = packetConn.WriteTo(rawMessage, destination)
if err != nil {
return nil, E.Cause(err, "mdns: write request on ", target.iface.Name, " ", target.family)
}
err = packetConn.SetReadDeadline(deadline)
if err != nil {
return nil, E.Cause(err, "mdns: set deadline on ", target.iface.Name, " ", target.family)
}
response := newResponseFromQuestion(question)
seenRecords := make(map[string]bool)
buffer := buf.Get(buf.UDPBufferSize)
defer buf.Put(buffer)
for {
n, source, readErr := packetConn.ReadFrom(buffer)
if readErr != nil {
if E.IsTimeout(readErr) {
if len(response.Answer) > 0 || len(response.Ns) > 0 || len(response.Extra) > 0 {
return response, nil
}
return nil, nil
}
return nil, E.Cause(readErr, "mdns: read response on ", target.iface.Name, " ", target.family)
}
if !validSource(source, target) {
continue
}
var candidate mDNS.Msg
err = candidate.Unpack(buffer[:n])
if err != nil {
t.logger.TraceContext(ctx, "mdns: unpack response: ", err)
continue
}
if !validResponse(&candidate, question) {
continue
}
normalizeResponse(&candidate, question)
mergeResponse(response, &candidate, seenRecords)
}
}
func (t *Transport) listenPacket(ctx context.Context, target queryTarget) (net.PacketConn, net.Addr, error) {
var listenConfig net.ListenConfig
listenConfig.Control = control.Append(listenConfig.Control, control.BindToInterface(t.networkManager.InterfaceFinder(), target.iface.Name, target.iface.Index))
netInterface := target.iface.NetInterface()
switch target.family {
case "udp4":
packetConn, err := listenConfig.ListenPacket(ctx, "udp4", "0.0.0.0:0")
if err != nil {
return nil, nil, E.Cause(err, "mdns: listen on ", target.iface.Name, " udp4")
}
ipv4Conn := ipv4.NewPacketConn(packetConn)
err = ipv4Conn.SetMulticastInterface(&netInterface)
if err != nil {
packetConn.Close()
return nil, nil, E.Cause(err, "mdns: set multicast interface on ", target.iface.Name, " udp4")
}
_ = ipv4Conn.SetMulticastTTL(255)
return packetConn, &net.UDPAddr{IP: mdnsGroupIPv4, Port: mdnsPort}, nil
case "udp6":
packetConn, err := listenConfig.ListenPacket(ctx, "udp6", "[::]:0")
if err != nil {
return nil, nil, E.Cause(err, "mdns: listen on ", target.iface.Name, " udp6")
}
ipv6Conn := ipv6.NewPacketConn(packetConn)
err = ipv6Conn.SetMulticastInterface(&netInterface)
if err != nil {
packetConn.Close()
return nil, nil, E.Cause(err, "mdns: set multicast interface on ", target.iface.Name, " udp6")
}
_ = ipv6Conn.SetMulticastHopLimit(255)
return packetConn, &net.UDPAddr{IP: mdnsGroupIPv6, Port: mdnsPort, Zone: target.iface.Name}, nil
default:
return nil, nil, E.New("mdns: unknown network: ", target.family)
}
}
func (t *Transport) queryTargets() ([]queryTarget, error) {
interfaces, err := t.fetchInterfaces()
if err != nil {
return nil, err
}
var targets []queryTarget
for _, iface := range interfaces {
supports4, supports6 := interfaceFamilies(iface)
if supports4 {
targets = append(targets, queryTarget{
iface: iface,
family: "udp4",
})
}
if supports6 {
targets = append(targets, queryTarget{
iface: iface,
family: "udp6",
})
}
}
if len(targets) == 0 {
return nil, E.New("missing usable mDNS interfaces")
}
return targets, nil
}
func (t *Transport) fetchInterfaces() ([]control.Interface, error) {
finder := t.networkManager.InterfaceFinder()
var interfaces []control.Interface
if len(t.interfaceNames) > 0 {
for _, interfaceName := range t.interfaceNames {
iface, err := finder.ByName(interfaceName)
if err != nil {
t.logger.Warn("mdns: interface ", interfaceName, " not found")
continue
}
if !isUsableInterface(*iface) {
t.logger.Warn("mdns: interface ", interfaceName, " is not usable")
continue
}
interfaces = append(interfaces, *iface)
}
} else {
interfaces = common.Filter(finder.Interfaces(), isUsableInterface)
}
if len(interfaces) == 0 {
return nil, E.New("mdns: missing usable interface")
}
return interfaces, nil
}
func isUsableInterface(iface control.Interface) bool {
return iface.Flags&net.FlagUp != 0 &&
iface.Flags&net.FlagMulticast != 0 &&
iface.Flags&net.FlagLoopback == 0
}
func interfaceFamilies(iface control.Interface) (supports4, supports6 bool) {
for _, prefix := range iface.Addresses {
addr := prefix.Addr()
if addr.IsLoopback() {
continue
}
if addr.Is4() {
supports4 = true
} else if addr.Is6() && !addr.Is4In6() {
supports6 = true
}
if supports4 && supports6 {
return
}
}
return
}
func makeQueryMessage(message *mDNS.Msg) *mDNS.Msg {
request := &mDNS.Msg{
Question: slices.Clone(message.Question),
}
for i := range request.Question {
stripQuestionClass(&request.Question[i])
}
return request
}
func newResponse(message *mDNS.Msg) *mDNS.Msg {
response := newResponseFromQuestion(message.Question[0])
response.Id = message.Id
return response
}
func newResponseFromQuestion(question mDNS.Question) *mDNS.Msg {
stripQuestionClass(&question)
return &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Response: true,
Authoritative: true,
Rcode: mDNS.RcodeSuccess,
},
Question: []mDNS.Question{question},
}
}
func validSource(source net.Addr, target queryTarget) bool {
sourceUDP, isUDP := source.(*net.UDPAddr)
if !isUDP || sourceUDP.Port != mdnsPort {
return false
}
sourceAddr, loaded := netip.AddrFromSlice(sourceUDP.IP)
if !loaded {
return false
}
sourceAddr = sourceAddr.Unmap()
if (target.family == "udp4" && !sourceAddr.Is4()) || (target.family == "udp6" && !sourceAddr.Is6()) {
return false
}
for _, prefix := range target.iface.Addresses {
if prefix.Contains(sourceAddr) {
return true
}
}
return false
}
func validResponse(response *mDNS.Msg, question mDNS.Question) bool {
if !response.Response ||
response.Opcode != mDNS.OpcodeQuery ||
response.Rcode != mDNS.RcodeSuccess {
return false
}
for _, responseQuestion := range response.Question {
if questionMatches(responseQuestion, question) {
return true
}
}
return responseHasMatchingRecord(response, question)
}
func responseHasMatchingRecord(response *mDNS.Msg, question mDNS.Question) bool {
for _, recordList := range [][]mDNS.RR{response.Answer, response.Ns, response.Extra} {
for _, record := range recordList {
if recordMatchesQuestion(record, question) {
return true
}
}
}
return false
}
func questionMatches(left mDNS.Question, right mDNS.Question) bool {
stripQuestionClass(&left)
stripQuestionClass(&right)
return left.Qtype == right.Qtype &&
left.Qclass == right.Qclass &&
strings.EqualFold(left.Name, right.Name)
}
func recordMatchesQuestion(record mDNS.RR, question mDNS.Question) bool {
header := record.Header()
return strings.EqualFold(header.Name, question.Name) &&
(question.Qtype == mDNS.TypeANY ||
header.Rrtype == question.Qtype ||
header.Rrtype == mDNS.TypeCNAME)
}
func normalizeResponse(response *mDNS.Msg, question mDNS.Question) {
response.Id = 0
response.Question = []mDNS.Question{question}
for i := range response.Question {
stripQuestionClass(&response.Question[i])
}
for _, recordList := range [][]mDNS.RR{response.Answer, response.Ns, response.Extra} {
for _, record := range recordList {
stripRecordClass(record)
}
}
}
func mergeResponse(destination *mDNS.Msg, source *mDNS.Msg, seenRecords map[string]bool) {
appendRecords := func(destinationRecords *[]mDNS.RR, sourceRecords []mDNS.RR) {
for _, record := range sourceRecords {
key := record.String()
if seenRecords[key] {
continue
}
seenRecords[key] = true
*destinationRecords = append(*destinationRecords, record)
}
}
appendRecords(&destination.Answer, source.Answer)
appendRecords(&destination.Ns, source.Ns)
appendRecords(&destination.Extra, source.Extra)
}
func stripQuestionClass(question *mDNS.Question) {
question.Qclass &^= mdnsClassTopBit
}
func stripRecordClass(record mDNS.RR) {
record.Header().Class &^= mdnsClassTopBit
}
+441
View File
@@ -0,0 +1,441 @@
package transport
import (
"context"
"errors"
"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 {
options queryMultiplexerOptions
connection *ConnPool[*multiplexConn]
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 {
net.Conn
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
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 {
return &queryMultiplexer{
options: options,
queries: make(map[uint16]*pendingQuery),
connection: NewConnPool(ConnPoolOptions[*multiplexConn]{
Mode: ConnPoolSingle,
IsAlive: func(conn *multiplexConn) bool {
return conn != nil
},
Close: func(conn *multiplexConn, cause error) {
conn.Close()
},
}),
}
}
func (m *queryMultiplexer) Close() error {
return m.connection.Close()
}
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()
}
func (m *queryMultiplexer) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
done := make(chan struct{})
var (
response *mDNS.Msg
err error
)
m.ExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
response = callbackResponse
err = callbackErr
close(done)
})
<-done
return response, err
}
func (m *queryMultiplexer) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
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) {
for firstAttempt := true; ; firstAttempt = false {
conn, connCtx, created, err := m.connection.AcquireShared(ctx, m.dialConn)
if err != nil {
callback(nil, err)
return
}
if created {
go m.recvLoop(conn)
}
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)
return
}
writeErr := m.options.write(conn, message, queryId)
if writeErr == nil {
return
}
pending := m.take(queryId)
m.connection.Invalidate(conn, writeErr)
if pending == nil {
return
}
if !created && firstAttempt {
continue
}
callback(nil, E.Cause(writeErr, "write request"))
return
}
}
func (m *queryMultiplexer) dialConn(ctx context.Context) (*multiplexConn, error) {
conn, err := m.options.dial(ctx)
if err != nil {
return nil, err
}
return &multiplexConn{Conn: conn}, nil
}
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
for {
m.queryId++
if _, exists := m.queries[m.queryId]; !exists {
break
}
if m.queryId == start {
return 0, E.New("no available query ID")
}
}
queryId := m.queryId
pending := &pendingQuery{
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.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.dispatch(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]
if !loaded {
m.queryAccess.Unlock()
return nil
}
delete(m.queries, queryId)
m.queryAccess.Unlock()
pending.stopContext()
pending.stopConn()
return pending
}
func (m *queryMultiplexer) complete(queryId uint16, response *mDNS.Msg, err error, releaseConn bool) {
pending := m.take(queryId)
if pending == nil {
return
}
if releaseConn {
m.connection.Release(pending.conn, true)
}
if response != nil {
response.Id = pending.message.Id
}
pending.callback(response, err)
}
func (m *queryMultiplexer) completeContextDone(queryId uint16, ctx context.Context) {
pending := m.take(queryId)
if pending == nil {
return
}
err := ctx.Err()
if errors.Is(err, context.DeadlineExceeded) && pending.conn.readEpoch.Load() == pending.readEpoch {
m.connection.Invalidate(pending.conn, err)
} else {
m.connection.Release(pending.conn, true)
}
pending.callback(nil, err)
}
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
}
conn.readEpoch.Add(1)
if message == nil {
continue
}
m.complete(message.Id, message, nil, true)
}
}
+546
View File
@@ -0,0 +1,546 @@
package transport
import (
"context"
"errors"
"io"
"net"
"sync/atomic"
"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)
}()
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 = multiplexer.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 = multiplexer.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 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")
if err != nil {
t.Fatal(err)
}
defer listener.Close()
accepted := make(chan net.Conn, 16)
go func() {
for {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
return
}
accepted <- conn
}
}()
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)
},
})
defer multiplexer.Close()
message := new(mDNS.Msg)
message.SetQuestion("example.com.", mDNS.TypeA)
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
start := time.Now()
_, err = multiplexer.Exchange(ctx, message)
elapsed := time.Since(start)
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatal("expected deadline exceeded, got ", err)
}
if elapsed > 2*time.Second {
t.Fatal("timeout not enforced, took ", elapsed)
}
firstConn := <-accepted
firstConn.SetReadDeadline(time.Now().Add(2 * time.Second))
_, err = io.Copy(io.Discard, firstConn)
if err != nil {
t.Fatal("expected the client side to close the connection, got ", err)
}
ctx2, cancel2 := context.WithTimeout(context.Background(), time.Second)
defer cancel2()
multiplexer.Exchange(ctx2, message)
select {
case <-accepted:
case <-time.After(time.Second):
t.Fatal("expected a fresh connection for the second query")
}
}
func TestMultiplexerSlowQueryKeepsActiveConn(t *testing.T) {
t.Parallel()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer listener.Close()
accepted := make(chan net.Conn, 16)
go func() {
for {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
return
}
accepted <- conn
go func() {
for {
request, readErr := ReadMessage(conn)
if readErr != nil {
return
}
if request.Question[0].Name == "slow.example.com." {
continue
}
response := new(mDNS.Msg)
response.SetReply(request)
WriteMessage(conn, request.Id, response)
}
}()
}
}()
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)
},
})
defer multiplexer.Close()
slowMessage := new(mDNS.Msg)
slowMessage.SetQuestion("slow.example.com.", mDNS.TypeA)
slowCtx, slowCancel := context.WithTimeout(context.Background(), time.Second)
defer slowCancel()
slowDone := make(chan error, 1)
go func() {
_, slowErr := multiplexer.Exchange(slowCtx, slowMessage)
slowDone <- slowErr
}()
select {
case <-accepted:
case <-time.After(time.Second):
t.Fatal("expected a connection for the slow query")
}
fastMessage := new(mDNS.Msg)
fastMessage.SetQuestion("fast.example.com.", mDNS.TypeA)
exchangeFast := func() {
fastCtx, fastCancel := context.WithTimeout(context.Background(), time.Second)
defer fastCancel()
_, fastErr := multiplexer.Exchange(fastCtx, fastMessage)
if fastErr != nil {
t.Fatal("fast query failed: ", fastErr)
}
}
deadline := time.Now().Add(3 * time.Second)
for {
if !time.Now().Before(deadline) {
t.Fatal("slow query did not complete")
}
exchangeFast()
select {
case slowErr := <-slowDone:
if !errors.Is(slowErr, context.DeadlineExceeded) {
t.Fatal("expected deadline exceeded for slow query, got ", slowErr)
}
exchangeFast()
if len(accepted) > 0 {
t.Fatal("slow query timeout must not replace the active connection")
}
return
case <-time.After(50 * time.Millisecond):
}
}
}
+7 -2
View File
@@ -22,7 +22,6 @@ import (
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
@@ -121,7 +120,7 @@ func (t *HTTP3Transport) newTransport() *http3.Transport {
if dialErr != nil {
return nil, dialErr
}
quicConn, dialErr := quic.DialEarly(ctx, bufio.NewUnbindPacketConn(conn), conn.RemoteAddr(), tlsCfg, cfg)
quicConn, dialErr := quic.DialEarlyConn(ctx, conn, tlsCfg, cfg)
if dialErr != nil {
conn.Close()
return nil, dialErr
@@ -209,3 +208,9 @@ func (t *HTTP3Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS
}
return &responseMessage, nil
}
func (t *HTTP3Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
go func() {
callback(t.Exchange(ctx, message))
}()
}
+7 -3
View File
@@ -17,7 +17,6 @@ import (
"github.com/sagernet/sing-box/option"
sQUIC "github.com/sagernet/sing-quic"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/bufio"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
@@ -109,8 +108,7 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
}
earlyConnection, err := sQUIC.DialEarly(
ctx,
bufio.NewUnbindPacketConn(rawConn),
t.serverAddr.UDPAddr(),
rawConn,
t.tlsConfig,
nil,
)
@@ -145,6 +143,12 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
return nil, err
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
go func() {
callback(t.Exchange(ctx, message))
}()
}
func (t *Transport) exchange(ctx context.Context, message *mDNS.Msg, conn *quic.Conn) (*mDNS.Msg, error) {
stream, err := conn.OpenStreamSync(ctx)
if err != nil {
+4
View File
@@ -71,3 +71,7 @@ func (t *SDNSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.
}
return t.client.Exchange(message, resolverInfo)
}
func (t *SDNSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
callback(t.Exchange(ctx, message))
}
+36 -23
View File
@@ -15,7 +15,6 @@ import (
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio/deadline"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
@@ -31,8 +30,9 @@ func RegisterTCP(registry *dns.TransportRegistry) {
type TCPTransport struct {
dns.TransportAdapter
dialer N.Dialer
serverAddr M.Socksaddr
dialer N.Dialer
serverAddr M.Socksaddr
multiplexer *queryMultiplexer
}
func NewTCP(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteDNSServerOptions) (adapter.DNSTransport, error) {
@@ -47,11 +47,33 @@ func NewTCP(ctx context.Context, logger log.ContextLogger, tag string, options o
if !serverAddr.IsValid() {
return nil, E.New("invalid server address: ", serverAddr)
}
return &TCPTransport{
TransportAdapter: dns.NewTransportAdapterWithRemoteOptions(C.DNSTypeTCP, tag, options),
dialer: transportDialer,
return NewTCPRaw(dns.NewTransportAdapterWithRemoteOptions(C.DNSTypeTCP, tag, options), transportDialer, serverAddr), nil
}
func NewTCPRaw(adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socksaddr) *TCPTransport {
t := &TCPTransport{
TransportAdapter: adapter,
dialer: dialer,
serverAddr: serverAddr,
}, nil
}
t.multiplexer = newQueryMultiplexer(queryMultiplexerOptions{
dial: func(ctx context.Context) (net.Conn, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial TCP connection")
}
return conn, nil
},
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,
probeReuse: true,
})
return t
}
func (t *TCPTransport) Start(stage adapter.StartStage) error {
@@ -62,28 +84,19 @@ func (t *TCPTransport) Start(stage adapter.StartStage) error {
}
func (t *TCPTransport) Close() error {
return nil
return t.multiplexer.Close()
}
func (t *TCPTransport) Reset() {
t.multiplexer.Reset()
}
func (t *TCPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial TCP connection")
}
defer conn.Close()
defer setConnDeadline(ctx, conn, deadline.NeedAdditionalReadDeadline(conn))()
err = WriteMessage(conn, 0, message)
if err != nil {
return nil, E.Cause(err, "write request")
}
response, err := ReadMessage(conn)
if err != nil {
return nil, E.Cause(err, "read response")
}
return response, nil
return t.multiplexer.Exchange(ctx, message)
}
func (t *TCPTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
t.multiplexer.ExchangeAsync(ctx, message, callback)
}
func setConnDeadline(ctx context.Context, conn net.Conn, needClose bool) func() {
+27 -67
View File
@@ -2,6 +2,7 @@ package transport
import (
"context"
"net"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/dialer"
@@ -11,7 +12,6 @@ import (
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/bufio/deadline"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
@@ -22,26 +22,16 @@ import (
var _ adapter.DNSTransport = (*TLSTransport)(nil)
const tlsDNSMaxInflight = 8
func RegisterTLS(registry *dns.TransportRegistry) {
dns.RegisterTransport[option.RemoteTLSDNSServerOptions](registry, C.DNSTypeTLS, NewTLS)
}
type TLSTransport struct {
dns.TransportAdapter
logger logger.ContextLogger
logger logger.ContextLogger
dialer tls.Dialer
serverAddr M.Socksaddr
tlsConfig tls.Config
connections *ConnPool[*tlsDNSConn]
}
type tlsDNSConn struct {
tls.Conn
queryId uint16
needDeadlineClose bool
multiplexer *queryMultiplexer
}
func NewTLS(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteTLSDNSServerOptions) (adapter.DNSTransport, error) {
@@ -66,23 +56,30 @@ func NewTLS(ctx context.Context, logger log.ContextLogger, tag string, options o
}
func NewTLSRaw(logger logger.ContextLogger, adapter dns.TransportAdapter, dialer N.Dialer, serverAddr M.Socksaddr, tlsConfig tls.Config) *TLSTransport {
return &TLSTransport{
t := &TLSTransport{
TransportAdapter: adapter,
logger: logger,
dialer: tls.NewDialer(dialer, tlsConfig),
serverAddr: serverAddr,
tlsConfig: tlsConfig,
connections: NewConnPool(ConnPoolOptions[*tlsDNSConn]{
Mode: ConnPoolOrdered,
MaxInflight: tlsDNSMaxInflight,
IsAlive: func(conn *tlsDNSConn) bool {
return conn != nil
},
Close: func(conn *tlsDNSConn, _ error) {
conn.Close()
},
}),
}
t.multiplexer = newQueryMultiplexer(queryMultiplexerOptions{
dial: func(ctx context.Context) (net.Conn, error) {
conn, err := t.dialer.DialTLSContext(ctx, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial TLS connection")
}
return conn, nil
},
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,
probeReuse: true,
})
return t
}
func (t *TLSTransport) Start(stage adapter.StartStage) error {
@@ -93,54 +90,17 @@ func (t *TLSTransport) Start(stage adapter.StartStage) error {
}
func (t *TLSTransport) Close() error {
return t.connections.Close()
return t.multiplexer.Close()
}
func (t *TLSTransport) Reset() {
t.connections.Reset()
t.multiplexer.Reset()
}
func (t *TLSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
var lastErr error
for range 2 {
conn, created, err := t.connections.Acquire(ctx, func(ctx context.Context) (*tlsDNSConn, error) {
tlsConn, err := t.dialer.DialTLSContext(ctx, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial TLS connection")
}
return &tlsDNSConn{
Conn: tlsConn,
needDeadlineClose: deadline.NeedAdditionalReadDeadline(tlsConn.NetConn()),
}, nil
})
if err != nil {
return nil, err
}
response, err := t.exchange(ctx, message, conn)
if err == nil {
t.connections.Release(conn, true)
return response, nil
}
lastErr = err
t.logger.DebugContext(ctx, "discarded pooled connection: ", err)
t.connections.Release(conn, false)
if created {
return nil, err
}
}
return nil, lastErr
return t.multiplexer.Exchange(ctx, message)
}
func (t *TLSTransport) exchange(ctx context.Context, message *mDNS.Msg, conn *tlsDNSConn) (*mDNS.Msg, error) {
defer setConnDeadline(ctx, conn, conn.needDeadlineClose)()
conn.queryId++
err := WriteMessage(conn, conn.queryId, message)
if err != nil {
return nil, E.Cause(err, "write request")
}
response, err := ReadMessage(conn)
if err != nil {
return nil, E.Cause(err, "read response")
}
return response, nil
func (t *TLSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
t.multiplexer.ExchangeAsync(ctx, message, callback)
}
+78 -156
View File
@@ -3,7 +3,6 @@ package transport
import (
"context"
"net"
"sync"
"sync/atomic"
"github.com/sagernet/sing-box/adapter"
@@ -12,6 +11,7 @@ import (
"github.com/sagernet/sing-box/dns"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio/deadline"
E "github.com/sagernet/sing/common/exceptions"
@@ -36,17 +36,7 @@ type UDPTransport struct {
serverAddr M.Socksaddr
udpSize atomic.Int32
connection *ConnPool[net.Conn]
callbackAccess sync.RWMutex
queryId uint16
callbacks map[uint16]*udpCallback
}
type udpCallback struct {
access sync.Mutex
response *mDNS.Msg
done chan struct{}
multiplexer *queryMultiplexer
}
func NewUDP(ctx context.Context, logger log.ContextLogger, tag string, options option.RemoteDNSServerOptions) (adapter.DNSTransport, error) {
@@ -70,18 +60,19 @@ func NewUDPRaw(logger logger.ContextLogger, adapter dns.TransportAdapter, dialer
logger: logger,
dialer: dialerInstance,
serverAddr: serverAddr,
callbacks: make(map[uint16]*udpCallback),
connection: NewConnPool(ConnPoolOptions[net.Conn]{
Mode: ConnPoolSingle,
IsAlive: func(conn net.Conn) bool {
return conn != nil
},
Close: func(conn net.Conn, cause error) {
conn.Close()
},
}),
}
t.udpSize.Store(2048)
t.multiplexer = newQueryMultiplexer(queryMultiplexerOptions{
dial: func(ctx context.Context) (net.Conn, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkUDP, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial UDP connection")
}
return conn, nil
},
write: t.writeQuery,
readNext: t.readResponse,
})
return t
}
@@ -93,28 +84,16 @@ func (t *UDPTransport) Start(stage adapter.StartStage) error {
}
func (t *UDPTransport) Close() error {
return t.connection.Close()
return t.multiplexer.Close()
}
func (t *UDPTransport) Reset() {
t.connection.Reset()
}
func (t *UDPTransport) nextAvailableQueryId() (uint16, error) {
start := t.queryId
for {
t.queryId++
if _, exists := t.callbacks[t.queryId]; !exists {
return t.queryId, nil
}
if t.queryId == start {
return 0, E.New("no available query ID")
}
}
t.multiplexer.Reset()
}
func (t *UDPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
response, err := t.exchange(ctx, message)
t.updateUDPSize(message)
response, err := t.multiplexer.Exchange(ctx, message)
if err != nil {
return nil, err
}
@@ -125,6 +104,67 @@ func (t *UDPTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.M
return response, nil
}
func (t *UDPTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
t.updateUDPSize(message)
t.multiplexer.ExchangeAsync(ctx, message, func(response *mDNS.Msg, err error) {
if err == nil && response.Truncated {
t.logger.InfoContext(ctx, "response truncated, retrying with TCP")
go func() {
callback(t.exchangeTCP(ctx, message))
}()
return
}
callback(response, err)
})
}
func (t *UDPTransport) updateUDPSize(message *mDNS.Msg) {
edns0Opt := message.IsEdns0()
if edns0Opt == nil {
return
}
udpSize := int32(edns0Opt.UDPSize())
for {
current := t.udpSize.Load()
if udpSize <= current {
return
}
if t.udpSize.CompareAndSwap(current, udpSize) {
t.Reset()
return
}
}
}
func (t *UDPTransport) writeQuery(conn net.Conn, message *mDNS.Msg, queryId uint16) error {
buffer := buf.NewSize(1 + message.Len())
defer buffer.Release()
exMessage := *message
exMessage.Compress = true
exMessage.Id = queryId
rawMessage, err := exMessage.PackBuffer(buffer.FreeBytes())
if err != nil {
return err
}
return common.Error(conn.Write(rawMessage))
}
func (t *UDPTransport) readResponse(conn net.Conn) (*mDNS.Msg, error) {
buffer := buf.NewSize(int(t.udpSize.Load()))
defer buffer.Release()
_, err := buffer.ReadOnceFrom(conn)
if err != nil {
return nil, err
}
var message mDNS.Msg
err = message.Unpack(buffer.Bytes())
if err != nil {
t.logger.Debug("discarded malformed UDP response: ", err)
return nil, nil
}
return &message, nil
}
func (t *UDPTransport) exchangeTCP(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
conn, err := t.dialer.DialContext(ctx, N.NetworkTCP, t.serverAddr)
if err != nil {
@@ -142,121 +182,3 @@ func (t *UDPTransport) exchangeTCP(ctx context.Context, message *mDNS.Msg) (*mDN
}
return response, nil
}
func (t *UDPTransport) exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
if edns0Opt := message.IsEdns0(); edns0Opt != nil {
udpSize := int32(edns0Opt.UDPSize())
for {
current := t.udpSize.Load()
if udpSize <= current {
break
}
if t.udpSize.CompareAndSwap(current, udpSize) {
t.Reset()
break
}
}
}
conn, connCtx, created, err := t.connection.AcquireShared(ctx, func(ctx context.Context) (net.Conn, error) {
rawConn, err := t.dialer.DialContext(ctx, N.NetworkUDP, t.serverAddr)
if err != nil {
return nil, E.Cause(err, "dial UDP connection")
}
return rawConn, nil
})
if err != nil {
return nil, err
}
if created {
go t.recvLoop(conn)
}
callback := &udpCallback{
done: make(chan struct{}),
}
t.callbackAccess.Lock()
queryId, err := t.nextAvailableQueryId()
if err != nil {
t.callbackAccess.Unlock()
t.connection.Release(conn, true)
return nil, err
}
t.callbacks[queryId] = callback
t.callbackAccess.Unlock()
defer func() {
t.callbackAccess.Lock()
delete(t.callbacks, queryId)
t.callbackAccess.Unlock()
}()
buffer := buf.NewSize(1 + message.Len())
defer buffer.Release()
exMessage := *message
exMessage.Compress = true
originalId := message.Id
exMessage.Id = queryId
rawMessage, err := exMessage.PackBuffer(buffer.FreeBytes())
if err != nil {
return nil, err
}
_, err = conn.Write(rawMessage)
if err != nil {
t.connection.Invalidate(conn, err)
return nil, E.Cause(err, "write request")
}
select {
case <-callback.done:
t.connection.Release(conn, true)
callback.response.Id = originalId
return callback.response, nil
case <-connCtx.Done():
return nil, context.Cause(connCtx)
case <-ctx.Done():
t.connection.Release(conn, true)
return nil, ctx.Err()
}
}
func (t *UDPTransport) recvLoop(conn net.Conn) {
for {
buffer := buf.NewSize(int(t.udpSize.Load()))
_, err := buffer.ReadOnceFrom(conn)
if err != nil {
buffer.Release()
t.connection.Invalidate(conn, err)
return
}
var message mDNS.Msg
err = message.Unpack(buffer.Bytes())
buffer.Release()
if err != nil {
t.logger.Debug("discarded malformed UDP response: ", err)
continue
}
t.callbackAccess.RLock()
callback, loaded := t.callbacks[message.Id]
t.callbackAccess.RUnlock()
if !loaded {
continue
}
callback.access.Lock()
select {
case <-callback.done:
default:
callback.response = &message
close(callback.done)
}
callback.access.Unlock()
}
}
-23
View File
@@ -1,21 +1,13 @@
package dns
import (
"net/netip"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/option"
)
var _ adapter.LegacyDNSTransport = (*TransportAdapter)(nil)
type TransportAdapter struct {
transportType string
transportTag string
dependencies []string
strategy C.DomainStrategy
clientSubnet netip.Prefix
}
func NewTransportAdapter(transportType string, transportTag string, dependencies []string) TransportAdapter {
@@ -35,8 +27,6 @@ func NewTransportAdapterWithLocalOptions(transportType string, transportTag stri
transportType: transportType,
transportTag: transportTag,
dependencies: dependencies,
strategy: C.DomainStrategy(localOptions.LegacyStrategy),
clientSubnet: localOptions.LegacyClientSubnet,
}
}
@@ -45,15 +35,10 @@ func NewTransportAdapterWithRemoteOptions(transportType string, transportTag str
if remoteOptions.DomainResolver != nil && remoteOptions.DomainResolver.Server != "" {
dependencies = append(dependencies, remoteOptions.DomainResolver.Server)
}
if remoteOptions.LegacyAddressResolver != "" {
dependencies = append(dependencies, remoteOptions.LegacyAddressResolver)
}
return TransportAdapter{
transportType: transportType,
transportTag: transportTag,
dependencies: dependencies,
strategy: C.DomainStrategy(remoteOptions.LegacyStrategy),
clientSubnet: remoteOptions.LegacyClientSubnet,
}
}
@@ -68,11 +53,3 @@ func (a *TransportAdapter) Tag() string {
func (a *TransportAdapter) Dependencies() []string {
return a.dependencies
}
func (a *TransportAdapter) LegacyStrategy() C.DomainStrategy {
return a.strategy
}
func (a *TransportAdapter) LegacyClientSubnet() netip.Prefix {
return a.clientSubnet
}
+10 -89
View File
@@ -2,104 +2,25 @@ package dns
import (
"context"
"net"
"time"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/dialer"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/option"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/service"
)
func NewLocalDialer(ctx context.Context, options option.LocalDNSServerOptions) (N.Dialer, error) {
if options.LegacyDefaultDialer {
return dialer.NewDefaultOutbound(ctx), nil
} else {
return dialer.NewWithOptions(dialer.Options{
Context: ctx,
Options: options.DialerOptions,
DirectResolver: true,
LegacyDNSDialer: options.Legacy,
})
}
return dialer.NewWithOptions(dialer.Options{
Context: ctx,
Options: options.DialerOptions,
DirectResolver: true,
})
}
func NewRemoteDialer(ctx context.Context, options option.RemoteDNSServerOptions) (N.Dialer, error) {
if options.LegacyDefaultDialer {
transportDialer := dialer.NewDefaultOutbound(ctx)
if options.LegacyAddressResolver != "" {
transport := service.FromContext[adapter.DNSTransportManager](ctx)
resolverTransport, loaded := transport.Transport(options.LegacyAddressResolver)
if !loaded {
return nil, E.New("address resolver not found: ", options.LegacyAddressResolver)
}
transportDialer = newTransportDialer(transportDialer, service.FromContext[adapter.DNSRouter](ctx), resolverTransport, C.DomainStrategy(options.LegacyAddressStrategy), time.Duration(options.LegacyAddressFallbackDelay))
} else if options.ServerIsDomain() {
return nil, E.New("missing address resolver for server: ", options.Server)
}
return transportDialer, nil
} else {
return dialer.NewWithOptions(dialer.Options{
Context: ctx,
Options: options.DialerOptions,
RemoteIsDomain: options.ServerIsDomain(),
DirectResolver: true,
LegacyDNSDialer: options.Legacy,
})
}
}
type legacyTransportDialer struct {
dialer N.Dialer
dnsRouter adapter.DNSRouter
transport adapter.DNSTransport
strategy C.DomainStrategy
fallbackDelay time.Duration
}
func newTransportDialer(dialer N.Dialer, dnsRouter adapter.DNSRouter, transport adapter.DNSTransport, strategy C.DomainStrategy, fallbackDelay time.Duration) *legacyTransportDialer {
return &legacyTransportDialer{
dialer,
dnsRouter,
transport,
strategy,
fallbackDelay,
}
}
func (d *legacyTransportDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
if destination.IsIP() {
return d.dialer.DialContext(ctx, network, destination)
}
addresses, err := d.dnsRouter.Lookup(ctx, destination.Fqdn, adapter.DNSQueryOptions{
Transport: d.transport,
Strategy: d.strategy,
return dialer.NewWithOptions(dialer.Options{
Context: ctx,
Options: options.DialerOptions,
RemoteIsDomain: options.ServerIsDomain(),
DirectResolver: true,
})
if err != nil {
return nil, err
}
return N.DialParallel(ctx, d.dialer, network, destination, addresses, d.strategy == C.DomainStrategyPreferIPv6, d.fallbackDelay)
}
func (d *legacyTransportDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
if destination.IsIP() {
return d.dialer.ListenPacket(ctx, destination)
}
addresses, err := d.dnsRouter.Lookup(ctx, destination.Fqdn, adapter.DNSQueryOptions{
Transport: d.transport,
Strategy: d.strategy,
})
if err != nil {
return nil, err
}
conn, _, err := N.ListenSerial(ctx, d.dialer, destination, addresses)
return conn, err
}
func (d *legacyTransportDialer) Upstream() any {
return d.dialer
}
+8
View File
@@ -2,6 +2,8 @@ package dns
import (
"context"
"maps"
"slices"
"sync"
"github.com/sagernet/sing-box/adapter"
@@ -44,6 +46,12 @@ func NewTransportRegistry() *TransportRegistry {
}
}
func (r *TransportRegistry) OptionTypes() []string {
r.access.Lock()
defer r.access.Unlock()
return slices.Sorted(maps.Keys(r.optionsType))
}
func (r *TransportRegistry) CreateOptions(transportType string) (any, bool) {
r.access.Lock()
defer r.access.Unlock()