diff --git a/adapter/network.go b/adapter/network.go index 14fe46c8..846d4392 100644 --- a/adapter/network.go +++ b/adapter/network.go @@ -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 } diff --git a/dns/client.go b/dns/client.go index 7a239713..99cb0ce1 100644 --- a/dns/client.go +++ b/dns/client.go @@ -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() } diff --git a/experimental/libbox/platform.go b/experimental/libbox/platform.go index 9e5eed71..87040da9 100644 --- a/experimental/libbox/platform.go +++ b/experimental/libbox/platform.go @@ -106,6 +106,7 @@ type NetworkInterface struct { Type int32 DNSServer StringIterator + Gateway StringIterator Metered bool } diff --git a/experimental/libbox/service.go b/experimental/libbox/service.go index 52b0b01c..00ed91d8 100644 --- a/experimental/libbox/service.go +++ b/experimental/libbox/service.go @@ -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, }) diff --git a/route/network.go b/route/network.go index 4fbcf22e..2f4a49b9 100644 --- a/route/network.go +++ b/route/network.go @@ -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 diff --git a/route/network_environment.go b/route/network_environment.go new file mode 100644 index 00000000..68e5b969 --- /dev/null +++ b/route/network_environment.go @@ -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") + } + } +} diff --git a/route/network_environment_darwin.go b/route/network_environment_darwin.go new file mode 100644 index 00000000..92633878 --- /dev/null +++ b/route/network_environment_darwin.go @@ -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{} + } +} diff --git a/route/network_environment_linux.go b/route/network_environment_linux.go new file mode 100644 index 00000000..702b3d92 --- /dev/null +++ b/route/network_environment_linux.go @@ -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 +} diff --git a/route/network_environment_stub.go b/route/network_environment_stub.go new file mode 100644 index 00000000..5c5778a7 --- /dev/null +++ b/route/network_environment_stub.go @@ -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 +} diff --git a/route/network_environment_windows.go b/route/network_environment_windows.go new file mode 100644 index 00000000..6fc00cb4 --- /dev/null +++ b/route/network_environment_windows.go @@ -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 +}