mirror of
https://github.com/shtorm-7/sing-box-extended.git
synced 2026-09-15 21:00:27 +00:00
Partition local DNS caches by interface signature
This commit is contained in:
@@ -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)
|
||||
|
||||
+60
-7
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
|
||||
@@ -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)}
|
||||
}
|
||||
@@ -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)}
|
||||
}
|
||||
+22
-22
@@ -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
|
||||
}
|
||||
+37
-37
@@ -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
|
||||
}
|
||||
+20
-20
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user