Partition local DNS caches by interface signature

This commit is contained in:
世界
2026-08-30 17:41:45 +08:00
parent 28b25598ed
commit 0e75f5f45c
16 changed files with 523 additions and 312 deletions
+60 -7
View File
@@ -3,8 +3,10 @@ package dns
import (
"context"
"errors"
"hash/fnv"
"net"
"net/netip"
"strconv"
"time"
"github.com/sagernet/sing-box/adapter"
@@ -87,13 +89,55 @@ type dnsCacheKey struct {
dns.Question
transportTag string
clientSubnet netip.Prefix
environment uint64
}
func (k dnsCacheKey) persistentName() string {
if !k.clientSubnet.IsValid() {
return k.transportTag
name := k.transportTag
if k.clientSubnet.IsValid() {
name += "\x00" + k.clientSubnet.String()
}
return k.transportTag + "\x00" + k.clientSubnet.String()
if k.environment != 0 {
name += "\x01" + strconv.FormatUint(k.environment, 36)
}
return name
}
func (c *Client) newCacheKey(transport adapter.DNSTransport, question dns.Question, message *dns.Msg, options adapter.DNSQueryOptions) dnsCacheKey {
cacheKey := dnsCacheKey{Question: question, transportTag: transport.Tag(), clientSubnet: c.effectiveClientSubnet(message, options)}
environmentTransport, withEnvironment := transport.(adapter.DNSTransportWithEnvironment)
if withEnvironment {
cacheKey.environment = environmentHash(environmentTransport.Environment())
}
return cacheKey
}
func (c *Client) finishCacheKey(transport adapter.DNSTransport, key dnsCacheKey) (dnsCacheKey, bool) {
environmentTransport, withEnvironment := transport.(adapter.DNSTransportWithEnvironment)
if !withEnvironment {
return key, true
}
environment := environmentHash(environmentTransport.Environment())
if environment == key.environment {
return key, true
}
if key.environment == 0 {
key.environment = environment
return key, true
}
return key, false
}
func environmentHash(environment []string) uint64 {
if len(environment) == 0 {
return 0
}
digest := fnv.New64a()
for _, entry := range environment {
digest.Write([]byte(entry))
digest.Write([]byte{0})
}
return digest.Sum64()
}
func (c *Client) effectiveClientSubnet(message *dns.Msg, options adapter.DNSQueryOptions) netip.Prefix {
@@ -231,7 +275,7 @@ func (c *Client) beginExchange(ctx context.Context, transport adapter.DNSTranspo
disableCache: disableCache,
}
if !disableCache {
cacheKey := dnsCacheKey{Question: question, transportTag: transport.Tag(), clientSubnet: c.effectiveClientSubnet(message, options)}
cacheKey := c.newCacheKey(transport, question, message, options)
operation.cacheKey = cacheKey
cond, loaded := c.cacheLock.LoadOrStore(cacheKey, make(chan struct{}))
if loaded {
@@ -243,6 +287,8 @@ func (c *Client) beginExchange(ctx context.Context, transport adapter.DNSTranspo
case <-ctx.Done():
return nil, nil, exchangeDone, ctx.Err()
}
cacheKey = c.newCacheKey(transport, question, message, options)
operation.cacheKey = cacheKey
} else {
operation.releaseCond = func() {
c.cacheLock.Delete(cacheKey)
@@ -303,7 +349,10 @@ func (c *Client) finishExchange(transport adapter.DNSTransport, operation *excha
}
timeToLive := applyResponseOptions(question, response, operation.options)
if !disableCache {
c.storeCache(operation.cacheKey, response, timeToLive)
cacheKey, storable := c.finishCacheKey(transport, operation.cacheKey)
if storable {
c.storeCache(cacheKey, response, timeToLive)
}
}
response.Id = operation.messageId
requestEDNSOpt := operation.message.IsEdns0()
@@ -476,7 +525,7 @@ func (c *Client) lookupToExchange(ctx context.Context, transport adapter.DNSTran
func (c *Client) questionCache(ctx context.Context, transport adapter.DNSTransport, message *dns.Msg, options adapter.DNSQueryOptions, responseChecker func(response *dns.Msg) bool) ([]netip.Addr, error) {
question := message.Question[0]
cacheKey := dnsCacheKey{Question: question, transportTag: transport.Tag(), clientSubnet: c.effectiveClientSubnet(message, options)}
cacheKey := c.newCacheKey(transport, question, message, options)
response, _, isStale := c.loadResponse(cacheKey)
if response == nil {
return nil, ErrNotCached
@@ -613,8 +662,12 @@ func (c *Client) backgroundRefreshDNS(transport adapter.DNSTransport, key dnsCac
} else if response.Rcode != dns.RcodeSuccess && response.Rcode != dns.RcodeNameError {
return
}
storeKey, storable := c.finishCacheKey(transport, key)
if !storable {
return
}
timeToLive := applyResponseOptions(key.Question, response, options)
c.storeCache(key, response, timeToLive)
c.storeCache(storeKey, response, timeToLive)
logRefreshedResponse(c.logger, ctx, response, timeToLive)
}()
}
+142 -100
View File
@@ -39,7 +39,10 @@ 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")
@@ -52,18 +55,21 @@ type Transport struct {
platformInterface adapter.PlatformInterface
interfaceName string
interfaceCallback *list.Element[tun.DefaultInterfaceUpdateCallback]
transportLock sync.RWMutex
updatedAt time.Time
lastError error
servers []M.Socksaddr
serverTransports []adapter.DNSTransport
refreshing atomic.Bool
search []string
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) {
transportDialer, err := dns.NewLocalDialer(ctx, options.LocalDNSServerOptions)
if err != nil {
@@ -120,26 +126,40 @@ func (t *Transport) Close() error {
if t.interfaceCallback != nil {
t.networkManager.InterfaceMonitor().UnregisterCallback(t.interfaceCallback)
}
t.transportLock.Lock()
defer t.transportLock.Unlock()
t.closeServerTransports()
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.closeServerTransports()
t.transportLock.Unlock()
t.refreshAccess.Lock()
defer t.refreshAccess.Unlock()
state := t.savedState.Swap(nil)
if state != nil {
closeServerTransports(state.serverTransports)
}
}
func (t *Transport) closeServerTransports() {
for _, serverTransport := range t.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()
}
t.serverTransports = nil
}
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
@@ -158,23 +178,23 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
}
func (t *Transport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
t.transportLock.RLock()
updatedAt := t.updatedAt
lastError := t.lastError
serverTransports := t.serverTransports
t.transportLock.RUnlock()
if lastError != nil {
callback(nil, E.Cause(lastError, "dhcp: fetch DNS servers"))
return
}
if len(serverTransports) == 0 {
state := t.savedState.Load()
if state == nil {
go t.exchangeCold(ctx, message, callback)
return
}
if time.Since(updatedAt) >= C.DHCPTTL {
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, serverTransports, callback)
t.exchangeWithTransports(ctx, message, state, callback)
}
func (t *Transport) exchangeCold(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) {
@@ -183,62 +203,60 @@ func (t *Transport) exchangeCold(ctx context.Context, message *mDNS.Msg, callbac
callback(nil, E.Cause(err, "dhcp: fetch DNS servers"))
return
}
t.transportLock.RLock()
serverTransports := t.serverTransports
t.transportLock.RUnlock()
if len(serverTransports) == 0 {
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, serverTransports, callback)
t.exchangeWithTransports(ctx, message, state, callback)
}
func (t *Transport) Fetch() []M.Socksaddr {
t.transportLock.RLock()
updatedAt := t.updatedAt
lastError := t.lastError
servers := t.servers
t.transportLock.RUnlock()
if lastError != nil {
state := t.savedState.Load()
if state == nil || state.lastError != nil {
return nil
}
if len(servers) > 0 && time.Since(updatedAt) >= C.DHCPTTL {
if len(state.servers) > 0 && time.Since(state.updatedAt) >= C.DHCPTTL {
t.startRefresh()
}
return servers
return state.servers
}
func (t *Transport) fetch() error {
t.transportLock.RLock()
updatedAt := t.updatedAt
lastError := t.lastError
t.transportLock.RUnlock()
if lastError != nil {
return lastError
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 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 nil
}
return t.updateServers()
return t.updateServersLocked()
}
func (t *Transport) startRefresh() {
if !t.refreshing.CompareAndSwap(false, true) {
if !t.refreshAccess.TryLock() {
return
}
go func() {
defer t.refreshing.Store(false)
t.transportLock.Lock()
defer t.transportLock.Unlock()
if time.Since(t.updatedAt) < C.DHCPTTL {
defer t.refreshAccess.Unlock()
state := t.savedState.Load()
if state != nil && time.Since(state.updatedAt) < C.DHCPTTL {
return
}
err := t.updateServers()
err := t.updateServersLocked()
if err != nil {
if errors.Is(err, errInterfaceIsCellular) && t.optional {
t.logger.Debug(E.Cause(err, "dhcp: refresh DNS servers"))
@@ -275,34 +293,47 @@ func (t *Transport) fetchInterface() (*control.Interface, error) {
}
}
func (t *Transport) updateServers() error {
func (t *Transport) updateServersLocked() error {
iface, err := t.fetchInterface()
if err != nil {
t.lastError = err
t.updatedAt = time.Now()
t.storeFailureLocked(err)
return E.Cause(err, "prepare interface")
}
t.logger.Info("dhcp: query DNS servers on ", iface.Name)
fetchCtx, cancel := context.WithTimeout(t.ctx, C.DHCPTimeout)
err = t.fetchServers0(fetchCtx, iface)
cancel()
t.updatedAt = time.Now()
if err != nil {
t.lastError = 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) {
t.transportLock.Lock()
err := t.updateServers()
t.transportLock.Unlock()
t.refreshAccess.Lock()
err := t.updateServersLocked()
t.refreshAccess.Unlock()
if err != nil {
if errors.Is(err, errInterfaceIsCellular) && t.optional {
t.logger.Debug(E.Cause(errInterfaceIsCellular, "dhcp: update DNS servers"))
@@ -390,44 +421,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.Info("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.Info("dhcp: updated DNS servers from ", iface.Name, ": [", strings.Join(common.Map(newState.servers, M.Socksaddr.String), ","), "], search: [", strings.Join(newState.search, ","), "]")
}
if !slices.Equal(t.servers, serverAddrs) || t.serverTransports == nil {
t.closeServerTransports()
serverTransports := make([]adapter.DNSTransport, 0, len(serverAddrs))
for _, serverAddr := range serverAddrs {
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)
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()
}
serverTransports = append(serverTransports, serverTransport)
return E.Cause(err, "initialize transport for ", serverAddr)
}
t.serverTransports = serverTransports
serverTransports = append(serverTransports, serverTransport)
}
newState.serverTransports = serverTransports
t.savedState.Store(newState)
if previousState != nil {
closeServerTransports(previousState.serverTransports)
}
t.servers = serverAddrs
return nil
}
+7 -7
View File
@@ -12,19 +12,19 @@ import (
mDNS "github.com/miekg/dns"
)
func (t *Transport) exchangeWithTransports(ctx context.Context, message *mDNS.Msg, serverTransports []adapter.DNSTransport, callback func(response *mDNS.Msg, err error)) {
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(domain)
names := t.nameList(state.search, domain)
if len(names) == 0 {
callback(nil, E.New("invalid domain: ", domain))
return
}
nameExchangers := make([]transport.AsyncExchanger, 0, len(names))
for _, fqdn := range names {
nameExchangers = append(nameExchangers, t.newNameExchanger(message, fqdn, serverTransports))
nameExchangers = append(nameExchangers, t.newNameExchanger(message, fqdn, state.serverTransports))
}
if len(serverTransports) == 1 || !(question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
if len(state.serverTransports) == 1 || !(question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
transport.ExchangeSequential(ctx, nameExchangers, nil, callback)
} else {
transport.ExchangeRace(ctx, nameExchangers, callback)
@@ -50,7 +50,7 @@ func (t *Transport) newNameExchanger(message *mDNS.Msg, fqdn string, serverTrans
}
}
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 {
@@ -68,11 +68,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)
+11 -2
View File
@@ -8,6 +8,7 @@ import (
"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/dns/transport/mdns"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
@@ -26,6 +27,7 @@ func RegisterTransport(registry *dns.TransportRegistry) {
var (
_ adapter.DNSTransport = (*Transport)(nil)
_ adapter.DNSTransportWithPreferredDomain = (*Transport)(nil)
_ adapter.DNSTransportWithEnvironment = (*Transport)(nil)
)
type Transport struct {
@@ -37,7 +39,7 @@ type Transport struct {
preferGo bool
resolved ResolvedResolver
mdnsTransport adapter.DNSTransport
configSource *systemConfigSource
configSource *systemconfig.Source
system systemResolver
serverSet atomic.Pointer[localServerSet]
serverSetAccess sync.Mutex
@@ -59,7 +61,7 @@ func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, opt
preferredResolver: preferredResolver,
dialer: transportDialer,
preferGo: options.PreferGo,
configSource: newSystemConfigSource(ctx),
configSource: systemconfig.NewSource(ctx),
}, nil
}
@@ -124,6 +126,13 @@ 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) {
done := make(chan struct{})
var (
+1
View File
@@ -10,6 +10,7 @@ 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))
}
+14 -1
View File
@@ -19,6 +19,7 @@ import (
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing-box/service/resolved"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/control"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
@@ -61,7 +62,8 @@ type DBusResolvedResolver struct {
}
type resolvedServerSet struct {
servers []resolvedServer
servers []resolvedServer
signature []string
}
type resolvedServer struct {
@@ -147,6 +149,14 @@ func (t *DBusResolvedResolver) 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 {
@@ -359,6 +369,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)
+14 -13
View File
@@ -7,13 +7,14 @@ import (
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/dns/transport/local/systemconfig"
E "github.com/sagernet/sing/common/exceptions"
mDNS "github.com/miekg/dns"
)
type localServerSet struct {
config *dnsConfig
config *systemconfig.Config
transports []adapter.DNSTransport
}
@@ -23,7 +24,7 @@ func (s *localServerSet) Close() {
}
}
func (t *Transport) serverSetFor(systemConfig *dnsConfig) (*localServerSet, error) {
func (t *Transport) serverSetFor(systemConfig *systemconfig.Config) (*localServerSet, error) {
serverSet := t.serverSet.Load()
if serverSet != nil && serverSet.config == systemConfig {
return serverSet, nil
@@ -34,10 +35,10 @@ func (t *Transport) serverSetFor(systemConfig *dnsConfig) (*localServerSet, erro
if serverSet != nil && serverSet.config == systemConfig {
return serverSet, nil
}
transports := make([]adapter.DNSTransport, 0, len(systemConfig.servers))
for _, serverAddr := range systemConfig.servers {
transports := make([]adapter.DNSTransport, 0, len(systemConfig.Servers))
for _, serverAddr := range systemConfig.Servers {
var serverTransport adapter.DNSTransport
if systemConfig.useTCP {
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)
@@ -69,7 +70,7 @@ func (t *Transport) exchangeAsync(ctx context.Context, message *mDNS.Msg, domain
callback(nil, err)
return
}
names := systemConfig.nameList(domain)
names := systemConfig.NameList(domain)
if len(names) == 0 {
callback(nil, E.New("invalid domain: ", domain))
return
@@ -79,23 +80,23 @@ func (t *Transport) exchangeAsync(ctx context.Context, message *mDNS.Msg, domain
nameExchangers = append(nameExchangers, newNameExchanger(systemConfig, serverSet, message, fqdn))
}
question := message.Question[0]
if systemConfig.singleRequest || !(question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
if systemConfig.SingleRequest || !(question.Qtype == mDNS.TypeA || question.Qtype == mDNS.TypeAAAA) {
transport.ExchangeSequential(ctx, nameExchangers, nil, callback)
} else {
transport.ExchangeRace(ctx, nameExchangers, callback)
}
}
func newNameExchanger(systemConfig *dnsConfig, serverSet *localServerSet, message *mDNS.Msg, fqdn string) transport.AsyncExchanger {
serverOffset := systemConfig.serverOffset()
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++ {
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) {
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)
})
-103
View File
@@ -1,103 +0,0 @@
package local
import (
"net/netip"
"os"
"slices"
"strings"
"sync/atomic"
"time"
M "github.com/sagernet/sing/common/metadata"
mDNS "github.com/miekg/dns"
)
var defaultNS = []M.Socksaddr{
M.SocksaddrFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 53),
M.SocksaddrFrom(netip.IPv6Loopback(), 53),
}
type dnsConfig 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 *dnsConfig) equal(other *dnsConfig) 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 *dnsConfig) serverOffset() uint32 {
if c.rotate {
return atomic.AddUint32(&c.soffset, 1) - 1
}
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 += "."
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 dnsDefaultSearch() []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)}
}
+113
View File
@@ -0,0 +1,113 @@
package systemconfig
import (
"net/netip"
"os"
"slices"
"strconv"
"strings"
"sync/atomic"
"time"
M "github.com/sagernet/sing/common/metadata"
mDNS "github.com/miekg/dns"
)
var defaultServers = []M.Socksaddr{
M.SocksaddrFrom(netip.AddrFrom4([4]byte{127, 0, 0, 1}), 53),
M.SocksaddrFrom(netip.IPv6Loopback(), 53),
}
type Config struct {
Servers []M.Socksaddr
Search []string
Ndots int
Timeout time.Duration
Attempts int
Rotate bool
soffset uint32
SingleRequest bool
UseTCP bool
TrustAD bool
}
func (c *Config) Equal(other *Config) bool {
return slices.Equal(c.Servers, other.Servers) &&
slices.Equal(c.Search, other.Search) &&
c.Ndots == other.Ndots &&
c.Timeout == other.Timeout &&
c.Attempts == other.Attempts &&
c.Rotate == other.Rotate &&
c.SingleRequest == other.SingleRequest &&
c.UseTCP == other.UseTCP &&
c.TrustAD == other.TrustAD
}
func (c *Config) Signature() []string {
signature := make([]string, 0, len(c.Servers)+len(c.Search)+1)
for _, server := range c.Servers {
signature = append(signature, server.String())
}
signature = append(signature, c.Search...)
return append(signature, "ndots:"+strconv.Itoa(c.Ndots))
}
func (c *Config) ServerOffset() uint32 {
if c.Rotate {
return atomic.AddUint32(&c.soffset, 1) - 1
}
return 0
}
func (c *Config) NameList(name string) []string {
l := len(name)
rooted := l > 0 && name[l-1] == '.'
if l > 254 || l == 254 && !rooted {
return nil
}
if rooted {
if avoidDNS(name) {
return nil
}
return []string{name}
}
hasNdots := strings.Count(name, ".") >= c.Ndots
name += "."
names := make([]string, 0, 1+len(c.Search))
if hasNdots && !avoidDNS(name) {
names = append(names, name)
}
for _, suffix := range c.Search {
fqdn := name + suffix
if !avoidDNS(fqdn) && len(fqdn) <= 254 {
names = append(names, fqdn)
}
}
if !hasNdots && !avoidDNS(name) {
names = append(names, name)
}
return names
}
func avoidDNS(name string) bool {
if name == "" {
return true
}
return strings.HasSuffix(strings.TrimSuffix(name, "."), ".onion")
}
func defaultSearch() []string {
hostname, err := os.Hostname()
if err != nil {
return nil
}
_, domain, found := strings.Cut(hostname, ".")
if !found || domain == "" {
return nil
}
return []string{mDNS.Fqdn(domain)}
}
@@ -1,6 +1,6 @@
//go:build cgo
package local
package systemconfig
/*
#include <dlfcn.h>
@@ -132,18 +132,18 @@ import (
mDNS "github.com/miekg/dns"
)
type systemConfigSource struct {
type Source struct {
interfaceMonitor tun.DefaultInterfaceMonitor
access sync.Mutex
notifyToken C.int
notifyValid bool
stale bool
interfaceIndex int
config *dnsConfig
config *Config
}
func newSystemConfigSource(ctx context.Context) *systemConfigSource {
source := &systemConfigSource{
func NewSource(ctx context.Context) *Source {
source := &Source{
interfaceMonitor: service.FromContext[adapter.NetworkManager](ctx).InterfaceMonitor(),
}
if C.box_dnsinfo_load() != 0 {
@@ -156,7 +156,7 @@ func newSystemConfigSource(ctx context.Context) *systemConfigSource {
return source
}
func (s *systemConfigSource) Configuration() *dnsConfig {
func (s *Source) Configuration() *Config {
interfaceIndex := s.defaultInterfaceIndex()
s.access.Lock()
defer s.access.Unlock()
@@ -175,14 +175,14 @@ func (s *systemConfigSource) Configuration() *dnsConfig {
return s.config
}
config := systemInfo.build(interfaceIndex)
if s.config != nil && config.equal(s.config) {
if s.config != nil && config.Equal(s.config) {
return s.config
}
s.config = config
return config
}
func (s *systemConfigSource) changedLocked() bool {
func (s *Source) changedLocked() bool {
if !s.notifyValid {
return true
}
@@ -194,13 +194,13 @@ func (s *systemConfigSource) changedLocked() bool {
return changed != 0
}
func (s *systemConfigSource) Reset() {
func (s *Source) Reset() {
s.access.Lock()
s.stale = true
s.access.Unlock()
}
func (s *systemConfigSource) Close() error {
func (s *Source) Close() error {
s.access.Lock()
defer s.access.Unlock()
if s.notifyValid {
@@ -210,7 +210,7 @@ func (s *systemConfigSource) Close() error {
return nil
}
func (s *systemConfigSource) defaultInterfaceIndex() int {
func (s *Source) defaultInterfaceIndex() int {
if s.interfaceMonitor == nil {
return 0
}
@@ -234,7 +234,7 @@ type dnsInfoConfig struct {
scopedResolvers []dnsInfoResolver
}
func (c *dnsInfoConfig) build(interfaceIndex int) *dnsConfig {
func (c *dnsInfoConfig) build(interfaceIndex int) *Config {
var selected dnsInfoResolver
if interfaceIndex != 0 {
selected = common.Find(c.scopedResolvers, func(it dnsInfoResolver) bool {
@@ -246,24 +246,24 @@ func (c *dnsInfoConfig) build(interfaceIndex int) *dnsConfig {
return it.domain == "" && len(it.servers) > 0
})
}
config := &dnsConfig{
ndots: 1,
timeout: 5 * time.Second,
attempts: 2,
config := &Config{
Ndots: 1,
Timeout: 5 * time.Second,
Attempts: 2,
}
if len(selected.servers) == 0 {
config.servers = defaultNS
config.search = dnsDefaultSearch()
config.Servers = defaultServers
config.Search = defaultSearch()
return config
}
config.servers = selected.servers
config.Servers = selected.servers
if len(selected.search) > 0 {
config.search = selected.search
config.Search = selected.search
} else {
config.search = dnsDefaultSearch()
config.Search = defaultSearch()
}
if selected.timeout > 0 {
config.timeout = selected.timeout
config.Timeout = selected.timeout
}
return config
}
@@ -1,6 +1,6 @@
//go:build !windows && !(darwin && cgo)
package local
package systemconfig
import (
"bufio"
@@ -20,30 +20,30 @@ import (
const resolvConfPath = "/etc/resolv.conf"
type systemConfigSource struct {
type Source struct {
updateAccess sync.Mutex
lastChecked time.Time
current atomic.Pointer[resolvConfig]
}
type resolvConfig struct {
config *dnsConfig
config *Config
mtime time.Time
noReload bool
}
func newSystemConfigSource(_ context.Context) *systemConfigSource {
source := &systemConfigSource{lastChecked: time.Now()}
source.current.Store(dnsReadConfig(resolvConfPath))
func NewSource(_ context.Context) *Source {
source := &Source{lastChecked: time.Now()}
source.current.Store(readResolvConfig(resolvConfPath))
return source
}
func (s *systemConfigSource) Configuration() *dnsConfig {
func (s *Source) Configuration() *Config {
s.tryUpdate()
return s.current.Load().config
}
func (s *systemConfigSource) tryUpdate() {
func (s *Source) tryUpdate() {
if s.current.Load().noReload {
return
}
@@ -65,41 +65,41 @@ func (s *systemConfigSource) tryUpdate() {
if mtime.Equal(current.mtime) {
return
}
updated := dnsReadConfig(resolvConfPath)
if updated.config.equal(current.config) {
updated := readResolvConfig(resolvConfPath)
if updated.config.Equal(current.config) {
updated.config = current.config
}
s.current.Store(updated)
}
func (s *systemConfigSource) Reset() {
func (s *Source) Reset() {
s.updateAccess.Lock()
s.lastChecked = time.Time{}
s.updateAccess.Unlock()
}
func (s *systemConfigSource) Close() error {
func (s *Source) Close() error {
return nil
}
func dnsReadConfig(path string) *resolvConfig {
config := &dnsConfig{
ndots: 1,
timeout: 5 * time.Second,
attempts: 2,
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 = defaultNS
config.search = dnsDefaultSearch()
config.Servers = defaultServers
config.Search = defaultSearch()
return result
}
defer file.Close()
fileInfo, err := file.Stat()
if err != nil {
config.servers = defaultNS
config.search = dnsDefaultSearch()
config.Servers = defaultServers
config.Search = defaultSearch()
return result
}
result.mtime = fileInfo.ModTime()
@@ -115,24 +115,24 @@ func dnsReadConfig(path string) *resolvConfig {
}
switch fields[0] {
case "nameserver":
if len(fields) > 1 && len(config.servers) < 3 {
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))
config.Servers = append(config.Servers, M.SocksaddrFrom(serverAddr, 53))
}
}
case "domain":
if len(fields) > 1 {
config.search = []string{mDNS.Fqdn(fields[1])}
config.Search = []string{mDNS.Fqdn(fields[1])}
}
case "search":
config.search = make([]string, 0, len(fields)-1)
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)
config.Search = append(config.Search, name)
}
case "options":
for _, option := range fields[1:] {
@@ -140,37 +140,37 @@ func dnsReadConfig(path string) *resolvConfig {
case strings.HasPrefix(option, "ndots:"):
value, parseErr := strconv.Atoi(option[len("ndots:"):])
if parseErr == nil {
config.ndots = min(max(value, 0), 15)
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
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)
config.Attempts = max(value, 1)
}
case option == "rotate":
config.rotate = true
config.Rotate = true
case option == "single-request" || option == "single-request-reopen":
config.singleRequest = true
config.SingleRequest = true
case option == "use-vc" || option == "usevc" || option == "tcp":
config.useTCP = true
config.UseTCP = true
case option == "trust-ad":
config.trustAD = true
config.TrustAD = true
case option == "no-reload":
result.noReload = true
}
}
}
}
if len(config.servers) == 0 {
config.servers = defaultNS
if len(config.Servers) == 0 {
config.Servers = defaultServers
}
if len(config.search) == 0 {
config.search = dnsDefaultSearch()
if len(config.Search) == 0 {
config.Search = defaultSearch()
}
return result
}
@@ -1,4 +1,4 @@
package local
package systemconfig
import (
"context"
@@ -22,16 +22,16 @@ import (
"golang.org/x/sys/windows"
)
type systemConfigSource struct {
type Source struct {
interfaceMonitor tun.DefaultInterfaceMonitor
access sync.Mutex
updateCallback *list.Element[tun.DefaultInterfaceUpdateCallback]
stale bool
config *dnsConfig
config *Config
}
func newSystemConfigSource(ctx context.Context) *systemConfigSource {
source := &systemConfigSource{}
func NewSource(ctx context.Context) *Source {
source := &Source{}
interfaceMonitor := service.FromContext[adapter.NetworkManager](ctx).InterfaceMonitor()
if interfaceMonitor != nil {
source.interfaceMonitor = interfaceMonitor
@@ -40,7 +40,7 @@ func newSystemConfigSource(ctx context.Context) *systemConfigSource {
return source
}
func (s *systemConfigSource) Configuration() *dnsConfig {
func (s *Source) Configuration() *Config {
s.access.Lock()
defer s.access.Unlock()
if s.config != nil && !s.stale && s.updateCallback != nil {
@@ -48,26 +48,26 @@ func (s *systemConfigSource) Configuration() *dnsConfig {
}
s.stale = false
config := s.readConfig()
if s.config != nil && config.equal(s.config) {
if s.config != nil && config.Equal(s.config) {
return s.config
}
s.config = config
return config
}
func (s *systemConfigSource) interfaceUpdated(defaultInterface *control.Interface, flags int) {
func (s *Source) interfaceUpdated(defaultInterface *control.Interface, flags int) {
s.access.Lock()
s.stale = true
s.access.Unlock()
}
func (s *systemConfigSource) Reset() {
func (s *Source) Reset() {
s.access.Lock()
s.stale = true
s.access.Unlock()
}
func (s *systemConfigSource) Close() error {
func (s *Source) Close() error {
s.access.Lock()
updateCallback := s.updateCallback
s.updateCallback = nil
@@ -78,18 +78,18 @@ func (s *systemConfigSource) Close() error {
return nil
}
func (s *systemConfigSource) readConfig() *dnsConfig {
config := &dnsConfig{
ndots: 1,
timeout: 5 * time.Second,
attempts: 2,
func (s *Source) readConfig() *Config {
config := &Config{
Ndots: 1,
Timeout: 5 * time.Second,
Attempts: 2,
}
defer func() {
if len(config.servers) == 0 {
config.servers = defaultNS
if len(config.Servers) == 0 {
config.Servers = defaultServers
}
if len(config.search) == 0 {
config.search = dnsDefaultSearch()
if len(config.Search) == 0 {
config.Search = defaultSearch()
}
}()
addresses, err := adapterAddresses()
@@ -149,7 +149,7 @@ func (s *systemConfigSource) readConfig() *dnsConfig {
}
servers = append(servers, M.SocksaddrFrom(address.Addr, 53))
}
config.servers = common.Uniq(servers)
config.Servers = common.Uniq(servers)
return config
}
+17
View File
@@ -11,6 +11,7 @@ import (
"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"
@@ -60,6 +61,7 @@ func RegisterTransport(registry *dns.TransportRegistry) {
var (
_ adapter.DNSTransport = (*Transport)(nil)
_ adapter.DNSTransportWithPreferredDomain = (*Transport)(nil)
_ adapter.DNSTransportWithEnvironment = (*Transport)(nil)
)
type Transport struct {
@@ -68,6 +70,7 @@ type Transport struct {
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) {
@@ -77,6 +80,7 @@ func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, opt
logger: logger,
networkManager: service.FromContext[adapter.NetworkManager](ctx),
interfaceNames: options.Interface,
configSource: systemconfig.NewSource(ctx),
}, nil
}
@@ -94,16 +98,29 @@ func (t *Transport) Start(stage adapter.StartStage) error {
}
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 {