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
}
+6 -2
View File
@@ -139,8 +139,12 @@ func (w *platformInterfaceWrapper) NetworkInterfaces() ([]adapter.NetworkInterfa
Addresses: common.Map(iteratorToArray[string](netInterface.Addresses), netip.MustParsePrefix),
Flags: linkFlags(uint32(netInterface.Flags)),
},
Type: C.InterfaceType(netInterface.Type),
DNSServers: iteratorToArray[string](netInterface.DNSServer),
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,
})
+40 -28
View File
@@ -34,29 +34,32 @@ import (
var _ adapter.NetworkManager = (*NetworkManager)(nil)
type NetworkManager struct {
ctx context.Context
logger logger.ContextLogger
router adapter.Router
interfaceFinder *control.DefaultInterfaceFinder
networkInterfaces common.TypedValue[[]adapter.NetworkInterface]
autoDetectInterface bool
defaultOptions adapter.NetworkOptions
autoRedirectOutputMark uint32
networkMonitor tun.NetworkUpdateMonitor
interfaceMonitor tun.DefaultInterfaceMonitor
packageManager tun.PackageManager
powerListener winpowrprof.EventListener
pauseManager pause.Manager
platformInterface adapter.PlatformInterface
connectionManager adapter.ConnectionManager
endpoint adapter.EndpointManager
inbound adapter.InboundManager
outbound adapter.OutboundManager
needWIFIState bool
wifiMonitor settings.WIFIMonitor
wifiState adapter.WIFIState
wifiStateMutex sync.RWMutex
started bool
ctx context.Context
logger logger.ContextLogger
router adapter.Router
interfaceFinder *control.DefaultInterfaceFinder
networkInterfaces common.TypedValue[[]adapter.NetworkInterface]
autoDetectInterface bool
defaultOptions adapter.NetworkOptions
autoRedirectOutputMark uint32
networkMonitor tun.NetworkUpdateMonitor
interfaceMonitor tun.DefaultInterfaceMonitor
packageManager tun.PackageManager
powerListener winpowrprof.EventListener
pauseManager pause.Manager
platformInterface adapter.PlatformInterface
connectionManager adapter.ConnectionManager
endpoint adapter.EndpointManager
inbound adapter.InboundManager
outbound adapter.OutboundManager
needWIFIState bool
wifiMonitor settings.WIFIMonitor
wifiState adapter.WIFIState
networkEnvironment uint64
stateAccess sync.RWMutex
environmentUpdateAccess sync.Mutex
environmentUpdateTimer *time.Timer
started bool
}
func NewNetworkManager(ctx context.Context, logger logger.ContextLogger, options option.RouteOptions, dnsOptions option.DNSOptions) (*NetworkManager, error) {
@@ -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
}