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

631
flow_dispatch.go Normal file
View file

@ -0,0 +1,631 @@
package tun
import (
"net/netip"
"sync/atomic"
"time"
"github.com/sagernet/sing-tun/gtcpip"
"github.com/sagernet/sing-tun/gtcpip/header"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
)
const (
tcpEstablishedTimeout = 2*time.Hour + 4*time.Minute
tcpTransitoryTimeout = 4 * time.Minute
defaultUDPTimeout = 5 * time.Minute
defaultICMPTimeout = time.Minute
flowTombstoneTimeout = 4 * time.Minute
flowTableCapacity = 16384
flowSweepInterval = 30 * time.Second
flowSweepLimit = flowTableCapacity / int(flowTombstoneTimeout/flowSweepInterval)
)
type ForwardWriteback interface {
ReturnHeadroom() int
WriteReturnPackets(packets [][]byte) error
}
type flowEntry struct {
action FlowAction
deadline int64
idle time.Duration
flow *forwardFlow
}
type forwardFlow struct {
nat *portNAT
reverseKey flowKey
forwardRule rewriteRule
reverseRule rewriteRule
effectiveMTU uint32
protocol uint8
clientAddress netip.Addr
clientSelector uint16
clientDestinationAddress netip.Addr
clientDestinationPort uint16
serverAddress netip.Addr
dnatAddress bool
dnatPort bool
finForward bool
established atomic.Bool
finReverse atomic.Bool
closed atomic.Bool
lastReverse atomic.Int64
}
func (f *forwardFlow) observeReverse(packet *forwardPacket, now int64) {
f.lastReverse.Store(now)
if packet.protocol != uint8(header.TCPProtocolNumber) {
return
}
f.established.Store(true)
if packet.tcpFlags&header.TCPFlagRst != 0 {
f.closed.Store(true)
return
}
if packet.tcpFlags&header.TCPFlagFin != 0 {
f.finReverse.Store(true)
}
}
type ForwardDispatcher struct {
epoch time.Time
handler Handler
writeback ForwardWriteback
logger logger.Logger
udpTimeout time.Duration
icmpTimeout time.Duration
table map[flowKey]*flowEntry
lastSweep int64
ports map[Port]*portNAT
natList atomic.Pointer[[]*portNAT]
activeNATs []*portNAT
writebackBatch [][]byte
returnPath forwardReturn
segmentBuffers [][]byte
segmentSizes []int
}
func NewForwardDispatcher(handler Handler, writeback ForwardWriteback, logger logger.Logger, udpTimeout time.Duration, icmpTimeout time.Duration) *ForwardDispatcher {
dispatcher := &ForwardDispatcher{
epoch: time.Now(),
handler: handler,
writeback: writeback,
logger: logger,
udpTimeout: udpTimeout,
icmpTimeout: icmpTimeout,
table: make(map[flowKey]*flowEntry),
ports: make(map[Port]*portNAT),
}
if dispatcher.udpTimeout <= 0 {
dispatcher.udpTimeout = defaultUDPTimeout
}
if dispatcher.icmpTimeout <= 0 {
dispatcher.icmpTimeout = defaultICMPTimeout
}
dispatcher.returnPath.dispatcher = dispatcher
return dispatcher
}
func (d *ForwardDispatcher) now() int64 {
return int64(time.Since(d.epoch))
}
func (d *ForwardDispatcher) Close() {
if d == nil {
return
}
d.returnPath.closed.Store(true)
for port, nat := range d.ports {
if nat != nil {
port.DetachReturn(&d.returnPath)
}
}
}
func (d *ForwardDispatcher) Dispatch(packet []byte) bool {
if d == nil {
return false
}
parsed, ok := parseForwardPacket(packet)
if !ok || parsed.fragment || !parsed.hasFlow {
return false
}
key := parsed.flowKey()
now := d.now()
entry, loaded := d.table[key]
if loaded && d.entryExpired(entry, now) {
d.removeEntry(key, entry)
loaded = false
}
if loaded {
return d.handleHit(key, entry, &parsed, packet, now)
}
if parsed.protocol == uint8(header.TCPProtocolNumber) &&
(parsed.tcpFlags&header.TCPFlagSyn == 0 || parsed.tcpFlags&header.TCPFlagAck != 0) {
return false
}
return d.judgeAndInstall(key, &parsed, packet, now)
}
func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *forwardPacket, raw []byte, now int64) bool {
switch entry.action {
case ActionFlow:
flow := entry.flow
if flow.closed.Load() {
d.tombstoneEntry(entry, now)
return true
}
if packet.protocol == uint8(header.TCPProtocolNumber) {
if packet.tcpFlags&header.TCPFlagRst != 0 {
d.forwardToPort(flow, packet, raw)
flow.closed.Store(true)
d.tombstoneEntry(entry, now)
return true
}
if packet.tcpFlags&header.TCPFlagFin != 0 {
flow.finForward = true
}
}
entry.idle = d.flowIdle(flow)
entry.deadline = now + int64(entry.idle)
d.forwardToPort(flow, packet, raw)
return true
case ActionAccept:
entry.deadline = now + int64(entry.idle)
if packet.protocol == uint8(header.TCPProtocolNumber) && packet.tcpFlags&header.TCPFlagRst != 0 {
d.removeEntry(key, entry)
}
return false
case ActionReject:
entry.deadline = now + int64(entry.idle)
d.stageReject(packet)
return true
default:
entry.deadline = now + int64(entry.idle)
return true
}
}
func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte, now int64) bool {
verdict := d.handler.JudgeFlow(packet.protocol, packet.source, packet.destination)
switch verdict.Action {
case ActionFlow:
if verdict.Port != nil {
flow, created := d.createFlow(packet, verdict)
if created {
entry := &flowEntry{action: ActionFlow, flow: flow, idle: d.flowIdle(flow)}
entry.deadline = now + int64(entry.idle)
d.insertEntry(key, entry, now)
d.forwardToPort(flow, packet, raw)
return true
}
}
d.installSimple(key, ActionAccept, packet.protocol, now)
return false
case ActionReject:
d.installSimple(key, ActionReject, packet.protocol, now)
d.stageReject(packet)
return true
case ActionDrop:
d.installSimple(key, ActionDrop, packet.protocol, now)
return true
default:
d.installSimple(key, ActionAccept, packet.protocol, now)
return false
}
}
func (d *ForwardDispatcher) installSimple(key flowKey, action FlowAction, protocol uint8, now int64) {
entry := &flowEntry{action: action, idle: d.idleTimeout(protocol, false)}
entry.deadline = now + int64(entry.idle)
d.insertEntry(key, entry, now)
}
func (d *ForwardDispatcher) idleTimeout(protocol uint8, established bool) time.Duration {
switch protocol {
case uint8(header.TCPProtocolNumber):
if established {
return tcpEstablishedTimeout
}
return tcpTransitoryTimeout
case uint8(header.UDPProtocolNumber):
return d.udpTimeout
default:
return d.icmpTimeout
}
}
func (d *ForwardDispatcher) flowIdle(flow *forwardFlow) time.Duration {
established := flow.established.Load() && !flow.finForward && !flow.finReverse.Load()
return d.idleTimeout(flow.protocol, established)
}
func (d *ForwardDispatcher) createFlow(packet *forwardPacket, verdict FlowVerdict) (*forwardFlow, bool) {
var portAddress netip.Addr
inet4Address, inet6Address := verdict.Port.PortAddresses()
if packet.ipVersion == 6 {
portAddress = inet6Address
} else {
portAddress = inet4Address
}
if !portAddress.IsValid() {
return nil, false
}
effectiveMTU := verdict.Port.PortMTU()
if packet.ipVersion == 6 && effectiveMTU != 0 && effectiveMTU < header.IPv6MinimumMTU {
return nil, false
}
isICMP := isICMPProtocol(packet.protocol)
clientDestinationAddress := packet.destination.Addr()
clientDestinationPort := packet.destination.Port()
serverAddress := clientDestinationAddress
serverPort := clientDestinationPort
if verdict.Destination.Addr().IsValid() {
serverAddress = verdict.Destination.Addr()
}
if verdict.Destination.Port() != 0 && !isICMP {
serverPort = verdict.Destination.Port()
}
nat := d.natFor(verdict.Port)
if nat == nil {
return nil, false
}
selector, reverseKey, allocated := nat.allocateSelector(packet.protocol, portAddress, serverAddress, serverPort, packet.source.Port())
if !allocated {
return nil, false
}
flow := &forwardFlow{
nat: nat,
reverseKey: reverseKey,
effectiveMTU: effectiveMTU,
protocol: packet.protocol,
clientAddress: packet.source.Addr(),
clientSelector: packet.source.Port(),
clientDestinationAddress: clientDestinationAddress,
clientDestinationPort: clientDestinationPort,
serverAddress: serverAddress,
dnatAddress: serverAddress != clientDestinationAddress,
dnatPort: serverPort != clientDestinationPort && !isICMP,
}
flow.forwardRule = rewriteRule{
sourceAddress: tcpip.AddrFromSlice(portAddress.AsSlice()),
sourcePort: selector,
rewriteSourcePort: true,
}
if flow.dnatAddress {
flow.forwardRule.destinationAddress = tcpip.AddrFromSlice(serverAddress.AsSlice())
}
if flow.dnatPort {
flow.forwardRule.destinationPort = serverPort
flow.forwardRule.rewriteDestinationPort = true
}
flow.reverseRule = rewriteRule{
destinationAddress: tcpip.AddrFromSlice(flow.clientAddress.AsSlice()),
destinationPort: flow.clientSelector,
rewriteDestinationPort: true,
}
if flow.dnatAddress {
flow.reverseRule.sourceAddress = tcpip.AddrFromSlice(clientDestinationAddress.AsSlice())
}
if flow.dnatPort {
flow.reverseRule.sourcePort = clientDestinationPort
flow.reverseRule.rewriteSourcePort = true
}
nat.insert(reverseKey, flow)
return flow, true
}
func (d *ForwardDispatcher) natFor(port Port) *portNAT {
nat, loaded := d.ports[port]
if loaded {
return nat
}
err := port.AttachReturn(&d.returnPath)
if err != nil {
d.logger.Trace(E.Cause(err, "attach return path"))
d.ports[port] = nil
return nil
}
nat = newPortNAT(port)
d.ports[port] = nat
var natList []*portNAT
current := d.natList.Load()
if current != nil {
natList = append(natList, *current...)
}
natList = append(natList, nat)
d.natList.Store(&natList)
return nat
}
func (d *ForwardDispatcher) forwardToPort(flow *forwardFlow, packet *forwardPacket, raw []byte) {
if flow.effectiveMTU != 0 && uint32(len(raw)) > flow.effectiveMTU {
if packet.protocol == uint8(header.TCPProtocolNumber) {
d.rewriteForward(flow, packet)
d.resegmentTCP(flow, packet, raw)
return
}
if packet.ipVersion == 4 {
ipHdr := packet.network.(header.IPv4)
if ipHdr.Flags()&header.IPv4FlagDontFragment == 0 {
d.rewriteForward(flow, packet)
fragments, ok := fragmentIPv4Packet(ipHdr, flow.effectiveMTU)
if ok {
for _, fragment := range fragments {
d.stagePort(flow.nat, fragment)
}
}
return
}
reply, ok := buildFragmentationNeeded(ipHdr, flow.effectiveMTU, d.writeback.ReturnHeadroom())
if ok {
d.writebackBatch = append(d.writebackBatch, reply)
}
return
}
reply, ok := buildPacketTooBig(packet.network.(header.IPv6), flow.effectiveMTU, d.writeback.ReturnHeadroom())
if ok {
d.writebackBatch = append(d.writebackBatch, reply)
}
return
}
d.rewriteForward(flow, packet)
d.stagePort(flow.nat, raw)
}
func (d *ForwardDispatcher) rewriteForward(flow *forwardFlow, packet *forwardPacket) {
if packet.isTCPSyn() {
applyRewriteRaw(packet, &flow.forwardRule)
clampTCPMSS(packet, flow.effectiveMTU)
recomputeChecksums(packet)
} else {
applyRewrite(packet, &flow.forwardRule)
}
}
func (d *ForwardDispatcher) stagePort(nat *portNAT, packet []byte) {
if len(nat.pending) == 0 {
d.activeNATs = append(d.activeNATs, nat)
}
nat.pending = append(nat.pending, packet)
}
func (d *ForwardDispatcher) flushPort(nat *portNAT) {
if len(nat.pending) == 0 {
return
}
err := nat.port.WritePackets(nat.pending)
if err != nil {
d.logger.Trace(E.Cause(err, "forward packets"))
}
nat.pending = nat.pending[:0]
}
func (d *ForwardDispatcher) stageReject(packet *forwardPacket) {
reply, ok := buildReject(packet, d.writeback.ReturnHeadroom())
if ok {
d.writebackBatch = append(d.writebackBatch, reply)
}
}
func (d *ForwardDispatcher) Flush() {
if d == nil {
return
}
for _, nat := range d.activeNATs {
d.flushPort(nat)
}
d.activeNATs = d.activeNATs[:0]
if len(d.writebackBatch) > 0 {
err := d.writeback.WriteReturnPackets(d.writebackBatch)
if err != nil {
d.logger.Trace(E.Cause(err, "write back packets"))
}
d.writebackBatch = d.writebackBatch[:0]
}
d.maybeSweep(d.now())
}
func (d *ForwardDispatcher) entryExpired(entry *flowEntry, now int64) bool {
if now <= entry.deadline {
return false
}
if entry.action == ActionFlow {
lastReverse := entry.flow.lastReverse.Load()
reverseDeadline := lastReverse + int64(entry.idle)
if lastReverse != 0 && now <= reverseDeadline {
entry.deadline = reverseDeadline
return false
}
}
return true
}
func (d *ForwardDispatcher) tombstoneEntry(entry *flowEntry, now int64) {
entry.action = ActionDrop
entry.idle = flowTombstoneTimeout
entry.deadline = now + int64(entry.idle)
}
func (d *ForwardDispatcher) removeEntry(key flowKey, entry *flowEntry) {
delete(d.table, key)
if entry.flow != nil {
entry.flow.closed.Store(true)
entry.flow.nat.delete(entry.flow.reverseKey)
}
}
func (d *ForwardDispatcher) insertEntry(key flowKey, entry *flowEntry, now int64) {
if len(d.table) >= flowTableCapacity {
d.evictEntries(now)
}
d.table[key] = entry
}
func (d *ForwardDispatcher) evictEntries(now int64) {
var (
freed int
visited int
oldestKey flowKey
oldest *flowEntry
)
for key, entry := range d.table {
if d.entryExpired(entry, now) {
d.removeEntry(key, entry)
freed++
} else if oldest == nil || entry.deadline < oldest.deadline {
oldestKey = key
oldest = entry
}
visited++
if visited >= flowSweepLimit {
break
}
}
if freed == 0 && oldest != nil {
d.removeEntry(oldestKey, oldest)
}
}
func (d *ForwardDispatcher) maybeSweep(now int64) {
if now-d.lastSweep < int64(flowSweepInterval) {
return
}
d.lastSweep = now
visited := 0
for key, entry := range d.table {
if entry.action == ActionFlow && entry.flow.closed.Load() {
d.tombstoneEntry(entry, now)
} else if d.entryExpired(entry, now) {
d.removeEntry(key, entry)
}
visited++
if visited >= flowSweepLimit {
break
}
}
}
func isICMPProtocol(protocol uint8) bool {
return protocol == uint8(header.ICMPv4ProtocolNumber) || protocol == uint8(header.ICMPv6ProtocolNumber)
}
var _ Return = (*forwardReturn)(nil)
type forwardReturn struct {
dispatcher *ForwardDispatcher
closed atomic.Bool
}
func (r *forwardReturn) ReturnHeadroom() int {
return r.dispatcher.writeback.ReturnHeadroom()
}
func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte {
if r.closed.Load() {
return packets
}
natListPtr := r.dispatcher.natList.Load()
if natListPtr == nil {
return packets
}
natList := *natListPtr
headroom := r.dispatcher.writeback.ReturnHeadroom()
unconsumed := packets[:0]
var writeBatch [][]byte
now := r.dispatcher.now()
for _, raw := range packets {
if len(raw) < headroom+header.IPv4MinimumSize {
unconsumed = append(unconsumed, raw)
continue
}
parsed, ok := parseForwardPacket(raw[headroom:])
if !ok || parsed.fragment {
unconsumed = append(unconsumed, raw)
continue
}
if !parsed.hasFlow {
if parsed.isICMPError() && returnICMPError(natList, &parsed) {
writeBatch = append(writeBatch, raw)
} else {
unconsumed = append(unconsumed, raw)
}
continue
}
var flow *forwardFlow
for _, nat := range natList {
flow = nat.lookup(parsed.flowKey())
if flow != nil {
break
}
}
if flow == nil {
unconsumed = append(unconsumed, raw)
continue
}
if flow.closed.Load() {
continue
}
flow.observeReverse(&parsed, now)
if parsed.isTCPSyn() {
applyRewriteRaw(&parsed, &flow.reverseRule)
clampTCPMSS(&parsed, flow.effectiveMTU)
recomputeChecksums(&parsed)
} else {
applyRewrite(&parsed, &flow.reverseRule)
}
writeBatch = append(writeBatch, raw)
}
if len(writeBatch) > 0 {
err := r.dispatcher.writeback.WriteReturnPackets(writeBatch)
if err != nil {
r.dispatcher.logger.Trace(E.Cause(err, "write return packets"))
}
}
return unconsumed
}
func returnICMPError(natList []*portNAT, parsed *forwardPacket) bool {
inner, ok := parsed.icmpErrorInner()
if !ok {
return false
}
embedded, parsedInner := parseEmbedded(inner)
if !parsedInner {
return false
}
key := embedded.flowKey().reversed()
var flow *forwardFlow
for _, nat := range natList {
flow = nat.lookup(key)
if flow != nil {
break
}
}
if flow == nil || flow.closed.Load() {
return false
}
rewriteEmbeddedSource(&embedded, tcpip.AddrFromSlice(flow.clientAddress.AsSlice()), flow.clientSelector, true)
if flow.dnatAddress || flow.dnatPort {
rewriteEmbeddedDestination(&embedded, tcpip.AddrFromSlice(flow.clientDestinationAddress.AsSlice()), flow.clientDestinationPort, flow.dnatPort)
}
parsed.network.SetDestinationAddr(flow.clientAddress)
if parsed.network.SourceAddr() == flow.serverAddress {
parsed.network.SetSourceAddr(flow.clientDestinationAddress)
}
recomputeChecksums(parsed)
return true
}