Fix windows bridge forwarding to connected subnets

This commit is contained in:
世界
2026-07-08 19:45:13 +08:00
parent 24607d331b
commit 5fee7552e7
2 changed files with 184 additions and 44 deletions
+131 -21
View File
@@ -6,7 +6,9 @@ import (
"context"
"encoding/binary"
"errors"
"net"
"net/netip"
"slices"
"sync"
"sync/atomic"
"time"
@@ -16,6 +18,7 @@ import (
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing-tun/gtcpip"
"github.com/sagernet/sing-tun/gtcpip/header"
"github.com/sagernet/sing/common/control"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
@@ -43,10 +46,62 @@ const (
divertICMPError
)
type localSegment struct {
prefix netip.Prefix
address netip.Addr
}
type egressState struct {
inet4 netip.Addr
inet6 netip.Addr
mtu uint32
inet4 netip.Addr
inet6 netip.Addr
mtu uint32
inet4Segments []localSegment
inet6Segments []localSegment
}
func (s *egressState) equal(other *egressState) bool {
return s.inet4 == other.inet4 && s.inet6 == other.inet6 && s.mtu == other.mtu &&
slices.Equal(s.inet4Segments, other.inet4Segments) &&
slices.Equal(s.inet6Segments, other.inet6Segments)
}
// sourceAddress picks the translated source for an outbound packet. Windows
// routes with the strong host model: the route lookup is constrained to the
// interface owning the source address, so choosing the source is what steers
// the packet. Destinations in a connected subnet take that subnet's own
// address; everything else takes the egress address.
func (s *egressState) sourceAddress(destination netip.Addr, isV6 bool) netip.Addr {
segments := s.inet4Segments
egressAddress := s.inet4
if isV6 {
segments = s.inet6Segments
egressAddress = s.inet6
}
for _, segment := range segments {
if segment.prefix.Contains(destination) {
return segment.address
}
}
return egressAddress
}
func (s *egressState) divertAddresses(isV6 bool) []netip.Addr {
egressAddress := s.inet4
segments := s.inet4Segments
if isV6 {
egressAddress = s.inet6
segments = s.inet6Segments
}
if !egressAddress.IsValid() {
return nil
}
addresses := []netip.Addr{egressAddress}
for _, segment := range segments {
if !slices.Contains(addresses, segment.address) {
addresses = append(addresses, segment.address)
}
}
return addresses
}
type diverter struct {
@@ -177,13 +232,13 @@ func (b *backendWindows) rebuildDivertersLocked(state *egressState) error {
b.diverters = nil
if b.inet4Port.IsValid() && state.inet4.IsValid() {
err := b.openFamilyDiverters(state.inet4, false)
err := b.openFamilyDiverters(state.divertAddresses(false), false)
if err != nil {
return err
}
}
if b.inet6Port.IsValid() && state.inet6.IsValid() {
err := b.openFamilyDiverters(state.inet6, true)
err := b.openFamilyDiverters(state.divertAddresses(true), true)
if err != nil {
return err
}
@@ -191,7 +246,7 @@ func (b *backendWindows) rebuildDivertersLocked(state *egressState) error {
return nil
}
func (b *backendWindows) openFamilyDiverters(egressAddr netip.Addr, isV6 bool) error {
func (b *backendWindows) openFamilyDiverters(addresses []netip.Addr, isV6 bool) error {
portHigh := uint16(uint32(b.reservedStart) + uint32(bridgeReservedPortCount) - 1)
entries := []struct {
what string
@@ -199,16 +254,16 @@ func (b *backendWindows) openFamilyDiverters(egressAddr netip.Addr, isV6 bool) e
build func() (*windivert.Filter, error)
}{
{"TCP", divertTransport, func() (*windivert.Filter, error) {
return windivert.InboundTCPPortRange(egressAddr, b.reservedStart, portHigh)
return windivert.InboundTCPPortRange(addresses, b.reservedStart, portHigh)
}},
{"UDP", divertTransport, func() (*windivert.Filter, error) {
return windivert.InboundUDPPortRange(egressAddr, b.reservedStart, portHigh)
return windivert.InboundUDPPortRange(addresses, b.reservedStart, portHigh)
}},
{"ICMP echo", divertICMPEcho, func() (*windivert.Filter, error) {
return windivert.InboundICMPEchoReply(egressAddr)
return windivert.InboundICMPEchoReply(addresses)
}},
{"ICMP error", divertICMPError, func() (*windivert.Filter, error) {
return windivert.InboundICMPError(egressAddr)
return windivert.InboundICMPError(addresses)
}},
}
for _, entry := range entries {
@@ -577,28 +632,24 @@ func (b *backendWindows) flushOutboundLocked() {
}
func (b *backendWindows) prepareOutbound(packet []byte, state *egressState) bool {
var (
isV6 bool
egressAddr netip.Addr
)
var isV6 bool
switch header.IPVersion(packet) {
case header.IPv4Version:
egressAddr = state.inet4
case header.IPv6Version:
isV6 = true
egressAddr = state.inet6
default:
return false
}
if !egressAddr.IsValid() {
return false
}
// The batched injection ioctl walks the buffer by IP total length; a
// packet with trailing bytes would desynchronize the walk and fail the
// whole batch.
if ipPacketLength(packet) != len(packet) {
return false
}
egressAddr := state.sourceAddress(packetRemoteAddress(packet, isV6, false), isV6)
if !egressAddr.IsValid() {
return false
}
info, valid := parseTransport(packet, isV6)
if !valid {
return false
@@ -642,8 +693,10 @@ func (b *backendWindows) syncEgress() {
}
state := b.currentEgressState()
previous := b.egress.Load()
if previous != nil && previous.inet4 == state.inet4 && previous.inet6 == state.inet6 {
if *previous != *state {
if previous != nil &&
slices.Equal(previous.divertAddresses(false), state.divertAddresses(false)) &&
slices.Equal(previous.divertAddresses(true), state.divertAddresses(true)) {
if !previous.equal(state) {
b.egress.Store(state)
}
return
@@ -688,9 +741,66 @@ func (b *backendWindows) currentEgressState() *egressState {
state.inet6 = address
}
}
b.collectLocalSegments(state, finder)
return state
}
// collectLocalSegments gathers the connected subnets whose destinations must
// bypass the egress pin so they leave on their own interface, mirroring the
// Linux backend (routing table) and the Darwin backend (pf pass-in rules).
// With a pinned egress only its own subnets are considered.
func (b *backendWindows) collectLocalSegments(state *egressState, finder control.InterfaceFinder) {
for _, localInterface := range finder.Interfaces() {
if b.boundInterface != "" && localInterface.Name != b.boundInterface {
continue
}
if localInterface.Flags&net.FlagUp == 0 || localInterface.Flags&net.FlagBroadcast == 0 ||
localInterface.Flags&net.FlagLoopback != 0 || localInterface.Flags&net.FlagPointToPoint != 0 {
continue
}
for _, prefix := range localInterface.Addresses {
address := prefix.Addr().Unmap()
if !address.IsGlobalUnicast() {
continue
}
segment := localSegment{
prefix: netip.PrefixFrom(address, prefix.Bits()).Masked(),
address: address,
}
if address.Is4() {
if state.inet4.IsValid() {
state.inet4Segments = appendSegment(state.inet4Segments, segment)
}
} else if state.inet6.IsValid() {
state.inet6Segments = appendSegment(state.inet6Segments, segment)
}
}
}
sortSegments(state.inet4Segments)
sortSegments(state.inet6Segments)
}
func appendSegment(segments []localSegment, segment localSegment) []localSegment {
for _, existing := range segments {
if existing.prefix == segment.prefix {
return segments
}
}
return append(segments, segment)
}
func sortSegments(segments []localSegment) {
slices.SortFunc(segments, func(a, b localSegment) int {
if a.prefix.Bits() != b.prefix.Bits() {
return b.prefix.Bits() - a.prefix.Bits()
}
if result := a.prefix.Addr().Compare(b.prefix.Addr()); result != 0 {
return result
}
return a.address.Compare(b.address)
})
}
func (b *backendWindows) Close() error {
b.closeOnce.Do(func() {
if b.closed != nil {