Merge tag 'v1.14.0'

This commit is contained in:
Shtorm
2026-09-04 01:08:30 +03:00
987 changed files with 141481 additions and 11302 deletions
+12
View File
@@ -15,6 +15,7 @@ import (
"github.com/sagernet/sing-box/common/dialer"
"github.com/sagernet/sing-box/common/sniff"
tf "github.com/sagernet/sing-box/common/tlsfragment"
"github.com/sagernet/sing-box/common/tlsspoof"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
@@ -129,6 +130,17 @@ func (m *ConnectionManager) NewConnection(ctx context.Context, this N.Dialer, co
if metadata.TLSFragment || metadata.TLSRecordFragment {
remoteConn = tf.NewConn(remoteConn, ctx, metadata.TLSFragment, metadata.TLSRecordFragment, metadata.TLSFragmentFallbackDelay)
}
if metadata.TLSSpoof != "" {
spoofConn, spoofErr := tlsspoof.NewConn(remoteConn, metadata.TLSSpoofMethod, metadata.TLSSpoof)
if spoofErr != nil {
spoofErr = E.Cause(spoofErr, "tls_spoof setup")
remoteConn.Close()
N.CloseOnHandshakeFailure(conn, onClose, spoofErr)
m.logger.ErrorContext(ctx, spoofErr)
return
}
remoteConn = spoofConn
}
serverFirst := sniff.Skip(&metadata)
var done atomic.Bool
if m.kickWriteHandshake(ctx, conn, remoteConn, serverFirst, false, &done, onClose) {
+22 -59
View File
@@ -10,17 +10,15 @@ import (
"github.com/sagernet/sing-box/dns"
dnsOutbound "github.com/sagernet/sing-box/protocol/dns"
R "github.com/sagernet/sing-box/route/rule"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/udpnat2"
mDNS "github.com/miekg/dns"
)
func (r *Router) hijackDNSStream(ctx context.Context, conn net.Conn, metadata adapter.InboundContext) error {
r.searchProcessInfo(ctx, &metadata)
metadata.Destination = M.Socksaddr{}
for {
conn.SetReadDeadline(time.Now().Add(C.DNSTimeout))
@@ -36,24 +34,7 @@ func (r *Router) hijackDNSStream(ctx context.Context, conn net.Conn, metadata ad
}
func (r *Router) hijackDNSPacket(ctx context.Context, conn N.PacketConn, packetBuffers []*N.PacketBuffer, metadata adapter.InboundContext, onClose N.CloseHandlerFunc) error {
if natConn, isNatConn := conn.(udpnat.Conn); isNatConn {
metadata.Destination = M.Socksaddr{}
for _, packet := range packetBuffers {
buffer := packet.Buffer
destination := packet.Destination
N.PutPacketBuffer(packet)
go ExchangeDNSPacket(ctx, r.dns, r.logger, natConn, buffer, metadata, destination)
}
natConn.SetHandler(&dnsHijacker{
router: r.dns,
logger: r.logger,
conn: conn,
ctx: ctx,
metadata: metadata,
onClose: onClose,
})
return nil
}
r.searchProcessInfo(ctx, &metadata)
err := dnsOutbound.NewDNSPacketConnection(ctx, r.dns, conn, packetBuffers, metadata)
N.CloseOnHandshakeFailure(conn, onClose, err)
if err != nil && !E.IsClosedOrCanceled(err) {
@@ -62,48 +43,30 @@ func (r *Router) hijackDNSPacket(ctx context.Context, conn N.PacketConn, packetB
return nil
}
func ExchangeDNSPacket(ctx context.Context, router adapter.DNSRouter, logger logger.ContextLogger, conn N.PacketConn, buffer *buf.Buffer, metadata adapter.InboundContext, destination M.Socksaddr) {
err := exchangeDNSPacket(ctx, router, conn, buffer, metadata, destination)
if err != nil && !R.IsRejected(err) && !E.IsClosedOrCanceled(err) {
logger.ErrorContext(ctx, E.Cause(err, "process DNS packet"))
}
}
func exchangeDNSPacket(ctx context.Context, router adapter.DNSRouter, conn N.PacketConn, buffer *buf.Buffer, metadata adapter.InboundContext, destination M.Socksaddr) error {
func (r *Router) HijackDNSPacket(ctx context.Context, payload []byte, writer N.PacketWriter, metadata adapter.InboundContext) {
var message mDNS.Msg
err := message.Unpack(buffer.Bytes())
buffer.Release()
err := message.Unpack(payload)
if err != nil {
return E.Cause(err, "unpack request")
r.logger.ErrorContext(ctx, E.Cause(err, "process DNS packet: unpack request"))
return
}
response, err := router.Exchange(adapter.WithContext(ctx, &metadata), &message, adapter.DNSQueryOptions{})
r.searchProcessInfo(ctx, &metadata)
destination := metadata.Destination
metadata.Destination = M.Socksaddr{}
r.dns.ExchangeAsync(adapter.WithContext(ctx, &metadata), &message, adapter.DNSQueryOptions{}, func(response *mDNS.Msg, exchangeErr error) {
if exchangeErr == nil {
exchangeErr = r.writeDNSPacketResponse(&message, response, writer, destination)
}
if exchangeErr != nil && !R.IsRejected(exchangeErr) && !E.IsClosedOrCanceled(exchangeErr) {
r.logger.ErrorContext(ctx, E.Cause(exchangeErr, "process DNS packet"))
}
})
}
func (r *Router) writeDNSPacketResponse(message *mDNS.Msg, response *mDNS.Msg, writer N.PacketWriter, destination M.Socksaddr) error {
responseBuffer, err := dns.TruncateDNSMessage(message, response, N.CalculateFrontHeadroom(writer), N.CalculateRearHeadroom(writer))
if err != nil {
return err
}
responseBuffer, err := dns.TruncateDNSMessage(&message, response, 1024)
if err != nil {
return err
}
err = conn.WritePacket(responseBuffer, destination)
return err
}
type dnsHijacker struct {
router adapter.DNSRouter
logger logger.ContextLogger
conn N.PacketConn
ctx context.Context
metadata adapter.InboundContext
onClose N.CloseHandlerFunc
}
func (h *dnsHijacker) NewPacketEx(buffer *buf.Buffer, destination M.Socksaddr) {
go ExchangeDNSPacket(h.ctx, h.router, h.logger, h.conn, buffer, h.metadata, destination)
}
func (h *dnsHijacker) Close() error {
if h.onClose != nil {
h.onClose(nil)
}
return nil
return writer.WritePacket(responseBuffer, destination)
}
+221
View File
@@ -0,0 +1,221 @@
package route
import (
"net"
"net/netip"
"os"
"sync"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
type fakeIPNATPacketConn struct {
N.NetPacketConn
origin M.Socksaddr
destination M.Socksaddr
directAccess sync.RWMutex
directDestinations map[netip.Addr]struct{}
}
func newFakeIPNATPacketConn(conn N.NetPacketConn, origin M.Socksaddr, destination M.Socksaddr) *fakeIPNATPacketConn {
return &fakeIPNATPacketConn{
NetPacketConn: conn,
origin: origin,
destination: destination,
directDestinations: make(map[netip.Addr]struct{}),
}
}
func (c *fakeIPNATPacketConn) rewriteReadDestination(destination M.Socksaddr) M.Socksaddr {
if destination.Addr == c.origin.Addr {
return M.Socksaddr{
Addr: c.destination.Addr,
Fqdn: c.destination.Fqdn,
Port: destination.Port,
}
}
if destination.Addr.IsValid() {
c.directAccess.RLock()
_, recorded := c.directDestinations[destination.Addr]
c.directAccess.RUnlock()
if !recorded {
c.directAccess.Lock()
c.directDestinations[destination.Addr] = struct{}{}
c.directAccess.Unlock()
}
}
return destination
}
func (c *fakeIPNATPacketConn) rewriteWriteDestination(destination M.Socksaddr) M.Socksaddr {
if destination.Addr.IsValid() {
c.directAccess.RLock()
_, direct := c.directDestinations[destination.Addr]
c.directAccess.RUnlock()
if direct {
return destination
}
}
return M.Socksaddr{
Addr: c.origin.Addr,
Port: destination.Port,
}
}
func (c *fakeIPNATPacketConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
n, addr, err = c.NetPacketConn.ReadFrom(p)
if err != nil {
return
}
addr = c.rewriteReadDestination(M.SocksaddrFromNet(addr)).UDPAddr()
return
}
func (c *fakeIPNATPacketConn) WriteTo(p []byte, addr net.Addr) (n int, err error) {
return c.NetPacketConn.WriteTo(p, c.rewriteWriteDestination(M.SocksaddrFromNet(addr)).UDPAddr())
}
func (c *fakeIPNATPacketConn) ReadPacket(buffer *buf.Buffer) (destination M.Socksaddr, err error) {
destination, err = c.NetPacketConn.ReadPacket(buffer)
if err != nil {
return
}
destination = c.rewriteReadDestination(destination)
return
}
func (c *fakeIPNATPacketConn) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error {
return c.NetPacketConn.WritePacket(buffer, c.rewriteWriteDestination(destination))
}
func (c *fakeIPNATPacketConn) CreateReadWaiter() (N.PacketReadWaiter, bool) {
waiter, created := bufio.CreatePacketReadWaiter(c.NetPacketConn)
if !created {
return nil, false
}
return &waitFakeIPNATPacketConn{c, waiter}, true
}
func (c *fakeIPNATPacketConn) CreatePacketBatchReadWaiter() (N.PacketBatchReadWaiter, bool) {
waiter, created := bufio.CreatePacketBatchReadWaiter(c.NetPacketConn)
if !created {
return nil, false
}
return &batchWaitFakeIPNATPacketConn{c, waiter}, true
}
func (c *fakeIPNATPacketConn) CreateConnectedPacketBatchReadWaiter() (N.ConnectedPacketBatchReadWaiter, bool) {
waiter, created := bufio.CreateConnectedPacketBatchReadWaiter(c.NetPacketConn)
if !created {
return nil, false
}
return &connectedBatchWaitFakeIPNATPacketConn{c, waiter}, true
}
func (c *fakeIPNATPacketConn) CreatePacketBatchWriter() (N.PacketBatchWriter, bool) {
writer, created := bufio.CreatePacketBatchWriter(c.NetPacketConn)
if !created {
return nil, false
}
return &fakeIPNATPacketBatchWriter{c, writer}, true
}
func (c *fakeIPNATPacketConn) CreateConnectedPacketBatchWriter() (N.ConnectedPacketBatchWriter, bool) {
return bufio.CreateConnectedPacketBatchWriter(c.NetPacketConn)
}
func (c *fakeIPNATPacketConn) UpdateDestination(destinationAddress netip.Addr) {
c.destination = M.SocksaddrFrom(destinationAddress, c.destination.Port)
}
func (c *fakeIPNATPacketConn) RemoteAddr() net.Addr {
return c.destination.UDPAddr()
}
func (c *fakeIPNATPacketConn) Upstream() any {
return c.NetPacketConn
}
type waitFakeIPNATPacketConn struct {
*fakeIPNATPacketConn
readWaiter N.PacketReadWaiter
}
func (c *waitFakeIPNATPacketConn) InitializeReadWaiter(options N.ReadWaitOptions) (needCopy bool) {
return c.readWaiter.InitializeReadWaiter(options)
}
func (c *waitFakeIPNATPacketConn) WaitReadPacket() (buffer *buf.Buffer, destination M.Socksaddr, err error) {
buffer, destination, err = c.readWaiter.WaitReadPacket()
if err != nil {
return
}
destination = c.rewriteReadDestination(destination)
return
}
type batchWaitFakeIPNATPacketConn struct {
*fakeIPNATPacketConn
readWaiter N.PacketBatchReadWaiter
}
func (c *batchWaitFakeIPNATPacketConn) InitializeReadWaiter(options N.ReadWaitOptions) (needCopy bool) {
return c.readWaiter.InitializeReadWaiter(options)
}
func (c *batchWaitFakeIPNATPacketConn) WaitReadPackets() (buffers []*buf.Buffer, destinations []M.Socksaddr, err error) {
buffers, destinations, err = c.readWaiter.WaitReadPackets()
if err != nil {
return
}
for index, destination := range destinations {
destinations[index] = c.rewriteReadDestination(destination)
}
return
}
type connectedBatchWaitFakeIPNATPacketConn struct {
*fakeIPNATPacketConn
readWaiter N.ConnectedPacketBatchReadWaiter
}
func (c *connectedBatchWaitFakeIPNATPacketConn) InitializeReadWaiter(options N.ReadWaitOptions) (needCopy bool) {
return c.readWaiter.InitializeReadWaiter(options)
}
func (c *connectedBatchWaitFakeIPNATPacketConn) WaitReadConnectedPackets() (buffers []*buf.Buffer, destination M.Socksaddr, err error) {
buffers, destination, err = c.readWaiter.WaitReadConnectedPackets()
if err != nil {
return
}
destination = c.rewriteReadDestination(destination)
return
}
type fakeIPNATPacketBatchWriter struct {
*fakeIPNATPacketConn
writer N.PacketBatchWriter
}
func (w *fakeIPNATPacketBatchWriter) WritePacketBatch(buffers []*buf.Buffer, destinations []M.Socksaddr) error {
if len(buffers) == 0 || len(buffers) != len(destinations) {
buf.ReleaseMulti(buffers)
return os.ErrInvalid
}
for index, destination := range destinations {
destinations[index] = w.rewriteWriteDestination(destination)
}
return w.writer.WritePacketBatch(buffers, destinations)
}
var (
_ bufio.NATPacketConn = (*fakeIPNATPacketConn)(nil)
_ N.PacketReadWaitCreator = (*fakeIPNATPacketConn)(nil)
_ N.PacketBatchReadWaitCreator = (*fakeIPNATPacketConn)(nil)
_ N.ConnectedPacketBatchReadWaitCreator = (*fakeIPNATPacketConn)(nil)
_ N.PacketBatchWriteCreator = (*fakeIPNATPacketConn)(nil)
_ N.ConnectedPacketBatchWriteCreator = (*fakeIPNATPacketConn)(nil)
)
+101
View File
@@ -0,0 +1,101 @@
package route
import (
"context"
"sync/atomic"
"time"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing/common/byteformats"
N "github.com/sagernet/sing/common/network"
)
var (
_ tun.FlowTracker = (*flowLogger)(nil)
_ tun.FlowTracker = multiFlowTracker(nil)
)
type flowLogger struct {
ctx context.Context
logger log.ContextLogger
network string
source string
destination string
outbound adapter.Outbound
createdAt time.Time
upload atomic.Int64
download atomic.Int64
}
func newFlowLogger(ctx context.Context, logger log.ContextLogger, metadata adapter.InboundContext, outbound adapter.Outbound) *flowLogger {
var source, destination string
if metadata.Network == N.NetworkICMP {
source = metadata.Source.AddrString()
destination = metadata.Destination.AddrString()
} else {
source = metadata.Source.String()
destination = metadata.Destination.String()
}
return &flowLogger{
ctx: ctx,
logger: logger,
network: metadata.Network,
source: source,
destination: destination,
outbound: outbound,
}
}
func (l *flowLogger) AttachFlow(handle tun.FlowHandle) {
l.createdAt = time.Now()
}
func (l *flowLogger) CountForward(n int) {
l.upload.Add(int64(n))
}
func (l *flowLogger) CountReverse(n int) {
l.download.Add(int64(n))
}
func (l *flowLogger) FlowEstablished() {
}
func (l *flowLogger) CloseFlow(reason tun.FlowCloseReason) {
l.logger.DebugContext(l.ctx, "flow closed: ", reason,
", upload ", byteformats.FormatBytes(uint64(l.upload.Load())), ", download ", byteformats.FormatBytes(uint64(l.download.Load())))
}
type multiFlowTracker []tun.FlowTracker
func (t multiFlowTracker) AttachFlow(handle tun.FlowHandle) {
for _, tracker := range t {
tracker.AttachFlow(handle)
}
}
func (t multiFlowTracker) CountForward(n int) {
for _, tracker := range t {
tracker.CountForward(n)
}
}
func (t multiFlowTracker) CountReverse(n int) {
for _, tracker := range t {
tracker.CountReverse(n)
}
}
func (t multiFlowTracker) FlowEstablished() {
for _, tracker := range t {
tracker.FlowEstablished()
}
}
func (t multiFlowTracker) CloseFlow(reason tun.FlowCloseReason) {
for _, tracker := range t {
tracker.CloseFlow(reason)
}
}
+245
View File
@@ -0,0 +1,245 @@
//go:build darwin
package route
import (
"net"
"net/netip"
"os"
"sync"
"time"
"github.com/sagernet/fswatch"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
"golang.org/x/net/route"
"golang.org/x/sys/unix"
)
var defaultLeaseFiles = []string{
"/var/db/dhcpd_leases",
"/tmp/dhcp.leases",
}
type neighborResolver struct {
logger logger.ContextLogger
leaseFiles []string
access sync.RWMutex
neighborIPToMAC map[netip.Addr]net.HardwareAddr
leaseIPToMAC map[netip.Addr]net.HardwareAddr
ipToHostname map[netip.Addr]string
macToHostname map[string]string
watcher *fswatch.Watcher
done chan struct{}
}
func newNeighborResolver(resolverLogger logger.ContextLogger, leaseFiles []string) (adapter.NeighborResolver, error) {
if len(leaseFiles) == 0 {
for _, path := range defaultLeaseFiles {
info, err := os.Stat(path)
if err == nil && info.Size() > 0 {
leaseFiles = append(leaseFiles, path)
}
}
}
return &neighborResolver{
logger: resolverLogger,
leaseFiles: leaseFiles,
neighborIPToMAC: make(map[netip.Addr]net.HardwareAddr),
leaseIPToMAC: make(map[netip.Addr]net.HardwareAddr),
ipToHostname: make(map[netip.Addr]string),
macToHostname: make(map[string]string),
done: make(chan struct{}),
}, nil
}
func (r *neighborResolver) Start() error {
err := r.loadNeighborTable()
if err != nil {
r.logger.Warn(E.Cause(err, "load neighbor table"))
}
r.doReloadLeaseFiles()
go r.subscribeNeighborUpdates()
if len(r.leaseFiles) > 0 {
watcher, err := fswatch.NewWatcher(fswatch.Options{
Path: r.leaseFiles,
Logger: r.logger,
Callback: func(_ string) {
r.doReloadLeaseFiles()
},
})
if err != nil {
r.logger.Warn(E.Cause(err, "create lease file watcher"))
} else {
r.watcher = watcher
err = watcher.Start()
if err != nil {
r.logger.Warn(E.Cause(err, "start lease file watcher"))
}
}
}
return nil
}
func (r *neighborResolver) Close() error {
close(r.done)
if r.watcher != nil {
return r.watcher.Close()
}
return nil
}
func (r *neighborResolver) LookupMAC(address netip.Addr) (net.HardwareAddr, bool) {
r.access.RLock()
defer r.access.RUnlock()
mac, found := r.neighborIPToMAC[address]
if found {
return mac, true
}
mac, found = r.leaseIPToMAC[address]
if found {
return mac, true
}
mac, found = extractMACFromEUI64(address)
if found {
return mac, true
}
return nil, false
}
func (r *neighborResolver) LookupAddresses(hostname string) []netip.Addr {
r.access.RLock()
defer r.access.RUnlock()
return lookupAddressesByHostname(hostname, r.ipToHostname, r.macToHostname, r.neighborIPToMAC, r.leaseIPToMAC)
}
func (r *neighborResolver) LookupHostname(address netip.Addr) (string, bool) {
r.access.RLock()
defer r.access.RUnlock()
hostname, found := r.ipToHostname[address]
if found {
return hostname, true
}
mac, macFound := r.neighborIPToMAC[address]
if !macFound {
mac, macFound = r.leaseIPToMAC[address]
}
if !macFound {
mac, macFound = extractMACFromEUI64(address)
}
if macFound {
hostname, found = r.macToHostname[mac.String()]
if found {
return hostname, true
}
}
return "", false
}
func (r *neighborResolver) loadNeighborTable() error {
entries, err := ReadNeighborEntries()
if err != nil {
return err
}
r.access.Lock()
defer r.access.Unlock()
for _, entry := range entries {
r.neighborIPToMAC[entry.Address] = entry.MACAddress
}
return nil
}
func (r *neighborResolver) subscribeNeighborUpdates() {
routeSocket, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, 0)
if err != nil {
r.logger.Warn(E.Cause(err, "subscribe neighbor updates"))
return
}
err = unix.SetNonblock(routeSocket, true)
if err != nil {
unix.Close(routeSocket)
r.logger.Warn(E.Cause(err, "set route socket nonblock"))
return
}
routeSocketFile := os.NewFile(uintptr(routeSocket), "route")
defer routeSocketFile.Close()
buffer := buf.NewPacket()
defer buffer.Release()
for {
select {
case <-r.done:
return
default:
}
err = setReadDeadline(routeSocketFile, 3*time.Second)
if err != nil {
r.logger.Warn(E.Cause(err, "set route socket read deadline"))
return
}
n, err := routeSocketFile.Read(buffer.FreeBytes())
if err != nil {
if nerr, ok := err.(net.Error); ok && nerr.Timeout() {
continue
}
select {
case <-r.done:
return
default:
}
r.logger.Warn(E.Cause(err, "receive neighbor update"))
continue
}
messages, err := route.ParseRIB(route.RIBTypeRoute, buffer.FreeBytes()[:n])
if err != nil {
continue
}
for _, message := range messages {
routeMessage, isRouteMessage := message.(*route.RouteMessage)
if !isRouteMessage {
continue
}
if routeMessage.Flags&unix.RTF_LLINFO == 0 {
continue
}
address, mac, isDelete, ok := ParseRouteNeighborMessage(routeMessage)
if !ok {
continue
}
r.access.Lock()
if isDelete {
delete(r.neighborIPToMAC, address)
} else {
r.neighborIPToMAC[address] = mac
}
r.access.Unlock()
}
}
}
func (r *neighborResolver) doReloadLeaseFiles() {
leaseIPToMAC, ipToHostname, macToHostname := ReloadLeaseFiles(r.leaseFiles)
r.access.Lock()
r.leaseIPToMAC = leaseIPToMAC
r.ipToHostname = ipToHostname
r.macToHostname = macToHostname
r.access.Unlock()
}
func setReadDeadline(file *os.File, timeout time.Duration) error {
rawConn, err := file.SyscallConn()
if err != nil {
return err
}
var controlErr error
err = rawConn.Control(func(fd uintptr) {
tv := unix.NsecToTimeval(int64(timeout))
controlErr = unix.SetsockoptTimeval(int(fd), unix.SOL_SOCKET, unix.SO_RCVTIMEO, &tv)
})
if err != nil {
return err
}
return controlErr
}
+56
View File
@@ -0,0 +1,56 @@
package route
import (
"net"
"net/netip"
"strings"
"github.com/sagernet/sing-box/dns"
)
func lookupAddressesByHostname(
hostname string,
ipToHostname map[netip.Addr]string,
macToHostname map[string]string,
ipToMACTables ...map[netip.Addr]net.HardwareAddr,
) []netip.Addr {
hostname = dns.FqdnToDomain(hostname)
if hostname == "" {
return nil
}
resultSet := make(map[netip.Addr]struct{})
var result []netip.Addr
addAddress := func(address netip.Addr) {
if isScopedIPv6Address(address) {
return
}
if _, exists := resultSet[address]; exists {
return
}
resultSet[address] = struct{}{}
result = append(result, address)
}
for address, entryHostname := range ipToHostname {
if strings.EqualFold(entryHostname, hostname) {
addAddress(address)
}
}
for mac, entryHostname := range macToHostname {
if !strings.EqualFold(entryHostname, hostname) {
continue
}
for _, table := range ipToMACTables {
for address, entryMAC := range table {
if entryMAC.String() == mac {
addAddress(address)
}
}
}
}
return result
}
func isScopedIPv6Address(address netip.Addr) bool {
// DNS AAAA records cannot carry an interface zone.
return address.Is6() && (address.IsLinkLocalUnicast() || address.Zone() != "")
}
+386
View File
@@ -0,0 +1,386 @@
package route
import (
"bufio"
"encoding/hex"
"net"
"net/netip"
"os"
"strconv"
"strings"
"time"
)
func parseLeaseFile(path string, ipToMAC map[netip.Addr]net.HardwareAddr, ipToHostname map[netip.Addr]string, macToHostname map[string]string) {
file, err := os.Open(path)
if err != nil {
return
}
defer file.Close()
if strings.HasSuffix(path, "dhcpd_leases") {
parseBootpdLeases(file, ipToMAC, ipToHostname, macToHostname)
return
}
if strings.HasSuffix(path, "kea-leases4.csv") {
parseKeaCSV4(file, ipToMAC, ipToHostname, macToHostname)
return
}
if strings.HasSuffix(path, "kea-leases6.csv") {
parseKeaCSV6(file, ipToMAC, ipToHostname, macToHostname)
return
}
if strings.HasSuffix(path, "dhcpd.leases") {
parseISCDhcpd(file, ipToMAC, ipToHostname, macToHostname)
return
}
parseDnsmasqOdhcpd(file, ipToMAC, ipToHostname, macToHostname)
}
func ReloadLeaseFiles(leaseFiles []string) (leaseIPToMAC map[netip.Addr]net.HardwareAddr, ipToHostname map[netip.Addr]string, macToHostname map[string]string) {
leaseIPToMAC = make(map[netip.Addr]net.HardwareAddr)
ipToHostname = make(map[netip.Addr]string)
macToHostname = make(map[string]string)
for _, path := range leaseFiles {
parseLeaseFile(path, leaseIPToMAC, ipToHostname, macToHostname)
}
return
}
func parseDnsmasqOdhcpd(file *os.File, ipToMAC map[netip.Addr]net.HardwareAddr, ipToHostname map[netip.Addr]string, macToHostname map[string]string) {
now := time.Now().Unix()
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := scanner.Text()
if strings.HasPrefix(line, "duid ") {
continue
}
if strings.HasPrefix(line, "# ") {
parseOdhcpdLine(line[2:], ipToMAC, ipToHostname, macToHostname)
continue
}
fields := strings.Fields(line)
if len(fields) < 4 {
continue
}
expiry, err := strconv.ParseInt(fields[0], 10, 64)
if err != nil {
continue
}
if expiry != 0 && expiry < now {
continue
}
if strings.Contains(fields[1], ":") {
mac, macErr := net.ParseMAC(fields[1])
if macErr != nil {
continue
}
address, addrOK := netip.AddrFromSlice(net.ParseIP(fields[2]))
if !addrOK {
continue
}
address = address.Unmap()
ipToMAC[address] = mac
hostname := fields[3]
if hostname != "*" {
ipToHostname[address] = hostname
macToHostname[mac.String()] = hostname
}
} else {
var mac net.HardwareAddr
if len(fields) >= 5 {
duid, duidErr := parseDUID(fields[4])
if duidErr == nil {
mac, _ = extractMACFromDUID(duid)
}
}
address, addrOK := netip.AddrFromSlice(net.ParseIP(fields[2]))
if !addrOK {
continue
}
address = address.Unmap()
if mac != nil {
ipToMAC[address] = mac
}
hostname := fields[3]
if hostname != "*" {
ipToHostname[address] = hostname
if mac != nil {
macToHostname[mac.String()] = hostname
}
}
}
}
}
func parseOdhcpdLine(line string, ipToMAC map[netip.Addr]net.HardwareAddr, ipToHostname map[netip.Addr]string, macToHostname map[string]string) {
fields := strings.Fields(line)
if len(fields) < 5 {
return
}
validTime, err := strconv.ParseInt(fields[4], 10, 64)
if err != nil {
return
}
if validTime == 0 {
return
}
if validTime > 0 && validTime < time.Now().Unix() {
return
}
hostname := fields[3]
if hostname == "-" || strings.HasPrefix(hostname, `broken\x20`) {
hostname = ""
}
if len(fields) >= 8 && fields[2] == "ipv4" {
mac, macErr := net.ParseMAC(fields[1])
if macErr != nil {
return
}
addressField := fields[7]
slashIndex := strings.IndexByte(addressField, '/')
if slashIndex >= 0 {
addressField = addressField[:slashIndex]
}
address, addrOK := netip.AddrFromSlice(net.ParseIP(addressField))
if !addrOK {
return
}
address = address.Unmap()
ipToMAC[address] = mac
if hostname != "" {
ipToHostname[address] = hostname
macToHostname[mac.String()] = hostname
}
return
}
var mac net.HardwareAddr
duidHex := fields[1]
duidBytes, hexErr := hex.DecodeString(duidHex)
if hexErr == nil {
mac, _ = extractMACFromDUID(duidBytes)
}
for i := 7; i < len(fields); i++ {
addressField := fields[i]
slashIndex := strings.IndexByte(addressField, '/')
if slashIndex >= 0 {
addressField = addressField[:slashIndex]
}
address, addrOK := netip.AddrFromSlice(net.ParseIP(addressField))
if !addrOK {
continue
}
address = address.Unmap()
if mac != nil {
ipToMAC[address] = mac
}
if hostname != "" {
ipToHostname[address] = hostname
if mac != nil {
macToHostname[mac.String()] = hostname
}
}
}
}
func parseISCDhcpd(file *os.File, ipToMAC map[netip.Addr]net.HardwareAddr, ipToHostname map[netip.Addr]string, macToHostname map[string]string) {
scanner := bufio.NewScanner(file)
var currentIP netip.Addr
var currentMAC net.HardwareAddr
var currentHostname string
var currentActive bool
var inLease bool
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if strings.HasPrefix(line, "lease ") && strings.HasSuffix(line, "{") {
ipString := strings.TrimSuffix(strings.TrimPrefix(line, "lease "), " {")
parsed, addrOK := netip.AddrFromSlice(net.ParseIP(ipString))
if addrOK {
currentIP = parsed.Unmap()
inLease = true
currentMAC = nil
currentHostname = ""
currentActive = false
}
continue
}
if line == "}" && inLease {
if currentActive && currentMAC != nil {
ipToMAC[currentIP] = currentMAC
if currentHostname != "" {
ipToHostname[currentIP] = currentHostname
macToHostname[currentMAC.String()] = currentHostname
}
} else {
delete(ipToMAC, currentIP)
delete(ipToHostname, currentIP)
}
inLease = false
continue
}
if !inLease {
continue
}
if rest, ok := strings.CutPrefix(line, "hardware ethernet "); ok {
macString := strings.TrimSuffix(rest, ";")
parsed, macErr := net.ParseMAC(macString)
if macErr == nil {
currentMAC = parsed
}
} else if rest, ok := strings.CutPrefix(line, "client-hostname "); ok {
hostname := strings.TrimSuffix(rest, ";")
hostname = strings.Trim(hostname, "\"")
if hostname != "" {
currentHostname = hostname
}
} else if rest, ok := strings.CutPrefix(line, "binding state "); ok {
state := strings.TrimSuffix(rest, ";")
currentActive = state == "active"
}
}
}
func parseKeaCSV4(file *os.File, ipToMAC map[netip.Addr]net.HardwareAddr, ipToHostname map[netip.Addr]string, macToHostname map[string]string) {
scanner := bufio.NewScanner(file)
firstLine := true
for scanner.Scan() {
if firstLine {
firstLine = false
continue
}
fields := strings.Split(scanner.Text(), ",")
if len(fields) < 10 {
continue
}
if fields[9] != "0" {
continue
}
address, addrOK := netip.AddrFromSlice(net.ParseIP(fields[0]))
if !addrOK {
continue
}
address = address.Unmap()
mac, macErr := net.ParseMAC(fields[1])
if macErr != nil {
continue
}
ipToMAC[address] = mac
hostname := ""
if len(fields) > 8 {
hostname = fields[8]
}
if hostname != "" {
ipToHostname[address] = hostname
macToHostname[mac.String()] = hostname
}
}
}
func parseKeaCSV6(file *os.File, ipToMAC map[netip.Addr]net.HardwareAddr, ipToHostname map[netip.Addr]string, macToHostname map[string]string) {
scanner := bufio.NewScanner(file)
firstLine := true
for scanner.Scan() {
if firstLine {
firstLine = false
continue
}
fields := strings.Split(scanner.Text(), ",")
if len(fields) < 14 {
continue
}
if fields[13] != "0" {
continue
}
address, addrOK := netip.AddrFromSlice(net.ParseIP(fields[0]))
if !addrOK {
continue
}
address = address.Unmap()
var mac net.HardwareAddr
if fields[12] != "" {
mac, _ = net.ParseMAC(fields[12])
}
if mac == nil {
duid, duidErr := hex.DecodeString(strings.ReplaceAll(fields[1], ":", ""))
if duidErr == nil {
mac, _ = extractMACFromDUID(duid)
}
}
hostname := ""
if len(fields) > 11 {
hostname = fields[11]
}
if mac != nil {
ipToMAC[address] = mac
}
if hostname != "" {
ipToHostname[address] = hostname
if mac != nil {
macToHostname[mac.String()] = hostname
}
}
}
}
func parseBootpdLeases(file *os.File, ipToMAC map[netip.Addr]net.HardwareAddr, ipToHostname map[netip.Addr]string, macToHostname map[string]string) {
now := time.Now().Unix()
scanner := bufio.NewScanner(file)
var currentName string
var currentIP netip.Addr
var currentMAC net.HardwareAddr
var currentLease int64
var inBlock bool
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "{" {
inBlock = true
currentName = ""
currentIP = netip.Addr{}
currentMAC = nil
currentLease = 0
continue
}
if line == "}" && inBlock {
if currentMAC != nil && currentIP.IsValid() {
if currentLease == 0 || currentLease >= now {
ipToMAC[currentIP] = currentMAC
if currentName != "" {
ipToHostname[currentIP] = currentName
macToHostname[currentMAC.String()] = currentName
}
}
}
inBlock = false
continue
}
if !inBlock {
continue
}
key, value, found := strings.Cut(line, "=")
if !found {
continue
}
switch key {
case "name":
currentName = value
case "ip_address":
parsed, addrOK := netip.AddrFromSlice(net.ParseIP(value))
if addrOK {
currentIP = parsed.Unmap()
}
case "hw_address":
typeAndMAC, hasSep := strings.CutPrefix(value, "1,")
if hasSep {
mac, macErr := net.ParseMAC(typeAndMAC)
if macErr == nil {
currentMAC = mac
}
}
case "lease":
leaseHex := strings.TrimPrefix(value, "0x")
parsed, parseErr := strconv.ParseInt(leaseHex, 16, 64)
if parseErr == nil {
currentLease = parsed
}
}
}
}
+230
View File
@@ -0,0 +1,230 @@
//go:build linux
package route
import (
"net"
"net/netip"
"os"
"slices"
"sync"
"time"
"github.com/sagernet/fswatch"
"github.com/sagernet/sing-box/adapter"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
"github.com/jsimonetti/rtnetlink"
"github.com/mdlayher/netlink"
"golang.org/x/sys/unix"
)
var defaultLeaseFiles = []string{
"/tmp/dhcp.leases",
"/var/lib/dhcp/dhcpd.leases",
"/var/lib/dhcpd/dhcpd.leases",
"/var/lib/kea/kea-leases4.csv",
"/var/lib/kea/kea-leases6.csv",
}
type neighborResolver struct {
logger logger.ContextLogger
leaseFiles []string
access sync.RWMutex
neighborIPToMAC map[netip.Addr]net.HardwareAddr
leaseIPToMAC map[netip.Addr]net.HardwareAddr
ipToHostname map[netip.Addr]string
macToHostname map[string]string
watcher *fswatch.Watcher
done chan struct{}
}
func newNeighborResolver(resolverLogger logger.ContextLogger, leaseFiles []string) (adapter.NeighborResolver, error) {
if len(leaseFiles) == 0 {
for _, path := range defaultLeaseFiles {
info, err := os.Stat(path)
if err == nil && info.Size() > 0 {
leaseFiles = append(leaseFiles, path)
}
}
}
return &neighborResolver{
logger: resolverLogger,
leaseFiles: leaseFiles,
neighborIPToMAC: make(map[netip.Addr]net.HardwareAddr),
leaseIPToMAC: make(map[netip.Addr]net.HardwareAddr),
ipToHostname: make(map[netip.Addr]string),
macToHostname: make(map[string]string),
done: make(chan struct{}),
}, nil
}
func (r *neighborResolver) Start() error {
err := r.loadNeighborTable()
if err != nil {
r.logger.Warn(E.Cause(err, "load neighbor table"))
}
r.doReloadLeaseFiles()
go r.subscribeNeighborUpdates()
if len(r.leaseFiles) > 0 {
watcher, err := fswatch.NewWatcher(fswatch.Options{
Path: r.leaseFiles,
Logger: r.logger,
Callback: func(_ string) {
r.doReloadLeaseFiles()
},
})
if err != nil {
r.logger.Warn(E.Cause(err, "create lease file watcher"))
} else {
r.watcher = watcher
err = watcher.Start()
if err != nil {
r.logger.Warn(E.Cause(err, "start lease file watcher"))
}
}
}
return nil
}
func (r *neighborResolver) Close() error {
close(r.done)
if r.watcher != nil {
return r.watcher.Close()
}
return nil
}
func (r *neighborResolver) LookupMAC(address netip.Addr) (net.HardwareAddr, bool) {
r.access.RLock()
defer r.access.RUnlock()
mac, found := r.neighborIPToMAC[address]
if found {
return mac, true
}
mac, found = r.leaseIPToMAC[address]
if found {
return mac, true
}
mac, found = extractMACFromEUI64(address)
if found {
return mac, true
}
return nil, false
}
func (r *neighborResolver) LookupAddresses(hostname string) []netip.Addr {
r.access.RLock()
defer r.access.RUnlock()
return lookupAddressesByHostname(hostname, r.ipToHostname, r.macToHostname, r.neighborIPToMAC, r.leaseIPToMAC)
}
func (r *neighborResolver) LookupHostname(address netip.Addr) (string, bool) {
r.access.RLock()
defer r.access.RUnlock()
hostname, found := r.ipToHostname[address]
if found {
return hostname, true
}
mac, macFound := r.neighborIPToMAC[address]
if !macFound {
mac, macFound = r.leaseIPToMAC[address]
}
if !macFound {
mac, macFound = extractMACFromEUI64(address)
}
if macFound {
hostname, found = r.macToHostname[mac.String()]
if found {
return hostname, true
}
}
return "", false
}
func (r *neighborResolver) loadNeighborTable() error {
connection, err := rtnetlink.Dial(nil)
if err != nil {
return E.Cause(err, "dial rtnetlink")
}
defer connection.Close()
neighbors, err := connection.Neigh.List()
if err != nil {
return E.Cause(err, "list neighbors")
}
r.access.Lock()
defer r.access.Unlock()
for _, neigh := range neighbors {
if neigh.Attributes == nil {
continue
}
if neigh.Attributes.LLAddress == nil || len(neigh.Attributes.Address) == 0 {
continue
}
address, ok := netip.AddrFromSlice(neigh.Attributes.Address)
if !ok {
continue
}
r.neighborIPToMAC[address] = slices.Clone(neigh.Attributes.LLAddress)
}
return nil
}
func (r *neighborResolver) subscribeNeighborUpdates() {
connection, err := netlink.Dial(unix.NETLINK_ROUTE, &netlink.Config{
Groups: 1 << (unix.RTNLGRP_NEIGH - 1),
})
if err != nil {
r.logger.Warn(E.Cause(err, "subscribe neighbor updates"))
return
}
defer connection.Close()
for {
select {
case <-r.done:
return
default:
}
err = connection.SetReadDeadline(time.Now().Add(3 * time.Second))
if err != nil {
r.logger.Warn(E.Cause(err, "set netlink read deadline"))
return
}
messages, err := connection.Receive()
if err != nil {
if nerr, ok := err.(net.Error); ok && nerr.Timeout() {
continue
}
select {
case <-r.done:
return
default:
}
r.logger.Warn(E.Cause(err, "receive neighbor update"))
continue
}
for _, message := range messages {
address, mac, isDelete, ok := ParseNeighborMessage(message)
if !ok {
continue
}
r.access.Lock()
if isDelete {
delete(r.neighborIPToMAC, address)
} else {
r.neighborIPToMAC[address] = mac
}
r.access.Unlock()
}
}
}
func (r *neighborResolver) doReloadLeaseFiles() {
leaseIPToMAC, ipToHostname, macToHostname := ReloadLeaseFiles(r.leaseFiles)
r.access.Lock()
r.leaseIPToMAC = leaseIPToMAC
r.ipToHostname = ipToHostname
r.macToHostname = macToHostname
r.access.Unlock()
}
+50
View File
@@ -0,0 +1,50 @@
package route
import (
"encoding/binary"
"encoding/hex"
"net"
"net/netip"
"slices"
"strings"
)
func extractMACFromDUID(duid []byte) (net.HardwareAddr, bool) {
if len(duid) < 4 {
return nil, false
}
duidType := binary.BigEndian.Uint16(duid[0:2])
hwType := binary.BigEndian.Uint16(duid[2:4])
if hwType != 1 {
return nil, false
}
switch duidType {
case 1:
if len(duid) < 14 {
return nil, false
}
return net.HardwareAddr(slices.Clone(duid[8:14])), true
case 3:
if len(duid) < 10 {
return nil, false
}
return net.HardwareAddr(slices.Clone(duid[4:10])), true
}
return nil, false
}
func extractMACFromEUI64(address netip.Addr) (net.HardwareAddr, bool) {
if !address.Is6() {
return nil, false
}
b := address.As16()
if b[11] != 0xff || b[12] != 0xfe {
return nil, false
}
return net.HardwareAddr{b[8] ^ 0x02, b[9], b[10], b[13], b[14], b[15]}, true
}
func parseDUID(s string) ([]byte, error) {
cleaned := strings.ReplaceAll(s, ":", "")
return hex.DecodeString(cleaned)
}
+90
View File
@@ -0,0 +1,90 @@
package route
import (
"net"
"net/netip"
"sync"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing/common/logger"
)
type platformNeighborResolver struct {
logger logger.ContextLogger
platform adapter.PlatformInterface
access sync.RWMutex
ipToMAC map[netip.Addr]net.HardwareAddr
ipToHostname map[netip.Addr]string
macToHostname map[string]string
}
func newPlatformNeighborResolver(resolverLogger logger.ContextLogger, platform adapter.PlatformInterface) adapter.NeighborResolver {
return &platformNeighborResolver{
logger: resolverLogger,
platform: platform,
ipToMAC: make(map[netip.Addr]net.HardwareAddr),
ipToHostname: make(map[netip.Addr]string),
macToHostname: make(map[string]string),
}
}
func (r *platformNeighborResolver) Start() error {
return r.platform.StartNeighborMonitor(r)
}
func (r *platformNeighborResolver) Close() error {
return r.platform.CloseNeighborMonitor(r)
}
func (r *platformNeighborResolver) LookupMAC(address netip.Addr) (net.HardwareAddr, bool) {
r.access.RLock()
defer r.access.RUnlock()
mac, found := r.ipToMAC[address]
if found {
return mac, true
}
return extractMACFromEUI64(address)
}
func (r *platformNeighborResolver) LookupAddresses(hostname string) []netip.Addr {
r.access.RLock()
defer r.access.RUnlock()
return lookupAddressesByHostname(hostname, r.ipToHostname, r.macToHostname, r.ipToMAC)
}
func (r *platformNeighborResolver) LookupHostname(address netip.Addr) (string, bool) {
r.access.RLock()
defer r.access.RUnlock()
hostname, found := r.ipToHostname[address]
if found {
return hostname, true
}
mac, found := r.ipToMAC[address]
if !found {
mac, found = extractMACFromEUI64(address)
}
if !found {
return "", false
}
hostname, found = r.macToHostname[mac.String()]
return hostname, found
}
func (r *platformNeighborResolver) UpdateNeighborTable(entries []adapter.NeighborEntry) {
ipToMAC := make(map[netip.Addr]net.HardwareAddr)
ipToHostname := make(map[netip.Addr]string)
macToHostname := make(map[string]string)
for _, entry := range entries {
ipToMAC[entry.Address] = entry.MACAddress
if entry.Hostname != "" {
ipToHostname[entry.Address] = entry.Hostname
macToHostname[entry.MACAddress.String()] = entry.Hostname
}
}
r.access.Lock()
r.ipToMAC = ipToMAC
r.ipToHostname = ipToHostname
r.macToHostname = macToHostname
r.access.Unlock()
r.logger.Info("updated neighbor table: ", len(entries), " entries")
}
+14
View File
@@ -0,0 +1,14 @@
//go:build !linux && !darwin
package route
import (
"os"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing/common/logger"
)
func newNeighborResolver(_ logger.ContextLogger, _ []string) (adapter.NeighborResolver, error) {
return nil, os.ErrInvalid
}
+104
View File
@@ -0,0 +1,104 @@
//go:build darwin
package route
import (
"net"
"net/netip"
"syscall"
"github.com/sagernet/sing-box/adapter"
E "github.com/sagernet/sing/common/exceptions"
"golang.org/x/net/route"
"golang.org/x/sys/unix"
)
func ReadNeighborEntries() ([]adapter.NeighborEntry, error) {
var entries []adapter.NeighborEntry
ipv4Entries, err := readNeighborEntriesAF(syscall.AF_INET)
if err != nil {
return nil, E.Cause(err, "read IPv4 neighbors")
}
entries = append(entries, ipv4Entries...)
ipv6Entries, err := readNeighborEntriesAF(syscall.AF_INET6)
if err != nil {
return nil, E.Cause(err, "read IPv6 neighbors")
}
entries = append(entries, ipv6Entries...)
return entries, nil
}
func readNeighborEntriesAF(addressFamily int) ([]adapter.NeighborEntry, error) {
rib, err := route.FetchRIB(addressFamily, route.RIBType(syscall.NET_RT_FLAGS), syscall.RTF_LLINFO)
if err != nil {
return nil, err
}
messages, err := route.ParseRIB(route.RIBType(syscall.NET_RT_FLAGS), rib)
if err != nil {
return nil, err
}
var entries []adapter.NeighborEntry
for _, message := range messages {
routeMessage, isRouteMessage := message.(*route.RouteMessage)
if !isRouteMessage {
continue
}
address, macAddress, ok := parseRouteNeighborEntry(routeMessage)
if !ok {
continue
}
entries = append(entries, adapter.NeighborEntry{
Address: address,
MACAddress: macAddress,
})
}
return entries, nil
}
func parseRouteNeighborEntry(message *route.RouteMessage) (address netip.Addr, macAddress net.HardwareAddr, ok bool) {
if len(message.Addrs) <= unix.RTAX_GATEWAY {
return
}
gateway, isLinkAddr := message.Addrs[unix.RTAX_GATEWAY].(*route.LinkAddr)
if !isLinkAddr || len(gateway.Addr) < 6 {
return
}
switch destination := message.Addrs[unix.RTAX_DST].(type) {
case *route.Inet4Addr:
address = netip.AddrFrom4(destination.IP)
case *route.Inet6Addr:
address = netip.AddrFrom16(destination.IP)
default:
return
}
macAddress = net.HardwareAddr(make([]byte, len(gateway.Addr)))
copy(macAddress, gateway.Addr)
ok = true
return
}
func ParseRouteNeighborMessage(message *route.RouteMessage) (address netip.Addr, macAddress net.HardwareAddr, isDelete bool, ok bool) {
isDelete = message.Type == unix.RTM_DELETE
if len(message.Addrs) <= unix.RTAX_GATEWAY {
return
}
switch destination := message.Addrs[unix.RTAX_DST].(type) {
case *route.Inet4Addr:
address = netip.AddrFrom4(destination.IP)
case *route.Inet6Addr:
address = netip.AddrFrom16(destination.IP)
default:
return
}
if !isDelete {
gateway, isLinkAddr := message.Addrs[unix.RTAX_GATEWAY].(*route.LinkAddr)
if !isLinkAddr || len(gateway.Addr) < 6 {
return
}
macAddress = net.HardwareAddr(make([]byte, len(gateway.Addr)))
copy(macAddress, gateway.Addr)
}
ok = true
return
}
+68
View File
@@ -0,0 +1,68 @@
//go:build linux
package route
import (
"net"
"net/netip"
"slices"
"github.com/sagernet/sing-box/adapter"
E "github.com/sagernet/sing/common/exceptions"
"github.com/jsimonetti/rtnetlink"
"github.com/mdlayher/netlink"
"golang.org/x/sys/unix"
)
func ReadNeighborEntries() ([]adapter.NeighborEntry, error) {
connection, err := rtnetlink.Dial(nil)
if err != nil {
return nil, E.Cause(err, "dial rtnetlink")
}
defer connection.Close()
neighbors, err := connection.Neigh.List()
if err != nil {
return nil, E.Cause(err, "list neighbors")
}
var entries []adapter.NeighborEntry
for _, neighbor := range neighbors {
if neighbor.Attributes == nil {
continue
}
if neighbor.Attributes.LLAddress == nil || len(neighbor.Attributes.Address) == 0 {
continue
}
address, ok := netip.AddrFromSlice(neighbor.Attributes.Address)
if !ok {
continue
}
entries = append(entries, adapter.NeighborEntry{
Address: address,
MACAddress: slices.Clone(neighbor.Attributes.LLAddress),
})
}
return entries, nil
}
func ParseNeighborMessage(message netlink.Message) (address netip.Addr, macAddress net.HardwareAddr, isDelete bool, ok bool) {
var neighMessage rtnetlink.NeighMessage
err := neighMessage.UnmarshalBinary(message.Data)
if err != nil {
return
}
if neighMessage.Attributes == nil || len(neighMessage.Attributes.Address) == 0 {
return
}
address, ok = netip.AddrFromSlice(neighMessage.Attributes.Address)
if !ok {
return
}
isDelete = message.Header.Type == unix.RTM_DELNEIGH
if !isDelete && neighMessage.Attributes.LLAddress == nil {
ok = false
return
}
macAddress = slices.Clone(neighMessage.Attributes.LLAddress)
return
}
+124 -46
View File
@@ -34,28 +34,37 @@ import (
var _ adapter.NetworkManager = (*NetworkManager)(nil)
type NetworkManager struct {
logger logger.ContextLogger
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
interfaceUpdateAccess sync.Mutex
interfaceUpdateCancel context.CancelFunc
interfaceUpdateRunAccess sync.Mutex
powerUpdateAccess sync.Mutex
powerUpdateCancel context.CancelFunc
started bool
}
func NewNetworkManager(ctx context.Context, logger logger.ContextLogger, options option.RouteOptions, dnsOptions option.DNSOptions) (*NetworkManager, error) {
@@ -70,6 +79,7 @@ func NewNetworkManager(ctx context.Context, logger logger.ContextLogger, options
return nil, E.New("`default_mark` is only supported on linux")
}
nm := &NetworkManager{
ctx: ctx,
logger: logger,
interfaceFinder: control.NewDefaultInterfaceFinder(),
autoDetectInterface: options.AutoDetectInterface,
@@ -78,10 +88,12 @@ func NewNetworkManager(ctx context.Context, logger logger.ContextLogger, options
RoutingMark: uint32(options.DefaultMark),
DomainResolver: defaultDomainResolver.Server,
DomainResolveOptions: adapter.DNSQueryOptions{
Strategy: C.DomainStrategy(defaultDomainResolver.Strategy),
DisableCache: defaultDomainResolver.DisableCache,
RewriteTTL: defaultDomainResolver.RewriteTTL,
ClientSubnet: defaultDomainResolver.ClientSubnet.Build(netip.Prefix{}),
Strategy: C.DomainStrategy(defaultDomainResolver.Strategy),
Timeout: time.Duration(defaultDomainResolver.Timeout),
DisableCache: defaultDomainResolver.DisableCache,
DisableOptimisticCache: defaultDomainResolver.DisableOptimisticCache,
RewriteTTL: defaultDomainResolver.RewriteTTL,
ClientSubnet: defaultDomainResolver.ClientSubnet.Build(netip.Prefix{}),
},
NetworkStrategy: (*C.NetworkStrategy)(options.DefaultNetworkStrategy),
NetworkType: common.Map(options.DefaultNetworkType, option.InterfaceType.Build),
@@ -104,7 +116,7 @@ func NewNetworkManager(ctx context.Context, logger logger.ContextLogger, options
return nil, E.New("`auto_detect_interface` is required by `default_network_strategy`")
}
}
usePlatformDefaultInterfaceMonitor := nm.platformInterface != nil
usePlatformDefaultInterfaceMonitor := nm.platformInterface != nil && nm.platformInterface.UsePlatformDefaultInterfaceMonitor()
enforceInterfaceMonitor := options.AutoDetectInterface
if !usePlatformDefaultInterfaceMonitor {
networkMonitor, err := tun.NewNetworkUpdateMonitor(logger)
@@ -113,6 +125,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,
@@ -136,6 +149,7 @@ func (r *NetworkManager) Start(stage adapter.StartStage) error {
monitor := taskmonitor.New(r.logger, C.StartTimeout)
switch stage {
case adapter.StartStateInitialize:
r.router = service.FromContext[adapter.Router](r.ctx)
if r.networkMonitor != nil {
monitor.Start("initialize network monitor")
err := r.networkMonitor.Start()
@@ -242,6 +256,14 @@ func (r *NetworkManager) Close() error {
})
monitor.Finish()
}
r.interfaceUpdateAccess.Lock()
interfaceUpdateCancel := r.interfaceUpdateCancel
r.interfaceUpdateCancel = nil
r.interfaceUpdateAccess.Unlock()
if interfaceUpdateCancel != nil {
interfaceUpdateCancel()
}
r.cancelPowerUpdate()
if r.networkMonitor != nil {
monitor.Start("close network monitor")
err = E.Append(err, r.networkMonitor.Close(), func(err error) error {
@@ -249,6 +271,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 {
@@ -264,6 +291,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 {
@@ -418,40 +446,41 @@ 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.Notice("WIFI state changed: SSID=", state.SSID, ", BSSID=", state.BSSID)
} else {
r.logger.Notice("WIFI disconnected")
}
} else {
r.wifiStateMutex.Unlock()
r.stateAccess.Unlock()
}
}
func (r *NetworkManager) UpdateWIFIState() {
func (r *NetworkManager) UpdateWIFIState(ctx context.Context) {
var state adapter.WIFIState
if r.wifiMonitor != nil {
state = r.wifiMonitor.ReadWIFIState()
state = r.wifiMonitor.ReadWIFIState(ctx)
} else if r.platformInterface != nil && r.platformInterface.UsePlatformWIFIMonitor() {
state = r.platformInterface.ReadWIFIState()
state = r.platformInterface.ReadWIFIState(ctx)
} else {
return
}
r.onWIFIStateChanged(state)
}
func (r *NetworkManager) ResetNetwork() {
func (r *NetworkManager) ResetNetwork(ctx context.Context) {
if r.connectionManager != nil {
r.connectionManager.CloseAll()
}
@@ -459,23 +488,25 @@ func (r *NetworkManager) ResetNetwork() {
for _, endpoint := range r.endpoint.Endpoints() {
listener, isListener := endpoint.(adapter.InterfaceUpdateListener)
if isListener {
listener.InterfaceUpdated()
listener.InterfaceUpdated(ctx)
}
}
for _, inbound := range r.inbound.Inbounds() {
listener, isListener := inbound.(adapter.InterfaceUpdateListener)
if isListener {
listener.InterfaceUpdated()
listener.InterfaceUpdated(ctx)
}
}
for _, outbound := range r.outbound.Outbounds() {
listener, isListener := outbound.(adapter.InterfaceUpdateListener)
if isListener {
listener.InterfaceUpdated()
listener.InterfaceUpdated(ctx)
}
}
r.router.ResetNetwork()
}
func (r *NetworkManager) notifyInterfaceUpdate(defaultInterface *control.Interface, flags int) {
@@ -484,8 +515,27 @@ func (r *NetworkManager) notifyInterfaceUpdate(defaultInterface *control.Interfa
r.logger.Error("missing default interface")
return
}
r.pauseManager.NetworkWake()
updateContext, updateCancel := context.WithCancel(r.ctx)
r.interfaceUpdateAccess.Lock()
previousCancel := r.interfaceUpdateCancel
r.interfaceUpdateCancel = updateCancel
r.interfaceUpdateAccess.Unlock()
if previousCancel != nil {
previousCancel()
}
go func() {
defer updateCancel()
r.updateInterface(updateContext, defaultInterface)
}()
}
func (r *NetworkManager) updateInterface(ctx context.Context, defaultInterface *control.Interface) {
r.interfaceUpdateRunAccess.Lock()
defer r.interfaceUpdateRunAccess.Unlock()
if ctx.Err() != nil {
return
}
var options []string
options = append(options, F.ToString("index ", defaultInterface.Index))
if C.IsAndroid && r.platformInterface == nil {
@@ -496,7 +546,7 @@ func (r *NetworkManager) notifyInterfaceUpdate(defaultInterface *control.Interfa
vpnStatus = "disabled"
}
options = append(options, "vpn "+vpnStatus)
} else if r.platformInterface != nil {
} else if r.platformInterface != nil && r.platformInterface.UsePlatformNetworkInterfaces() {
networkInterface := common.Find(r.networkInterfaces.Load(), func(it adapter.NetworkInterface) bool {
return it.Interface.Index == defaultInterface.Index
})
@@ -513,19 +563,26 @@ func (r *NetworkManager) notifyInterfaceUpdate(defaultInterface *control.Interfa
}
}
r.logger.Notice("updated default interface ", defaultInterface.Name, ", ", strings.Join(options, ", "))
r.UpdateWIFIState()
r.UpdateWIFIState(ctx)
if ctx.Err() != nil {
return
}
r.updateNetworkEnvironment()
if ctx.Err() != nil {
return
}
if !r.started {
return
}
r.ResetNetwork()
r.ResetNetwork(ctx)
}
func (r *NetworkManager) notifyWindowsPowerEvent(event int) {
switch event {
case winpowrprof.EVENT_SUSPEND:
r.pauseManager.DevicePause()
r.ResetNetwork()
r.cancelPowerUpdate()
r.ResetNetwork(r.ctx)
case winpowrprof.EVENT_RESUME:
if !r.pauseManager.IsDevicePaused() {
return
@@ -533,7 +590,28 @@ func (r *NetworkManager) notifyWindowsPowerEvent(event int) {
fallthrough
case winpowrprof.EVENT_RESUME_AUTOMATIC:
r.pauseManager.DeviceWake()
r.ResetNetwork()
updateContext, updateCancel := context.WithCancel(r.ctx)
r.powerUpdateAccess.Lock()
previousCancel := r.powerUpdateCancel
r.powerUpdateCancel = updateCancel
r.powerUpdateAccess.Unlock()
if previousCancel != nil {
previousCancel()
}
go func() {
defer updateCancel()
r.ResetNetwork(updateContext)
}()
}
}
func (r *NetworkManager) cancelPowerUpdate() {
r.powerUpdateAccess.Lock()
previousCancel := r.powerUpdateCancel
r.powerUpdateCancel = nil
r.powerUpdateAccess.Unlock()
if previousCancel != nil {
previousCancel()
}
}
+100
View File
@@ -0,0 +1,100 @@
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 (
gatewayStrings []string
hardwareStrings []string
wifiSSID 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)
gatewayStrings = common.Map(gateways, netip.Addr.String)
wifiState := r.WIFIState()
if wifiState.SSID != "" {
wifiSSID = wifiState.SSID
} else if len(gateways) > 0 {
hardwareAddresses := systemNeighborHardwareAddresses(defaultInterface.Interface.Index, gateways)
for _, gateway := range gateways {
hardwareAddress := hardwareAddresses[gateway]
if len(hardwareAddress) > 0 {
hardwareStrings = append(hardwareStrings, hardwareAddress.String())
}
}
}
}
var options []string
if len(gatewayStrings) > 0 {
options = append(options, "gateway "+formatEnvironmentValues(gatewayStrings))
}
if wifiSSID != "" {
options = append(options, "ssid "+wifiSSID)
}
if len(hardwareStrings) > 0 {
options = append(options, "gateway_mac "+formatEnvironmentValues(hardwareStrings))
}
var environmentHash uint64
if len(options) > 0 {
digest := fnv.New64a()
for _, option := range options {
digest.Write([]byte(option))
digest.Write([]byte{0})
}
environmentHash = digest.Sum64()
}
r.stateAccess.Lock()
changed := environmentHash != r.networkEnvironment
r.networkEnvironment = environmentHash
r.stateAccess.Unlock()
if !changed || len(options) == 0 {
return
}
r.logger.Info("updated network environment: ", strings.Join(options, ", "))
}
func formatEnvironmentValues(values []string) string {
if len(values) == 1 {
return values[0]
}
return "[" + strings.Join(values, " ") + "]"
}
+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
}
+3
View File
@@ -44,6 +44,9 @@ func (s *platformSearcher) FindProcessInfo(ctx context.Context, network string,
return s.platform.FindConnectionOwner(request)
}
func (s *platformSearcher) ResetCache() {
}
func (s *platformSearcher) Close() error {
return nil
}
+290 -162
View File
@@ -11,11 +11,11 @@ import (
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/sniff"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/log"
R "github.com/sagernet/sing-box/route/rule"
mux "github.com/sagernet/sing-mux"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing-tun/ping"
vmess "github.com/sagernet/sing-vmess"
"github.com/sagernet/sing-vmess"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio"
@@ -29,6 +29,16 @@ import (
"golang.org/x/exp/slices"
)
var defaultPacketSniffers = []sniff.PacketSniffer{
sniff.DomainNameQuery,
sniff.QUICClientHello,
sniff.STUNMessage,
sniff.UTP,
sniff.UDPTracker,
sniff.DTLSRecord,
sniff.NTP,
}
// Deprecated: use RouteConnectionEx instead.
func (r *Router) RouteConnection(ctx context.Context, conn net.Conn, metadata adapter.InboundContext) error {
done := make(chan any)
@@ -81,7 +91,7 @@ func (r *Router) routeConnection(ctx context.Context, conn net.Conn, metadata ad
metadata.LastInbound = metadata.Inbound
metadata.Inbound = metadata.InboundDetour
metadata.InboundDetour = ""
injectable.NewConnectionEx(ctx, conn, metadata, onClose)
injectable.NewConnection(ctx, conn, metadata, onClose)
return nil
}
metadata.Network = N.NetworkTCP
@@ -95,10 +105,14 @@ func (r *Router) routeConnection(ctx context.Context, conn net.Conn, metadata ad
case uot.LegacyMagicAddress:
return E.New("global UoT (legacy) not supported since sing-box v1.7.0.")
}
if metadata.InboundType == C.TypeTun && metadata.Protocol == C.ProtocolDNS {
N.CloseOnHandshakeFailure(conn, onClose, r.hijackDNSStream(ctx, conn, metadata))
return nil
}
if deadline.NeedAdditionalReadDeadline(conn) {
conn = deadline.NewConn(conn)
}
selectedRule, _, buffers, _, err := r.matchRule(ctx, &metadata, false, false, conn, nil)
selectedRule, _, buffers, _, err := r.matchRule(ctx, &metadata, conn, nil)
if err != nil {
return err
}
@@ -155,11 +169,15 @@ func (r *Router) routeConnection(ctx context.Context, conn net.Conn, metadata ad
for _, buffer := range buffers {
conn = bufio.NewCachedConn(conn, buffer)
}
if selectedRule != nil {
metadata.RouteRule = selectedRule.String()
}
metadata.RouteOutbound = selectedOutbound.Tag()
for _, tracker := range r.trackers {
conn = tracker.RoutedConnection(ctx, conn, metadata, selectedRule, selectedOutbound)
}
if outboundHandler, isHandler := selectedOutbound.(adapter.ConnectionHandlerEx); isHandler {
outboundHandler.NewConnectionEx(ctx, conn, metadata, onClose)
if outboundHandler, isHandler := selectedOutbound.(adapter.ConnectionHandler); isHandler {
outboundHandler.NewConnection(ctx, conn, metadata, onClose)
} else {
r.connection.NewConnection(ctx, selectedOutbound, conn, metadata, onClose)
}
@@ -222,7 +240,7 @@ func (r *Router) routePacketConnection(ctx context.Context, conn N.PacketConn, m
metadata.LastInbound = metadata.Inbound
metadata.Inbound = metadata.InboundDetour
metadata.InboundDetour = ""
injectable.NewPacketConnectionEx(ctx, conn, metadata, onClose)
injectable.NewPacketConnection(ctx, conn, metadata, onClose)
return nil
}
// TODO: move to UoT
@@ -232,8 +250,10 @@ func (r *Router) routePacketConnection(ctx context.Context, conn N.PacketConn, m
/*if deadline.NeedAdditionalReadDeadline(conn) {
conn = deadline.NewPacketConn(bufio.NewNetPacketConn(conn))
}*/
selectedRule, _, _, packetBuffers, err := r.matchRule(ctx, &metadata, false, false, nil, conn)
if metadata.InboundType == C.TypeTun && metadata.Protocol == C.ProtocolDNS {
return r.hijackDNSPacket(ctx, conn, nil, metadata, onClose)
}
selectedRule, _, _, packetBuffers, err := r.matchRule(ctx, &metadata, nil, conn)
if err != nil {
return err
}
@@ -287,144 +307,265 @@ func (r *Router) routePacketConnection(ctx context.Context, conn N.PacketConn, m
conn = bufio.NewCachedPacketConn(conn, buffer.Buffer, buffer.Destination)
N.PutPacketBuffer(buffer)
}
if selectedRule != nil {
metadata.RouteRule = selectedRule.String()
}
metadata.RouteOutbound = selectedOutbound.Tag()
for _, tracker := range r.trackers {
conn = tracker.RoutedPacketConnection(ctx, conn, metadata, selectedRule, selectedOutbound)
}
if metadata.FakeIP {
conn = bufio.NewNATPacketConn(bufio.NewNetPacketConn(conn), metadata.OriginDestination, metadata.Destination)
conn = newFakeIPNATPacketConn(bufio.NewNetPacketConn(conn), metadata.OriginDestination, metadata.Destination)
}
if outboundHandler, isHandler := selectedOutbound.(adapter.PacketConnectionHandlerEx); isHandler {
outboundHandler.NewPacketConnectionEx(ctx, conn, metadata, onClose)
if outboundHandler, isHandler := selectedOutbound.(adapter.PacketConnectionHandler); isHandler {
outboundHandler.NewPacketConnection(ctx, conn, metadata, onClose)
} else {
r.connection.NewPacketConnection(ctx, selectedOutbound, conn, metadata, onClose)
}
return nil
}
func (r *Router) PreMatch(metadata adapter.InboundContext, routeContext tun.DirectRouteContext, timeout time.Duration, supportBypass bool) (tun.DirectRouteDestination, error) {
selectedRule, _, _, _, err := r.matchRule(r.ctx, &metadata, true, supportBypass, nil, nil)
func (r *Router) PreMatch(metadata adapter.InboundContext, firstPacket []byte) adapter.PreMatchResult {
ctx := log.ContextWithNewID(r.ctx)
metadata.PreMatch = true
continueResult := adapter.PreMatchResult{Action: adapter.PreMatchContinue}
packetDestination := metadata.Destination
err := r.prepareMatchMetadata(ctx, &metadata)
if err != nil {
return nil, err
return continueResult
}
var directRouteOutbound adapter.DirectRouteOutbound
if selectedRule != nil {
switch action := selectedRule.Action().(type) {
case *R.RuleActionReject:
switch metadata.Network {
case N.NetworkTCP:
if action.Method == C.RuleActionRejectMethodReply {
return nil, E.New("reject method `reply` is not supported for TCP connections")
for currentRuleIndex, currentRule := range r.rules {
metadata.ResetRuleCache()
if !currentRule.Match(&metadata) {
continue
}
ruleDescription := currentRule.String()
if ruleDescription != "" {
r.logger.DebugContext(ctx, "pre-match[", currentRuleIndex, "] ", currentRule, " => ", currentRule.Action())
} else {
r.logger.DebugContext(ctx, "pre-match[", currentRuleIndex, "] => ", currentRule.Action())
}
switch action := currentRule.Action().(type) {
case *R.RuleActionSniff:
if metadata.Network == N.NetworkICMP {
continue
}
if metadata.Network != N.NetworkUDP || len(firstPacket) == 0 {
return continueResult
}
if sniff.Skip(&metadata) || metadata.Protocol != "" {
continue
}
if len(action.PacketSniffers) == 0 && len(action.StreamSniffers) > 0 {
continue
}
if slices.Equal(metadata.SnifferNames, action.SnifferNames) && metadata.SniffError != nil {
continue
}
packetSniffers := action.PacketSniffers
if len(packetSniffers) == 0 {
packetSniffers = defaultPacketSniffers
}
sniffErr := sniff.PeekPacket(ctx, &metadata, firstPacket, packetSniffers...)
metadata.SnifferNames = action.SnifferNames
metadata.SniffError = sniffErr
if sniffErr != nil {
if errors.Is(sniffErr, sniff.ErrNeedMoreData) {
return continueResult
}
case N.NetworkUDP:
if action.Method == C.RuleActionRejectMethodReply {
return nil, E.New("reject method `reply` is not supported for UDP connections")
continue
}
//goland:noinspection GoDeprecation
if action.OverrideDestination && M.IsDomainName(metadata.Domain) {
metadata.Destination = M.Socksaddr{
Fqdn: metadata.Domain,
Port: metadata.Destination.Port,
}
}
return nil, action.Error(context.Background())
case *R.RuleActionBypass:
if supportBypass {
return nil, &R.BypassedError{Cause: tun.ErrBypass}
if metadata.Domain != "" && metadata.Client != "" {
r.logger.DebugContext(ctx, "sniffed packet protocol: ", metadata.Protocol, ", domain: ", metadata.Domain, ", client: ", metadata.Client)
} else if metadata.Domain != "" {
r.logger.DebugContext(ctx, "sniffed packet protocol: ", metadata.Protocol, ", domain: ", metadata.Domain)
} else if metadata.Client != "" {
r.logger.DebugContext(ctx, "sniffed packet protocol: ", metadata.Protocol, ", client: ", metadata.Client)
} else {
r.logger.DebugContext(ctx, "sniffed packet protocol: ", metadata.Protocol)
}
if routeContext == nil {
return nil, nil
}
outbound, loaded := r.outbound.Outbound(action.Outbound)
if !loaded {
return nil, E.New("outbound not found: ", action.Outbound)
}
if !common.Contains(outbound.Network(), metadata.Network) {
return nil, E.New(metadata.Network, " is not supported by outbound: ", action.Outbound)
}
directRouteOutbound = outbound.(adapter.DirectRouteOutbound)
case *R.RuleActionRouteOptions:
applyRouteOptionsOverride(&metadata, action)
case *R.RuleActionRoute:
if routeContext == nil {
return nil, nil
applyRouteOptionsOverride(&metadata, &action.RuleActionRouteOptions)
return r.preMatchFlow(ctx, &metadata, packetDestination, currentRule, action.Outbound)
case *R.RuleActionBypass:
applyRouteOptionsOverride(&metadata, &action.RuleActionRouteOptions)
if action.Outbound == "" {
if metadata.Destination.IsDomain() || metadata.Destination != packetDestination {
return continueResult
}
return adapter.PreMatchResult{Action: adapter.PreMatchBypass}
}
outbound, loaded := r.outbound.Outbound(action.Outbound)
if !loaded {
return nil, E.New("outbound not found: ", action.Outbound)
return r.preMatchFlow(ctx, &metadata, packetDestination, currentRule, action.Outbound)
case *R.RuleActionReject:
rejectErr := action.Error(r.ctx)
if errors.Is(rejectErr, R.ErrDrop) {
return adapter.PreMatchResult{Action: adapter.PreMatchDrop}
}
if !common.Contains(outbound.Network(), metadata.Network) {
return nil, E.New(metadata.Network, " is not supported by outbound: ", action.Outbound)
return adapter.PreMatchResult{Action: adapter.PreMatchReject}
case *R.RuleActionHijackDNS:
if metadata.Network != N.NetworkUDP {
return continueResult
}
directRouteOutbound = outbound.(adapter.DirectRouteOutbound)
return adapter.PreMatchResult{Action: adapter.PreMatchHijackDNS}
case *R.RuleActionResolve:
resolveErr := r.actionResolve(adapter.WithContext(ctx, &metadata), &metadata, action)
if resolveErr != nil {
r.logger.DebugContext(ctx, "pre-match[", currentRuleIndex, "] ", currentRule, " => ", action, ": ", resolveErr)
return adapter.PreMatchResult{Action: adapter.PreMatchReject}
}
default:
return continueResult
}
}
if directRouteOutbound == nil {
if selectedRule != nil || metadata.Network != N.NetworkICMP {
return nil, nil
return r.preMatchFlow(ctx, &metadata, packetDestination, nil, "")
}
func applyRouteOptionsOverride(metadata *adapter.InboundContext, routeOptions *R.RuleActionRouteOptions) {
if routeOptions.OverrideAddress.IsValid() {
metadata.Destination = M.Socksaddr{
Addr: routeOptions.OverrideAddress.Addr,
Port: metadata.Destination.Port,
Fqdn: routeOptions.OverrideAddress.Fqdn,
}
defaultOutbound := r.outbound.Default()
if !common.Contains(defaultOutbound.Network(), metadata.Network) {
return nil, E.New(metadata.Network, " is not supported by default outbound: ", defaultOutbound.Tag())
}
if routeOptions.OverridePort > 0 {
metadata.Destination = M.Socksaddr{
Addr: metadata.Destination.Addr,
Port: routeOptions.OverridePort,
Fqdn: metadata.Destination.Fqdn,
}
}
if routeOptions.UDPTimeout > 0 {
metadata.UDPTimeout = routeOptions.UDPTimeout
}
}
func (r *Router) preMatchFlow(ctx context.Context, metadata *adapter.InboundContext, packetDestination M.Socksaddr, matchedRule adapter.Rule, outboundTag string) adapter.PreMatchResult {
continueResult := adapter.PreMatchResult{Action: adapter.PreMatchContinue}
var outbound adapter.Outbound
if outboundTag == "" {
outbound = r.outbound.Default()
} else {
var loaded bool
outbound, loaded = r.outbound.Outbound(outboundTag)
if !loaded {
return continueResult
}
}
for range 8 {
group, isGroup := outbound.(adapter.OutboundGroup)
if !isGroup {
break
}
selectedOutbound, selectedLoaded := r.outbound.Outbound(group.Now())
if !selectedLoaded {
return continueResult
}
outbound = selectedOutbound
}
if !common.Contains(outbound.Network(), metadata.Network) {
return continueResult
}
flowOutbound, isFlowOutbound := outbound.(adapter.FlowOutbound)
if !isFlowOutbound {
return continueResult
}
flowAction := flowOutbound.PreMatchFlow(metadata.Network, metadata.Destination.Addr)
if flowAction != adapter.PreMatchFlow {
return adapter.PreMatchResult{Action: flowAction, Outbound: outbound}
}
result := adapter.PreMatchResult{Action: adapter.PreMatchFlow, Outbound: outbound}
if metadata.Network == N.NetworkUDP {
if metadata.UDPTimeout > 0 {
result.UDPTimeout = metadata.UDPTimeout
} else {
protocol := metadata.Protocol
if protocol == "" {
protocol = C.PortProtocols[metadata.Destination.Port]
}
if protocol != "" {
result.UDPTimeout = C.ProtocolTimeouts[protocol]
}
}
directRouteOutbound = defaultOutbound.(adapter.DirectRouteOutbound)
}
if metadata.Destination.IsDomain() {
if len(metadata.DestinationAddresses) == 0 {
var strategy C.DomainStrategy
if metadata.Source.IsIPv4() {
strategy = C.DomainStrategyIPv4Only
} else {
strategy = C.DomainStrategyIPv6Only
}
err = r.actionResolve(r.ctx, &metadata, &R.RuleActionResolve{
Strategy: strategy,
})
if err != nil {
return nil, err
}
if !metadata.FakeIP {
return continueResult
}
var newDestination netip.Addr
if metadata.Source.IsIPv4() {
for _, address := range metadata.DestinationAddresses {
if address.Is4() {
newDestination = address
break
}
}
} else {
for _, address := range metadata.DestinationAddresses {
if address.Is6() {
newDestination = address
break
}
for _, address := range metadata.DestinationAddresses {
if address.Is4() == packetDestination.IsIPv4() {
newDestination = address
break
}
}
if !newDestination.IsValid() {
if metadata.Source.IsIPv4() {
return nil, E.New("no IPv4 address found for domain: ", metadata.Destination.Fqdn)
if len(metadata.DestinationAddresses) == 0 {
r.logger.WarnContext(ctx, "pre-match: reject ", metadata.Network, " connection from ", metadata.Source.AddrString(), " to fake destination ", metadata.Destination.Fqdn, ": a resolve action is required before routing to outbound/", outbound.Type(), "[", outbound.Tag(), "]")
} else {
return nil, E.New("no IPv6 address found for domain: ", metadata.Destination.Fqdn)
r.logger.DebugContext(ctx, "pre-match: reject ", metadata.Network, " connection from ", metadata.Source.AddrString(), " to fake destination ", metadata.Destination.Fqdn, ": no resolved address for this address family")
}
return adapter.PreMatchResult{Action: adapter.PreMatchReject}
}
flowAction = flowOutbound.PreMatchFlow(metadata.Network, newDestination)
if flowAction != adapter.PreMatchFlow {
return adapter.PreMatchResult{Action: flowAction, Outbound: outbound}
}
result.Destination = netip.AddrPortFrom(newDestination, metadata.Destination.Port)
} else if metadata.Destination != packetDestination {
result.Destination = metadata.Destination.AddrPort()
}
r.logger.InfoContext(ctx, "pre-match: forward ", metadata.Network, " connection from ", metadata.Source.AddrString(), " to ", metadata.Destination.AddrString(), " via outbound/", outbound.Type(), "[", outbound.Tag(), "]")
metadataCopy := *metadata
result.NewTracker = func() tun.FlowTracker {
flowTrackers := make([]tun.FlowTracker, 0, len(r.trackers)+1)
flowTrackers = append(flowTrackers, newFlowLogger(ctx, r.logger, metadataCopy, outbound))
for _, tracker := range r.trackers {
flowTracker := tracker.RoutedFlow(ctx, metadataCopy, matchedRule, outbound)
if flowTracker != nil {
flowTrackers = append(flowTrackers, flowTracker)
}
}
metadata.Destination = M.Socksaddr{
Addr: newDestination,
if len(flowTrackers) == 1 {
return flowTrackers[0]
}
routeContext = ping.NewContextDestinationWriter(routeContext, metadata.OriginDestination.Addr)
var routeDestination tun.DirectRouteDestination
routeDestination, err = directRouteOutbound.NewDirectRouteConnection(metadata, routeContext, timeout)
if err != nil {
return nil, err
}
return ping.NewDestinationWriter(routeDestination, newDestination), nil
return multiFlowTracker(flowTrackers)
}
return directRouteOutbound.NewDirectRouteConnection(metadata, routeContext, timeout)
return result
}
func (r *Router) matchRule(
ctx context.Context, metadata *adapter.InboundContext, preMatch bool, supportBypass bool,
inputConn net.Conn, inputPacketConn N.PacketConn,
) (
selectedRule adapter.Rule, selectedRuleIndex int,
buffers []*buf.Buffer, packetBuffers []*N.PacketBuffer, fatalErr error,
) {
func (r *Router) prepareMatchMetadata(ctx context.Context, metadata *adapter.InboundContext) error {
r.searchProcessInfo(ctx, metadata)
if r.neighborResolver != nil && metadata.SourceMACAddress == nil && metadata.Source.Addr.IsValid() {
mac, macFound := r.neighborResolver.LookupMAC(metadata.Source.Addr)
if macFound {
metadata.SourceMACAddress = mac
}
hostname, hostnameFound := r.neighborResolver.LookupHostname(metadata.Source.Addr)
if hostnameFound {
metadata.SourceHostname = hostname
if macFound {
r.logger.InfoContext(ctx, "found neighbor: ", mac, ", hostname: ", hostname)
} else {
r.logger.InfoContext(ctx, "found neighbor hostname: ", hostname)
}
} else if macFound {
r.logger.InfoContext(ctx, "found neighbor: ", mac)
}
}
if metadata.Destination.Addr.IsValid() && r.dnsTransport.FakeIP() != nil && r.dnsTransport.FakeIP().Store().Contains(metadata.Destination.Addr) {
domain, loaded := r.dnsTransport.FakeIP().Store().Lookup(metadata.Destination.Addr)
if !loaded {
fatalErr = E.New("missing fakeip record, try enable `experimental.cache_file`")
return
return E.New("missing fakeip record, try enable `experimental.cache_file`")
}
if domain != "" {
metadata.OriginDestination = metadata.Destination
@@ -447,6 +588,20 @@ func (r *Router) matchRule(
} else if metadata.Destination.IsIPv6() {
metadata.IPVersion = 6
}
return nil
}
func (r *Router) matchRule(
ctx context.Context, metadata *adapter.InboundContext,
inputConn net.Conn, inputPacketConn N.PacketConn,
) (
selectedRule adapter.Rule, selectedRuleIndex int,
buffers []*buf.Buffer, packetBuffers []*N.PacketBuffer, fatalErr error,
) {
fatalErr = r.prepareMatchMetadata(ctx, metadata)
if fatalErr != nil {
return
}
match:
for currentRuleIndex, currentRule := range r.rules {
@@ -454,23 +609,11 @@ match:
if !currentRule.Match(metadata) {
continue
}
if !preMatch {
ruleDescription := currentRule.String()
if ruleDescription != "" {
r.logger.DebugContext(ctx, "match[", currentRuleIndex, "] ", currentRule, " => ", currentRule.Action())
} else {
r.logger.DebugContext(ctx, "match[", currentRuleIndex, "] => ", currentRule.Action())
}
ruleDescription := currentRule.String()
if ruleDescription != "" {
r.logger.DebugContext(ctx, "match[", currentRuleIndex, "] ", currentRule, " => ", currentRule.Action())
} else {
switch currentRule.Action().Type() {
case C.RuleActionTypeReject:
ruleDescription := currentRule.String()
if ruleDescription != "" {
r.logger.DebugContext(ctx, "pre-match[", currentRuleIndex, "] ", currentRule, " => ", currentRule.Action())
} else {
r.logger.DebugContext(ctx, "pre-match[", currentRuleIndex, "] => ", currentRule.Action())
}
}
r.logger.DebugContext(ctx, "match[", currentRuleIndex, "] => ", currentRule.Action())
}
var routeOptions *R.RuleActionRouteOptions
switch action := currentRule.Action().(type) {
@@ -478,6 +621,10 @@ match:
routeOptions = &action.RuleActionRouteOptions
case *R.RuleActionRouteOptions:
routeOptions = action
case *R.RuleActionBypass:
if action.Outbound != "" {
routeOptions = &action.RuleActionRouteOptions
}
}
if routeOptions != nil {
// TODO: add nat
@@ -485,20 +632,9 @@ match:
metadata.RouteOriginalDestination = metadata.Destination
}
if routeOptions.OverrideAddress.IsValid() {
metadata.Destination = M.Socksaddr{
Addr: routeOptions.OverrideAddress.Addr,
Port: metadata.Destination.Port,
Fqdn: routeOptions.OverrideAddress.Fqdn,
}
metadata.DestinationAddresses = nil
}
if routeOptions.OverridePort > 0 {
metadata.Destination = M.Socksaddr{
Addr: metadata.Destination.Addr,
Port: routeOptions.OverridePort,
Fqdn: metadata.Destination.Fqdn,
}
}
applyRouteOptionsOverride(metadata, routeOptions)
if routeOptions.OverrideGateway != nil && routeOptions.OverrideGateway.IsValid() {
metadata.Gateway = routeOptions.OverrideGateway
}
@@ -530,24 +666,22 @@ match:
if routeOptions.TLSRecordFragment {
metadata.TLSRecordFragment = true
}
if routeOptions.TLSSpoof != "" {
metadata.TLSSpoof = routeOptions.TLSSpoof
metadata.TLSSpoofMethod = routeOptions.TLSSpoofMethod
}
}
switch action := currentRule.Action().(type) {
case *R.RuleActionSniff:
if !preMatch {
newBuffer, newPacketBuffers, newErr := r.actionSniff(ctx, metadata, action, inputConn, inputPacketConn, buffers, packetBuffers)
if newBuffer != nil {
buffers = append(buffers, newBuffer)
} else if len(newPacketBuffers) > 0 {
packetBuffers = append(packetBuffers, newPacketBuffers...)
}
if newErr != nil {
fatalErr = newErr
return
}
} else if metadata.Network != N.NetworkICMP {
selectedRule = currentRule
selectedRuleIndex = currentRuleIndex
break match
newBuffer, newPacketBuffers, newErr := r.actionSniff(ctx, metadata, action, inputConn, inputPacketConn, buffers, packetBuffers)
if newBuffer != nil {
buffers = append(buffers, newBuffer)
} else if len(newPacketBuffers) > 0 {
packetBuffers = append(packetBuffers, newPacketBuffers...)
}
if newErr != nil {
fatalErr = newErr
return
}
case *R.RuleActionResolve:
fatalErr = r.actionResolve(ctx, metadata, action)
@@ -565,7 +699,7 @@ match:
}
if actionType == C.RuleActionTypeBypass {
bypassAction := currentRule.Action().(*R.RuleActionBypass)
if !supportBypass && bypassAction.Outbound == "" {
if bypassAction.Outbound == "" {
continue match
}
selectedRule = currentRule
@@ -654,15 +788,7 @@ func (r *Router) actionSniff(
if len(action.PacketSniffers) > 0 {
packetSniffers = action.PacketSniffers
} else {
packetSniffers = []sniff.PacketSniffer{
sniff.DomainNameQuery,
sniff.QUICClientHello,
sniff.STUNMessage,
sniff.UTP,
sniff.UDPTracker,
sniff.DTLSRecord,
sniff.NTP,
}
packetSniffers = defaultPacketSniffers
}
var err error
for _, packetBuffer := range inputPacketBuffers {
@@ -784,11 +910,13 @@ func (r *Router) actionResolve(ctx context.Context, metadata *adapter.InboundCon
}
}
addresses, err := r.dns.Lookup(adapter.WithContext(ctx, metadata), metadata.Destination.Fqdn, adapter.DNSQueryOptions{
Transport: transport,
Strategy: action.Strategy,
DisableCache: action.DisableCache,
RewriteTTL: action.RewriteTTL,
ClientSubnet: action.ClientSubnet,
Transport: transport,
Strategy: action.Strategy,
DisableCache: action.DisableCache,
DisableOptimisticCache: action.DisableOptimisticCache,
RewriteTTL: action.RewriteTTL,
Timeout: action.Timeout,
ClientSubnet: action.ClientSubnet,
})
if err != nil {
return err
+93 -25
View File
@@ -34,13 +34,18 @@ type Router struct {
connection adapter.ConnectionManager
network adapter.NetworkManager
defaultOutbound adapter.Outbound
httpClientManager adapter.HTTPClientManager
rules []adapter.Rule
final string
needFindProcess bool
needFindNeighbor bool
leaseFiles []string
ruleSets []adapter.RuleSet
ruleSetMap map[string]adapter.RuleSet
ruleSetUpdater *R.RuleSetUpdater
processSearcher process.Searcher
processCache freelru.Cache[processCacheKey, processCacheEntry]
processCache *freelru.Cache[processCacheKey, processCacheEntry]
neighborResolver adapter.NeighborResolver
pauseManager pause.Manager
trackers []adapter.ConnectionTracker
platformInterface adapter.PlatformInterface
@@ -57,10 +62,13 @@ func NewRouter(ctx context.Context, logFactory log.Factory, options option.Route
dnsTransport: service.FromContext[adapter.DNSTransportManager](ctx),
connection: service.FromContext[adapter.ConnectionManager](ctx),
network: service.FromContext[adapter.NetworkManager](ctx),
httpClientManager: service.FromContext[adapter.HTTPClientManager](ctx),
rules: make([]adapter.Rule, 0, len(options.Rules)),
final: options.Final,
ruleSetMap: make(map[string]adapter.RuleSet),
needFindProcess: hasRule(options.Rules, isProcessRule) || hasDNSRule(dnsOptions.Rules, isProcessDNSRule) || options.FindProcess,
needFindNeighbor: hasRule(options.Rules, isNeighborRule) || hasDNSRule(dnsOptions.Rules, isNeighborDNSRule) || hasLocalNeighborDNSServer(dnsOptions.Servers) || options.FindNeighbor,
leaseFiles: options.DHCPLeaseFiles,
pauseManager: service.FromContext[pause.Manager](ctx),
platformInterface: service.FromContext[adapter.PlatformInterface](ctx),
started: make(chan struct{}),
@@ -69,6 +77,10 @@ func NewRouter(ctx context.Context, logFactory log.Factory, options option.Route
func (r *Router) Initialize(rules []option.Rule, ruleSets []option.RuleSet) error {
for i, options := range rules {
err := R.ValidateNoNestedRuleActions(options)
if err != nil {
return E.Cause(err, "parse rule[", i, "]")
}
rule, err := R.NewRule(r.ctx, r.logger, options, false)
if err != nil {
return E.Cause(err, "parse rule[", i, "]")
@@ -76,15 +88,17 @@ func (r *Router) Initialize(rules []option.Rule, ruleSets []option.RuleSet) erro
r.rules = append(r.rules, rule)
}
for i, options := range ruleSets {
if _, exists := r.ruleSetMap[options.Tag]; exists {
return E.New("duplicate rule-set tag: ", options.Tag)
for _, tag := range options.Tag {
if _, exists := r.ruleSetMap[tag]; exists {
return E.New("duplicate rule-set tag: ", tag)
}
ruleSet, err := R.NewRuleSet(r.ctx, r.logger, tag, options)
if err != nil {
return E.Cause(err, "parse rule-set[", i, "]")
}
r.ruleSets = append(r.ruleSets, ruleSet)
r.ruleSetMap[tag] = ruleSet
}
ruleSet, err := R.NewRuleSet(r.ctx, r.logger, options)
if err != nil {
return E.Cause(err, "parse rule-set[", i, "]")
}
r.ruleSets = append(r.ruleSets, ruleSet)
r.ruleSetMap[options.Tag] = ruleSet
}
return nil
}
@@ -92,16 +106,46 @@ func (r *Router) Initialize(rules []option.Rule, ruleSets []option.RuleSet) erro
func (r *Router) Start(stage adapter.StartStage) error {
monitor := taskmonitor.New(r.logger, C.StartTimeout)
switch stage {
case adapter.StartStateInitialize:
if r.needFindNeighbor {
if r.platformInterface != nil && r.platformInterface.UsePlatformNeighborResolver() {
monitor.Start("initialize neighbor resolver")
resolver := newPlatformNeighborResolver(r.logger, r.platformInterface)
err := resolver.Start()
monitor.Finish()
if err != nil {
r.logger.Error(E.Cause(err, "start neighbor resolver"))
} else {
r.neighborResolver = resolver
}
} else {
monitor.Start("initialize neighbor resolver")
resolver, err := newNeighborResolver(r.logger, r.leaseFiles)
monitor.Finish()
if err != nil {
if err != os.ErrInvalid {
r.logger.Error(E.Cause(err, "create neighbor resolver"))
}
} else {
err = resolver.Start()
if err != nil {
r.logger.Error(E.Cause(err, "start neighbor resolver"))
} else {
r.neighborResolver = resolver
}
}
}
}
case adapter.StartStateStart:
var cacheContext *adapter.HTTPStartContext
var startContext *adapter.HTTPStartContext
if len(r.ruleSets) > 0 {
monitor.Start("initialize rule-set")
cacheContext = adapter.NewHTTPStartContext(r.ctx)
startContext = adapter.NewHTTPStartContext()
var ruleSetStartGroup task.Group
for i, ruleSet := range r.ruleSets {
ruleSetInPlace := ruleSet
ruleSetStartGroup.Append0(func(ctx context.Context) error {
err := ruleSetInPlace.StartContext(ctx, cacheContext)
err := ruleSetInPlace.StartContext(ctx, startContext)
if err != nil {
return E.Cause(err, "initialize rule-set[", i, "]")
}
@@ -116,9 +160,10 @@ func (r *Router) Start(stage adapter.StartStage) error {
return err
}
}
if cacheContext != nil {
cacheContext.Close()
if startContext != nil {
startContext.Close()
}
r.ruleSetUpdater = R.NewRuleSetUpdater(r.ctx, r.ruleSets)
r.network.Initialize(r.ruleSets)
needFindProcess := r.needFindProcess
for _, ruleSet := range r.ruleSets {
@@ -151,7 +196,7 @@ func (r *Router) Start(stage adapter.StartStage) error {
}
}
if r.processSearcher != nil {
processCache := common.Must1(freelru.NewSharded[processCacheKey, processCacheEntry](256, maphash.NewHasher[processCacheKey]().Hash32))
processCache := common.Must1(freelru.New[processCacheKey, processCacheEntry](256, maphash.NewHasher[processCacheKey]().Hash32, true))
processCache.SetLifetime(200 * time.Millisecond)
r.processCache = processCache
}
@@ -164,13 +209,8 @@ func (r *Router) Start(stage adapter.StartStage) error {
return E.Cause(err, "initialize rule[", i, "]")
}
}
for _, ruleSet := range r.ruleSets {
monitor.Start("post start rule_set[", ruleSet.Name(), "]")
err := ruleSet.PostStart()
monitor.Finish()
if err != nil {
return E.Cause(err, "post start rule_set[", ruleSet.Name(), "]")
}
if r.ruleSetUpdater != nil {
r.ruleSetUpdater.Start()
}
if r.final != "" {
defaultOutbound, loaded := r.outbound.Outbound(r.final)
@@ -195,6 +235,13 @@ func (r *Router) Start(stage adapter.StartStage) error {
func (r *Router) Close() error {
monitor := taskmonitor.New(r.logger, C.StopTimeout)
var err error
if r.neighborResolver != nil {
monitor.Start("close neighbor resolver")
err = E.Append(err, r.neighborResolver.Close(), func(closeErr error) error {
return E.Cause(closeErr, "close neighbor resolver")
})
monitor.Finish()
}
for i, rule := range r.rules {
monitor.Start("close rule[", i, "]")
err = E.Append(err, rule.Close(), func(err error) error {
@@ -202,6 +249,13 @@ func (r *Router) Close() error {
})
monitor.Finish()
}
if r.ruleSetUpdater != nil {
monitor.Start("close rule-set updater")
err = E.Append(err, r.ruleSetUpdater.Close(), func(err error) error {
return E.Cause(err, "close rule-set updater")
})
monitor.Finish()
}
for i, ruleSet := range r.ruleSets {
monitor.Start("close rule-set[", i, "]")
err = E.Append(err, ruleSet.Close(), func(err error) error {
@@ -236,7 +290,21 @@ func (r *Router) NeedFindProcess() bool {
return r.needFindProcess
}
func (r *Router) ResetNetwork() {
r.network.ResetNetwork()
r.dns.ResetNetwork()
func (r *Router) NeedFindNeighbor() bool {
return r.needFindNeighbor
}
func (r *Router) NeighborResolver() adapter.NeighborResolver {
return r.neighborResolver
}
func (r *Router) ResetNetwork() {
r.httpClientManager.ResetNetwork()
r.dns.ResetNetwork()
if r.processCache != nil {
r.processCache.Purge()
}
if r.processSearcher != nil {
r.processSearcher.ResetCache()
}
}
+9 -101
View File
@@ -1,7 +1,5 @@
package rule
import "github.com/sagernet/sing-box/adapter"
type ruleMatchState uint8
const (
@@ -11,108 +9,18 @@ const (
ruleMatchDestinationPort
)
type ruleMatchStateSet uint16
func singleRuleMatchState(state ruleMatchState) ruleMatchStateSet {
return 1 << state
type ruleGroupMatch struct {
required ruleMatchState
satisfied ruleMatchState
}
func emptyRuleMatchState() ruleMatchStateSet {
return singleRuleMatchState(0)
func (g ruleGroupMatch) done() bool {
return g.required&^g.satisfied == 0
}
func (s ruleMatchStateSet) isEmpty() bool {
return s == 0
}
func (s ruleMatchStateSet) contains(state ruleMatchState) bool {
return s&(1<<state) != 0
}
func (s ruleMatchStateSet) add(state ruleMatchState) ruleMatchStateSet {
return s | singleRuleMatchState(state)
}
func (s ruleMatchStateSet) merge(other ruleMatchStateSet) ruleMatchStateSet {
return s | other
}
func (s ruleMatchStateSet) combine(other ruleMatchStateSet) ruleMatchStateSet {
if s.isEmpty() || other.isEmpty() {
return 0
func (g ruleGroupMatch) mergeWith(other ruleGroupMatch) ruleGroupMatch {
return ruleGroupMatch{
required: g.required | other.required,
satisfied: g.satisfied | other.satisfied,
}
var combined ruleMatchStateSet
for left := range ruleMatchState(16) {
if !s.contains(left) {
continue
}
for right := range ruleMatchState(16) {
if !other.contains(right) {
continue
}
combined = combined.add(left | right)
}
}
return combined
}
func (s ruleMatchStateSet) withBase(base ruleMatchState) ruleMatchStateSet {
if s.isEmpty() {
return 0
}
var withBase ruleMatchStateSet
for state := range ruleMatchState(16) {
if !s.contains(state) {
continue
}
withBase = withBase.add(state | base)
}
return withBase
}
func (s ruleMatchStateSet) filter(allowed func(ruleMatchState) bool) ruleMatchStateSet {
var filtered ruleMatchStateSet
for state := range ruleMatchState(16) {
if !s.contains(state) {
continue
}
if allowed(state) {
filtered = filtered.add(state)
}
}
return filtered
}
type ruleStateMatcher interface {
matchStates(metadata *adapter.InboundContext) ruleMatchStateSet
}
type ruleStateMatcherWithBase interface {
matchStatesWithBase(metadata *adapter.InboundContext, base ruleMatchState) ruleMatchStateSet
}
func matchHeadlessRuleStatesWithBase(rule adapter.HeadlessRule, metadata *adapter.InboundContext, base ruleMatchState) ruleMatchStateSet {
if matcher, isStateMatcher := rule.(ruleStateMatcherWithBase); isStateMatcher {
return matcher.matchStatesWithBase(metadata, base)
}
if matcher, isStateMatcher := rule.(ruleStateMatcher); isStateMatcher {
return matcher.matchStates(metadata).withBase(base)
}
if rule.Match(metadata) {
return emptyRuleMatchState().withBase(base)
}
return 0
}
func matchRuleItemStatesWithBase(item RuleItem, metadata *adapter.InboundContext, base ruleMatchState) ruleMatchStateSet {
if matcher, isStateMatcher := item.(ruleStateMatcherWithBase); isStateMatcher {
return matcher.matchStatesWithBase(metadata, base)
}
if matcher, isStateMatcher := item.(ruleStateMatcher); isStateMatcher {
return matcher.matchStates(metadata).withBase(base)
}
if item.Match(metadata) {
return emptyRuleMatchState().withBase(base)
}
return 0
}
+89 -118
View File
@@ -18,7 +18,7 @@ type abstractDefaultRule struct {
destinationIPCIDRItems []RuleItem
destinationPortItems []RuleItem
allItems []RuleItem
ruleSetItem RuleItem
ruleSetItem *RuleSetItem
invert bool
action adapter.RuleAction
}
@@ -52,124 +52,99 @@ func (r *abstractDefaultRule) Close() error {
}
func (r *abstractDefaultRule) Match(metadata *adapter.InboundContext) bool {
return !r.matchStates(metadata).isEmpty()
if len(r.allItems) == 0 {
return true
}
matched := r.matchInner(metadata)
if r.invert {
if !matched {
metadata.DeferredIPCIDRMatchGroups = 0
return true
}
return metadata.DeferredIPCIDRMatchGroups != 0
}
return matched
}
func (r *abstractDefaultRule) matchInner(metadata *adapter.InboundContext) bool {
groups := r.evaluateGroups(metadata)
for _, item := range r.items {
if !item.Match(metadata) {
return false
}
}
var matched bool
if r.ruleSetItem != nil {
matched = r.ruleSetItem.matchWithOuterGroups(metadata, groups)
} else {
matched = groups.done()
}
if matched {
metadata.DeferredIPCIDRMatchGroups &^= uint8(groups.satisfied)
}
return matched
}
func (r *abstractDefaultRule) evaluateForMerge(metadata *adapter.InboundContext) (ruleGroupMatch, bool) {
groups := r.evaluateGroups(metadata)
for _, item := range r.items {
if !item.Match(metadata) {
return ruleGroupMatch{}, false
}
}
return groups, true
}
func (r *abstractDefaultRule) destinationIPCIDRMatchesSource(metadata *adapter.InboundContext) bool {
return !metadata.IgnoreDestinationIPCIDRMatch && metadata.IPCIDRMatchSource && len(r.destinationIPCIDRItems) > 0
return metadata.IPCIDRMatchSource && len(r.destinationIPCIDRItems) > 0
}
func (r *abstractDefaultRule) destinationIPCIDRMatchesDestination(metadata *adapter.InboundContext) bool {
return !metadata.IgnoreDestinationIPCIDRMatch && !metadata.IPCIDRMatchSource && len(r.destinationIPCIDRItems) > 0
}
func (r *abstractDefaultRule) requiresSourceAddressMatch(metadata *adapter.InboundContext) bool {
return len(r.sourceAddressItems) > 0 || r.destinationIPCIDRMatchesSource(metadata)
}
func (r *abstractDefaultRule) requiresDestinationAddressMatch(metadata *adapter.InboundContext) bool {
return len(r.destinationAddressItems) > 0 || r.destinationIPCIDRMatchesDestination(metadata)
}
func (r *abstractDefaultRule) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return r.matchStatesWithBase(metadata, 0)
}
func (r *abstractDefaultRule) matchStatesWithBase(metadata *adapter.InboundContext, inheritedBase ruleMatchState) ruleMatchStateSet {
if len(r.allItems) == 0 {
return emptyRuleMatchState().withBase(inheritedBase)
}
evaluationBase := inheritedBase
if r.invert {
evaluationBase = 0
}
baseState := evaluationBase
func (r *abstractDefaultRule) evaluateGroups(metadata *adapter.InboundContext) ruleGroupMatch {
var groups ruleGroupMatch
if len(r.sourceAddressItems) > 0 {
metadata.DidMatch = true
groups.required |= ruleMatchSourceAddress
if matchAnyItem(r.sourceAddressItems, metadata) {
baseState |= ruleMatchSourceAddress
groups.satisfied |= ruleMatchSourceAddress
}
}
if r.destinationIPCIDRMatchesSource(metadata) && !baseState.has(ruleMatchSourceAddress) {
metadata.DidMatch = true
if matchAnyItem(r.destinationIPCIDRItems, metadata) {
baseState |= ruleMatchSourceAddress
if r.destinationIPCIDRMatchesSource(metadata) {
groups.required |= ruleMatchSourceAddress
if !groups.satisfied.has(ruleMatchSourceAddress) && matchAnyItem(r.destinationIPCIDRItems, metadata) {
groups.satisfied |= ruleMatchSourceAddress
}
} else if r.destinationIPCIDRMatchesSource(metadata) {
metadata.DidMatch = true
}
if len(r.sourcePortItems) > 0 {
metadata.DidMatch = true
groups.required |= ruleMatchSourcePort
if matchAnyItem(r.sourcePortItems, metadata) {
baseState |= ruleMatchSourcePort
groups.satisfied |= ruleMatchSourcePort
}
}
if len(r.destinationAddressItems) > 0 {
metadata.DidMatch = true
groups.required |= ruleMatchDestinationAddress
if matchAnyItem(r.destinationAddressItems, metadata) {
baseState |= ruleMatchDestinationAddress
groups.satisfied |= ruleMatchDestinationAddress
}
}
if r.destinationIPCIDRMatchesDestination(metadata) && !baseState.has(ruleMatchDestinationAddress) {
metadata.DidMatch = true
if matchAnyItem(r.destinationIPCIDRItems, metadata) {
baseState |= ruleMatchDestinationAddress
if r.destinationIPCIDRMatchesDestination(metadata) {
groups.required |= ruleMatchDestinationAddress
if !groups.satisfied.has(ruleMatchDestinationAddress) && matchAnyItem(r.destinationIPCIDRItems, metadata) {
groups.satisfied |= ruleMatchDestinationAddress
}
} else if r.destinationIPCIDRMatchesDestination(metadata) {
metadata.DidMatch = true
}
if len(r.destinationPortItems) > 0 {
metadata.DidMatch = true
groups.required |= ruleMatchDestinationPort
if matchAnyItem(r.destinationPortItems, metadata) {
baseState |= ruleMatchDestinationPort
groups.satisfied |= ruleMatchDestinationPort
}
}
for _, item := range r.items {
metadata.DidMatch = true
if !item.Match(metadata) {
return r.invertedFailure(inheritedBase)
}
if metadata.IgnoreDestinationIPCIDRMatch && !metadata.IPCIDRMatchSource && len(r.destinationIPCIDRItems) > 0 && len(r.destinationAddressItems) == 0 {
metadata.DeferredIPCIDRMatchGroups |= uint8(ruleMatchDestinationAddress)
}
var stateSet ruleMatchStateSet
if r.ruleSetItem != nil {
metadata.DidMatch = true
stateSet = matchRuleItemStatesWithBase(r.ruleSetItem, metadata, baseState)
} else {
stateSet = singleRuleMatchState(baseState)
}
stateSet = stateSet.filter(func(state ruleMatchState) bool {
if r.requiresSourceAddressMatch(metadata) && !state.has(ruleMatchSourceAddress) {
return false
}
if len(r.sourcePortItems) > 0 && !state.has(ruleMatchSourcePort) {
return false
}
if r.requiresDestinationAddressMatch(metadata) && !state.has(ruleMatchDestinationAddress) {
return false
}
if len(r.destinationPortItems) > 0 && !state.has(ruleMatchDestinationPort) {
return false
}
return true
})
if stateSet.isEmpty() {
return r.invertedFailure(inheritedBase)
}
if r.invert {
// DNS pre-lookup defers destination address-limit checks until the response phase.
if metadata.IgnoreDestinationIPCIDRMatch && stateSet == emptyRuleMatchState() && !metadata.DidMatch && len(r.destinationIPCIDRItems) > 0 {
return emptyRuleMatchState().withBase(inheritedBase)
}
return 0
}
return stateSet
}
func (r *abstractDefaultRule) invertedFailure(base ruleMatchState) ruleMatchStateSet {
if r.invert {
return emptyRuleMatchState().withBase(base)
}
return 0
return groups
}
func (r *abstractDefaultRule) Action() adapter.RuleAction {
@@ -227,50 +202,46 @@ func (r *abstractLogicalRule) Close() error {
}
func (r *abstractLogicalRule) Match(metadata *adapter.InboundContext) bool {
return !r.matchStates(metadata).isEmpty()
}
func (r *abstractLogicalRule) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return r.matchStatesWithBase(metadata, 0)
}
func (r *abstractLogicalRule) matchStatesWithBase(metadata *adapter.InboundContext, base ruleMatchState) ruleMatchStateSet {
evaluationBase := base
if r.invert {
evaluationBase = 0
}
var stateSet ruleMatchStateSet
var (
matched bool
deferredGroups uint8
)
if r.mode == C.LogicalTypeAnd {
stateSet = emptyRuleMatchState().withBase(evaluationBase)
matched = true
for _, rule := range r.rules {
nestedMetadata := *metadata
nestedMetadata.ResetRuleCache()
nestedStateSet := matchHeadlessRuleStatesWithBase(rule, &nestedMetadata, evaluationBase)
if nestedStateSet.isEmpty() {
if r.invert {
return emptyRuleMatchState().withBase(base)
}
return 0
if !rule.Match(&nestedMetadata) {
matched = false
deferredGroups = 0
break
}
stateSet = stateSet.combine(nestedStateSet)
deferredGroups |= nestedMetadata.DeferredIPCIDRMatchGroups
}
} else {
for _, rule := range r.rules {
nestedMetadata := *metadata
nestedMetadata.ResetRuleCache()
stateSet = stateSet.merge(matchHeadlessRuleStatesWithBase(rule, &nestedMetadata, evaluationBase))
}
if stateSet.isEmpty() {
if r.invert {
return emptyRuleMatchState().withBase(base)
if rule.Match(&nestedMetadata) {
matched = true
if nestedMetadata.DeferredIPCIDRMatchGroups == 0 {
deferredGroups = 0
break
}
deferredGroups |= nestedMetadata.DeferredIPCIDRMatchGroups
}
return 0
}
}
if matched {
metadata.DeferredIPCIDRMatchGroups |= deferredGroups
}
if r.invert {
return 0
if !matched {
return true
}
return deferredGroups != 0
}
return stateSet
return matched
}
func (r *abstractLogicalRule) Action() adapter.RuleAction {
-4
View File
@@ -24,10 +24,6 @@ func (f *fakeRuleSet) StartContext(context.Context, *adapter.HTTPStartContext) e
return nil
}
func (f *fakeRuleSet) PostStart() error {
return nil
}
func (f *fakeRuleSet) Metadata() adapter.RuleSetMetadata {
return adapter.RuleSetMetadata{}
}
+181 -78
View File
@@ -11,9 +11,9 @@ import (
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/dialer"
"github.com/sagernet/sing-box/common/sniff"
"github.com/sagernet/sing-box/common/tlsspoof"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing/common"
E "github.com/sagernet/sing/common/exceptions"
F "github.com/sagernet/sing/common/format"
@@ -24,61 +24,65 @@ import (
"github.com/miekg/dns"
)
func newRuleActionRouteOptions(options option.RawRouteOptionsActionOptions) (RuleActionRouteOptions, error) {
spoof, spoofMethod, err := tlsspoof.ParseOptions(options.TLSSpoof, options.TLSSpoofMethod)
if err != nil {
return RuleActionRouteOptions{}, err
}
var overrideGateway *netip.Addr
if options.OverrideGateway != "" {
parsed := M.ParseAddr(options.OverrideGateway)
overrideGateway = &parsed
}
return RuleActionRouteOptions{
OverrideAddress: M.ParseSocksaddrHostPort(options.OverrideAddress, 0),
OverridePort: options.OverridePort,
OverrideGateway: overrideGateway,
NetworkStrategy: (*C.NetworkStrategy)(options.NetworkStrategy),
FallbackDelay: time.Duration(options.FallbackDelay),
UDPDisableDomainUnmapping: options.UDPDisableDomainUnmapping,
UDPConnect: options.UDPConnect,
UDPTimeout: time.Duration(options.UDPTimeout),
TLSFragment: options.TLSFragment,
TLSFragmentFallbackDelay: time.Duration(options.TLSFragmentFallbackDelay),
TLSRecordFragment: options.TLSRecordFragment,
TLSSpoof: spoof,
TLSSpoofMethod: spoofMethod,
}, nil
}
func NewRuleAction(ctx context.Context, logger logger.ContextLogger, action option.RuleAction) (adapter.RuleAction, error) {
switch action.Action {
case "":
return nil, nil
case C.RuleActionTypeRoute:
var overrideGateway *netip.Addr
if action.RouteOptions.OverrideGateway != "" {
parsed := M.ParseAddr(action.RouteOptions.OverrideGateway)
overrideGateway = &parsed
routeOptions, err := newRuleActionRouteOptions(action.RouteOptions.RawRouteOptionsActionOptions)
if err != nil {
return nil, err
}
return &RuleActionRoute{
Outbound: action.RouteOptions.Outbound,
RuleActionRouteOptions: RuleActionRouteOptions{
OverrideAddress: M.ParseSocksaddrHostPort(action.RouteOptions.OverrideAddress, 0),
OverridePort: action.RouteOptions.OverridePort,
OverrideGateway: overrideGateway,
NetworkStrategy: (*C.NetworkStrategy)(action.RouteOptions.NetworkStrategy),
FallbackDelay: time.Duration(action.RouteOptions.FallbackDelay),
UDPDisableDomainUnmapping: action.RouteOptions.UDPDisableDomainUnmapping,
UDPConnect: action.RouteOptions.UDPConnect,
TLSFragment: action.RouteOptions.TLSFragment,
TLSFragmentFallbackDelay: time.Duration(action.RouteOptions.TLSFragmentFallbackDelay),
TLSRecordFragment: action.RouteOptions.TLSRecordFragment,
},
Outbound: action.RouteOptions.Outbound,
RuleActionRouteOptions: routeOptions,
}, nil
case C.RuleActionTypeRouteOptions:
return &RuleActionRouteOptions{
OverrideAddress: M.ParseSocksaddrHostPort(action.RouteOptionsOptions.OverrideAddress, 0),
OverridePort: action.RouteOptionsOptions.OverridePort,
NetworkStrategy: (*C.NetworkStrategy)(action.RouteOptionsOptions.NetworkStrategy),
FallbackDelay: time.Duration(action.RouteOptionsOptions.FallbackDelay),
UDPDisableDomainUnmapping: action.RouteOptionsOptions.UDPDisableDomainUnmapping,
UDPConnect: action.RouteOptionsOptions.UDPConnect,
UDPTimeout: time.Duration(action.RouteOptionsOptions.UDPTimeout),
TLSFragment: action.RouteOptionsOptions.TLSFragment,
TLSFragmentFallbackDelay: time.Duration(action.RouteOptionsOptions.TLSFragmentFallbackDelay),
TLSRecordFragment: action.RouteOptionsOptions.TLSRecordFragment,
}, nil
routeOptions, err := newRuleActionRouteOptions(option.RawRouteOptionsActionOptions(action.RouteOptionsOptions))
if err != nil {
return nil, err
}
return &routeOptions, nil
case C.RuleActionTypeBypass:
routeOptions, err := newRuleActionRouteOptions(action.BypassOptions.RawRouteOptionsActionOptions)
if err != nil {
return nil, err
}
return &RuleActionBypass{
Outbound: action.BypassOptions.Outbound,
RuleActionRouteOptions: RuleActionRouteOptions{
OverrideAddress: M.ParseSocksaddrHostPort(action.BypassOptions.OverrideAddress, 0),
OverridePort: action.BypassOptions.OverridePort,
NetworkStrategy: (*C.NetworkStrategy)(action.BypassOptions.NetworkStrategy),
FallbackDelay: time.Duration(action.BypassOptions.FallbackDelay),
UDPDisableDomainUnmapping: action.BypassOptions.UDPDisableDomainUnmapping,
UDPConnect: action.BypassOptions.UDPConnect,
TLSFragment: action.BypassOptions.TLSFragment,
TLSFragmentFallbackDelay: time.Duration(action.BypassOptions.TLSFragmentFallbackDelay),
TLSRecordFragment: action.BypassOptions.TLSRecordFragment,
},
Outbound: action.BypassOptions.Outbound,
RuleActionRouteOptions: routeOptions,
}, nil
case C.RuleActionTypeDirect:
directDialer, err := dialer.New(ctx, option.DialerOptions(action.DirectOptions), false)
directDialer, err := dialer.New(ctx, option.DialerOptions{
AbstractDialerOptions: action.DirectOptions.AbstractDialerOptions,
}, false)
if err != nil {
return nil, err
}
@@ -113,11 +117,13 @@ func NewRuleAction(ctx context.Context, logger logger.ContextLogger, action opti
return sniffAction, sniffAction.build()
case C.RuleActionTypeResolve:
return &RuleActionResolve{
Server: action.ResolveOptions.Server,
Strategy: C.DomainStrategy(action.ResolveOptions.Strategy),
DisableCache: action.ResolveOptions.DisableCache,
RewriteTTL: action.ResolveOptions.RewriteTTL,
ClientSubnet: action.ResolveOptions.ClientSubnet.Build(netip.Prefix{}),
Server: action.ResolveOptions.Server,
Timeout: time.Duration(action.ResolveOptions.Timeout),
Strategy: C.DomainStrategy(action.ResolveOptions.Strategy),
DisableCache: action.ResolveOptions.DisableCache,
DisableOptimisticCache: action.ResolveOptions.DisableOptimisticCache,
RewriteTTL: action.ResolveOptions.RewriteTTL,
ClientSubnet: action.ResolveOptions.ClientSubnet.Build(netip.Prefix{}),
}, nil
default:
panic(F.ToString("unknown rule action: ", action.Action))
@@ -130,20 +136,44 @@ func NewDNSRuleAction(logger logger.ContextLogger, action option.DNSRuleAction)
return nil
case C.RuleActionTypeRoute:
return &RuleActionDNSRoute{
Server: action.RouteOptions.Server,
Server: action.RouteOptions.Server,
Speculative: action.RouteOptions.Speculative,
RuleActionDNSRouteOptions: RuleActionDNSRouteOptions{
Strategy: C.DomainStrategy(action.RouteOptions.Strategy),
DisableCache: action.RouteOptions.DisableCache,
RewriteTTL: action.RouteOptions.RewriteTTL,
ClientSubnet: netip.Prefix(common.PtrValueOrDefault(action.RouteOptions.ClientSubnet)),
Strategy: C.DomainStrategy(action.RouteOptions.Strategy),
Timeout: time.Duration(action.RouteOptions.Timeout),
DisableCache: action.RouteOptions.DisableCache,
DisableOptimisticCache: action.RouteOptions.DisableOptimisticCache,
RewriteTTL: action.RouteOptions.RewriteTTL,
ClientSubnet: netip.Prefix(common.PtrValueOrDefault(action.RouteOptions.ClientSubnet)),
RemoveClientSubnet: action.RouteOptions.RemoveClientSubnet,
},
}
case C.RuleActionTypeEvaluate:
return &RuleActionEvaluate{
Server: action.EvaluateOptions.Server,
Tag: action.EvaluateOptions.Tag,
Speculative: action.EvaluateOptions.Speculative,
RuleActionDNSRouteOptions: RuleActionDNSRouteOptions{
Strategy: C.DomainStrategy(action.EvaluateOptions.Strategy),
Timeout: time.Duration(action.EvaluateOptions.Timeout),
DisableCache: action.EvaluateOptions.DisableCache,
DisableOptimisticCache: action.EvaluateOptions.DisableOptimisticCache,
RewriteTTL: action.EvaluateOptions.RewriteTTL,
ClientSubnet: netip.Prefix(common.PtrValueOrDefault(action.EvaluateOptions.ClientSubnet)),
RemoveClientSubnet: action.EvaluateOptions.RemoveClientSubnet,
},
}
case C.RuleActionTypeRespond:
return &RuleActionRespond{}
case C.RuleActionTypeRouteOptions:
return &RuleActionDNSRouteOptions{
Strategy: C.DomainStrategy(action.RouteOptionsOptions.Strategy),
DisableCache: action.RouteOptionsOptions.DisableCache,
RewriteTTL: action.RouteOptionsOptions.RewriteTTL,
ClientSubnet: netip.Prefix(common.PtrValueOrDefault(action.RouteOptionsOptions.ClientSubnet)),
Strategy: C.DomainStrategy(action.RouteOptionsOptions.Strategy),
Timeout: time.Duration(action.RouteOptionsOptions.Timeout),
DisableCache: action.RouteOptionsOptions.DisableCache,
DisableOptimisticCache: action.RouteOptionsOptions.DisableOptimisticCache,
RewriteTTL: action.RouteOptionsOptions.RewriteTTL,
ClientSubnet: netip.Prefix(common.PtrValueOrDefault(action.RouteOptionsOptions.ClientSubnet)),
RemoveClientSubnet: action.RouteOptionsOptions.RemoveClientSubnet,
}
case C.RuleActionTypeReject:
return &RuleActionReject{
@@ -212,6 +242,8 @@ type RuleActionRouteOptions struct {
TLSFragment bool
TLSFragmentFallbackDelay time.Duration
TLSRecordFragment bool
TLSSpoof string
TLSSpoofMethod tlsspoof.Method
}
func (r *RuleActionRouteOptions) Type() string {
@@ -240,7 +272,7 @@ func (r *RuleActionRouteOptions) Descriptions() []string {
descriptions = append(descriptions, F.ToString("network-type=", strings.Join(common.Map(r.NetworkType, C.InterfaceType.String), ",")))
}
if r.FallbackNetworkType != nil {
descriptions = append(descriptions, F.ToString("fallback-network-type="+strings.Join(common.Map(r.NetworkType, C.InterfaceType.String), ",")))
descriptions = append(descriptions, F.ToString("fallback-network-type=", strings.Join(common.Map(r.FallbackNetworkType, C.InterfaceType.String), ",")))
}
if r.FallbackDelay > 0 {
descriptions = append(descriptions, F.ToString("fallback-delay=", r.FallbackDelay.String()))
@@ -263,11 +295,16 @@ func (r *RuleActionRouteOptions) Descriptions() []string {
if r.TLSRecordFragment {
descriptions = append(descriptions, "tls-record-fragment")
}
if r.TLSSpoof != "" {
descriptions = append(descriptions, F.ToString("tls-spoof=", r.TLSSpoof))
descriptions = append(descriptions, F.ToString("tls-spoof-method=", r.TLSSpoofMethod.String()))
}
return descriptions
}
type RuleActionDNSRoute struct {
Server string
Server string
Speculative bool
RuleActionDNSRouteOptions
}
@@ -276,25 +313,69 @@ func (r *RuleActionDNSRoute) Type() string {
}
func (r *RuleActionDNSRoute) String() string {
return formatDNSRouteAction("route", r.Server, r.Speculative, r.RuleActionDNSRouteOptions)
}
type RuleActionEvaluate struct {
Server string
Tag string
Speculative bool
RuleActionDNSRouteOptions
}
func (r *RuleActionEvaluate) Type() string {
return C.RuleActionTypeEvaluate
}
func (r *RuleActionEvaluate) String() string {
return formatDNSRouteAction("evaluate", r.Server, r.Speculative, r.RuleActionDNSRouteOptions)
}
type RuleActionRespond struct{}
func (r *RuleActionRespond) Type() string {
return C.RuleActionTypeRespond
}
func (r *RuleActionRespond) String() string {
return "respond"
}
func formatDNSRouteAction(action string, server string, speculative bool, options RuleActionDNSRouteOptions) string {
var descriptions []string
descriptions = append(descriptions, r.Server)
if r.DisableCache {
descriptions = append(descriptions, server)
if speculative {
descriptions = append(descriptions, "speculative")
}
if options.DisableCache {
descriptions = append(descriptions, "disable-cache")
}
if r.RewriteTTL != nil {
descriptions = append(descriptions, F.ToString("rewrite-ttl=", *r.RewriteTTL))
if options.DisableOptimisticCache {
descriptions = append(descriptions, "disable-optimistic-cache")
}
if r.ClientSubnet.IsValid() {
descriptions = append(descriptions, F.ToString("client-subnet=", r.ClientSubnet))
if options.RewriteTTL != nil {
descriptions = append(descriptions, F.ToString("rewrite-ttl=", *options.RewriteTTL))
}
return F.ToString("route(", strings.Join(descriptions, ","), ")")
if options.Timeout > 0 {
descriptions = append(descriptions, F.ToString("timeout=", options.Timeout.String()))
}
if options.ClientSubnet.IsValid() {
descriptions = append(descriptions, F.ToString("client-subnet=", options.ClientSubnet))
}
if options.RemoveClientSubnet {
descriptions = append(descriptions, "remove-client-subnet")
}
return F.ToString(action, "(", strings.Join(descriptions, ","), ")")
}
type RuleActionDNSRouteOptions struct {
Strategy C.DomainStrategy
DisableCache bool
RewriteTTL *uint32
ClientSubnet netip.Prefix
Strategy C.DomainStrategy
Timeout time.Duration
DisableCache bool
DisableOptimisticCache bool
RewriteTTL *uint32
ClientSubnet netip.Prefix
RemoveClientSubnet bool
}
func (r *RuleActionDNSRouteOptions) Type() string {
@@ -306,12 +387,21 @@ func (r *RuleActionDNSRouteOptions) String() string {
if r.DisableCache {
descriptions = append(descriptions, "disable-cache")
}
if r.DisableOptimisticCache {
descriptions = append(descriptions, "disable-optimistic-cache")
}
if r.RewriteTTL != nil {
descriptions = append(descriptions, F.ToString("rewrite-ttl=", *r.RewriteTTL))
}
if r.Timeout > 0 {
descriptions = append(descriptions, F.ToString("timeout=", r.Timeout.String()))
}
if r.ClientSubnet.IsValid() {
descriptions = append(descriptions, F.ToString("client-subnet=", r.ClientSubnet))
}
if r.RemoveClientSubnet {
descriptions = append(descriptions, "remove-client-subnet")
}
return F.ToString("route-options(", strings.Join(descriptions, ","), ")")
}
@@ -328,6 +418,11 @@ func (r *RuleActionDirect) String() string {
return "direct" + r.description
}
var (
ErrReset = E.New("connection reset")
ErrDrop = E.New("packet dropped")
)
type RejectedError struct {
Cause error
}
@@ -385,9 +480,9 @@ func (r *RuleActionReject) Error(ctx context.Context) error {
var returnErr error
switch r.Method {
case C.RuleActionRejectMethodDefault:
returnErr = &RejectedError{tun.ErrReset}
returnErr = &RejectedError{ErrReset}
case C.RuleActionRejectMethodDrop:
return &RejectedError{tun.ErrDrop}
return &RejectedError{ErrDrop}
case C.RuleActionRejectMethodReply:
return nil
default:
@@ -407,7 +502,7 @@ func (r *RuleActionReject) Error(ctx context.Context) error {
if ctx != nil {
r.logger.DebugContext(ctx, "dropped due to flooding")
}
return &RejectedError{tun.ErrDrop}
return &RejectedError{ErrDrop}
}
return returnErr
}
@@ -481,11 +576,13 @@ func (r *RuleActionSniff) String() string {
}
type RuleActionResolve struct {
Server string
Strategy C.DomainStrategy
DisableCache bool
RewriteTTL *uint32
ClientSubnet netip.Prefix
Server string
Timeout time.Duration
Strategy C.DomainStrategy
DisableCache bool
DisableOptimisticCache bool
RewriteTTL *uint32
ClientSubnet netip.Prefix
}
func (r *RuleActionResolve) Type() string {
@@ -503,9 +600,15 @@ func (r *RuleActionResolve) String() string {
if r.DisableCache {
options = append(options, "disable_cache")
}
if r.DisableOptimisticCache {
options = append(options, "disable_optimistic_cache")
}
if r.RewriteTTL != nil {
options = append(options, F.ToString("rewrite_ttl=", *r.RewriteTTL))
}
if r.Timeout > 0 {
options = append(options, F.ToString("timeout=", r.Timeout.String()))
}
if r.ClientSubnet.IsValid() {
options = append(options, F.ToString("client_subnet=", r.ClientSubnet))
}
+22 -8
View File
@@ -47,10 +47,6 @@ type DefaultRule struct {
abstractDefaultRule
}
func (r *DefaultRule) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return r.abstractDefaultRule.matchStates(metadata)
}
type RuleItem interface {
Match(metadata *adapter.InboundContext) bool
String() string
@@ -209,6 +205,14 @@ func NewDefaultRule(ctx context.Context, logger log.ContextLogger, options optio
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.PackageNameRegex) > 0 {
item, err := NewPackageNameRegexItem(options.PackageNameRegex)
if err != nil {
return nil, E.Cause(err, "package_name_regex")
}
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.User) > 0 {
item := NewUserItem(options.User)
rule.items = append(rule.items, item)
@@ -264,6 +268,16 @@ func NewDefaultRule(ctx context.Context, logger log.ContextLogger, options optio
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.SourceMACAddress) > 0 {
item := NewSourceMACAddressItem(options.SourceMACAddress)
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.SourceHostname) > 0 {
item := NewSourceHostnameItem(options.SourceHostname)
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.PreferredBy) > 0 {
item := NewPreferredByItem(ctx, options.PreferredBy)
rule.items = append(rule.items, item)
@@ -291,10 +305,6 @@ type LogicalRule struct {
abstractLogicalRule
}
func (r *LogicalRule) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return r.abstractLogicalRule.matchStates(metadata)
}
func NewLogicalRule(ctx context.Context, logger log.ContextLogger, options option.LogicalRule) (*LogicalRule, error) {
action, err := NewRuleAction(ctx, logger, options.RuleAction)
if err != nil {
@@ -316,6 +326,10 @@ func NewLogicalRule(ctx context.Context, logger log.ContextLogger, options optio
return nil, E.New("unknown logical mode: ", options.Mode)
}
for i, subOptions := range options.Rules {
err = validateNoNestedRuleActions(subOptions, true)
if err != nil {
return nil, E.Cause(err, "sub rule[", i, "]")
}
subRule, err := NewRule(ctx, logger, subOptions, false)
if err != nil {
return nil, E.Cause(err, "sub rule[", i, "]")
+227 -44
View File
@@ -5,58 +5,117 @@ import (
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/experimental/deprecated"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/service"
"github.com/miekg/dns"
)
func NewDNSRule(ctx context.Context, logger log.ContextLogger, options option.DNSRule, checkServer bool) (adapter.DNSRule, error) {
func NewDNSRule(ctx context.Context, logger log.ContextLogger, options option.DNSRule, checkServer bool, legacyDNSMode bool) (adapter.DNSRule, error) {
switch options.Type {
case "", C.RuleTypeDefault:
if !options.DefaultOptions.IsValid() {
return nil, E.New("missing conditions")
}
if !checkServer && options.DefaultOptions.Action == C.RuleActionTypeEvaluate {
return nil, E.New(options.DefaultOptions.Action, " is only allowed on top-level DNS rules")
}
err := validateDNSRuleAction(options.DefaultOptions.DNSRuleAction)
if err != nil {
return nil, err
}
if options.DefaultOptions.Race && !options.DefaultOptions.MatchResponse.IsEnabled() {
return nil, E.New("`race` requires `match_response`")
}
switch options.DefaultOptions.Action {
case "", C.RuleActionTypeRoute:
if options.DefaultOptions.RouteOptions.Server == "" && checkServer {
return nil, E.New("missing server field")
}
case C.RuleActionTypeEvaluate:
if options.DefaultOptions.EvaluateOptions.Server == "" && checkServer {
return nil, E.New("missing server field")
}
}
return NewDefaultDNSRule(ctx, logger, options.DefaultOptions)
return NewDefaultDNSRule(ctx, logger, options.DefaultOptions, legacyDNSMode)
case C.RuleTypeLogical:
if !options.LogicalOptions.IsValid() {
return nil, E.New("missing conditions")
}
if !checkServer && options.LogicalOptions.Action == C.RuleActionTypeEvaluate {
return nil, E.New(options.LogicalOptions.Action, " is only allowed on top-level DNS rules")
}
err := validateDNSRuleAction(options.LogicalOptions.DNSRuleAction)
if err != nil {
return nil, err
}
switch options.LogicalOptions.Action {
case "", C.RuleActionTypeRoute:
if options.LogicalOptions.RouteOptions.Server == "" && checkServer {
return nil, E.New("missing server field")
}
case C.RuleActionTypeEvaluate:
if options.LogicalOptions.EvaluateOptions.Server == "" && checkServer {
return nil, E.New("missing server field")
}
}
return NewLogicalDNSRule(ctx, logger, options.LogicalOptions)
return NewLogicalDNSRule(ctx, logger, options.LogicalOptions, legacyDNSMode)
default:
return nil, E.New("unknown rule type: ", options.Type)
}
}
func validateDNSRuleAction(action option.DNSRuleAction) error {
if action.Action == C.RuleActionTypeReject && action.RejectOptions.Method == C.RuleActionRejectMethodReply {
return E.New("reject method `reply` is not supported for DNS rules")
}
var routeOptions option.AbstractDNSRouteActionOptions
switch action.Action {
case "", C.RuleActionTypeRoute:
routeOptions = action.RouteOptions.AbstractDNSRouteActionOptions
case C.RuleActionTypeEvaluate:
routeOptions = action.EvaluateOptions.AbstractDNSRouteActionOptions
case C.RuleActionTypeRouteOptions:
routeOptions = option.AbstractDNSRouteActionOptions(action.RouteOptionsOptions)
}
if routeOptions.RemoveClientSubnet && routeOptions.ClientSubnet != nil {
return E.New("`client_subnet` and `remove_client_subnet` are mutually exclusive")
}
if action.Race {
switch action.Action {
case "", C.RuleActionTypeRoute, C.RuleActionTypeRespond, C.RuleActionTypeReject, C.RuleActionTypePredefined:
default:
return E.New("`race` requires a final action")
}
if action.RouteOptions.Speculative {
return E.New("`race` and `speculative` cannot be combined on the same rule")
}
}
return nil
}
var _ adapter.DNSRule = (*DefaultDNSRule)(nil)
type DefaultDNSRule struct {
abstractDefaultRule
matchResponse bool
matchResponseTag string
race bool
}
func (r *DefaultDNSRule) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return r.abstractDefaultRule.matchStates(metadata)
}
func NewDefaultDNSRule(ctx context.Context, logger log.ContextLogger, options option.DefaultDNSRule) (*DefaultDNSRule, error) {
func NewDefaultDNSRule(ctx context.Context, logger log.ContextLogger, options option.DefaultDNSRule, legacyDNSMode bool) (*DefaultDNSRule, error) {
rule := &DefaultDNSRule{
abstractDefaultRule: abstractDefaultRule{
invert: options.Invert,
action: NewDNSRuleAction(logger, options.DNSRuleAction),
},
matchResponse: options.MatchResponse.IsEnabled(),
matchResponseTag: options.MatchResponse.ResponseTag(),
race: options.Race,
}
if len(options.Inbound) > 0 {
item := NewInboundRule(options.Inbound)
@@ -80,6 +139,16 @@ func NewDefaultDNSRule(ctx context.Context, logger log.ContextLogger, options op
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.QueryClientSubnet) > 0 {
item := NewQueryClientSubnetItem(options.QueryClientSubnet)
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if options.QueryDNSSEC {
item := NewQueryDNSSECItem()
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.Network) > 0 {
item := NewNetworkItem(options.Network)
rule.items = append(rule.items, item)
@@ -116,7 +185,7 @@ func NewDefaultDNSRule(ctx context.Context, logger log.ContextLogger, options op
rule.destinationAddressItems = append(rule.destinationAddressItems, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.Geosite) > 0 {
if len(options.Geosite) > 0 { //nolint:staticcheck
return nil, E.New("geosite database is deprecated in sing-box 1.8.0 and removed in sing-box 1.12.0")
}
if len(options.SourceGeoIP) > 0 {
@@ -156,6 +225,26 @@ func NewDefaultDNSRule(ctx context.Context, logger log.ContextLogger, options op
rule.destinationIPCIDRItems = append(rule.destinationIPCIDRItems, item)
rule.allItems = append(rule.allItems, item)
}
if options.ResponseRcode != nil {
item := NewDNSResponseRCodeItem(int(*options.ResponseRcode))
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.ResponseAnswer) > 0 {
item := NewDNSResponseRecordItem("response_answer", options.ResponseAnswer, dnsResponseAnswers)
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.ResponseNs) > 0 {
item := NewDNSResponseRecordItem("response_ns", options.ResponseNs, dnsResponseNS)
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.ResponseExtra) > 0 {
item := NewDNSResponseRecordItem("response_extra", options.ResponseExtra, dnsResponseExtra)
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.SourcePort) > 0 {
item := NewPortItem(true, options.SourcePort)
rule.sourcePortItems = append(rule.sourcePortItems, item)
@@ -205,6 +294,14 @@ func NewDefaultDNSRule(ctx context.Context, logger log.ContextLogger, options op
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.PackageNameRegex) > 0 {
item, err := NewPackageNameRegexItem(options.PackageNameRegex)
if err != nil {
return nil, E.Cause(err, "package_name_regex")
}
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.User) > 0 {
item := NewUserItem(options.User)
rule.items = append(rule.items, item)
@@ -265,6 +362,28 @@ func NewDefaultDNSRule(ctx context.Context, logger log.ContextLogger, options op
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.SourceMACAddress) > 0 {
item := NewSourceMACAddressItem(options.SourceMACAddress)
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.SourceHostname) > 0 {
item := NewSourceHostnameItem(options.SourceHostname)
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.PreferredBy) > 0 {
item := NewPreferredByDNSItem(ctx, options.PreferredBy)
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if options.RuleSetIPCIDRAcceptEmpty { //nolint:staticcheck
if legacyDNSMode {
deprecated.Report(ctx, deprecated.OptionRuleSetIPCIDRAcceptEmpty)
} else {
return nil, E.New(deprecated.OptionRuleSetIPCIDRAcceptEmpty.MessageWithLink())
}
}
if len(options.RuleSet) > 0 {
//nolint:staticcheck
if options.Deprecated_RulesetIPCIDRMatchSource {
@@ -274,7 +393,7 @@ func NewDefaultDNSRule(ctx context.Context, logger log.ContextLogger, options op
if options.RuleSetIPCIDRMatchSource {
matchSource = true
}
item := NewRuleSetItem(router, options.RuleSet, matchSource, options.RuleSetIPCIDRAcceptEmpty)
item := NewRuleSetItem(router, options.RuleSet, matchSource, options.RuleSetIPCIDRAcceptEmpty) //nolint:staticcheck
rule.ruleSetItem = item
rule.allItems = append(rule.allItems, item)
}
@@ -289,44 +408,99 @@ func (r *DefaultDNSRule) WithAddressLimit() bool {
if len(r.destinationIPCIDRItems) > 0 {
return true
}
if r.ruleSetItem != nil {
ruleSet, isRuleSet := r.ruleSetItem.(*RuleSetItem)
if isRuleSet && ruleSet.ContainsDestinationIPCIDRRule() {
return true
}
}
return false
return r.ruleSetItem != nil && r.ruleSetItem.ContainsDestinationIPCIDRRule()
}
func (r *DefaultDNSRule) Match(metadata *adapter.InboundContext) bool {
metadata.IgnoreDestinationIPCIDRMatch = true
defer func() {
metadata.IgnoreDestinationIPCIDRMatch = false
}()
return !r.matchStates(metadata).isEmpty()
return r.matchForMatch(metadata)
}
func (r *DefaultDNSRule) MatchAddressLimit(metadata *adapter.InboundContext) bool {
return !r.matchStates(metadata).isEmpty()
func (r *DefaultDNSRule) LegacyPreMatch(metadata *adapter.InboundContext) bool {
if r.matchResponse {
return false
}
metadata.IgnoreDestinationIPCIDRMatch = true
defer func() { metadata.IgnoreDestinationIPCIDRMatch = false }()
return r.abstractDefaultRule.Match(metadata)
}
func (r *DefaultDNSRule) MatchResponseTag() string {
return r.matchResponseTag
}
func (r *DefaultDNSRule) MatchResponseTags() []string {
if r.matchResponseTag == "" {
return nil
}
return []string{r.matchResponseTag}
}
func (r *DefaultDNSRule) MatchResponseAnonymous() bool {
return r.matchResponse && r.matchResponseTag == ""
}
func (r *DefaultDNSRule) Race() bool {
return r.race
}
func (r *DefaultDNSRule) matchForMatch(metadata *adapter.InboundContext) bool {
if r.matchResponse {
response := metadata.DNSResponse
if r.matchResponseTag != "" {
response = metadata.NamedDNSResponses[r.matchResponseTag]
}
if response == nil {
return r.invert
}
matchMetadata := *metadata
matchMetadata.DNSResponse = response
matchMetadata.DestinationAddressMatchFromResponse = true
return r.abstractDefaultRule.Match(&matchMetadata)
}
return r.abstractDefaultRule.Match(metadata)
}
func (r *DefaultDNSRule) MatchAddressLimit(metadata *adapter.InboundContext, response *dns.Msg) bool {
matchMetadata := *metadata
matchMetadata.ResetRuleCache()
matchMetadata.DNSResponse = response
matchMetadata.DestinationAddressMatchFromResponse = true
return r.abstractDefaultRule.Match(&matchMetadata)
}
var _ adapter.DNSRule = (*LogicalDNSRule)(nil)
type LogicalDNSRule struct {
abstractLogicalRule
matchResponseTags []string
matchResponseAnonymous bool
race bool
}
func (r *LogicalDNSRule) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return r.abstractLogicalRule.matchStates(metadata)
func (r *LogicalDNSRule) MatchResponseTag() string {
return ""
}
func NewLogicalDNSRule(ctx context.Context, logger log.ContextLogger, options option.LogicalDNSRule) (*LogicalDNSRule, error) {
func (r *LogicalDNSRule) MatchResponseTags() []string {
return r.matchResponseTags
}
func (r *LogicalDNSRule) MatchResponseAnonymous() bool {
return r.matchResponseAnonymous
}
func (r *LogicalDNSRule) Race() bool {
return r.race
}
func NewLogicalDNSRule(ctx context.Context, logger log.ContextLogger, options option.LogicalDNSRule, legacyDNSMode bool) (*LogicalDNSRule, error) {
r := &LogicalDNSRule{
abstractLogicalRule: abstractLogicalRule{
rules: make([]adapter.HeadlessRule, len(options.Rules)),
invert: options.Invert,
action: NewDNSRuleAction(logger, options.DNSRuleAction),
},
race: options.Race,
}
switch options.Mode {
case C.LogicalTypeAnd:
@@ -337,12 +511,26 @@ func NewLogicalDNSRule(ctx context.Context, logger log.ContextLogger, options op
return nil, E.New("unknown logical mode: ", options.Mode)
}
for i, subRule := range options.Rules {
rule, err := NewDNSRule(ctx, logger, subRule, false)
err := validateNoNestedDNSRuleActions(subRule, true)
if err != nil {
return nil, E.Cause(err, "sub rule[", i, "]")
}
rule, err := NewDNSRule(ctx, logger, subRule, false, legacyDNSMode)
if err != nil {
return nil, E.Cause(err, "sub rule[", i, "]")
}
r.rules[i] = rule
}
for _, subRule := range r.rules {
if dnsRule, isDNSRule := subRule.(adapter.DNSRule); isDNSRule {
r.matchResponseTags = append(r.matchResponseTags, dnsRule.MatchResponseTags()...)
r.matchResponseAnonymous = r.matchResponseAnonymous || dnsRule.MatchResponseAnonymous()
}
}
r.matchResponseTags = common.Uniq(r.matchResponseTags)
if r.race && len(r.matchResponseTags) == 0 && !r.matchResponseAnonymous {
return nil, E.New("`race` requires `match_response` in sub-rules")
}
return r, nil
}
@@ -352,28 +540,23 @@ func (r *LogicalDNSRule) Action() adapter.RuleAction {
func (r *LogicalDNSRule) WithAddressLimit() bool {
for _, rawRule := range r.rules {
switch rule := rawRule.(type) {
case *DefaultDNSRule:
if rule.WithAddressLimit() {
return true
}
case *LogicalDNSRule:
if rule.WithAddressLimit() {
return true
}
if dnsRule, isDNSRule := rawRule.(adapter.DNSRule); isDNSRule && dnsRule.WithAddressLimit() {
return true
}
}
return false
}
func (r *LogicalDNSRule) Match(metadata *adapter.InboundContext) bool {
func (r *LogicalDNSRule) LegacyPreMatch(metadata *adapter.InboundContext) bool {
metadata.IgnoreDestinationIPCIDRMatch = true
defer func() {
metadata.IgnoreDestinationIPCIDRMatch = false
}()
return !r.matchStates(metadata).isEmpty()
defer func() { metadata.IgnoreDestinationIPCIDRMatch = false }()
return r.abstractLogicalRule.Match(metadata)
}
func (r *LogicalDNSRule) MatchAddressLimit(metadata *adapter.InboundContext) bool {
return !r.matchStates(metadata).isEmpty()
func (r *LogicalDNSRule) MatchAddressLimit(metadata *adapter.InboundContext, response *dns.Msg) bool {
matchMetadata := *metadata
matchMetadata.ResetRuleCache()
matchMetadata.DNSResponse = response
matchMetadata.DestinationAddressMatchFromResponse = true
return r.abstractLogicalRule.Match(&matchMetadata)
}
+386
View File
@@ -0,0 +1,386 @@
package rule
import (
"context"
"net"
"testing"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common/json"
M "github.com/sagernet/sing/common/metadata"
"github.com/sagernet/sing/service"
"github.com/miekg/dns"
"github.com/stretchr/testify/require"
)
type addressFilterRouter struct {
adapter.Router
ruleSets map[string]adapter.RuleSet
}
func (r *addressFilterRouter) RuleSet(tag string) (adapter.RuleSet, bool) {
ruleSet, loaded := r.ruleSets[tag]
return ruleSet, loaded
}
func addressFilterContext(t *testing.T, ruleSetConfigs map[string]string) context.Context {
t.Helper()
router := &addressFilterRouter{ruleSets: make(map[string]adapter.RuleSet)}
ctx := service.ContextWith[adapter.Router](context.Background(), router)
for tag, config := range ruleSetConfigs {
var plainOptions option.PlainRuleSetCompat
err := json.UnmarshalContext(ctx, []byte(config), &plainOptions)
require.NoError(t, err)
ruleSet, err := NewLocalRuleSet(ctx, log.NewNOPFactory().Logger(), tag, option.RuleSet{
Type: C.RuleSetTypeInline,
InlineOptions: plainOptions.Options,
})
require.NoError(t, err)
router.ruleSets[tag] = ruleSet
}
return ctx
}
func addressFilterDNSRule(t *testing.T, ctx context.Context, config string) adapter.DNSRule {
t.Helper()
var ruleOptions option.DNSRule
err := json.UnmarshalContext(ctx, []byte(config), &ruleOptions)
require.NoError(t, err)
rule, err := NewDNSRule(ctx, log.NewNOPFactory().NewLogger("test"), ruleOptions, true, true)
require.NoError(t, err)
require.NoError(t, rule.Start())
return rule
}
func addressFilterResponse(address string) *dns.Msg {
response := &dns.Msg{}
response.Rcode = dns.RcodeSuccess
response.Answer = append(response.Answer, &dns.A{
Hdr: dns.RR_Header{Rrtype: dns.TypeA, Class: dns.ClassINET},
A: net.ParseIP(address).To4(),
})
return response
}
// addressFilterFlow mirrors dns/router.go: LegacyPreMatch under
// IgnoreDestinationIPCIDRMatch, then addressLimitResponseCheck against the
// response.
func addressFilterFlow(rule adapter.DNSRule, domain string, responseAddress string) (preMatched bool, routed bool) {
metadata := adapter.InboundContext{
Domain: domain,
QueryType: dns.TypeA,
Source: M.ParseSocksaddrHostPort("192.168.1.10", 5353),
}
metadata.ResetRuleCache()
preMatched = rule.LegacyPreMatch(&metadata)
if !preMatched {
return false, false
}
if !rule.WithAddressLimit() {
return true, true
}
checkMetadata := metadata
return true, rule.MatchAddressLimit(&checkMetadata, addressFilterResponse(responseAddress))
}
func TestDNSAddressFilterInvert(t *testing.T) {
t.Parallel()
ctx := addressFilterContext(t, map[string]string{
"mixed": `{"version": 3, "rules": [{"domain_suffix": ["ads.example"]}, {"ip_cidr": ["1.1.1.0/24"]}]}`,
"cn-ip": `{"version": 3, "rules": [{"ip_cidr": ["1.1.1.0/24"]}]}`,
"lan-ip": `{"version": 3, "rules": [{"ip_cidr": ["192.168.0.0/16"]}]}`,
"other-net": `{"version": 3, "rules": [{"ip_cidr": ["10.99.0.0/16"]}]}`,
"cn-domain": `{"version": 3, "rules": [{"domain_suffix": ["cn.example"]}]}`,
})
testCases := []struct {
name string
rule string
domain string
responseAddress string
expectPreMatch bool
expectRouted bool
}{
{
name: "direct mixed invert, domain hit",
rule: `{"domain_suffix": ["lookup.example"], "ip_cidr": ["1.1.1.0/24"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
},
{
name: "direct mixed invert, both miss",
rule: `{"domain_suffix": ["lookup.example"], "ip_cidr": ["1.1.1.0/24"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "other.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "direct mixed invert, ip hit",
rule: `{"domain_suffix": ["lookup.example"], "ip_cidr": ["1.1.1.0/24"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "other.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
{
name: "ip rule-set invert, ip miss",
rule: `{"rule_set": ["cn-ip"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "ip rule-set invert, ip hit",
rule: `{"rule_set": ["cn-ip"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
{
name: "mixed rule-set invert, both miss",
rule: `{"rule_set": ["mixed"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "mixed rule-set invert, domain hit skips pre-lookup",
rule: `{"rule_set": ["mixed"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "ads.example",
responseAddress: "8.8.8.8",
},
{
name: "mixed rule-set invert, ip hit",
rule: `{"rule_set": ["mixed"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
{
name: "logical invert, ip miss",
rule: `{"type": "logical", "mode": "or", "invert": true, "rules": [{"rule_set": ["cn-ip"]}], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "logical invert, ip hit",
rule: `{"type": "logical", "mode": "or", "invert": true, "rules": [{"rule_set": ["cn-ip"]}], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
{
name: "logical and with inverted ip rule-set, ip miss",
rule: `{"type": "logical", "mode": "and", "rules": [{"domain_suffix": ["lookup.example"]}, {"rule_set": ["cn-ip"], "invert": true}], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "logical and with inverted ip rule-set, ip hit",
rule: `{"type": "logical", "mode": "and", "rules": [{"domain_suffix": ["lookup.example"]}, {"rule_set": ["cn-ip"], "invert": true}], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
{
name: "logical or invert, domain hit skips pre-lookup",
rule: `{"type": "logical", "mode": "or", "invert": true, "rules": [{"domain_suffix": ["ads.example"]}, {"rule_set": ["cn-ip"]}], "action": "route", "server": "proxy"}`,
domain: "ads.example",
responseAddress: "8.8.8.8",
},
{
name: "logical or invert, both miss",
rule: `{"type": "logical", "mode": "or", "invert": true, "rules": [{"domain_suffix": ["ads.example"]}, {"rule_set": ["cn-ip"]}], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "logical or invert, ip hit",
rule: `{"type": "logical", "mode": "or", "invert": true, "rules": [{"domain_suffix": ["ads.example"]}, {"rule_set": ["cn-ip"]}], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
{
name: "outer domain with ip rule-set invert, domain hit skips pre-lookup",
rule: `{"domain_suffix": ["cn.example"], "rule_set": ["cn-ip"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "cn.example",
responseAddress: "8.8.8.8",
},
{
name: "outer domain with ip rule-set invert, both miss",
rule: `{"domain_suffix": ["cn.example"], "rule_set": ["cn-ip"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "other.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "source ip with ip invert, source hit ip miss",
rule: `{"source_ip_cidr": ["192.168.1.0/24"], "ip_cidr": ["1.1.1.0/24"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "source ip with ip invert, source hit ip hit",
rule: `{"source_ip_cidr": ["192.168.1.0/24"], "ip_cidr": ["1.1.1.0/24"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
{
name: "source ip with ip invert, source miss",
rule: `{"source_ip_cidr": ["10.99.0.0/16"], "ip_cidr": ["1.1.1.0/24"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "ip rule-set without invert, ip hit",
rule: `{"rule_set": ["cn-ip"], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
expectRouted: true,
},
{
name: "ip rule-set without invert, ip miss",
rule: `{"rule_set": ["cn-ip"], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
},
{
name: "mixed rule-set without invert, ip hit",
rule: `{"rule_set": ["mixed"], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
expectRouted: true,
},
{
name: "mixed rule-set without invert, both miss",
rule: `{"rule_set": ["mixed"], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
},
{
name: "direct ip invert, ip miss",
rule: `{"ip_cidr": ["1.1.1.0/24"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "direct ip invert, ip hit",
rule: `{"ip_cidr": ["1.1.1.0/24"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
{
name: "nested logical invert, ip miss",
rule: `{"type": "logical", "mode": "or", "invert": true, "rules": [{"type": "logical", "mode": "or", "rules": [{"rule_set": ["cn-ip"]}]}], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "nested logical invert, ip hit",
rule: `{"type": "logical", "mode": "or", "invert": true, "rules": [{"type": "logical", "mode": "or", "rules": [{"rule_set": ["cn-ip"]}]}], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
{
name: "nested logical and invert, domain hit ip miss",
rule: `{"type": "logical", "mode": "or", "invert": true, "rules": [{"type": "logical", "mode": "and", "rules": [{"domain_suffix": ["lookup.example"]}, {"rule_set": ["cn-ip"]}]}], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "nested logical and invert, domain hit ip hit",
rule: `{"type": "logical", "mode": "or", "invert": true, "rules": [{"type": "logical", "mode": "and", "rules": [{"domain_suffix": ["lookup.example"]}, {"rule_set": ["cn-ip"]}]}], "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
{
name: "match-source rule-set invert, source in set",
rule: `{"rule_set": ["lan-ip"], "rule_set_ip_cidr_match_source": true, "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
},
{
name: "match-source rule-set invert, source not in set",
rule: `{"rule_set": ["other-net"], "rule_set_ip_cidr_match_source": true, "invert": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "match-source rule-set, source in set",
rule: `{"rule_set": ["lan-ip"], "rule_set_ip_cidr_match_source": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "match-source rule-set, source not in set",
rule: `{"rule_set": ["other-net"], "rule_set_ip_cidr_match_source": true, "action": "route", "server": "proxy"}`,
domain: "lookup.example",
responseAddress: "8.8.8.8",
},
{
name: "direct ip with domain rule-set invert, domain hit skips pre-lookup",
rule: `{"ip_cidr": ["1.1.1.0/24"], "rule_set": ["cn-domain"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "cn.example",
responseAddress: "8.8.8.8",
},
{
name: "direct ip with domain rule-set invert, both miss",
rule: `{"ip_cidr": ["1.1.1.0/24"], "rule_set": ["cn-domain"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "other.example",
responseAddress: "8.8.8.8",
expectPreMatch: true,
expectRouted: true,
},
{
name: "direct ip with domain rule-set invert, ip hit",
rule: `{"ip_cidr": ["1.1.1.0/24"], "rule_set": ["cn-domain"], "invert": true, "action": "route", "server": "proxy"}`,
domain: "other.example",
responseAddress: "1.1.1.5",
expectPreMatch: true,
},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
t.Parallel()
rule := addressFilterDNSRule(t, ctx, testCase.rule)
preMatched, routed := addressFilterFlow(rule, testCase.domain, testCase.responseAddress)
require.Equal(t, testCase.expectPreMatch, preMatched, "pre-lookup match")
require.Equal(t, testCase.expectRouted, routed, "routed")
})
}
}
+8 -8
View File
@@ -34,10 +34,6 @@ type DefaultHeadlessRule struct {
abstractDefaultRule
}
func (r *DefaultHeadlessRule) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return r.abstractDefaultRule.matchStates(metadata)
}
func NewDefaultHeadlessRule(ctx context.Context, options option.DefaultHeadlessRule) (*DefaultHeadlessRule, error) {
networkManager := service.FromContext[adapter.NetworkManager](ctx)
rule := &DefaultHeadlessRule{
@@ -153,6 +149,14 @@ func NewDefaultHeadlessRule(ctx context.Context, options option.DefaultHeadlessR
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if len(options.PackageNameRegex) > 0 {
item, err := NewPackageNameRegexItem(options.PackageNameRegex)
if err != nil {
return nil, E.Cause(err, "package_name_regex")
}
rule.items = append(rule.items, item)
rule.allItems = append(rule.allItems, item)
}
if networkManager != nil {
if len(options.NetworkType) > 0 {
item := NewNetworkTypeItem(networkManager, common.Map(options.NetworkType, option.InterfaceType.Build))
@@ -208,10 +212,6 @@ type LogicalHeadlessRule struct {
abstractLogicalRule
}
func (r *LogicalHeadlessRule) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return r.abstractLogicalRule.matchStates(metadata)
}
func NewLogicalHeadlessRule(ctx context.Context, options option.LogicalHeadlessRule) (*LogicalHeadlessRule, error) {
r := &LogicalHeadlessRule{
abstractLogicalRule{
+11 -1
View File
@@ -6,6 +6,7 @@ import (
"strings"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing/common"
E "github.com/sagernet/sing/common/exceptions"
"go4.org/netipx"
@@ -77,11 +78,20 @@ func (r *IPCIDRItem) Match(metadata *adapter.InboundContext) bool {
if r.isSource || metadata.IPCIDRMatchSource {
return r.ipSet.Contains(metadata.Source.Addr)
}
if metadata.DestinationAddressMatchFromResponse {
addresses := metadata.DNSResponseAddressesForMatch()
if len(addresses) == 0 {
// Legacy rule_set_ip_cidr_accept_empty only applies when the DNS response
// does not expose any address answers for matching.
return metadata.IPCIDRAcceptEmpty
}
return slices.ContainsFunc(addresses, r.ipSet.Contains)
}
if metadata.Destination.IsIP() {
return r.ipSet.Contains(metadata.Destination.Addr)
}
if len(metadata.DestinationAddresses) > 0 {
return slices.ContainsFunc(metadata.DestinationAddresses, r.ipSet.Contains)
return common.Any(metadata.DestinationAddresses, r.ipSet.Contains)
}
return metadata.IPCIDRAcceptEmpty
}
+3
View File
@@ -13,6 +13,9 @@ func NewIPAcceptAnyItem() *IPAcceptAnyItem {
}
func (r *IPAcceptAnyItem) Match(metadata *adapter.InboundContext) bool {
if metadata.DestinationAddressMatchFromResponse {
return len(metadata.DNSResponseAddressesForMatch()) > 0
}
return len(metadata.DestinationAddresses) > 0
}
+12 -11
View File
@@ -1,8 +1,6 @@
package rule
import (
"net/netip"
"github.com/sagernet/sing-box/adapter"
N "github.com/sagernet/sing/common/network"
)
@@ -18,21 +16,24 @@ func NewIPIsPrivateItem(isSource bool) *IPIsPrivateItem {
}
func (r *IPIsPrivateItem) Match(metadata *adapter.InboundContext) bool {
var destination netip.Addr
if r.isSource {
destination = metadata.Source.Addr
} else {
destination = metadata.Destination.Addr
return !N.IsPublicAddr(metadata.Source.Addr)
}
if destination.IsValid() {
return !N.IsPublicAddr(destination)
}
if !r.isSource {
for _, destinationAddress := range metadata.DestinationAddresses {
if metadata.DestinationAddressMatchFromResponse {
for _, destinationAddress := range metadata.DNSResponseAddressesForMatch() {
if !N.IsPublicAddr(destinationAddress) {
return true
}
}
return false
}
if metadata.Destination.Addr.IsValid() {
return !N.IsPublicAddr(metadata.Destination.Addr)
}
for _, destinationAddress := range metadata.DestinationAddresses {
if !N.IsPublicAddr(destinationAddress) {
return true
}
}
return false
}
@@ -0,0 +1,55 @@
package rule
import (
"regexp"
"slices"
"strings"
"github.com/sagernet/sing-box/adapter"
E "github.com/sagernet/sing/common/exceptions"
F "github.com/sagernet/sing/common/format"
)
var _ RuleItem = (*PackageNameRegexItem)(nil)
type PackageNameRegexItem struct {
matchers []*regexp.Regexp
description string
}
func NewPackageNameRegexItem(expressions []string) (*PackageNameRegexItem, error) {
matchers := make([]*regexp.Regexp, 0, len(expressions))
for i, regex := range expressions {
matcher, err := regexp.Compile(regex)
if err != nil {
return nil, E.Cause(err, "parse expression ", i)
}
matchers = append(matchers, matcher)
}
description := "package_name_regex="
eLen := len(expressions)
if eLen == 1 {
description += expressions[0]
} else if eLen > 3 {
description += F.ToString("[", strings.Join(expressions[:3], " "), "]")
} else {
description += F.ToString("[", strings.Join(expressions, " "), "]")
}
return &PackageNameRegexItem{matchers, description}, nil
}
func (r *PackageNameRegexItem) Match(metadata *adapter.InboundContext) bool {
if metadata.ProcessInfo == nil || len(metadata.ProcessInfo.AndroidPackageNames) == 0 {
return false
}
for _, matcher := range r.matchers {
if slices.ContainsFunc(metadata.ProcessInfo.AndroidPackageNames, matcher.MatchString) {
return true
}
}
return false
}
func (r *PackageNameRegexItem) String() string {
return r.description
}
+3 -3
View File
@@ -50,14 +50,14 @@ func (r *PreferredByItem) Match(metadata *adapter.InboundContext) bool {
}
if domainHost != "" {
for _, outbound := range r.outbounds {
if outbound.PreferredDomain(domainHost) {
if outbound.PreferredDomain(metadata, domainHost) {
return true
}
}
}
if metadata.Destination.IsIP() {
for _, outbound := range r.outbounds {
if outbound.PreferredAddress(metadata.Destination.Addr) {
if outbound.PreferredAddress(metadata, metadata.Destination.Addr) {
return true
}
}
@@ -65,7 +65,7 @@ func (r *PreferredByItem) Match(metadata *adapter.InboundContext) bool {
if len(metadata.DestinationAddresses) > 0 {
for _, address := range metadata.DestinationAddresses {
for _, outbound := range r.outbounds {
if outbound.PreferredAddress(address) {
if outbound.PreferredAddress(metadata, address) {
return true
}
}
+74
View File
@@ -0,0 +1,74 @@
package rule
import (
"context"
"strings"
"github.com/sagernet/sing-box/adapter"
E "github.com/sagernet/sing/common/exceptions"
F "github.com/sagernet/sing/common/format"
"github.com/sagernet/sing/service"
mDNS "github.com/miekg/dns"
)
var _ RuleItem = (*PreferredByDNSItem)(nil)
type PreferredByDNSItem struct {
ctx context.Context
transportTags []string
transports []adapter.DNSTransportWithPreferredDomain
}
func NewPreferredByDNSItem(ctx context.Context, transportTags []string) *PreferredByDNSItem {
return &PreferredByDNSItem{
ctx: ctx,
transportTags: transportTags,
}
}
func (r *PreferredByDNSItem) Start() error {
transportManager := service.FromContext[adapter.DNSTransportManager](r.ctx)
for _, transportTag := range r.transportTags {
rawTransport, loaded := transportManager.Transport(transportTag)
if !loaded {
return E.New("DNS server not found: ", transportTag)
}
transportWithPreferredDomain, withPreferredDomain := rawTransport.(adapter.DNSTransportWithPreferredDomain)
if !withPreferredDomain {
return E.New("DNS server type does not support preferred_by: ", rawTransport.Type())
}
r.transports = append(r.transports, transportWithPreferredDomain)
}
return nil
}
func (r *PreferredByDNSItem) Match(metadata *adapter.InboundContext) bool {
var domainHost string
if metadata.Domain != "" {
domainHost = metadata.Domain
} else {
domainHost = metadata.Destination.Fqdn
}
if domainHost == "" {
return false
}
canonical := mDNS.CanonicalName(domainHost)
for _, transport := range r.transports {
if transport.PreferredDomain(canonical) {
return true
}
}
return false
}
func (r *PreferredByDNSItem) String() string {
description := "preferred_by="
pLen := len(r.transportTags)
if pLen == 1 {
description += F.ToString(r.transportTags[0])
} else {
description += "[" + strings.Join(F.MapToString(r.transportTags), " ") + "]"
}
return description
}
@@ -0,0 +1,42 @@
package rule
import (
"net/netip"
"slices"
"strings"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/json/badoption"
)
var _ RuleItem = (*QueryClientSubnetItem)(nil)
type QueryClientSubnetItem struct {
prefixes []netip.Prefix
}
func NewQueryClientSubnetItem(prefixables badoption.Listable[*badoption.Prefixable]) *QueryClientSubnetItem {
return &QueryClientSubnetItem{
prefixes: common.Map(prefixables, func(it *badoption.Prefixable) netip.Prefix {
return it.Build(netip.Prefix{})
}),
}
}
func (r *QueryClientSubnetItem) Match(metadata *adapter.InboundContext) bool {
clientSubnet := metadata.QueryClientSubnet
if !clientSubnet.IsValid() {
return false
}
return slices.ContainsFunc(r.prefixes, func(prefix netip.Prefix) bool {
return clientSubnet.Bits() >= prefix.Bits() && prefix.Contains(clientSubnet.Addr())
})
}
func (r *QueryClientSubnetItem) String() string {
if len(r.prefixes) == 1 {
return "query_client_subnet=" + r.prefixes[0].String()
}
return "query_client_subnet=[" + strings.Join(common.Map(r.prefixes, netip.Prefix.String), " ") + "]"
}
+21
View File
@@ -0,0 +1,21 @@
package rule
import (
"github.com/sagernet/sing-box/adapter"
)
var _ RuleItem = (*QueryDNSSECItem)(nil)
type QueryDNSSECItem struct{}
func NewQueryDNSSECItem() *QueryDNSSECItem {
return &QueryDNSSECItem{}
}
func (r *QueryDNSSECItem) Match(metadata *adapter.InboundContext) bool {
return metadata.QueryDNSSEC
}
func (r *QueryDNSSECItem) String() string {
return "query_dnssec=true"
}
+26
View File
@@ -0,0 +1,26 @@
package rule
import (
"github.com/sagernet/sing-box/adapter"
F "github.com/sagernet/sing/common/format"
"github.com/miekg/dns"
)
var _ RuleItem = (*DNSResponseRCodeItem)(nil)
type DNSResponseRCodeItem struct {
rcode int
}
func NewDNSResponseRCodeItem(rcode int) *DNSResponseRCodeItem {
return &DNSResponseRCodeItem{rcode: rcode}
}
func (r *DNSResponseRCodeItem) Match(metadata *adapter.InboundContext) bool {
return metadata.DNSResponse != nil && metadata.DNSResponse.Rcode == r.rcode
}
func (r *DNSResponseRCodeItem) String() string {
return F.ToString("response_rcode=", dns.RcodeToString[r.rcode])
}
+62
View File
@@ -0,0 +1,62 @@
package rule
import (
"slices"
"strings"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/option"
"github.com/miekg/dns"
)
var _ RuleItem = (*DNSResponseRecordItem)(nil)
type DNSResponseRecordItem struct {
field string
records []option.DNSRecordOptions
selector func(*dns.Msg) []dns.RR
}
func NewDNSResponseRecordItem(field string, records []option.DNSRecordOptions, selector func(*dns.Msg) []dns.RR) *DNSResponseRecordItem {
return &DNSResponseRecordItem{
field: field,
records: records,
selector: selector,
}
}
func (r *DNSResponseRecordItem) Match(metadata *adapter.InboundContext) bool {
if metadata.DNSResponse == nil {
return false
}
records := r.selector(metadata.DNSResponse)
for _, expected := range r.records {
if slices.ContainsFunc(records, expected.Match) {
return true
}
}
return false
}
func (r *DNSResponseRecordItem) String() string {
descriptions := make([]string, 0, len(r.records))
for _, record := range r.records {
if record.RR != nil {
descriptions = append(descriptions, record.RR.String())
}
}
return r.field + "=[" + strings.Join(descriptions, " ") + "]"
}
func dnsResponseAnswers(message *dns.Msg) []dns.RR {
return message.Answer
}
func dnsResponseNS(message *dns.Msg) []dns.RR {
return message.Ns
}
func dnsResponseExtra(message *dns.Msg) []dns.RR {
return message.Extra
}
+101 -14
View File
@@ -29,9 +29,11 @@ func NewRuleSetItem(router adapter.Router, tagList []string, ipCIDRMatchSource b
}
func (r *RuleSetItem) Start() error {
_ = r.Close()
for _, tag := range r.tagList {
ruleSet, loaded := r.router.RuleSet(tag)
if !loaded {
_ = r.Close()
return E.New("rule-set not found: ", tag)
}
ruleSet.IncRef()
@@ -40,24 +42,109 @@ func (r *RuleSetItem) Start() error {
return nil
}
func (r *RuleSetItem) Match(metadata *adapter.InboundContext) bool {
return !r.matchStates(metadata).isEmpty()
}
func (r *RuleSetItem) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return r.matchStatesWithBase(metadata, 0)
}
func (r *RuleSetItem) matchStatesWithBase(metadata *adapter.InboundContext, base ruleMatchState) ruleMatchStateSet {
var stateSet ruleMatchStateSet
func (r *RuleSetItem) Close() error {
for _, ruleSet := range r.setList {
ruleSet.DecRef()
}
clear(r.setList)
r.setList = nil
return nil
}
func (r *RuleSetItem) Match(metadata *adapter.InboundContext) bool {
for _, ruleSet := range r.setList {
nestedMetadata := r.nestedMetadata(metadata)
if ruleSet.Match(&nestedMetadata) {
return true
}
}
return false
}
func (r *RuleSetItem) matchWithOuterGroups(metadata *adapter.InboundContext, outerGroups ruleGroupMatch) bool {
outerDone := outerGroups.done()
var (
matched bool
deferredGroups uint8
)
for _, ruleSet := range r.setList {
nestedMetadata := r.nestedMetadata(metadata)
if provider, isProvider := ruleSet.(mergeableRuleProvider); isProvider {
branch := provider.mergeableRule()
if branch != nil {
branchGroups, branchMatched := branch.evaluateForMerge(&nestedMetadata)
if branchMatched {
merged := outerGroups.mergeWith(branchGroups)
if merged.done() {
branchDeferredGroups := nestedMetadata.DeferredIPCIDRMatchGroups &^ uint8(merged.satisfied)
if branchDeferredGroups == 0 {
metadata.DeferredIPCIDRMatchGroups &^= uint8(merged.satisfied)
return true
}
matched = true
deferredGroups |= branchDeferredGroups
}
}
continue
}
}
if outerDone && ruleSet.Match(&nestedMetadata) {
if nestedMetadata.DeferredIPCIDRMatchGroups == 0 {
return true
}
matched = true
deferredGroups |= nestedMetadata.DeferredIPCIDRMatchGroups
}
}
if matched {
metadata.DeferredIPCIDRMatchGroups |= deferredGroups
}
return matched
}
func (r *RuleSetItem) nestedMetadata(metadata *adapter.InboundContext) adapter.InboundContext {
nestedMetadata := *metadata
nestedMetadata.ResetRuleMatchCache()
nestedMetadata.IPCIDRMatchSource = r.ipCidrMatchSource
nestedMetadata.IPCIDRAcceptEmpty = r.ipCidrAcceptEmpty
return nestedMetadata
}
type mergeableRuleProvider interface {
mergeableRule() *DefaultHeadlessRule
}
func mergeableRuleIn(rules []adapter.HeadlessRule) *DefaultHeadlessRule {
if len(rules) != 1 {
return nil
}
rule, isDefault := rules[0].(*DefaultHeadlessRule)
if !isDefault || rule.invert || rule.ruleSetItem != nil {
return nil
}
return rule
}
func matchAnyHeadlessRule(rules []adapter.HeadlessRule, metadata *adapter.InboundContext) bool {
var (
matched bool
deferredGroups uint8
)
for _, rule := range rules {
nestedMetadata := *metadata
nestedMetadata.ResetRuleMatchCache()
nestedMetadata.IPCIDRMatchSource = r.ipCidrMatchSource
nestedMetadata.IPCIDRAcceptEmpty = r.ipCidrAcceptEmpty
stateSet = stateSet.merge(matchHeadlessRuleStatesWithBase(ruleSet, &nestedMetadata, base))
if rule.Match(&nestedMetadata) {
if nestedMetadata.DeferredIPCIDRMatchGroups == 0 {
return true
}
matched = true
deferredGroups |= nestedMetadata.DeferredIPCIDRMatchGroups
}
}
return stateSet
if matched {
metadata.DeferredIPCIDRMatchGroups |= deferredGroups
}
return matched
}
func (r *RuleSetItem) ContainsDestinationIPCIDRRule() bool {
+148
View File
@@ -0,0 +1,148 @@
package rule
import (
"context"
"net"
"sync/atomic"
"testing"
"github.com/sagernet/sing-box/adapter"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/x/list"
"github.com/stretchr/testify/require"
"go4.org/netipx"
)
type ruleSetItemTestRouter struct {
ruleSets map[string]adapter.RuleSet
}
func (r *ruleSetItemTestRouter) Start(adapter.StartStage) error { return nil }
func (r *ruleSetItemTestRouter) Close() error { return nil }
func (r *ruleSetItemTestRouter) PreMatch(adapter.InboundContext, []byte) adapter.PreMatchResult {
return adapter.PreMatchResult{}
}
func (r *ruleSetItemTestRouter) HijackDNSPacket(context.Context, []byte, N.PacketWriter, adapter.InboundContext) {
}
func (r *ruleSetItemTestRouter) RouteConnection(context.Context, net.Conn, adapter.InboundContext) error {
return nil
}
func (r *ruleSetItemTestRouter) RoutePacketConnection(context.Context, N.PacketConn, adapter.InboundContext) error {
return nil
}
func (r *ruleSetItemTestRouter) RouteConnectionEx(context.Context, net.Conn, adapter.InboundContext, N.CloseHandlerFunc) {
}
func (r *ruleSetItemTestRouter) RoutePacketConnectionEx(context.Context, N.PacketConn, adapter.InboundContext, N.CloseHandlerFunc) {
}
func (r *ruleSetItemTestRouter) RuleSet(tag string) (adapter.RuleSet, bool) {
ruleSet, loaded := r.ruleSets[tag]
return ruleSet, loaded
}
func (r *ruleSetItemTestRouter) Rules() []adapter.Rule { return nil }
func (r *ruleSetItemTestRouter) NeedFindProcess() bool { return false }
func (r *ruleSetItemTestRouter) NeedFindNeighbor() bool { return false }
func (r *ruleSetItemTestRouter) NeighborResolver() adapter.NeighborResolver { return nil }
func (r *ruleSetItemTestRouter) AppendTracker(adapter.ConnectionTracker) {}
func (r *ruleSetItemTestRouter) ResetNetwork() {}
type countingRuleSet struct {
name string
refs atomic.Int32
}
func (s *countingRuleSet) Name() string { return s.name }
func (s *countingRuleSet) StartContext(context.Context, *adapter.HTTPStartContext) error { return nil }
func (s *countingRuleSet) PostStart() error { return nil }
func (s *countingRuleSet) Metadata() adapter.RuleSetMetadata { return adapter.RuleSetMetadata{} }
func (s *countingRuleSet) ExtractIPSet() []*netipx.IPSet { return nil }
func (s *countingRuleSet) IncRef() { s.refs.Add(1) }
func (s *countingRuleSet) DecRef() {
if s.refs.Add(-1) < 0 {
panic("rule-set: negative refs")
}
}
func (s *countingRuleSet) Cleanup() {}
func (s *countingRuleSet) RegisterCallback(adapter.RuleSetUpdateCallback) *list.Element[adapter.RuleSetUpdateCallback] {
return nil
}
func (s *countingRuleSet) UnregisterCallback(*list.Element[adapter.RuleSetUpdateCallback]) {}
func (s *countingRuleSet) Close() error { return nil }
func (s *countingRuleSet) Match(*adapter.InboundContext) bool { return true }
func (s *countingRuleSet) String() string { return s.name }
func (s *countingRuleSet) RefCount() int32 { return s.refs.Load() }
func TestRuleSetItemCloseReleasesRefs(t *testing.T) {
t.Parallel()
firstSet := &countingRuleSet{name: "first"}
secondSet := &countingRuleSet{name: "second"}
item := NewRuleSetItem(&ruleSetItemTestRouter{
ruleSets: map[string]adapter.RuleSet{
"first": firstSet,
"second": secondSet,
},
}, []string{"first", "second"}, false, false)
require.NoError(t, item.Start())
require.EqualValues(t, 1, firstSet.RefCount())
require.EqualValues(t, 1, secondSet.RefCount())
require.NoError(t, item.Close())
require.Zero(t, firstSet.RefCount())
require.Zero(t, secondSet.RefCount())
require.NoError(t, item.Close())
require.Zero(t, firstSet.RefCount())
require.Zero(t, secondSet.RefCount())
}
func TestRuleSetItemStartRollbackOnFailure(t *testing.T) {
t.Parallel()
firstSet := &countingRuleSet{name: "first"}
item := NewRuleSetItem(&ruleSetItemTestRouter{
ruleSets: map[string]adapter.RuleSet{
"first": firstSet,
},
}, []string{"first", "missing"}, false, false)
err := item.Start()
require.ErrorContains(t, err, "rule-set not found: missing")
require.Zero(t, firstSet.RefCount())
}
func TestRuleSetItemRestartKeepsBalancedRefs(t *testing.T) {
t.Parallel()
firstSet := &countingRuleSet{name: "first"}
item := NewRuleSetItem(&ruleSetItemTestRouter{
ruleSets: map[string]adapter.RuleSet{
"first": firstSet,
},
}, []string{"first"}, false, false)
require.NoError(t, item.Start())
require.EqualValues(t, 1, firstSet.RefCount())
require.NoError(t, item.Start())
require.EqualValues(t, 1, firstSet.RefCount())
require.NoError(t, item.Close())
require.Zero(t, firstSet.RefCount())
}
+42
View File
@@ -0,0 +1,42 @@
package rule
import (
"strings"
"github.com/sagernet/sing-box/adapter"
)
var _ RuleItem = (*SourceHostnameItem)(nil)
type SourceHostnameItem struct {
hostnames []string
hostnameMap map[string]bool
}
func NewSourceHostnameItem(hostnameList []string) *SourceHostnameItem {
rule := &SourceHostnameItem{
hostnames: hostnameList,
hostnameMap: make(map[string]bool),
}
for _, hostname := range hostnameList {
rule.hostnameMap[hostname] = true
}
return rule
}
func (r *SourceHostnameItem) Match(metadata *adapter.InboundContext) bool {
if metadata.SourceHostname == "" {
return false
}
return r.hostnameMap[metadata.SourceHostname]
}
func (r *SourceHostnameItem) String() string {
var description string
if len(r.hostnames) == 1 {
description = "source_hostname=" + r.hostnames[0]
} else {
description = "source_hostname=[" + strings.Join(r.hostnames, " ") + "]"
}
return description
}
@@ -0,0 +1,48 @@
package rule
import (
"net"
"strings"
"github.com/sagernet/sing-box/adapter"
)
var _ RuleItem = (*SourceMACAddressItem)(nil)
type SourceMACAddressItem struct {
addresses []string
addressMap map[string]bool
}
func NewSourceMACAddressItem(addressList []string) *SourceMACAddressItem {
rule := &SourceMACAddressItem{
addresses: addressList,
addressMap: make(map[string]bool),
}
for _, address := range addressList {
parsed, err := net.ParseMAC(address)
if err == nil {
rule.addressMap[parsed.String()] = true
} else {
rule.addressMap[address] = true
}
}
return rule
}
func (r *SourceMACAddressItem) Match(metadata *adapter.InboundContext) bool {
if metadata.SourceMACAddress == nil {
return false
}
return r.addressMap[metadata.SourceMACAddress.String()]
}
func (r *SourceMACAddressItem) String() string {
var description string
if len(r.addresses) == 1 {
description = "source_mac_address=" + r.addresses[0]
} else {
description = "source_mac_address=[" + strings.Join(r.addresses, " ") + "]"
}
return description
}
+71
View File
@@ -0,0 +1,71 @@
package rule
import (
"reflect"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/option"
E "github.com/sagernet/sing/common/exceptions"
)
func ValidateNoNestedRuleActions(rule option.Rule) error {
return validateNoNestedRuleActions(rule, false)
}
func ValidateNoNestedDNSRuleActions(rule option.DNSRule) error {
return validateNoNestedDNSRuleActions(rule, false)
}
func validateNoNestedRuleActions(rule option.Rule, nested bool) error {
if nested && ruleHasConfiguredAction(rule) {
return E.New(option.RouteRuleActionNestedUnsupportedMessage)
}
if rule.Type != C.RuleTypeLogical {
return nil
}
for i, subRule := range rule.LogicalOptions.Rules {
err := validateNoNestedRuleActions(subRule, true)
if err != nil {
return E.Cause(err, "sub rule[", i, "]")
}
}
return nil
}
func validateNoNestedDNSRuleActions(rule option.DNSRule, nested bool) error {
if nested && dnsRuleHasConfiguredAction(rule) {
return E.New(option.DNSRuleActionNestedUnsupportedMessage)
}
if rule.Type != C.RuleTypeLogical {
return nil
}
for i, subRule := range rule.LogicalOptions.Rules {
err := validateNoNestedDNSRuleActions(subRule, true)
if err != nil {
return E.Cause(err, "sub rule[", i, "]")
}
}
return nil
}
func ruleHasConfiguredAction(rule option.Rule) bool {
switch rule.Type {
case "", C.RuleTypeDefault:
return !reflect.DeepEqual(rule.DefaultOptions.RuleAction, option.RuleAction{})
case C.RuleTypeLogical:
return !reflect.DeepEqual(rule.LogicalOptions.RuleAction, option.RuleAction{})
default:
return false
}
}
func dnsRuleHasConfiguredAction(rule option.DNSRule) bool {
switch rule.Type {
case "", C.RuleTypeDefault:
return !reflect.DeepEqual(rule.DefaultOptions.DNSRuleAction, option.DNSRuleAction{})
case C.RuleTypeLogical:
return !reflect.DeepEqual(rule.LogicalOptions.DNSRuleAction, option.DNSRuleAction{})
default:
return false
}
}
+88
View File
@@ -0,0 +1,88 @@
package rule
import (
"context"
"testing"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/stretchr/testify/require"
)
func TestNewRuleRejectsNestedRuleAction(t *testing.T) {
t.Parallel()
_, err := NewRule(context.Background(), log.NewNOPFactory().NewLogger("router"), option.Rule{
Type: C.RuleTypeLogical,
LogicalOptions: option.LogicalRule{
RawLogicalRule: option.RawLogicalRule{
Mode: C.LogicalTypeAnd,
Rules: []option.Rule{{
Type: C.RuleTypeDefault,
DefaultOptions: option.DefaultRule{
RuleAction: option.RuleAction{
Action: C.RuleActionTypeRoute,
RouteOptions: option.RouteActionOptions{
Outbound: "direct",
},
},
},
}},
},
},
}, false)
require.ErrorContains(t, err, option.RouteRuleActionNestedUnsupportedMessage)
}
func TestNewDNSRuleRejectsNestedRuleAction(t *testing.T) {
t.Parallel()
_, err := NewDNSRule(context.Background(), log.NewNOPFactory().NewLogger("dns"), option.DNSRule{
Type: C.RuleTypeLogical,
LogicalOptions: option.LogicalDNSRule{
RawLogicalDNSRule: option.RawLogicalDNSRule{
Mode: C.LogicalTypeAnd,
Rules: []option.DNSRule{{
Type: C.RuleTypeDefault,
DefaultOptions: option.DefaultDNSRule{
DNSRuleAction: option.DNSRuleAction{
Action: C.RuleActionTypeRoute,
RouteOptions: option.DNSRouteActionOptions{
Server: "default",
},
},
},
}},
},
DNSRuleAction: option.DNSRuleAction{
Action: C.RuleActionTypeRoute,
RouteOptions: option.DNSRouteActionOptions{
Server: "default",
},
},
},
}, true, false)
require.ErrorContains(t, err, option.DNSRuleActionNestedUnsupportedMessage)
}
func TestNewDNSRuleRejectsReplyRejectMethod(t *testing.T) {
t.Parallel()
_, err := NewDNSRule(context.Background(), log.NewNOPFactory().NewLogger("dns"), option.DNSRule{
Type: C.RuleTypeDefault,
DefaultOptions: option.DefaultDNSRule{
RawDefaultDNSRule: option.RawDefaultDNSRule{
Domain: []string{"example.com"},
},
DNSRuleAction: option.DNSRuleAction{
Action: C.RuleActionTypeReject,
RejectOptions: option.RejectActionOptions{
Method: C.RuleActionRejectMethodReply,
},
},
},
}, false, false)
require.ErrorContains(t, err, "reject method `reply` is not supported for DNS rules")
}
+37 -4
View File
@@ -2,6 +2,7 @@ package rule
import (
"context"
"reflect"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
@@ -9,16 +10,17 @@ import (
"github.com/sagernet/sing/common"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
"github.com/sagernet/sing/service"
"go4.org/netipx"
)
func NewRuleSet(ctx context.Context, logger logger.ContextLogger, options option.RuleSet) (adapter.RuleSet, error) {
func NewRuleSet(ctx context.Context, logger logger.ContextLogger, tag string, options option.RuleSet) (adapter.RuleSet, error) {
switch options.Type {
case C.RuleSetTypeInline, C.RuleSetTypeLocal, "":
return NewLocalRuleSet(ctx, logger, options)
return NewLocalRuleSet(ctx, logger, tag, options)
case C.RuleSetTypeRemote:
return NewRemoteRuleSet(ctx, logger, options), nil
return NewRemoteRuleSet(ctx, logger, tag, options)
default:
return nil, E.New("unknown rule-set type: ", options.Type)
}
@@ -59,7 +61,7 @@ func HasHeadlessRule(rules []option.HeadlessRule, cond func(rule option.DefaultH
}
func isProcessHeadlessRule(rule option.DefaultHeadlessRule) bool {
return len(rule.ProcessName) > 0 || len(rule.ProcessPath) > 0 || len(rule.ProcessPathRegex) > 0 || len(rule.PackageName) > 0
return len(rule.ProcessName) > 0 || len(rule.ProcessPath) > 0 || len(rule.ProcessPathRegex) > 0 || len(rule.PackageName) > 0 || len(rule.PackageNameRegex) > 0
}
func isWIFIHeadlessRule(rule option.DefaultHeadlessRule) bool {
@@ -69,3 +71,34 @@ func isWIFIHeadlessRule(rule option.DefaultHeadlessRule) bool {
func isIPCIDRHeadlessRule(rule option.DefaultHeadlessRule) bool {
return len(rule.IPCIDR) > 0 || rule.IPSet != nil
}
func isDNSQueryTypeHeadlessRule(rule option.DefaultHeadlessRule) bool {
return len(rule.QueryType) > 0
}
func isNonIPCIDRHeadlessRule(rule option.DefaultHeadlessRule) bool {
ipOnly := option.DefaultHeadlessRule{
IPCIDR: rule.IPCIDR,
IPSet: rule.IPSet,
Invert: rule.Invert,
}
return !reflect.DeepEqual(rule, ipOnly)
}
func buildRuleSetMetadata(headlessRules []option.HeadlessRule) adapter.RuleSetMetadata {
return adapter.RuleSetMetadata{
ContainsProcessRule: HasHeadlessRule(headlessRules, isProcessHeadlessRule),
ContainsWIFIRule: HasHeadlessRule(headlessRules, isWIFIHeadlessRule),
ContainsIPCIDRRule: HasHeadlessRule(headlessRules, isIPCIDRHeadlessRule),
ContainsDNSQueryTypeRule: HasHeadlessRule(headlessRules, isDNSQueryTypeHeadlessRule),
ContainsNonIPCIDRRule: HasHeadlessRule(headlessRules, isNonIPCIDRHeadlessRule),
}
}
func validateRuleSetMetadataUpdate(ctx context.Context, tag string, metadata adapter.RuleSetMetadata) error {
validator := service.FromContext[adapter.DNSRuleSetUpdateValidator](ctx)
if validator == nil {
return nil
}
return validator.ValidateRuleSetMetadataUpdate(tag, metadata)
}
+14 -28
View File
@@ -2,7 +2,6 @@ package rule
import (
"context"
"os"
"path/filepath"
"strings"
"sync"
@@ -39,11 +38,11 @@ type LocalRuleSet struct {
refs atomic.Int32
}
func NewLocalRuleSet(ctx context.Context, logger logger.Logger, options option.RuleSet) (*LocalRuleSet, error) {
func NewLocalRuleSet(ctx context.Context, logger logger.Logger, tag string, options option.RuleSet) (*LocalRuleSet, error) {
ruleSet := &LocalRuleSet{
ctx: ctx,
logger: logger,
tag: options.Tag,
tag: tag,
fileFormat: options.Format,
}
if options.Type == C.RuleSetTypeInline {
@@ -55,7 +54,7 @@ func NewLocalRuleSet(ctx context.Context, logger logger.Logger, options option.R
return nil, err
}
} else {
filePath := filemanager.BasePath(ctx, options.LocalOptions.Path)
filePath := filemanager.BasePath(ctx, strings.ReplaceAll(options.LocalOptions.Path, C.RuleSetTagPlaceholder, tag))
filePath, _ = filepath.Abs(filePath)
err := ruleSet.reloadFile(filePath)
if err != nil {
@@ -66,7 +65,7 @@ func NewLocalRuleSet(ctx context.Context, logger logger.Logger, options option.R
Callback: func(path string) {
uErr := ruleSet.reloadFile(path)
if uErr != nil {
logger.Error(E.Cause(uErr, "reload rule-set ", options.Tag))
logger.Error(E.Cause(uErr, "reload rule-set ", tag))
}
},
})
@@ -100,7 +99,7 @@ func (s *LocalRuleSet) reloadFile(path string) error {
var ruleSet option.PlainRuleSetCompat
switch s.fileFormat {
case C.RuleSetFormatSource, "":
content, err := os.ReadFile(path)
content, err := filemanager.ReadFile(s.ctx, path)
if err != nil {
return err
}
@@ -110,7 +109,7 @@ func (s *LocalRuleSet) reloadFile(path string) error {
}
case C.RuleSetFormatBinary:
setFile, err := os.Open(path)
setFile, err := filemanager.Open(s.ctx, path)
if err != nil {
return err
}
@@ -138,10 +137,11 @@ func (s *LocalRuleSet) reloadRules(headlessRules []option.HeadlessRule) error {
return E.Cause(err, "parse rule_set.rules.[", i, "]")
}
}
var metadata adapter.RuleSetMetadata
metadata.ContainsProcessRule = HasHeadlessRule(headlessRules, isProcessHeadlessRule)
metadata.ContainsWIFIRule = HasHeadlessRule(headlessRules, isWIFIHeadlessRule)
metadata.ContainsIPCIDRRule = HasHeadlessRule(headlessRules, isIPCIDRHeadlessRule)
metadata := buildRuleSetMetadata(headlessRules)
err = validateRuleSetMetadataUpdate(s.ctx, s.tag, metadata)
if err != nil {
return err
}
s.access.Lock()
s.rules = rules
s.metadata = metadata
@@ -153,10 +153,6 @@ func (s *LocalRuleSet) reloadRules(headlessRules []option.HeadlessRule) error {
return nil
}
func (s *LocalRuleSet) PostStart() error {
return nil
}
func (s *LocalRuleSet) Metadata() adapter.RuleSetMetadata {
s.access.RLock()
defer s.access.RUnlock()
@@ -203,19 +199,9 @@ func (s *LocalRuleSet) Close() error {
}
func (s *LocalRuleSet) Match(metadata *adapter.InboundContext) bool {
return !s.matchStates(metadata).isEmpty()
return matchAnyHeadlessRule(s.rules, metadata)
}
func (s *LocalRuleSet) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return s.matchStatesWithBase(metadata, 0)
}
func (s *LocalRuleSet) matchStatesWithBase(metadata *adapter.InboundContext, base ruleMatchState) ruleMatchStateSet {
var stateSet ruleMatchStateSet
for _, rule := range s.rules {
nestedMetadata := *metadata
nestedMetadata.ResetRuleMatchCache()
stateSet = stateSet.merge(matchHeadlessRuleStatesWithBase(rule, &nestedMetadata, base))
}
return stateSet
func (s *LocalRuleSet) mergeableRule() *DefaultHeadlessRule {
return mergeableRuleIn(s.rules)
}
+104 -107
View File
@@ -3,11 +3,10 @@ package rule
import (
"bytes"
"context"
"crypto/tls"
"crypto/sha256"
"io"
"net"
"net/http"
"runtime"
"path/filepath"
"strings"
"sync"
"sync/atomic"
@@ -16,17 +15,16 @@ import (
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/srs"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/experimental/deprecated"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing/common"
E "github.com/sagernet/sing/common/exceptions"
F "github.com/sagernet/sing/common/format"
"github.com/sagernet/sing/common/json"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/ntp"
"github.com/sagernet/sing/common/x/list"
"github.com/sagernet/sing/service"
"github.com/sagernet/sing/service/filemanager"
"github.com/sagernet/sing/service/pause"
"go4.org/netipx"
@@ -39,22 +37,25 @@ type RemoteRuleSet struct {
cancel context.CancelFunc
logger logger.ContextLogger
outbound adapter.OutboundManager
tag string
url string
urlHash [32]byte
initialPath string
options option.RuleSet
updateInterval time.Duration
dialer N.Dialer
httpClient *http.Client
access sync.RWMutex
rules []adapter.HeadlessRule
metadata adapter.RuleSetMetadata
lastUpdated time.Time
lastEtag string
updateTicker *time.Ticker
cacheFile adapter.CacheFile
pauseManager pause.Manager
callbacks list.List[adapter.RuleSetUpdateCallback]
refs atomic.Int32
}
func NewRemoteRuleSet(ctx context.Context, logger logger.ContextLogger, options option.RuleSet) *RemoteRuleSet {
func NewRemoteRuleSet(ctx context.Context, logger logger.ContextLogger, tag string, options option.RuleSet) (*RemoteRuleSet, error) {
ctx, cancel := context.WithCancel(ctx)
var updateInterval time.Duration
if options.RemoteOptions.UpdateInterval > 0 {
@@ -62,19 +63,29 @@ func NewRemoteRuleSet(ctx context.Context, logger logger.ContextLogger, options
} else {
updateInterval = 24 * time.Hour
}
var initialPath string
if options.RemoteOptions.InitialPath != "" {
initialPath = filemanager.BasePath(ctx, strings.ReplaceAll(options.RemoteOptions.InitialPath, C.RuleSetTagPlaceholder, tag))
initialPath, _ = filepath.Abs(initialPath)
}
url := strings.ReplaceAll(options.RemoteOptions.URL, C.RuleSetTagPlaceholder, tag)
return &RemoteRuleSet{
ctx: ctx,
cancel: cancel,
outbound: service.FromContext[adapter.OutboundManager](ctx),
logger: logger,
tag: tag,
url: url,
urlHash: sha256.Sum256([]byte(url)),
initialPath: initialPath,
options: options,
updateInterval: updateInterval,
pauseManager: service.FromContext[pause.Manager](ctx),
}
}, nil
}
func (s *RemoteRuleSet) Name() string {
return s.options.Tag
return s.tag
}
func (s *RemoteRuleSet) String() string {
@@ -83,40 +94,47 @@ func (s *RemoteRuleSet) String() string {
func (s *RemoteRuleSet) StartContext(ctx context.Context, startContext *adapter.HTTPStartContext) error {
s.cacheFile = service.FromContext[adapter.CacheFile](s.ctx)
var dialer N.Dialer
if s.options.RemoteOptions.DownloadDetour != "" {
outbound, loaded := s.outbound.Outbound(s.options.RemoteOptions.DownloadDetour)
if !loaded {
return E.New("download detour not found: ", s.options.RemoteOptions.DownloadDetour)
}
dialer = outbound
} else {
dialer = s.outbound.Default()
transport, err := s.resolveTransport()
if err != nil {
return E.Cause(err, "create rule-set http client")
}
s.dialer = dialer
startContext.Register(transport)
s.httpClient = &http.Client{Transport: transport}
if s.cacheFile != nil {
if savedSet := s.cacheFile.LoadRuleSet(s.options.Tag); savedSet != nil {
err := s.loadBytes(savedSet.Content)
if err != nil {
s.logger.Warn(E.Cause(err, "restore cached rule-set, will refetch"))
savedSet := s.cacheFile.LoadRuleSet(s.tag)
if savedSet != nil {
if len(savedSet.URLHash) > 0 && !bytes.Equal(savedSet.URLHash, s.urlHash[:]) {
s.logger.Info("cached rule-set was downloaded from another URL, will refetch")
} else {
s.lastUpdated = savedSet.LastUpdated
s.lastEtag = savedSet.LastEtag
err = s.loadBytes(savedSet.Content)
if err != nil {
s.logger.Warn(E.Cause(err, "restore cached rule-set, will refetch"))
} else {
s.lastUpdated = savedSet.LastUpdated
s.lastEtag = savedSet.LastEtag
}
}
}
}
if s.lastUpdated.IsZero() {
err := s.fetch(ctx, startContext)
var loadedFromInitialPath bool
if s.lastUpdated.IsZero() && s.initialPath != "" {
var content []byte
content, err = filemanager.ReadFile(s.ctx, s.initialPath)
if err == nil {
err = s.loadBytes(content)
}
if err != nil {
return E.Cause(err, "initial rule-set: ", s.options.Tag)
s.logger.Warn(E.Cause(err, "load initial rule-set from ", s.initialPath))
} else {
loadedFromInitialPath = true
}
}
if s.lastUpdated.IsZero() && !loadedFromInitialPath {
err = s.fetch(ctx, true)
if err != nil {
return E.Cause(err, "initial rule-set: ", s.tag)
}
}
s.updateTicker = time.NewTicker(s.updateInterval)
return nil
}
func (s *RemoteRuleSet) PostStart() error {
go s.loopUpdate()
return nil
}
@@ -190,10 +208,13 @@ func (s *RemoteRuleSet) loadBytes(content []byte) error {
return E.Cause(err, "parse rule_set.rules.[", i, "]")
}
}
metadata := buildRuleSetMetadata(plainRuleSet.Rules)
err = validateRuleSetMetadataUpdate(s.ctx, s.tag, metadata)
if err != nil {
return err
}
s.access.Lock()
s.metadata.ContainsProcessRule = HasHeadlessRule(plainRuleSet.Rules, isProcessHeadlessRule)
s.metadata.ContainsWIFIRule = HasHeadlessRule(plainRuleSet.Rules, isWIFIHeadlessRule)
s.metadata.ContainsIPCIDRRule = HasHeadlessRule(plainRuleSet.Rules, isIPCIDRHeadlessRule)
s.metadata = metadata
s.rules = rules
callbacks := s.callbacks.Array()
s.access.Unlock()
@@ -203,139 +224,115 @@ func (s *RemoteRuleSet) loadBytes(content []byte) error {
return nil
}
func (s *RemoteRuleSet) loopUpdate() {
if time.Since(s.lastUpdated) > s.updateInterval {
err := s.fetch(s.ctx, nil)
if err != nil {
s.logger.Error("fetch rule-set ", s.options.Tag, ": ", err)
} else if s.refs.Load() == 0 {
s.rules = nil
}
}
for {
runtime.GC()
select {
case <-s.ctx.Done():
return
case <-s.updateTicker.C:
s.updateOnce()
}
}
}
func (s *RemoteRuleSet) updateOnce() {
err := s.fetch(s.ctx, nil)
err := s.fetch(s.ctx, false)
if err != nil {
s.logger.Error("fetch rule-set ", s.options.Tag, ": ", err)
s.logger.Error("fetch rule-set ", s.tag, ": ", err)
} else if s.refs.Load() == 0 {
s.rules = nil
}
}
func (s *RemoteRuleSet) fetch(ctx context.Context, startContext *adapter.HTTPStartContext) error {
s.logger.Debug("updating rule-set ", s.options.Tag, " from URL: ", s.options.RemoteOptions.URL)
var httpClient *http.Client
if startContext != nil {
httpClient = startContext.HTTPClient(s.options.RemoteOptions.DownloadDetour, s.dialer)
} else {
httpClient = &http.Client{
Transport: &http.Transport{
ForceAttemptHTTP2: true,
TLSHandshakeTimeout: C.TCPTimeout,
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
return s.dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
},
TLSClientConfig: &tls.Config{
Time: ntp.TimeFuncFromContext(s.ctx),
RootCAs: adapter.RootPoolFromContext(s.ctx),
},
},
}
}
request, err := http.NewRequest("GET", s.options.RemoteOptions.URL, nil)
func (s *RemoteRuleSet) fetch(ctx context.Context, isStart bool) error {
s.logger.Debug("updating rule-set ", s.tag, " from URL: ", s.url)
request, err := http.NewRequest("GET", s.url, nil)
if err != nil {
return err
}
if s.lastEtag != "" {
request.Header.Set("If-None-Match", s.lastEtag)
}
response, err := httpClient.Do(request.WithContext(ctx))
if !isStart {
defer s.httpClient.CloseIdleConnections()
}
response, err := s.httpClient.Do(request.WithContext(ctx))
if err != nil {
return err
}
defer response.Body.Close()
switch response.StatusCode {
case http.StatusOK:
case http.StatusNotModified:
s.lastUpdated = time.Now()
if s.cacheFile != nil {
savedRuleSet := s.cacheFile.LoadRuleSet(s.options.Tag)
savedRuleSet := s.cacheFile.LoadRuleSet(s.tag)
if savedRuleSet != nil {
savedRuleSet.LastUpdated = s.lastUpdated
err = s.cacheFile.SaveRuleSet(s.options.Tag, savedRuleSet)
savedRuleSet.URLHash = s.urlHash[:]
err = s.cacheFile.SaveRuleSet(s.tag, savedRuleSet)
if err != nil {
s.logger.Error("save rule-set updated time: ", err)
return nil
}
}
}
s.logger.Notice("update rule-set ", s.options.Tag, ": not modified")
s.logger.Notice("update rule-set ", s.tag, ": not modified")
return nil
default:
return E.New("unexpected status: ", response.Status)
}
content, err := io.ReadAll(response.Body)
if err != nil {
response.Body.Close()
return err
}
err = s.loadBytes(content)
if err != nil {
response.Body.Close()
return err
}
response.Body.Close()
eTagHeader := response.Header.Get("Etag")
if eTagHeader != "" {
s.lastEtag = eTagHeader
}
s.lastUpdated = time.Now()
if s.cacheFile != nil {
err = s.cacheFile.SaveRuleSet(s.options.Tag, &adapter.SavedBinary{
err = s.cacheFile.SaveRuleSet(s.tag, &adapter.SavedBinary{
LastUpdated: s.lastUpdated,
Content: content,
LastEtag: s.lastEtag,
URLHash: s.urlHash[:],
})
if err != nil {
s.logger.Error("save rule-set cache: ", err)
}
}
s.logger.Notice("updated rule-set ", s.options.Tag)
s.logger.Notice("updated rule-set ", s.tag)
return nil
}
func (s *RemoteRuleSet) resolveTransport() (adapter.HTTPTransport, error) {
httpClientManager := service.FromContext[adapter.HTTPClientManager](s.ctx)
if s.options.RemoteOptions.HTTPClient != nil && !s.options.RemoteOptions.HTTPClient.IsEmpty() {
if s.options.RemoteOptions.DownloadDetour != "" { //nolint:staticcheck
return nil, E.New("http_client is conflict with deprecated download_detour field")
}
return httpClientManager.ResolveTransport(s.ctx, s.logger, *s.options.RemoteOptions.HTTPClient)
}
if s.options.RemoteOptions.DownloadDetour != "" { //nolint:staticcheck
deprecated.Report(s.ctx, deprecated.OptionLegacyRuleSetDownloadDetour)
return httpClientManager.ResolveTransport(s.ctx, s.logger, option.HTTPClientOptions{
DialerOptions: option.DialerOptions{
Detour: s.options.RemoteOptions.DownloadDetour, //nolint:staticcheck
},
DisableEmptyDirectCheck: true,
})
}
defaultTransport := httpClientManager.DefaultTransport()
if defaultTransport == nil {
return nil, E.New("default http client transport is not initialized")
}
return defaultTransport, nil
}
func (s *RemoteRuleSet) Close() error {
s.rules = nil
s.cancel()
if s.updateTicker != nil {
s.updateTicker.Stop()
}
return nil
}
func (s *RemoteRuleSet) Match(metadata *adapter.InboundContext) bool {
return !s.matchStates(metadata).isEmpty()
return matchAnyHeadlessRule(s.rules, metadata)
}
func (s *RemoteRuleSet) matchStates(metadata *adapter.InboundContext) ruleMatchStateSet {
return s.matchStatesWithBase(metadata, 0)
}
func (s *RemoteRuleSet) matchStatesWithBase(metadata *adapter.InboundContext, base ruleMatchState) ruleMatchStateSet {
var stateSet ruleMatchStateSet
for _, rule := range s.rules {
nestedMetadata := *metadata
nestedMetadata.ResetRuleMatchCache()
stateSet = stateSet.merge(matchHeadlessRuleStatesWithBase(rule, &nestedMetadata, base))
}
return stateSet
func (s *RemoteRuleSet) mergeableRule() *DefaultHeadlessRule {
return mergeableRuleIn(s.rules)
}
+690 -30
View File
@@ -2,6 +2,7 @@ package rule
import (
"context"
"net"
"net/netip"
"strings"
"testing"
@@ -9,11 +10,11 @@ import (
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/convertor/adguard"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/option"
slogger "github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
mDNS "github.com/miekg/dns"
"github.com/stretchr/testify/require"
)
@@ -294,11 +295,11 @@ func TestRouteRuleSetOrSemantics(t *testing.T) {
})
require.True(t, rule.Match(&metadata))
})
t.Run("later rule in same set can satisfy outer group", func(t *testing.T) {
t.Run("multi rule set does not satisfy outer group", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest(
"rule-set-or",
"rule-set-and",
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addOtherItem(rule, NewNetworkItem([]string{N.NetworkTCP}))
}),
@@ -310,7 +311,7 @@ func TestRouteRuleSetOrSemantics(t *testing.T) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
})
require.True(t, rule.Match(&metadata))
require.False(t, rule.Match(&metadata))
})
t.Run("cross ruleset union is not allowed", func(t *testing.T) {
t.Parallel()
@@ -332,7 +333,7 @@ func TestRouteRuleSetOrSemantics(t *testing.T) {
func TestRouteRuleSetLogicalSemantics(t *testing.T) {
t.Parallel()
t.Run("logical or keeps all successful branch states", func(t *testing.T) {
t.Run("logical set does not satisfy outer group", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest("logical-or", headlessLogicalRule(
@@ -349,9 +350,9 @@ func TestRouteRuleSetLogicalSemantics(t *testing.T) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
})
require.True(t, rule.Match(&metadata))
require.False(t, rule.Match(&metadata))
})
t.Run("logical and unions child states", func(t *testing.T) {
t.Run("logical branch does not lift outer group requirements", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest("logical-and", headlessLogicalRule(
@@ -369,9 +370,28 @@ func TestRouteRuleSetLogicalSemantics(t *testing.T) {
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
addSourcePortItem(rule, []uint16{2000})
})
require.False(t, rule.Match(&metadata))
})
t.Run("logical branch matches on its own conditions", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest("logical-and-self", headlessLogicalRule(
C.LogicalTypeAnd,
false,
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"example.com"})
}),
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addSourcePortItem(rule, []uint16{1000})
}),
))
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
addSourcePortItem(rule, []uint16{1000})
})
require.True(t, rule.Match(&metadata))
})
t.Run("invert success does not contribute positive state", func(t *testing.T) {
t.Run("inverted set does not satisfy outer group", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest("invert", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
@@ -386,9 +406,240 @@ func TestRouteRuleSetLogicalSemantics(t *testing.T) {
})
}
func TestRouteRuleSetInvertMergedBranchSemantics(t *testing.T) {
func TestRuleSetShapeBoundary(t *testing.T) {
t.Parallel()
t.Run("default invert keeps inherited group outside grouped predicate", func(t *testing.T) {
buildOuter := func(ruleSet *LocalRuleSet) *DefaultRule {
return routeRuleForTest(func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, []string{"extra.example.org"}, nil)
addDestinationPortItem(rule, []uint16{443})
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
}
singleShape := buildOuter(newLocalRuleSetForTest("flat-single", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"a.example.com", "b.example.com"})
})))
multiShape := buildOuter(newLocalRuleSetForTest(
"flat-multi",
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"a.example.com"})
}),
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"b.example.com"})
}),
))
queries := []struct {
name string
domain string
port uint16
singleResult bool
multiResult bool
}{
{"in set", "www.b.example.com", 443, true, false},
{"outer own domain", "extra.example.org", 443, true, false},
{"neither", "other.example.net", 443, false, false},
{"in set with wrong port", "www.b.example.com", 80, false, false},
}
for _, query := range queries {
t.Run(query.name, func(t *testing.T) {
t.Parallel()
singleMetadata := testMetadata(query.domain)
singleMetadata.Destination.Port = query.port
multiMetadata := testMetadata(query.domain)
multiMetadata.Destination.Port = query.port
require.Equal(t, query.singleResult, singleShape.Match(&singleMetadata))
require.Equal(t, query.multiResult, multiShape.Match(&multiMetadata))
})
}
}
func TestRuleSetCaseBoundary(t *testing.T) {
t.Parallel()
t.Run("single cross group rule merges into outer", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest("port-single", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationPortRangeItem(t, rule, []string{"400:500"})
}))
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
addDestinationPortItem(rule, []uint16{8080})
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
require.True(t, rule.Match(&metadata))
})
t.Run("multi rule set keeps outer group absolute", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest(
"port-multi",
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationPortRangeItem(t, rule, []string{"400:500"})
}),
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"never.example"})
}),
)
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
addDestinationPortItem(rule, []uint16{8080})
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
require.False(t, rule.Match(&metadata))
})
}
func TestRuleSetLogicalBranchSelfContained(t *testing.T) {
t.Parallel()
newRuleSet := func() *LocalRuleSet {
return newLocalRuleSetForTest("and-branch", headlessLogicalRule(
C.LogicalTypeAnd,
false,
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"b.example.com"})
}),
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationPortRangeItem(t, rule, []string{"800:900"})
}),
))
}
t.Run("branch matches only on its own conditions", func(t *testing.T) {
t.Parallel()
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{newRuleSet()}})
})
matchedMetadata := testMetadata("www.b.example.com")
matchedMetadata.Destination.Port = 850
require.True(t, rule.Match(&matchedMetadata))
unmatchedMetadata := testMetadata("www.b.example.com")
require.False(t, rule.Match(&unmatchedMetadata))
})
t.Run("outer condition and set are both required", func(t *testing.T) {
t.Parallel()
matchedRule := routeRuleForTest(func(rule *abstractDefaultRule) {
addSourcePortItem(rule, []uint16{1000})
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{newRuleSet()}})
})
metadata := testMetadata("www.b.example.com")
metadata.Destination.Port = 850
require.True(t, matchedRule.Match(&metadata))
setMissMetadata := testMetadata("other.example.net")
setMissMetadata.Destination.Port = 850
require.False(t, matchedRule.Match(&setMissMetadata))
outerMissRule := routeRuleForTest(func(rule *abstractDefaultRule) {
addSourcePortItem(rule, []uint16{2000})
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{newRuleSet()}})
})
outerMissMetadata := testMetadata("www.b.example.com")
outerMissMetadata.Destination.Port = 850
require.False(t, outerMissRule.Match(&outerMissMetadata))
})
}
func TestRuleSetMixedReference(t *testing.T) {
t.Parallel()
t.Run("single rule set merges as or", func(t *testing.T) {
t.Parallel()
ruleSet := newLocalRuleSetForTest("mixed-single", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"set.example.com"})
}))
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, []string{"extra.example.org"}, nil)
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
extraMetadata := testMetadata("extra.example.org")
require.True(t, rule.Match(&extraMetadata))
setMetadata := testMetadata("www.set.example.com")
require.True(t, rule.Match(&setMetadata))
otherMetadata := testMetadata("other.example.net")
require.False(t, rule.Match(&otherMetadata))
})
t.Run("multi rule set is an independent condition", func(t *testing.T) {
t.Parallel()
ruleSet := newLocalRuleSetForTest(
"mixed-multi",
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"set.example.com"})
}),
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"set2.example.com"})
}),
)
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
addDestinationPortItem(rule, []uint16{443})
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
matchedMetadata := testMetadata("www.set2.example.com")
require.True(t, rule.Match(&matchedMetadata))
portMissMetadata := testMetadata("www.set2.example.com")
portMissMetadata.Destination.Port = 80
require.False(t, rule.Match(&portMissMetadata))
setMissMetadata := testMetadata("other.example.net")
require.False(t, rule.Match(&setMissMetadata))
})
}
func TestRuleSetStandaloneReference(t *testing.T) {
t.Parallel()
ruleSet := newLocalRuleSetForTest(
"standalone",
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"a.example.net"})
}),
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationPortRangeItem(t, rule, []string{"400:500"})
}),
)
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
domainMetadata := testMetadata("www.a.example.net")
domainMetadata.Destination.Port = 80
require.True(t, rule.Match(&domainMetadata))
portMetadata := testMetadata("other.example.org")
require.True(t, rule.Match(&portMetadata))
missMetadata := testMetadata("other.example.org")
missMetadata.Destination.Port = 80
require.False(t, rule.Match(&missMetadata))
}
func TestRuleSetInvertedSingleRuleIsBoolean(t *testing.T) {
t.Parallel()
ruleSet := newLocalRuleSetForTest("inverted-single", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
rule.invert = true
addDestinationAddressItem(t, rule, nil, []string{"blocked.example"})
}))
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
allowedMetadata := testMetadata("good.example.org")
require.True(t, rule.Match(&allowedMetadata))
blockedMetadata := testMetadata("www.blocked.example")
require.False(t, rule.Match(&blockedMetadata))
}
func TestRuleSetEmptySetNeverMatches(t *testing.T) {
t.Parallel()
emptySet := newLocalRuleSetForTest("empty")
t.Run("outer own group does not bypass the set", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"example.com"})
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{emptySet}})
})
require.False(t, rule.Match(&metadata))
})
t.Run("standalone reference does not match", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{emptySet}})
})
require.False(t, rule.Match(&metadata))
})
}
func TestRouteRuleSetInvertBranchSemantics(t *testing.T) {
t.Parallel()
t.Run("inverted default branch acts as boolean term", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest("invert-grouped", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
@@ -401,7 +652,7 @@ func TestRouteRuleSetInvertMergedBranchSemantics(t *testing.T) {
})
require.True(t, rule.Match(&metadata))
})
t.Run("default invert keeps inherited group after negation succeeds", func(t *testing.T) {
t.Run("inverted default branch with non grouped condition acts as boolean term", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest("invert-network", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
@@ -414,7 +665,7 @@ func TestRouteRuleSetInvertMergedBranchSemantics(t *testing.T) {
})
require.True(t, rule.Match(&metadata))
})
t.Run("logical invert keeps inherited group outside grouped predicate", func(t *testing.T) {
t.Run("inverted logical branch acts as boolean term", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest("logical-invert-grouped", headlessLogicalRule(
@@ -430,7 +681,7 @@ func TestRouteRuleSetInvertMergedBranchSemantics(t *testing.T) {
})
require.True(t, rule.Match(&metadata))
})
t.Run("logical invert keeps inherited group after negation succeeds", func(t *testing.T) {
t.Run("inverted logical branch with non grouped condition acts as boolean term", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest("logical-invert-network", headlessLogicalRule(
@@ -497,21 +748,26 @@ func TestDefaultRuleDoesNotReuseGroupedMatchCacheAcrossEvaluations(t *testing.T)
func TestRouteRuleSetRemoteUsesSameSemantics(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newRemoteRuleSetForTest(
"remote",
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addOtherItem(rule, NewNetworkItem([]string{N.NetworkTCP}))
addOtherItem(rule, NewNetworkItem([]string{N.NetworkUDP}))
}),
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"example.com"})
}),
)
rule := routeRuleForTest(func(rule *abstractDefaultRule) {
standaloneRule := routeRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
standaloneMetadata := testMetadata("www.example.com")
require.True(t, standaloneRule.Match(&standaloneMetadata))
combinedRule := routeRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
})
require.True(t, rule.Match(&metadata))
combinedMetadata := testMetadata("www.example.com")
require.False(t, combinedRule.Match(&combinedMetadata))
}
func TestDNSRuleSetSemantics(t *testing.T) {
@@ -540,7 +796,7 @@ func TestDNSRuleSetSemantics(t *testing.T) {
})
require.False(t, rule.Match(&metadata))
})
t.Run("outer destination group stays outside inverted grouped branch", func(t *testing.T) {
t.Run("inverted branch acts as boolean term", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.baidu.com")
ruleSet := newLocalRuleSetForTest("dns-invert-grouped", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
@@ -553,7 +809,7 @@ func TestDNSRuleSetSemantics(t *testing.T) {
})
require.True(t, rule.Match(&metadata))
})
t.Run("outer destination group stays outside inverted logical branch", func(t *testing.T) {
t.Run("inverted logical branch acts as boolean term", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest("dns-logical-invert-network", headlessLogicalRule(
@@ -579,7 +835,7 @@ func TestDNSRuleSetSemantics(t *testing.T) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
})
require.True(t, rule.MatchAddressLimit(&metadata))
require.True(t, rule.MatchAddressLimit(&metadata, dnsResponseForTest(netip.MustParseAddr("203.0.113.1"))))
})
t.Run("dns keeps ruleset or semantics", func(t *testing.T) {
t.Parallel()
@@ -594,7 +850,7 @@ func TestDNSRuleSetSemantics(t *testing.T) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{emptyStateSet, destinationStateSet}})
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
})
require.True(t, rule.MatchAddressLimit(&metadata))
require.True(t, rule.MatchAddressLimit(&metadata, dnsResponseForTest(netip.MustParseAddr("203.0.113.1"))))
})
t.Run("ruleset ip cidr flags stay scoped", func(t *testing.T) {
t.Parallel()
@@ -608,10 +864,381 @@ func TestDNSRuleSetSemantics(t *testing.T) {
ipCidrAcceptEmpty: true,
})
})
require.True(t, rule.MatchAddressLimit(&metadata))
require.True(t, rule.MatchAddressLimit(&metadata, dnsResponseForTest(netip.MustParseAddr("203.0.113.1"))))
require.False(t, rule.MatchAddressLimit(&metadata, dnsResponseForTest(netip.MustParseAddr("8.8.8.8"))))
require.True(t, rule.MatchAddressLimit(&metadata, dnsResponseForTest()))
require.False(t, metadata.IPCIDRMatchSource)
require.False(t, metadata.IPCIDRAcceptEmpty)
})
t.Run("pre lookup ruleset only deferred fields fail closed", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("lookup.example")
ruleSet := newLocalRuleSetForTest("dns-prelookup-ipcidr", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
}))
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
// This is accepted without match_response so mixed rule_set deployments keep
// working; the destination-IP-only branch simply cannot match before a DNS
// response is available.
require.False(t, rule.Match(&metadata))
})
t.Run("pre lookup ruleset destination cidr does not fall back to other predicates", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("lookup.example")
ruleSet := newLocalRuleSetForTest("dns-prelookup-network-and-ipcidr", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addOtherItem(rule, NewNetworkItem([]string{N.NetworkTCP}))
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
}))
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
require.False(t, rule.Match(&metadata))
})
t.Run("pre lookup mixed ruleset still matches non response branch", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("www.example.com")
ruleSet := newLocalRuleSetForTest(
"dns-prelookup-mixed",
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addOtherItem(rule, NewNetworkItem([]string{N.NetworkTCP}))
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
}),
headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationAddressItem(t, rule, nil, []string{"example.com"})
}),
)
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
// Destination-IP predicates inside rule_set fail closed before the DNS response,
// but they must not force validation errors or suppress sibling non-response
// branches.
require.True(t, rule.Match(&metadata))
})
}
func TestDNSMatchResponseRuleSetDestinationCIDRUsesDNSResponse(t *testing.T) {
t.Parallel()
ruleSet := newLocalRuleSetForTest("dns-response-ipcidr", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
}))
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
rule.matchResponse = true
matchedMetadata := testMetadata("lookup.example")
matchedMetadata.DNSResponse = dnsResponseForTest(netip.MustParseAddr("203.0.113.1"))
require.True(t, rule.Match(&matchedMetadata))
require.Empty(t, matchedMetadata.DestinationAddresses)
unmatchedMetadata := testMetadata("lookup.example")
unmatchedMetadata.DNSResponse = dnsResponseForTest(netip.MustParseAddr("8.8.8.8"))
require.False(t, rule.Match(&unmatchedMetadata))
}
func TestDNSMatchResponseMissingResponseUsesBooleanSemantics(t *testing.T) {
t.Parallel()
t.Run("plain rule remains false", func(t *testing.T) {
t.Parallel()
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {})
rule.matchResponse = true
metadata := testMetadata("lookup.example")
require.False(t, rule.Match(&metadata))
})
t.Run("invert rule becomes true", func(t *testing.T) {
t.Parallel()
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
rule.invert = true
})
rule.matchResponse = true
metadata := testMetadata("lookup.example")
require.True(t, rule.Match(&metadata))
})
t.Run("logical wrapper respects inverted child", func(t *testing.T) {
t.Parallel()
nestedRule := dnsRuleForTest(func(rule *abstractDefaultRule) {
rule.invert = true
})
nestedRule.matchResponse = true
logicalRule := &LogicalDNSRule{
abstractLogicalRule: abstractLogicalRule{
rules: []adapter.HeadlessRule{nestedRule},
mode: C.LogicalTypeAnd,
},
}
metadata := testMetadata("lookup.example")
require.True(t, logicalRule.Match(&metadata))
})
}
func TestDNSAddressLimitIgnoresDestinationAddresses(t *testing.T) {
t.Parallel()
testCases := []struct {
name string
build func(*testing.T, *abstractDefaultRule)
matchedResponse *mDNS.Msg
unmatchedResponse *mDNS.Msg
}{
{
name: "ip_cidr",
build: func(t *testing.T, rule *abstractDefaultRule) {
t.Helper()
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
},
matchedResponse: dnsResponseForTest(netip.MustParseAddr("203.0.113.1")),
unmatchedResponse: dnsResponseForTest(netip.MustParseAddr("8.8.8.8")),
},
{
name: "ip_is_private",
build: func(t *testing.T, rule *abstractDefaultRule) {
t.Helper()
addDestinationIPIsPrivateItem(rule)
},
matchedResponse: dnsResponseForTest(netip.MustParseAddr("10.0.0.1")),
unmatchedResponse: dnsResponseForTest(netip.MustParseAddr("8.8.8.8")),
},
{
name: "ip_accept_any",
build: func(t *testing.T, rule *abstractDefaultRule) {
t.Helper()
addDestinationIPAcceptAnyItem(rule)
},
matchedResponse: dnsResponseForTest(netip.MustParseAddr("203.0.113.1")),
unmatchedResponse: dnsResponseForTest(),
},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
t.Parallel()
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
testCase.build(t, rule)
})
mismatchMetadata := testMetadata("lookup.example")
mismatchMetadata.DestinationAddresses = []netip.Addr{netip.MustParseAddr("203.0.113.1")}
require.False(t, rule.MatchAddressLimit(&mismatchMetadata, testCase.unmatchedResponse))
matchMetadata := testMetadata("lookup.example")
matchMetadata.DestinationAddresses = []netip.Addr{netip.MustParseAddr("8.8.8.8")}
require.True(t, rule.MatchAddressLimit(&matchMetadata, testCase.matchedResponse))
})
}
}
func TestDNSLegacyAddressLimitPreLookupDefersDirectRules(t *testing.T) {
t.Parallel()
testCases := []struct {
name string
build func(*testing.T, *abstractDefaultRule)
matchedResponse *mDNS.Msg
unmatchedResponse *mDNS.Msg
}{
{
name: "ip_cidr",
build: func(t *testing.T, rule *abstractDefaultRule) {
t.Helper()
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
},
matchedResponse: dnsResponseForTest(netip.MustParseAddr("203.0.113.1")),
unmatchedResponse: dnsResponseForTest(netip.MustParseAddr("8.8.8.8")),
},
{
name: "ip_is_private",
build: func(t *testing.T, rule *abstractDefaultRule) {
t.Helper()
addDestinationIPIsPrivateItem(rule)
},
matchedResponse: dnsResponseForTest(netip.MustParseAddr("10.0.0.1")),
unmatchedResponse: dnsResponseForTest(netip.MustParseAddr("8.8.8.8")),
},
{
name: "ip_accept_any",
build: func(t *testing.T, rule *abstractDefaultRule) {
t.Helper()
addDestinationIPAcceptAnyItem(rule)
},
matchedResponse: dnsResponseForTest(netip.MustParseAddr("203.0.113.1")),
unmatchedResponse: dnsResponseForTest(),
},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
t.Parallel()
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
testCase.build(t, rule)
})
preLookupMetadata := testMetadata("lookup.example")
require.True(t, rule.LegacyPreMatch(&preLookupMetadata))
matchedMetadata := testMetadata("lookup.example")
require.True(t, rule.MatchAddressLimit(&matchedMetadata, testCase.matchedResponse))
unmatchedMetadata := testMetadata("lookup.example")
require.False(t, rule.MatchAddressLimit(&unmatchedMetadata, testCase.unmatchedResponse))
})
}
}
func TestDNSLegacyAddressLimitPreLookupDefersRuleSetDestinationCIDR(t *testing.T) {
t.Parallel()
ruleSet := newLocalRuleSetForTest("dns-legacy-ipcidr", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
}))
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
preLookupMetadata := testMetadata("lookup.example")
require.True(t, rule.LegacyPreMatch(&preLookupMetadata))
matchedMetadata := testMetadata("lookup.example")
require.True(t, rule.MatchAddressLimit(&matchedMetadata, dnsResponseForTest(netip.MustParseAddr("203.0.113.1"))))
unmatchedMetadata := testMetadata("lookup.example")
require.False(t, rule.MatchAddressLimit(&unmatchedMetadata, dnsResponseForTest(netip.MustParseAddr("8.8.8.8"))))
}
func TestDNSLegacyLogicalAddressLimitPreLookupDefersNestedRules(t *testing.T) {
t.Parallel()
nestedRule := dnsRuleForTest(func(rule *abstractDefaultRule) {
addDestinationIPIsPrivateItem(rule)
})
logicalRule := &LogicalDNSRule{
abstractLogicalRule: abstractLogicalRule{
rules: []adapter.HeadlessRule{nestedRule},
mode: C.LogicalTypeAnd,
},
}
preLookupMetadata := testMetadata("lookup.example")
require.True(t, logicalRule.LegacyPreMatch(&preLookupMetadata))
matchedMetadata := testMetadata("lookup.example")
require.True(t, logicalRule.MatchAddressLimit(&matchedMetadata, dnsResponseForTest(netip.MustParseAddr("10.0.0.1"))))
unmatchedMetadata := testMetadata("lookup.example")
require.False(t, logicalRule.MatchAddressLimit(&unmatchedMetadata, dnsResponseForTest(netip.MustParseAddr("8.8.8.8"))))
}
func TestDNSLegacyInvertAddressLimitPreLookupRegression(t *testing.T) {
t.Parallel()
testCases := []struct {
name string
build func(*testing.T, *abstractDefaultRule)
matchedAddrs []netip.Addr
unmatchedAddrs []netip.Addr
}{
{
name: "ip_cidr",
build: func(t *testing.T, rule *abstractDefaultRule) {
t.Helper()
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
},
matchedAddrs: []netip.Addr{netip.MustParseAddr("203.0.113.1")},
unmatchedAddrs: []netip.Addr{netip.MustParseAddr("8.8.8.8")},
},
{
name: "ip_is_private",
build: func(t *testing.T, rule *abstractDefaultRule) {
t.Helper()
addDestinationIPIsPrivateItem(rule)
},
matchedAddrs: []netip.Addr{netip.MustParseAddr("10.0.0.1")},
unmatchedAddrs: []netip.Addr{netip.MustParseAddr("8.8.8.8")},
},
{
name: "ip_accept_any",
build: func(t *testing.T, rule *abstractDefaultRule) {
t.Helper()
addDestinationIPAcceptAnyItem(rule)
},
matchedAddrs: []netip.Addr{netip.MustParseAddr("203.0.113.1")},
},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
t.Parallel()
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
rule.invert = true
testCase.build(t, rule)
})
preLookupMetadata := testMetadata("lookup.example")
require.True(t, rule.LegacyPreMatch(&preLookupMetadata))
matchedMetadata := testMetadata("lookup.example")
require.False(t, rule.MatchAddressLimit(&matchedMetadata, dnsResponseForTest(testCase.matchedAddrs...)))
unmatchedMetadata := testMetadata("lookup.example")
require.True(t, rule.MatchAddressLimit(&unmatchedMetadata, dnsResponseForTest(testCase.unmatchedAddrs...)))
})
}
}
func TestDNSLegacyInvertLogicalAddressLimitPreLookupRegression(t *testing.T) {
t.Parallel()
t.Run("inverted deferred child does not suppress branch", func(t *testing.T) {
t.Parallel()
logicalRule := &LogicalDNSRule{
abstractLogicalRule: abstractLogicalRule{
rules: []adapter.HeadlessRule{
dnsRuleForTest(func(rule *abstractDefaultRule) {
rule.invert = true
addDestinationIPIsPrivateItem(rule)
}),
},
mode: C.LogicalTypeAnd,
},
}
preLookupMetadata := testMetadata("lookup.example")
require.True(t, logicalRule.LegacyPreMatch(&preLookupMetadata))
})
}
func TestDNSLegacyInvertRuleSetAddressLimitPreLookupRegression(t *testing.T) {
t.Parallel()
ruleSet := newLocalRuleSetForTest("dns-legacy-invert-ipcidr", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
rule.invert = true
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
}))
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
preLookupMetadata := testMetadata("lookup.example")
require.True(t, rule.LegacyPreMatch(&preLookupMetadata))
matchedMetadata := testMetadata("lookup.example")
require.False(t, rule.MatchAddressLimit(&matchedMetadata, dnsResponseForTest(netip.MustParseAddr("203.0.113.1"))))
unmatchedMetadata := testMetadata("lookup.example")
require.True(t, rule.MatchAddressLimit(&unmatchedMetadata, dnsResponseForTest(netip.MustParseAddr("8.8.8.8"))))
}
func TestDNSInvertAddressLimitPreLookupRegression(t *testing.T) {
@@ -662,14 +1289,14 @@ func TestDNSInvertAddressLimitPreLookupRegression(t *testing.T) {
matchedMetadata := testMetadata("lookup.example")
matchedMetadata.DestinationAddresses = testCase.matchedAddrs
require.False(t, rule.MatchAddressLimit(&matchedMetadata))
require.False(t, rule.MatchAddressLimit(&matchedMetadata, dnsResponseForTest(testCase.matchedAddrs...)))
unmatchedMetadata := testMetadata("lookup.example")
unmatchedMetadata.DestinationAddresses = testCase.unmatchedAddrs
require.True(t, rule.MatchAddressLimit(&unmatchedMetadata))
require.True(t, rule.MatchAddressLimit(&unmatchedMetadata, dnsResponseForTest(testCase.unmatchedAddrs...)))
})
}
t.Run("mixed resolved and deferred fields keep old pre lookup false", func(t *testing.T) {
t.Run("mixed resolved and deferred fields invert matches pre lookup", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("lookup.example")
rule := dnsRuleForTest(func(rule *abstractDefaultRule) {
@@ -677,9 +1304,9 @@ func TestDNSInvertAddressLimitPreLookupRegression(t *testing.T) {
addOtherItem(rule, NewNetworkItem([]string{N.NetworkTCP}))
addDestinationIPCIDRItem(t, rule, []string{"203.0.113.0/24"})
})
require.False(t, rule.Match(&metadata))
require.True(t, rule.Match(&metadata))
})
t.Run("ruleset only deferred fields keep old pre lookup false", func(t *testing.T) {
t.Run("ruleset only deferred fields invert matches pre lookup", func(t *testing.T) {
t.Parallel()
metadata := testMetadata("lookup.example")
ruleSet := newLocalRuleSetForTest("dns-ruleset-ipcidr", headlessDefaultRule(t, func(rule *abstractDefaultRule) {
@@ -689,7 +1316,7 @@ func TestDNSInvertAddressLimitPreLookupRegression(t *testing.T) {
rule.invert = true
addRuleSetItem(rule, &RuleSetItem{setList: []adapter.RuleSet{ruleSet}})
})
require.False(t, rule.Match(&metadata))
require.True(t, rule.Match(&metadata))
})
}
@@ -731,8 +1358,8 @@ func newLocalRuleSetForTest(tag string, rules ...adapter.HeadlessRule) *LocalRul
func newRemoteRuleSetForTest(tag string, rules ...adapter.HeadlessRule) *RemoteRuleSet {
return &RemoteRuleSet{
options: option.RuleSet{Tag: tag},
rules: rules,
tag: tag,
rules: rules,
}
}
@@ -760,6 +1387,39 @@ func testMetadata(domain string) adapter.InboundContext {
}
}
func dnsResponseForTest(addresses ...netip.Addr) *mDNS.Msg {
response := &mDNS.Msg{
MsgHdr: mDNS.MsgHdr{
Response: true,
Rcode: mDNS.RcodeSuccess,
},
}
for _, address := range addresses {
if address.Is4() {
response.Answer = append(response.Answer, &mDNS.A{
Hdr: mDNS.RR_Header{
Name: mDNS.Fqdn("lookup.example"),
Rrtype: mDNS.TypeA,
Class: mDNS.ClassINET,
Ttl: 60,
},
A: net.IP(append([]byte(nil), address.AsSlice()...)),
})
} else {
response.Answer = append(response.Answer, &mDNS.AAAA{
Hdr: mDNS.RR_Header{
Name: mDNS.Fqdn("lookup.example"),
Rrtype: mDNS.TypeAAAA,
Class: mDNS.ClassINET,
Ttl: 60,
},
AAAA: net.IP(append([]byte(nil), address.AsSlice()...)),
})
}
}
return response
}
func addRuleSetItem(rule *abstractDefaultRule, item *RuleSetItem) {
rule.ruleSetItem = item
rule.allItems = append(rule.allItems, item)
@@ -0,0 +1,111 @@
package rule
import (
"context"
"sync/atomic"
"testing"
"github.com/sagernet/sing-box/adapter"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/option"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/json/badoption"
"github.com/sagernet/sing/common/x/list"
"github.com/sagernet/sing/service"
"github.com/stretchr/testify/require"
)
type fakeDNSRuleSetUpdateValidator struct {
validate func(tag string, metadata adapter.RuleSetMetadata) error
}
func (v *fakeDNSRuleSetUpdateValidator) ValidateRuleSetMetadataUpdate(tag string, metadata adapter.RuleSetMetadata) error {
if v.validate == nil {
return nil
}
return v.validate(tag, metadata)
}
func TestLocalRuleSetReloadRulesRejectsInvalidUpdateBeforeCommit(t *testing.T) {
t.Parallel()
var callbackCount atomic.Int32
ctx := service.ContextWith[adapter.DNSRuleSetUpdateValidator](context.Background(), &fakeDNSRuleSetUpdateValidator{
validate: func(tag string, metadata adapter.RuleSetMetadata) error {
require.Equal(t, "dynamic-set", tag)
if metadata.ContainsDNSQueryTypeRule {
return E.New("dns conflict")
}
return nil
},
})
ruleSet := &LocalRuleSet{
ctx: ctx,
tag: "dynamic-set",
fileFormat: C.RuleSetFormatSource,
}
_ = ruleSet.callbacks.PushBack(func(adapter.RuleSet) {
callbackCount.Add(1)
})
err := ruleSet.reloadRules([]option.HeadlessRule{{
Type: C.RuleTypeDefault,
DefaultOptions: option.DefaultHeadlessRule{
Domain: badoption.Listable[string]{"example.com"},
},
}})
require.NoError(t, err)
require.Equal(t, int32(1), callbackCount.Load())
require.False(t, ruleSet.metadata.ContainsDNSQueryTypeRule)
require.True(t, ruleSet.Match(&adapter.InboundContext{Domain: "example.com"}))
err = ruleSet.reloadRules([]option.HeadlessRule{{
Type: C.RuleTypeDefault,
DefaultOptions: option.DefaultHeadlessRule{
QueryType: badoption.Listable[option.DNSQueryType]{option.DNSQueryType(1)},
},
}})
require.ErrorContains(t, err, "dns conflict")
require.Equal(t, int32(1), callbackCount.Load())
require.False(t, ruleSet.metadata.ContainsDNSQueryTypeRule)
require.True(t, ruleSet.Match(&adapter.InboundContext{Domain: "example.com"}))
}
func TestRemoteRuleSetLoadBytesRejectsInvalidUpdateBeforeCommit(t *testing.T) {
t.Parallel()
var callbackCount atomic.Int32
ctx := service.ContextWith[adapter.DNSRuleSetUpdateValidator](context.Background(), &fakeDNSRuleSetUpdateValidator{
validate: func(tag string, metadata adapter.RuleSetMetadata) error {
require.Equal(t, "dynamic-set", tag)
if metadata.ContainsDNSQueryTypeRule {
return E.New("dns conflict")
}
return nil
},
})
ruleSet := &RemoteRuleSet{
ctx: ctx,
tag: "dynamic-set",
options: option.RuleSet{
Format: C.RuleSetFormatSource,
},
callbacks: list.List[adapter.RuleSetUpdateCallback]{},
}
_ = ruleSet.callbacks.PushBack(func(adapter.RuleSet) {
callbackCount.Add(1)
})
err := ruleSet.loadBytes([]byte(`{"version":4,"rules":[{"domain":["example.com"]}]}`))
require.NoError(t, err)
require.Equal(t, int32(1), callbackCount.Load())
require.False(t, ruleSet.metadata.ContainsDNSQueryTypeRule)
require.True(t, ruleSet.Match(&adapter.InboundContext{Domain: "example.com"}))
err = ruleSet.loadBytes([]byte(`{"version":4,"rules":[{"query_type":["A"]}]}`))
require.ErrorContains(t, err, "dns conflict")
require.Equal(t, int32(1), callbackCount.Load())
require.False(t, ruleSet.metadata.ContainsDNSQueryTypeRule)
require.True(t, ruleSet.Match(&adapter.InboundContext{Domain: "example.com"}))
}
+87
View File
@@ -0,0 +1,87 @@
package rule
import (
"context"
"runtime"
"time"
"github.com/sagernet/sing-box/adapter"
)
type RuleSetUpdater struct {
ctx context.Context
cancel context.CancelFunc
ruleSets []*RemoteRuleSet
}
func NewRuleSetUpdater(ctx context.Context, ruleSets []adapter.RuleSet) *RuleSetUpdater {
var remoteRuleSets []*RemoteRuleSet
for _, ruleSet := range ruleSets {
remoteRuleSet, isRemote := ruleSet.(*RemoteRuleSet)
if isRemote {
remoteRuleSets = append(remoteRuleSets, remoteRuleSet)
}
}
if len(remoteRuleSets) == 0 {
return nil
}
ctx, cancel := context.WithCancel(ctx)
return &RuleSetUpdater{
ctx: ctx,
cancel: cancel,
ruleSets: remoteRuleSets,
}
}
func (u *RuleSetUpdater) Start() {
go u.loopUpdate()
}
func (u *RuleSetUpdater) Close() error {
u.cancel()
return nil
}
func (u *RuleSetUpdater) loopUpdate() {
nextUpdates := make([]time.Time, len(u.ruleSets))
for i, ruleSet := range u.ruleSets {
nextUpdates[i] = ruleSet.lastUpdated.Add(ruleSet.updateInterval)
}
timer := time.NewTimer(0)
defer timer.Stop()
for {
select {
case <-u.ctx.Done():
return
case <-timer.C:
}
now := time.Now()
var updated bool
for i, ruleSet := range u.ruleSets {
if now.Before(nextUpdates[i]) {
continue
}
ruleSet.updateOnce()
nextUpdates[i] = now.Add(ruleSet.updateInterval)
updated = true
}
if updated {
runtime.GC()
}
timer.Reset(waitUntilNext(nextUpdates))
}
}
func waitUntilNext(nextUpdates []time.Time) time.Duration {
next := nextUpdates[0]
for _, nextUpdate := range nextUpdates[1:] {
if nextUpdate.Before(next) {
next = nextUpdate
}
}
wait := time.Until(next)
if wait < 0 {
return 0
}
return wait
}
+26 -2
View File
@@ -38,11 +38,19 @@ func hasDNSRule(rules []option.DNSRule, cond func(rule option.DefaultDNSRule) bo
}
func isProcessRule(rule option.DefaultRule) bool {
return len(rule.ProcessName) > 0 || len(rule.ProcessPath) > 0 || len(rule.ProcessPathRegex) > 0 || len(rule.PackageName) > 0 || len(rule.User) > 0 || len(rule.UserID) > 0
return len(rule.ProcessName) > 0 || len(rule.ProcessPath) > 0 || len(rule.ProcessPathRegex) > 0 || len(rule.PackageName) > 0 || len(rule.PackageNameRegex) > 0 || len(rule.User) > 0 || len(rule.UserID) > 0
}
func isProcessDNSRule(rule option.DefaultDNSRule) bool {
return len(rule.ProcessName) > 0 || len(rule.ProcessPath) > 0 || len(rule.ProcessPathRegex) > 0 || len(rule.PackageName) > 0 || len(rule.User) > 0 || len(rule.UserID) > 0
return len(rule.ProcessName) > 0 || len(rule.ProcessPath) > 0 || len(rule.ProcessPathRegex) > 0 || len(rule.PackageName) > 0 || len(rule.PackageNameRegex) > 0 || len(rule.User) > 0 || len(rule.UserID) > 0
}
func isNeighborRule(rule option.DefaultRule) bool {
return len(rule.SourceMACAddress) > 0 || len(rule.SourceHostname) > 0
}
func isNeighborDNSRule(rule option.DefaultDNSRule) bool {
return len(rule.SourceMACAddress) > 0 || len(rule.SourceHostname) > 0
}
func isWIFIRule(rule option.DefaultRule) bool {
@@ -52,3 +60,19 @@ func isWIFIRule(rule option.DefaultRule) bool {
func isWIFIDNSRule(rule option.DefaultDNSRule) bool {
return len(rule.WIFISSID) > 0 || len(rule.WIFIBSSID) > 0
}
func hasLocalNeighborDNSServer(servers []option.DNSServerOptions) bool {
for _, server := range servers {
if server.Type != C.DNSTypeLocal {
continue
}
localOptions, isLocal := server.Options.(*option.LocalDNSServerOptions)
if !isLocal {
continue
}
if len(localOptions.NeighborDomain) > 0 {
return true
}
}
return false
}