From 0e75f5f45c9a39de11520fd116ade022a5a7652c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Tue, 11 Aug 2026 09:42:55 +0800 Subject: [PATCH] Partition local DNS caches by interface signature --- adapter/dns.go | 5 + dns/client.go | 67 ++++- dns/transport/dhcp/dhcp.go | 242 ++++++++++-------- dns/transport/dhcp/dhcp_shared.go | 14 +- dns/transport/local/local.go | 13 +- dns/transport/local/local_resolved.go | 1 + dns/transport/local/local_resolved_linux.go | 15 +- dns/transport/local/local_shared.go | 27 +- dns/transport/local/system_config.go | 103 -------- dns/transport/local/systemconfig/config.go | 113 ++++++++ .../source_darwin.go} | 44 ++-- .../source_resolv.go} | 74 +++--- .../source_windows.go} | 40 +-- dns/transport/mdns/mdns.go | 17 ++ experimental/libbox/dns.go | 15 ++ service/resolved/transport.go | 45 ++++ 16 files changed, 523 insertions(+), 312 deletions(-) delete mode 100644 dns/transport/local/system_config.go create mode 100644 dns/transport/local/systemconfig/config.go rename dns/transport/local/{system_config_darwin.go => systemconfig/source_darwin.go} (91%) rename dns/transport/local/{system_config_resolv.go => systemconfig/source_resolv.go} (65%) rename dns/transport/local/{system_config_windows.go => systemconfig/source_windows.go} (83%) diff --git a/adapter/dns.go b/adapter/dns.go index b399d6b4..977cf11d 100644 --- a/adapter/dns.go +++ b/adapter/dns.go @@ -94,6 +94,11 @@ type DNSTransportWithPreferredDomain interface { PreferredDomain(domain string) bool } +type DNSTransportWithEnvironment interface { + DNSTransport + Environment() []string +} + type DNSTransportRegistry interface { option.DNSTransportOptionsRegistry CreateDNSTransport(ctx context.Context, logger log.ContextLogger, tag string, transportType string, options any) (DNSTransport, error) diff --git a/dns/client.go b/dns/client.go index 904c99e5..7a239713 100644 --- a/dns/client.go +++ b/dns/client.go @@ -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) }() } diff --git a/dns/transport/dhcp/dhcp.go b/dns/transport/dhcp/dhcp.go index 3abc7cf5..113d308b 100644 --- a/dns/transport/dhcp/dhcp.go +++ b/dns/transport/dhcp/dhcp.go @@ -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 } diff --git a/dns/transport/dhcp/dhcp_shared.go b/dns/transport/dhcp/dhcp_shared.go index 3123d8f9..f840846b 100644 --- a/dns/transport/dhcp/dhcp_shared.go +++ b/dns/transport/dhcp/dhcp_shared.go @@ -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) diff --git a/dns/transport/local/local.go b/dns/transport/local/local.go index 7e33a55c..b7bf8347 100644 --- a/dns/transport/local/local.go +++ b/dns/transport/local/local.go @@ -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 ( diff --git a/dns/transport/local/local_resolved.go b/dns/transport/local/local_resolved.go index 13b2a434..b698f6a9 100644 --- a/dns/transport/local/local_resolved.go +++ b/dns/transport/local/local_resolved.go @@ -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)) } diff --git a/dns/transport/local/local_resolved_linux.go b/dns/transport/local/local_resolved_linux.go index 7f835946..cbbd3f36 100644 --- a/dns/transport/local/local_resolved_linux.go +++ b/dns/transport/local/local_resolved_linux.go @@ -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) diff --git a/dns/transport/local/local_shared.go b/dns/transport/local/local_shared.go index d26e5f80..271776c3 100644 --- a/dns/transport/local/local_shared.go +++ b/dns/transport/local/local_shared.go @@ -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) }) diff --git a/dns/transport/local/system_config.go b/dns/transport/local/system_config.go deleted file mode 100644 index 00f5aa3b..00000000 --- a/dns/transport/local/system_config.go +++ /dev/null @@ -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)} -} diff --git a/dns/transport/local/systemconfig/config.go b/dns/transport/local/systemconfig/config.go new file mode 100644 index 00000000..22a94f1f --- /dev/null +++ b/dns/transport/local/systemconfig/config.go @@ -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)} +} diff --git a/dns/transport/local/system_config_darwin.go b/dns/transport/local/systemconfig/source_darwin.go similarity index 91% rename from dns/transport/local/system_config_darwin.go rename to dns/transport/local/systemconfig/source_darwin.go index 5b59f761..e8bda779 100644 --- a/dns/transport/local/system_config_darwin.go +++ b/dns/transport/local/systemconfig/source_darwin.go @@ -1,6 +1,6 @@ //go:build cgo -package local +package systemconfig /* #include @@ -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 } diff --git a/dns/transport/local/system_config_resolv.go b/dns/transport/local/systemconfig/source_resolv.go similarity index 65% rename from dns/transport/local/system_config_resolv.go rename to dns/transport/local/systemconfig/source_resolv.go index ad2a530f..cc14a213 100644 --- a/dns/transport/local/system_config_resolv.go +++ b/dns/transport/local/systemconfig/source_resolv.go @@ -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 } diff --git a/dns/transport/local/system_config_windows.go b/dns/transport/local/systemconfig/source_windows.go similarity index 83% rename from dns/transport/local/system_config_windows.go rename to dns/transport/local/systemconfig/source_windows.go index ebeeb6c7..c45f06ea 100644 --- a/dns/transport/local/system_config_windows.go +++ b/dns/transport/local/systemconfig/source_windows.go @@ -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 } diff --git a/dns/transport/mdns/mdns.go b/dns/transport/mdns/mdns.go index 76851ebc..18542c55 100644 --- a/dns/transport/mdns/mdns.go +++ b/dns/transport/mdns/mdns.go @@ -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 { diff --git a/experimental/libbox/dns.go b/experimental/libbox/dns.go index 38e68740..7fcef2ce 100644 --- a/experimental/libbox/dns.go +++ b/experimental/libbox/dns.go @@ -15,6 +15,7 @@ import ( "github.com/sagernet/sing/common" E "github.com/sagernet/sing/common/exceptions" M "github.com/sagernet/sing/common/metadata" + "github.com/sagernet/sing/service" mDNS "github.com/miekg/dns" ) @@ -29,6 +30,7 @@ type platformTransport struct { dns.TransportAdapter iif LocalDNSTransport preferredResolver *local.PreferredDomainResolver + networkManager adapter.NetworkManager } func newPlatformTransport(ctx context.Context, logger log.ContextLogger, iif LocalDNSTransport, tag string, options option.LocalDNSServerOptions) (*platformTransport, error) { @@ -40,6 +42,7 @@ func newPlatformTransport(ctx context.Context, logger log.ContextLogger, iif Loc TransportAdapter: dns.NewTransportAdapterWithLocalOptions(C.DNSTypeLocal, tag, options), iif: iif, preferredResolver: preferredResolver, + networkManager: service.FromContext[adapter.NetworkManager](ctx), }, nil } @@ -59,6 +62,17 @@ func (p *platformTransport) PreferredDomain(domain string) bool { return p.preferredResolver.PreferredDomain(domain) } +func (p *platformTransport) Environment() []string { + if p.networkManager == nil { + return nil + } + defaultInterface := p.networkManager.DefaultNetworkInterface() + if defaultInterface == nil { + return nil + } + return defaultInterface.DNSServers +} + func (p *platformTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) { localResponse := p.preferredResolver.Lookup(message) if localResponse != nil { @@ -170,4 +184,5 @@ func (c *ExchangeContext) ErrnoCode(code int32) { var ( _ adapter.DNSTransport = (*platformTransport)(nil) _ adapter.DNSTransportWithPreferredDomain = (*platformTransport)(nil) + _ adapter.DNSTransportWithEnvironment = (*platformTransport)(nil) ) diff --git a/service/resolved/transport.go b/service/resolved/transport.go index 067017d6..4e53896f 100644 --- a/service/resolved/transport.go +++ b/service/resolved/transport.go @@ -6,6 +6,8 @@ import ( "context" "net/netip" "os" + "slices" + "strconv" "sync" "sync/atomic" "time" @@ -34,6 +36,7 @@ func RegisterTransport(registry *dns.TransportRegistry) { var ( _ adapter.DNSTransport = (*Transport)(nil) _ adapter.DNSTransportWithPreferredDomain = (*Transport)(nil) + _ adapter.DNSTransportWithEnvironment = (*Transport)(nil) ) type Transport struct { @@ -122,6 +125,48 @@ func (t *Transport) Reset() { } } +func (t *Transport) Environment() []string { + if t.service == nil { + return nil + } + t.service.linkAccess.RLock() + defer t.service.linkAccess.RUnlock() + linkIndexes := make([]int32, 0, len(t.service.links)) + for linkIndex := range t.service.links { + linkIndexes = append(linkIndexes, linkIndex) + } + slices.Sort(linkIndexes) + var environment []string + for _, linkIndex := range linkIndexes { + link := t.service.links[linkIndex] + linkEntry := "link:" + strconv.Itoa(int(linkIndex)) + if link.dnsOverTLS { + linkEntry += ":tls" + } + environment = append(environment, linkEntry) + for _, address := range link.address { + serverAddr, ok := netip.AddrFromSlice(address.Address) + if ok { + environment = append(environment, serverAddr.String()) + } + } + for _, address := range link.addressEx { + serverAddr, ok := netip.AddrFromSlice(address.Address) + if ok { + environment = append(environment, M.SocksaddrFrom(serverAddr, address.Port).String()+"/"+address.Name) + } + } + for _, domain := range link.domain { + if domain.RoutingOnly { + environment = append(environment, "routing-only:"+domain.Domain) + } else { + environment = append(environment, domain.Domain) + } + } + } + return environment +} + func (t *Transport) updateTransports(link *TransportLink) error { t.linkAccess.Lock() defer t.linkAccess.Unlock()