Refactor local DNS cache partitioning

This commit is contained in:
世界
2026-08-30 17:41:45 +08:00
parent 96fe4f7317
commit 4e91d92c5f
10 changed files with 434 additions and 56 deletions
+3
View File
@@ -3,6 +3,7 @@ package adapter
import (
"encoding/hex"
"net"
"net/netip"
"strings"
"time"
@@ -18,6 +19,7 @@ type NetworkManager interface {
UpdateInterfaces() error
DefaultNetworkInterface() *NetworkInterface
NetworkInterfaces() []NetworkInterface
NetworkEnvironment() uint64
AutoDetectInterface() bool
AutoDetectInterfaceFunc() control.Func
ProtectFunc() control.Func
@@ -76,6 +78,7 @@ type NetworkInterface struct {
control.Interface
Type C.InterfaceType
DNSServers []string
Gateways []netip.Addr
Expensive bool
Constrained bool
}
+32 -26
View File
@@ -2,6 +2,7 @@ package dns
import (
"context"
"encoding/binary"
"errors"
"hash/fnv"
"net"
@@ -18,6 +19,7 @@ import (
"github.com/sagernet/sing/common/task"
"github.com/sagernet/sing/contrab/freelru"
"github.com/sagernet/sing/contrab/maphash"
"github.com/sagernet/sing/service"
"github.com/miekg/dns"
)
@@ -43,6 +45,7 @@ type Client struct {
initRDRCFunc func() adapter.RDRCStore
dnsCache adapter.DNSCacheStore
initDNSCacheFunc func() adapter.DNSCacheStore
networkManager adapter.NetworkManager
logger logger.ContextLogger
cache *freelru.Cache[dnsCacheKey, *dns.Msg]
cacheLock compatible.Map[dnsCacheKey, chan struct{}]
@@ -104,53 +107,56 @@ func (k dnsCacheKey) persistentName() string {
}
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())
clientSubnet := options.ClientSubnet
if !clientSubnet.IsValid() {
clientSubnet = c.clientSubnet
}
if !clientSubnet.IsValid() {
clientSubnet = clientSubnetFromMessage(message)
}
return dnsCacheKey{
Question: question,
transportTag: transport.Tag(),
clientSubnet: clientSubnet,
environment: c.environmentHash(transport),
}
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 {
environment := c.environmentHash(transport)
if environment == key.environment || key.environment == 0 {
key.environment = environment
return key, true
}
return key, false
}
func environmentHash(environment []string) uint64 {
if len(environment) == 0 {
func (c *Client) environmentHash(transport adapter.DNSTransport) uint64 {
environmentTransport, withEnvironment := transport.(adapter.DNSTransportWithEnvironment)
if !withEnvironment {
return 0
}
var networkEnvironment uint64
if c.networkManager != nil {
networkEnvironment = c.networkManager.NetworkEnvironment()
}
environment := environmentTransport.Environment()
if len(environment) == 0 {
return networkEnvironment
}
digest := fnv.New64a()
for _, entry := range environment {
digest.Write([]byte(entry))
digest.Write([]byte{0})
}
var hashBytes [8]byte
binary.BigEndian.PutUint64(hashBytes[:], networkEnvironment)
digest.Write(hashBytes[:])
return digest.Sum64()
}
func (c *Client) effectiveClientSubnet(message *dns.Msg, options adapter.DNSQueryOptions) netip.Prefix {
if options.ClientSubnet.IsValid() {
return options.ClientSubnet
}
if c.clientSubnet.IsValid() {
return c.clientSubnet
}
return clientSubnetFromMessage(message)
}
func (c *Client) Start() {
c.networkManager = service.FromContext[adapter.NetworkManager](c.ctx)
if c.initRDRCFunc != nil {
c.rdrc = c.initRDRCFunc()
}
+1
View File
@@ -106,6 +106,7 @@ type NetworkInterface struct {
Type int32
DNSServer StringIterator
Gateway StringIterator
Metered bool
}
+4
View File
@@ -141,6 +141,10 @@ func (w *platformInterfaceWrapper) NetworkInterfaces() ([]adapter.NetworkInterfa
},
Type: C.InterfaceType(netInterface.Type),
DNSServers: iteratorToArray[string](netInterface.DNSServer),
Gateways: common.Filter(common.Map(iteratorToArray[string](netInterface.Gateway), func(it string) netip.Addr {
gateway, _ := netip.ParseAddr(it)
return gateway.Unmap().WithZone("")
}), netip.Addr.IsValid),
Expensive: netInterface.Metered || isDefault && w.isExpensive,
Constrained: isDefault && w.isConstrained,
})
+18 -6
View File
@@ -55,7 +55,10 @@ type NetworkManager struct {
needWIFIState bool
wifiMonitor settings.WIFIMonitor
wifiState adapter.WIFIState
wifiStateMutex sync.RWMutex
networkEnvironment uint64
stateAccess sync.RWMutex
environmentUpdateAccess sync.Mutex
environmentUpdateTimer *time.Timer
started bool
}
@@ -117,6 +120,7 @@ func NewNetworkManager(ctx context.Context, logger logger.ContextLogger, options
return nil, E.Cause(err, "create network monitor")
}
nm.networkMonitor = networkMonitor
networkMonitor.RegisterCallback(nm.postUpdateNetworkEnvironment)
interfaceMonitor, err := tun.NewDefaultInterfaceMonitor(nm.networkMonitor, logger, tun.DefaultInterfaceMonitorOptions{
InterfaceFinder: nm.interfaceFinder,
OverrideAndroidVPN: options.OverrideAndroidVPN,
@@ -254,6 +258,11 @@ func (r *NetworkManager) Close() error {
})
monitor.Finish()
}
r.environmentUpdateAccess.Lock()
if r.environmentUpdateTimer != nil {
r.environmentUpdateTimer.Stop()
}
r.environmentUpdateAccess.Unlock()
if r.wifiMonitor != nil {
monitor.Start("close WIFI monitor")
err = E.Append(err, r.wifiMonitor.Close(), func(err error) error {
@@ -269,6 +278,7 @@ func (r *NetworkManager) InterfaceFinder() control.InterfaceFinder {
}
func (r *NetworkManager) UpdateInterfaces() error {
defer r.updateNetworkEnvironment()
if r.platformInterface == nil || !r.platformInterface.UsePlatformNetworkInterfaces() {
return r.interfaceFinder.Update()
} else {
@@ -423,24 +433,25 @@ func (r *NetworkManager) NeedWIFIState() bool {
}
func (r *NetworkManager) WIFIState() adapter.WIFIState {
r.wifiStateMutex.RLock()
defer r.wifiStateMutex.RUnlock()
r.stateAccess.RLock()
defer r.stateAccess.RUnlock()
return r.wifiState
}
func (r *NetworkManager) onWIFIStateChanged(state adapter.WIFIState) {
state.BSSID = adapter.NormalizeWIFIBSSID(state.BSSID)
r.wifiStateMutex.Lock()
r.stateAccess.Lock()
if state != r.wifiState {
r.wifiState = state
r.wifiStateMutex.Unlock()
r.stateAccess.Unlock()
r.postUpdateNetworkEnvironment()
if state.SSID != "" {
r.logger.Info("WIFI state changed: SSID=", state.SSID, ", BSSID=", state.BSSID)
} else {
r.logger.Info("WIFI disconnected")
}
} else {
r.wifiStateMutex.Unlock()
r.stateAccess.Unlock()
}
}
@@ -521,6 +532,7 @@ func (r *NetworkManager) notifyInterfaceUpdate(defaultInterface *control.Interfa
}
r.logger.Info("updated default interface ", defaultInterface.Name, ", ", strings.Join(options, ", "))
r.UpdateWIFIState()
r.updateNetworkEnvironment()
if !r.started {
return
+84
View File
@@ -0,0 +1,84 @@
package route
import (
"hash/fnv"
"net/netip"
"slices"
"strings"
"time"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing/common"
)
func (r *NetworkManager) NetworkEnvironment() uint64 {
r.stateAccess.RLock()
defer r.stateAccess.RUnlock()
return r.networkEnvironment
}
func (r *NetworkManager) postUpdateNetworkEnvironment() {
r.environmentUpdateAccess.Lock()
defer r.environmentUpdateAccess.Unlock()
if r.environmentUpdateTimer == nil {
r.environmentUpdateTimer = time.AfterFunc(time.Second, r.updateNetworkEnvironment)
} else {
r.environmentUpdateTimer.Reset(time.Second)
}
}
func (r *NetworkManager) updateNetworkEnvironment() {
r.environmentUpdateAccess.Lock()
defer r.environmentUpdateAccess.Unlock()
if r.environmentUpdateTimer != nil {
r.environmentUpdateTimer.Stop()
}
var defaultInterface *adapter.NetworkInterface
if r.interfaceMonitor != nil {
defaultInterface = r.DefaultNetworkInterface()
}
var environment []string
if defaultInterface != nil {
gateways := defaultInterface.Gateways
if len(gateways) == 0 {
gateways = systemGateways(defaultInterface.Interface.Index)
}
gateways = common.Uniq(gateways)
slices.SortFunc(gateways, netip.Addr.Compare)
for _, gateway := range gateways {
environment = append(environment, "gateway:"+gateway.String())
}
wifiState := r.WIFIState()
if wifiState.SSID != "" {
environment = append(environment, "ssid:"+wifiState.SSID)
} else if len(gateways) > 0 {
hardwareAddresses := systemNeighborHardwareAddresses(defaultInterface.Interface.Index, gateways)
for _, gateway := range gateways {
hardwareAddress := hardwareAddresses[gateway]
if len(hardwareAddress) > 0 {
environment = append(environment, "gateway_mac:"+hardwareAddress.String())
}
}
}
}
var environmentHash uint64
if len(environment) > 0 {
digest := fnv.New64a()
for _, entry := range environment {
digest.Write([]byte(entry))
digest.Write([]byte{0})
}
environmentHash = digest.Sum64()
}
r.stateAccess.Lock()
changed := environmentHash != r.networkEnvironment
r.networkEnvironment = environmentHash
r.stateAccess.Unlock()
if changed {
if len(environment) > 0 {
r.logger.Info("updated network environment: ", strings.Join(environment, ", "))
} else {
r.logger.Info("updated network environment: empty")
}
}
}
+82
View File
@@ -0,0 +1,82 @@
package route
import (
"net"
"net/netip"
"slices"
"golang.org/x/net/route"
"golang.org/x/sys/unix"
)
func systemGateways(interfaceIndex int) []netip.Addr {
rib, err := route.FetchRIB(unix.AF_UNSPEC, route.RIBTypeRoute, 0)
if err != nil {
return nil
}
messages, err := route.ParseRIB(route.RIBTypeRoute, rib)
if err != nil {
return nil
}
var gateways []netip.Addr
for _, message := range messages {
routeMessage, isRouteMessage := message.(*route.RouteMessage)
if !isRouteMessage || routeMessage.Index != interfaceIndex || routeMessage.Flags&unix.RTF_GATEWAY == 0 {
continue
}
destination := routeAddressAt(routeMessage.Addrs, unix.RTAX_DST)
if !destination.IsValid() || !destination.IsUnspecified() {
continue
}
gateway := routeAddressAt(routeMessage.Addrs, unix.RTAX_GATEWAY)
if gateway.IsValid() {
gateways = append(gateways, gateway)
}
}
return gateways
}
func systemNeighborHardwareAddresses(interfaceIndex int, addresses []netip.Addr) map[netip.Addr]net.HardwareAddr {
rib, err := route.FetchRIB(unix.AF_UNSPEC, route.RIBType(unix.NET_RT_FLAGS), unix.RTF_LLINFO)
if err != nil {
return nil
}
messages, err := route.ParseRIB(route.RIBTypeRoute, rib)
if err != nil {
return nil
}
hardwareAddresses := make(map[netip.Addr]net.HardwareAddr)
for _, message := range messages {
routeMessage, isRouteMessage := message.(*route.RouteMessage)
if !isRouteMessage || routeMessage.Index != interfaceIndex {
continue
}
destination := routeAddressAt(routeMessage.Addrs, unix.RTAX_DST)
if !slices.Contains(addresses, destination) {
continue
}
if len(routeMessage.Addrs) <= unix.RTAX_GATEWAY {
continue
}
linkAddress, isLinkAddress := routeMessage.Addrs[unix.RTAX_GATEWAY].(*route.LinkAddr)
if !isLinkAddress || len(linkAddress.Addr) == 0 {
continue
}
hardwareAddresses[destination] = net.HardwareAddr(linkAddress.Addr)
}
return hardwareAddresses
}
func routeAddressAt(addresses []route.Addr, index int) netip.Addr {
if len(addresses) <= index {
return netip.Addr{}
}
switch address := addresses[index].(type) {
case *route.Inet4Addr:
return netip.AddrFrom4(address.IP)
case *route.Inet6Addr:
return netip.AddrFrom16(address.IP)
default:
return netip.Addr{}
}
}
+59
View File
@@ -0,0 +1,59 @@
package route
import (
"net"
"net/netip"
"slices"
"github.com/sagernet/netlink"
)
func systemGateways(interfaceIndex int) []netip.Addr {
routes, err := netlink.RouteListFiltered(netlink.FAMILY_ALL, &netlink.Route{LinkIndex: interfaceIndex}, netlink.RT_FILTER_OIF)
if err != nil {
return nil
}
var gateways []netip.Addr
for _, currentRoute := range routes {
if currentRoute.Gw == nil {
continue
}
if currentRoute.Dst != nil {
ones, _ := currentRoute.Dst.Mask.Size()
if ones != 0 {
continue
}
}
gateway, valid := netip.AddrFromSlice(currentRoute.Gw)
if valid {
gateways = append(gateways, gateway.Unmap())
}
}
return gateways
}
func systemNeighborHardwareAddresses(interfaceIndex int, addresses []netip.Addr) map[netip.Addr]net.HardwareAddr {
neighbors, err := netlink.NeighList(interfaceIndex, netlink.FAMILY_ALL)
if err != nil {
return nil
}
hardwareAddresses := make(map[netip.Addr]net.HardwareAddr)
for _, neighbor := range neighbors {
if neighbor.State&(netlink.NUD_INCOMPLETE|netlink.NUD_FAILED) != 0 {
continue
}
if len(neighbor.HardwareAddr) == 0 {
continue
}
neighborAddress, valid := netip.AddrFromSlice(neighbor.IP)
if !valid {
continue
}
neighborAddress = neighborAddress.Unmap()
if !slices.Contains(addresses, neighborAddress) {
continue
}
hardwareAddresses[neighborAddress] = neighbor.HardwareAddr
}
return hardwareAddresses
}
+16
View File
@@ -0,0 +1,16 @@
//go:build !darwin && !linux && !windows
package route
import (
"net"
"net/netip"
)
func systemGateways(interfaceIndex int) []netip.Addr {
return nil
}
func systemNeighborHardwareAddresses(interfaceIndex int, addresses []netip.Addr) map[netip.Addr]net.HardwareAddr {
return nil
}
+111
View File
@@ -0,0 +1,111 @@
package route
import (
"net"
"net/netip"
"slices"
"syscall"
"unsafe"
"golang.org/x/sys/windows"
)
func systemGateways(interfaceIndex int) []netip.Addr {
bufferSize := uint32(15000)
var buffer []byte
for {
buffer = make([]byte, bufferSize)
const flags = windows.GAA_FLAG_INCLUDE_GATEWAYS |
windows.GAA_FLAG_SKIP_ANYCAST |
windows.GAA_FLAG_SKIP_MULTICAST |
windows.GAA_FLAG_SKIP_DNS_SERVER
err := windows.GetAdaptersAddresses(syscall.AF_UNSPEC, flags, 0, (*windows.IpAdapterAddresses)(unsafe.Pointer(&buffer[0])), &bufferSize)
if err == nil {
break
}
if err != windows.ERROR_BUFFER_OVERFLOW || bufferSize <= uint32(len(buffer)) {
return nil
}
}
var gateways []netip.Addr
for adapter := (*windows.IpAdapterAddresses)(unsafe.Pointer(&buffer[0])); adapter != nil; adapter = adapter.Next {
if int(adapter.IfIndex) != interfaceIndex && int(adapter.Ipv6IfIndex) != interfaceIndex {
continue
}
for gatewayAddress := adapter.FirstGatewayAddress; gatewayAddress != nil; gatewayAddress = gatewayAddress.Next {
gateway, valid := netip.AddrFromSlice(gatewayAddress.Address.IP())
if valid {
gateways = append(gateways, gateway.Unmap().WithZone(""))
}
}
}
return gateways
}
var (
modiphlpapi = windows.NewLazySystemDLL("iphlpapi.dll")
procGetIpNetTable2 = modiphlpapi.NewProc("GetIpNetTable2")
procFreeMibTable = modiphlpapi.NewProc("FreeMibTable")
)
const (
neighborStateUnreachable = 0
neighborStateIncomplete = 1
)
type mibIPNetRow2 struct {
Address windows.RawSockaddrInet6
InterfaceIndex uint32
InterfaceLUID uint64
PhysicalAddress [32]byte
PhysicalAddressLength uint32
State uint32
Flags uint8
_ [3]byte
ReachabilityTime uint32
}
type mibIPNetTable2 struct {
NumEntries uint32
_ [4]byte
Table [1]mibIPNetRow2
}
func systemNeighborHardwareAddresses(interfaceIndex int, addresses []netip.Addr) map[netip.Addr]net.HardwareAddr {
var table *mibIPNetTable2
result, _, _ := procGetIpNetTable2.Call(uintptr(syscall.AF_UNSPEC), uintptr(unsafe.Pointer(&table)))
if result != 0 || table == nil {
return nil
}
defer procFreeMibTable.Call(uintptr(unsafe.Pointer(table)))
rows := unsafe.Slice(&table.Table[0], table.NumEntries)
hardwareAddresses := make(map[netip.Addr]net.HardwareAddr)
for i := range rows {
row := &rows[i]
if int(row.InterfaceIndex) != interfaceIndex {
continue
}
if row.State == neighborStateUnreachable || row.State == neighborStateIncomplete {
continue
}
if row.PhysicalAddressLength == 0 || row.PhysicalAddressLength > uint32(len(row.PhysicalAddress)) {
continue
}
var rowAddress netip.Addr
switch row.Address.Family {
case windows.AF_INET:
rowAddress = netip.AddrFrom4((*windows.RawSockaddrInet4)(unsafe.Pointer(&row.Address)).Addr)
case windows.AF_INET6:
rowAddress = netip.AddrFrom16(row.Address.Addr)
default:
continue
}
if !slices.Contains(addresses, rowAddress) {
continue
}
hardwareAddress := make(net.HardwareAddr, row.PhysicalAddressLength)
copy(hardwareAddress, row.PhysicalAddress[:row.PhysicalAddressLength])
hardwareAddresses[rowAddress] = hardwareAddress
}
return hardwareAddresses
}