mirror of
https://github.com/shtorm-7/sing-box-extended.git
synced 2026-09-15 21:00:27 +00:00
Merge tag 'v1.14.0'
This commit is contained in:
+509
-320
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,4 +2,6 @@
|
||||
|
||||
package hosts
|
||||
|
||||
var DefaultPath = "/etc/hosts"
|
||||
func defaultPath() (string, error) {
|
||||
return "/etc/hosts", nil
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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):
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}()
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user