Add flow dispatcher

This commit is contained in:
世界 2026-07-06 11:49:18 +08:00
parent 47bdde06c3
commit ed63adda33
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
27 changed files with 2469 additions and 963 deletions

View file

@ -5,6 +5,7 @@ import (
"errors"
"net"
"net/netip"
"slices"
"syscall"
"time"
@ -46,7 +47,7 @@ type System struct {
tcpPort6 uint16
tcpNat *TCPNat
udpNat *udpnat.Service
directNat *DirectRouteMapping
dispatcher *ForwardDispatcher
bindInterface bool
interfaceFinder control.InterfaceFinder
frontHeadroom int
@ -101,6 +102,7 @@ func NewSystem(options StackOptions) (Stack, error) {
}
func (s *System) Close() error {
s.dispatcher.Close()
return common.Close(
s.tcpListener,
s.tcpListener6,
@ -162,7 +164,13 @@ func (s *System) start() error {
}
s.tcpNat = NewNat(s.ctx, s.udpTimeout)
s.udpNat = udpnat.New(s.handler, s.preparePacketConnection, s.udpTimeout, false)
s.directNat = NewDirectRouteMapping(s.icmpTimeout)
if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN {
s.frontHeadroom = linuxTUN.FrontHeadroom()
s.txChecksumOffload = linuxTUN.TXChecksumOffload()
}
if s.handler != nil {
s.dispatcher = NewForwardDispatcher(s.handler, newSystemWriteback(s.tun, s.frontHeadroom), s.logger, s.udpTimeout, s.icmpTimeout)
}
return nil
}
@ -172,8 +180,6 @@ func (s *System) tunLoop() {
return
}
if linuxTUN, isLinuxTUN := s.tun.(LinuxTUN); isLinuxTUN {
s.frontHeadroom = linuxTUN.FrontHeadroom()
s.txChecksumOffload = linuxTUN.TXChecksumOffload()
batchSize := linuxTUN.BatchSize()
if batchSize > 1 {
s.batchLoopLinux(linuxTUN, batchSize)
@ -204,6 +210,7 @@ func (s *System) tunLoop() {
s.logger.Trace(E.Cause(err, "write packet"))
}
}
s.dispatcher.Flush()
}
}
@ -223,6 +230,7 @@ func (s *System) wintunLoop(winTun WinTun) {
s.logger.Trace(E.Cause(err, "write packet"))
}
}
s.dispatcher.Flush()
release()
}
}
@ -263,11 +271,13 @@ func (s *System) batchLoopLinux(linuxTUN LinuxTUN, batchSize int) {
}
writeBuffers = writeBuffers[:0]
}
s.dispatcher.Flush()
}
}
func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
var writeBuffers []*buf.Buffer
var releaseBuffers []*buf.Buffer
for {
buffers, err := darwinTUN.BatchRead()
if err != nil {
@ -280,6 +290,7 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
continue
}
writeBuffers = writeBuffers[:0]
releaseBuffers = releaseBuffers[:0]
for _, buffer := range buffers {
packetSize := buffer.Len()
if packetSize < header.IPv4MinimumSize {
@ -289,7 +300,7 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
if s.processPacket(buffer.Bytes()) {
writeBuffers = append(writeBuffers, buffer)
} else {
buffer.Release()
releaseBuffers = append(releaseBuffers, buffer)
}
}
if len(writeBuffers) > 0 {
@ -299,6 +310,8 @@ func (s *System) batchLoopDarwin(darwinTUN DarwinTUN) {
}
buf.ReleaseMulti(writeBuffers)
}
s.dispatcher.Flush()
buf.ReleaseMulti(releaseBuffers)
}
}
@ -338,11 +351,53 @@ func (s *System) acceptLoop(listener net.Listener) {
}
}
func (s *System) dispatchIPv4(ipHdr header.IPv4, destination netip.Addr) bool {
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
if slices.Contains(s.inet4LoopbackAddress, destination) {
return false
}
if ipHdr.SourceAddr() == s.inet4Address &&
ipHdr.FragmentOffset() == 0 &&
len(ipHdr.Payload()) >= header.TCPMinimumSize &&
header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort {
return false
}
case header.ICMPv4ProtocolNumber:
if destination == s.inet4Address {
return false
}
}
return s.dispatcher.Dispatch(ipHdr)
}
func (s *System) dispatchIPv6(ipHdr header.IPv6, destination netip.Addr) bool {
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
if slices.Contains(s.inet6LoopbackAddress, destination) {
return false
}
if ipHdr.SourceAddr() == s.inet6Address &&
len(ipHdr.Payload()) >= header.TCPMinimumSize &&
header.TCP(ipHdr.Payload()).SourcePort() == s.tcpPort6 {
return false
}
case header.ICMPv6ProtocolNumber:
if destination == s.inet6Address {
return false
}
}
return s.dispatcher.Dispatch(ipHdr)
}
func (s *System) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
destination := ipHdr.DestinationAddr()
if destination == s.broadcastAddr || !destination.IsGlobalUnicast() {
return
}
if s.dispatchIPv4(ipHdr, destination) {
return false, nil
}
writeBack = true
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
@ -360,9 +415,13 @@ func (s *System) processIPv4(ipHdr header.IPv4) (writeBack bool, err error) {
}
func (s *System) processIPv6(ipHdr header.IPv6) (writeBack bool, err error) {
if !ipHdr.DestinationAddr().IsGlobalUnicast() {
destination := ipHdr.DestinationAddr()
if !destination.IsGlobalUnicast() {
return
}
if s.dispatchIPv6(ipHdr, destination) {
return false, nil
}
writeBack = true
switch ipHdr.TransportProtocol() {
case header.TCPProtocolNumber:
@ -404,14 +463,7 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
}
}
if !loopback {
natPort, err := s.tcpNat.Lookup(source, destination, s.handler)
if err != nil {
if errors.Is(err, ErrDrop) {
return false, nil
} else {
return false, s.resetIPv4TCP(ipHdr, tcpHdr)
}
}
natPort := s.tcpNat.Lookup(source, destination)
ipHdr.SetSourceAddr(s.inet4NextAddress)
tcpHdr.SetSourcePort(natPort)
ipHdr.SetDestinationAddr(s.inet4Address)
@ -429,51 +481,6 @@ func (s *System) processIPv4TCP(ipHdr header.IPv4, tcpHdr header.TCP) (bool, err
return true, nil
}
func (s *System) resetIPv4TCP(origIPHdr header.IPv4, origTCPHdr header.TCP) error {
frontHeadroom := s.frontHeadroom + PacketOffset
newPacket := buf.NewSize(frontHeadroom + header.IPv4MinimumSize + header.TCPMinimumSize)
defer newPacket.Release()
newPacket.Resize(frontHeadroom, header.IPv4MinimumSize+header.TCPMinimumSize)
ipHdr := header.IPv4(newPacket.Bytes())
ipHdr.Encode(&header.IPv4Fields{
TotalLength: uint16(newPacket.Len()),
Protocol: uint8(header.TCPProtocolNumber),
SrcAddr: origIPHdr.DestinationAddr(),
DstAddr: origIPHdr.SourceAddr(),
})
tcpHdr := header.TCP(ipHdr.Payload())
fields := header.TCPFields{
SrcPort: origTCPHdr.DestinationPort(),
DstPort: origTCPHdr.SourcePort(),
DataOffset: header.TCPMinimumSize,
Flags: header.TCPFlagRst,
}
if origTCPHdr.Flags()&header.TCPFlagAck != 0 {
fields.SeqNum = origTCPHdr.AckNumber()
} else {
fields.Flags |= header.TCPFlagAck
ackNum := origTCPHdr.SequenceNumber() + uint32(len(origTCPHdr.Payload()))
if origTCPHdr.Flags()&header.TCPFlagSyn != 0 {
ackNum++
}
if origTCPHdr.Flags()&header.TCPFlagFin != 0 {
ackNum++
}
fields.AckNum = ackNum
}
tcpHdr.Encode(&fields)
if !s.txChecksumOffload {
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum(header.TCPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), header.TCPMinimumSize)))
}
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version)
} else {
newPacket.Advance(-s.frontHeadroom)
}
return common.Error(s.tun.Write(newPacket.Bytes()))
}
func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, error) {
source := netip.AddrPortFrom(ipHdr.SourceAddr(), tcpHdr.SourcePort())
destination := netip.AddrPortFrom(ipHdr.DestinationAddr(), tcpHdr.DestinationPort())
@ -499,14 +506,7 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
}
}
if !loopback {
natPort, err := s.tcpNat.Lookup(source, destination, s.handler)
if err != nil {
if errors.Is(err, ErrDrop) {
return false, nil
} else {
return false, s.resetIPv6TCP(ipHdr, tcpHdr)
}
}
natPort := s.tcpNat.Lookup(source, destination)
ipHdr.SetSourceAddr(s.inet6NextAddress)
tcpHdr.SetSourcePort(natPort)
ipHdr.SetDestinationAddr(s.inet6Address)
@ -523,50 +523,6 @@ func (s *System) processIPv6TCP(ipHdr header.IPv6, tcpHdr header.TCP) (bool, err
return true, nil
}
func (s *System) resetIPv6TCP(origIPHdr header.IPv6, origTCPHdr header.TCP) error {
frontHeadroom := s.frontHeadroom + PacketOffset
newPacket := buf.NewSize(frontHeadroom + header.IPv6MinimumSize + header.TCPMinimumSize)
defer newPacket.Release()
newPacket.Resize(frontHeadroom, header.IPv6MinimumSize+header.TCPMinimumSize)
ipHdr := header.IPv6(newPacket.Bytes())
ipHdr.Encode(&header.IPv6Fields{
PayloadLength: uint16(header.TCPMinimumSize),
TransportProtocol: header.TCPProtocolNumber,
SrcAddr: origIPHdr.DestinationAddr(),
DstAddr: origIPHdr.SourceAddr(),
})
tcpHdr := header.TCP(ipHdr.Payload())
fields := header.TCPFields{
SrcPort: origTCPHdr.DestinationPort(),
DstPort: origTCPHdr.SourcePort(),
DataOffset: header.TCPMinimumSize,
Flags: header.TCPFlagRst,
}
if origTCPHdr.Flags()&header.TCPFlagAck != 0 {
fields.SeqNum = origTCPHdr.AckNumber()
} else {
fields.Flags |= header.TCPFlagAck
ackNum := origTCPHdr.SequenceNumber() + uint32(len(origTCPHdr.Payload()))
if origTCPHdr.Flags()&header.TCPFlagSyn != 0 {
ackNum++
}
if origTCPHdr.Flags()&header.TCPFlagFin != 0 {
ackNum++
}
fields.AckNum = ackNum
}
tcpHdr.Encode(&fields)
if !s.txChecksumOffload {
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum(header.TCPProtocolNumber, ipHdr.SourceAddressSlice(), ipHdr.DestinationAddressSlice(), header.TCPMinimumSize)))
}
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
} else {
newPacket.Advance(-s.frontHeadroom)
}
return common.Error(s.tun.Write(newPacket.Bytes()))
}
func (s *System) processIPv4UDP(ipHdr header.IPv4, udpHdr header.UDP) error {
if ipHdr.Flags()&header.IPv4FlagMoreFragments != 0 {
return E.New("ipv4: fragment dropped")
@ -594,19 +550,6 @@ func (s *System) processIPv6UDP(ipHdr header.IPv6, udpHdr header.UDP) error {
}
func (s *System) preparePacketConnection(source M.Socksaddr, destination M.Socksaddr, userData any) (bool, context.Context, N.PacketWriter, N.CloseHandlerFunc) {
_, pErr := s.handler.PrepareConnection(N.NetworkUDP, source, destination, nil, 0)
if pErr != nil {
if !errors.Is(pErr, ErrDrop) {
if source.IsIPv4() {
ipHdr := userData.(header.IPv4)
s.rejectIPv4WithICMP(ipHdr, header.ICMPv4PortUnreachable)
} else {
ipHdr := userData.(header.IPv6)
s.rejectIPv6WithICMP(ipHdr, header.ICMPv6PortUnreachable)
}
}
return false, nil, nil, nil
}
var writer N.PacketWriter
if source.IsIPv4() {
packet := userData.(header.IPv4)
@ -640,29 +583,6 @@ func (s *System) processIPv4ICMP(ipHdr header.IPv4, icmpHdr header.ICMPv4) (bool
if icmpHdr.Type() != header.ICMPv4Echo || icmpHdr.Code() != 0 {
return false, nil
}
sourceAddr := ipHdr.SourceAddr()
destinationAddr := ipHdr.DestinationAddr()
if destinationAddr != s.inet4Address {
action, err := s.directNat.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) {
return s.handler.PrepareConnection(
N.NetworkICMP,
M.SocksaddrFrom(sourceAddr, 0),
M.SocksaddrFrom(destinationAddr, 0),
&systemICMPDirectPacketWriter4{s.tun, s.frontHeadroom + PacketOffset, sourceAddr},
timeout,
)
})
if err != nil {
if errors.Is(err, ErrReset) {
return false, s.rejectIPv4WithICMP(ipHdr, header.ICMPv4HostUnreachable)
} else if errors.Is(err, ErrDrop) {
return false, nil
}
}
if action != nil {
return false, action.WritePacket(buf.As(ipHdr).ToOwned())
}
}
icmpHdr.SetType(header.ICMPv4EchoReply)
sourceAddress := ipHdr.SourceAddr()
ipHdr.SetSourceAddr(ipHdr.DestinationAddr())
@ -672,70 +592,10 @@ func (s *System) processIPv4ICMP(ipHdr header.IPv4, icmpHdr header.ICMPv4) (bool
return true, nil
}
func (s *System) rejectIPv4WithICMP(ipHdr header.IPv4, code header.ICMPv4Code) error {
frontHeadroom := s.frontHeadroom + PacketOffset
mtu := s.mtu
const maxIPData = header.IPv4MinimumProcessableDatagramSize - header.IPv4MinimumSize
if mtu > maxIPData {
mtu = maxIPData
}
available := mtu - header.ICMPv4MinimumSize
if available < len(ipHdr)+header.ICMPv4MinimumErrorPayloadSize {
return nil
}
payload := ipHdr
if len(payload) > available {
payload = payload[:available]
}
newPacket := buf.NewSize(frontHeadroom + header.IPv4MinimumSize + header.ICMPv4MinimumSize + len(payload))
defer newPacket.Release()
newPacket.Resize(frontHeadroom, header.IPv4MinimumSize+header.ICMPv4MinimumSize+len(payload))
newIPHdr := header.IPv4(newPacket.Bytes())
newIPHdr.Encode(&header.IPv4Fields{
TotalLength: uint16(newPacket.Len()),
Protocol: uint8(header.ICMPv4ProtocolNumber),
SrcAddr: ipHdr.DestinationAddr(),
DstAddr: ipHdr.SourceAddr(),
})
newIPHdr.SetChecksum(^newIPHdr.CalculateChecksum())
icmpHdr := header.ICMPv4(newIPHdr.Payload())
icmpHdr.SetType(header.ICMPv4DstUnreachable)
icmpHdr.SetCode(code)
icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr[:header.ICMPv4MinimumSize], checksum.Checksum(ipHdr.Payload(), 0)))
copy(icmpHdr.Payload(), payload)
if PacketOffset > 0 {
newPacket.ExtendHeader(PacketOffset)[3] = syscall.AF_INET
} else {
newPacket.Advance(-s.frontHeadroom)
}
return common.Error(s.tun.Write(newPacket.Bytes()))
}
func (s *System) processIPv6ICMP(ipHdr header.IPv6, icmpHdr header.ICMPv6) (bool, error) {
if icmpHdr.Type() != header.ICMPv6EchoRequest || icmpHdr.Code() != 0 {
return false, nil
}
sourceAddr := ipHdr.SourceAddr()
destinationAddr := ipHdr.DestinationAddr()
if destinationAddr != s.inet6Address {
action, err := s.directNat.Lookup(DirectRouteSession{Source: sourceAddr, Destination: destinationAddr}, func(timeout time.Duration) (DirectRouteDestination, error) {
return s.handler.PrepareConnection(
N.NetworkICMP,
M.SocksaddrFrom(sourceAddr, 0),
M.SocksaddrFrom(destinationAddr, 0),
&systemICMPDirectPacketWriter6{s.tun, s.frontHeadroom + PacketOffset, sourceAddr},
timeout,
)
})
if errors.Is(err, ErrReset) {
return false, s.rejectIPv6WithICMP(ipHdr, header.ICMPv6AddressUnreachable)
} else if errors.Is(err, ErrDrop) {
return false, nil
}
if action != nil {
return false, action.WritePacket(buf.As(ipHdr).ToOwned())
}
}
icmpHdr.SetType(header.ICMPv6EchoReply)
sourceAddress := ipHdr.SourceAddr()
ipHdr.SetSourceAddr(ipHdr.DestinationAddr())
@ -748,50 +608,6 @@ func (s *System) processIPv6ICMP(ipHdr header.IPv6, icmpHdr header.ICMPv6) (bool
return true, nil
}
func (s *System) rejectIPv6WithICMP(ipHdr header.IPv6, code header.ICMPv6Code) error {
frontHeadroom := s.frontHeadroom + PacketOffset
mtu := s.mtu
const maxIPv6Data = header.IPv6MinimumMTU - header.IPv6FixedHeaderSize
if mtu > maxIPv6Data {
mtu = maxIPv6Data
}
available := mtu - header.ICMPv6ErrorHeaderSize
if available < header.IPv6MinimumSize {
return nil
}
payload := ipHdr
if len(payload) > available {
payload = payload[:available]
}
newPacket := buf.NewSize(frontHeadroom + header.IPv6MinimumSize + header.ICMPv6DstUnreachableMinimumSize + len(payload))
defer newPacket.Release()
newPacket.Resize(frontHeadroom, header.IPv6MinimumSize+header.ICMPv6DstUnreachableMinimumSize+len(payload))
newIPHdr := header.IPv6(newPacket.Bytes())
newIPHdr.Encode(&header.IPv6Fields{
PayloadLength: uint16(header.ICMPv6DstUnreachableMinimumSize + len(payload)),
TransportProtocol: header.ICMPv6ProtocolNumber,
SrcAddr: ipHdr.DestinationAddr(),
DstAddr: ipHdr.SourceAddr(),
})
icmpHdr := header.ICMPv6(newIPHdr.Payload())
icmpHdr.SetType(header.ICMPv6DstUnreachable)
icmpHdr.SetCode(code)
icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: icmpHdr[:header.ICMPv6DstUnreachableMinimumSize],
Src: newIPHdr.SourceAddressSlice(),
Dst: newIPHdr.DestinationAddressSlice(),
PayloadCsum: checksum.Checksum(payload, 0),
PayloadLen: len(payload),
}))
copy(icmpHdr.Payload(), payload)
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
} else {
newPacket.Advance(-s.frontHeadroom)
}
return common.Error(s.tun.Write(newPacket.Bytes()))
}
type systemUDPPacketWriter4 struct {
tun Tun
frontHeadroom int
@ -868,45 +684,37 @@ func (w *systemUDPPacketWriter6) WritePacket(buffer *buf.Buffer, destination M.S
return common.Error(w.tun.Write(newPacket.Bytes()))
}
type systemICMPDirectPacketWriter4 struct {
type systemWriteback struct {
tun Tun
linuxTUN LinuxTUN
frontHeadroom int
source netip.Addr
}
func (w *systemICMPDirectPacketWriter4) WritePacket(p []byte) error {
newPacket := buf.NewSize(w.frontHeadroom + len(p))
defer newPacket.Release()
newPacket.Resize(w.frontHeadroom, 0)
newPacket.Write(p)
ipHdr := header.IPv4(newPacket.Bytes())
ipHdr.SetDestinationAddr(w.source)
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv4Version)
} else {
newPacket.Advance(-w.frontHeadroom)
func newSystemWriteback(tunInterface Tun, frontHeadroom int) *systemWriteback {
writeback := &systemWriteback{tun: tunInterface, frontHeadroom: frontHeadroom}
if linuxTUN, isLinuxTUN := tunInterface.(LinuxTUN); isLinuxTUN {
writeback.linuxTUN = linuxTUN
}
return common.Error(w.tun.Write(newPacket.Bytes()))
return writeback
}
type systemICMPDirectPacketWriter6 struct {
tun Tun
frontHeadroom int
source netip.Addr
func (w *systemWriteback) ReturnHeadroom() int {
return w.frontHeadroom + PacketOffset
}
func (w *systemICMPDirectPacketWriter6) WritePacket(p []byte) error {
newPacket := buf.NewSize(w.frontHeadroom + len(p))
defer newPacket.Release()
newPacket.Resize(w.frontHeadroom, 0)
newPacket.Write(p)
ipHdr := header.IPv6(newPacket.Bytes())
ipHdr.SetDestinationAddr(w.source)
if PacketOffset > 0 {
PacketFillHeader(newPacket.ExtendHeader(PacketOffset), header.IPv6Version)
} else {
newPacket.Advance(-w.frontHeadroom)
func (w *systemWriteback) WriteReturnPackets(packets [][]byte) error {
if w.linuxTUN != nil {
return common.Error(w.linuxTUN.BatchWrite(packets, w.frontHeadroom))
}
return common.Error(w.tun.Write(newPacket.Bytes()))
var writeErrors []error
for _, packet := range packets {
if PacketOffset > 0 {
PacketFillHeader(packet, header.IPVersion(packet[PacketOffset:]))
}
_, err := w.tun.Write(packet)
if err != nil {
writeErrors = append(writeErrors, err)
}
}
return E.Errors(writeErrors...)
}