From 14c8f75f7a7685db9a507251d1f062a72d6b6a92 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Mon, 6 Jul 2026 23:37:38 +0800 Subject: [PATCH] Add flow tracking --- flow.go | 43 +++++++++++++ flow_dispatch.go | 72 +++++++++++++++++---- nfqueue_linux.go | 1 + stack_gvisor_icmp.go | 148 +++++++++++++++++++++++++++++++++++-------- stack_gvisor_tcp.go | 2 +- stack_gvisor_udp.go | 11 +++- tun.go | 2 +- 7 files changed, 240 insertions(+), 39 deletions(-) diff --git a/flow.go b/flow.go index a920e93..db2cd1d 100644 --- a/flow.go +++ b/flow.go @@ -6,6 +6,7 @@ type FlowVerdict struct { Action FlowAction Port Port Destination netip.AddrPort + NewTracker func() FlowTracker } type FlowAction uint8 @@ -18,6 +19,48 @@ const ( ActionBypass ) +type FlowTracker interface { + AttachFlow(handle FlowHandle) + CountForward(n int) + CountReverse(n int) + FlowEstablished() + CloseFlow(reason FlowCloseReason) +} + +type FlowHandle interface { + CloseFlow() +} + +type FlowCloseReason uint8 + +const ( + FlowCloseReset FlowCloseReason = iota + FlowCloseFinished + FlowCloseTimeout + FlowCloseEvicted + FlowCloseShutdown + FlowCloseInterrupted +) + +func (r FlowCloseReason) String() string { + switch r { + case FlowCloseReset: + return "connection reset" + case FlowCloseFinished: + return "finished" + case FlowCloseTimeout: + return "idle timeout" + case FlowCloseEvicted: + return "evicted" + case FlowCloseShutdown: + return "stack closed" + case FlowCloseInterrupted: + return "interrupted" + default: + return "unknown" + } +} + type Port interface { PortAddresses() (v4 netip.Addr, v6 netip.Addr) PortMTU() uint32 diff --git a/flow_dispatch.go b/flow_dispatch.go index d9ec9a3..b70dc1f 100644 --- a/flow_dispatch.go +++ b/flow_dispatch.go @@ -14,6 +14,7 @@ import ( const ( tcpEstablishedTimeout = 2*time.Hour + 4*time.Minute tcpTransitoryTimeout = 4 * time.Minute + tcpClosingTimeout = 10 * time.Second defaultUDPTimeout = 5 * time.Minute @@ -46,6 +47,7 @@ type forwardFlow struct { reverseRule rewriteRule effectiveMTU uint32 protocol uint8 + tracker FlowTracker clientAddress netip.Addr clientSelector uint16 @@ -62,14 +64,29 @@ type forwardFlow struct { lastReverse atomic.Int64 } +func (f *forwardFlow) close(reason FlowCloseReason) { + if !f.closed.CompareAndSwap(false, true) { + return + } + if f.tracker != nil { + f.tracker.CloseFlow(reason) + } +} + +func (f *forwardFlow) CloseFlow() { + f.close(FlowCloseInterrupted) +} + func (f *forwardFlow) observeReverse(packet *forwardPacket, now int64) { f.lastReverse.Store(now) if packet.protocol != uint8(header.TCPProtocolNumber) { return } - f.established.Store(true) + if f.established.CompareAndSwap(false, true) && f.tracker != nil { + f.tracker.FlowEstablished() + } if packet.tcpFlags&header.TCPFlagRst != 0 { - f.closed.Store(true) + f.close(FlowCloseReset) return } if packet.tcpFlags&header.TCPFlagFin != 0 { @@ -128,6 +145,11 @@ func (d *ForwardDispatcher) Close() { return } d.returnPath.closed.Store(true) + for _, entry := range d.table { + if entry.flow != nil { + entry.flow.close(FlowCloseShutdown) + } + } for port, nat := range d.ports { if nat != nil { port.DetachReturn(&d.returnPath) @@ -147,7 +169,7 @@ func (d *ForwardDispatcher) Dispatch(packet []byte) bool { now := d.now() entry, loaded := d.table[key] if loaded && d.entryExpired(entry, now) { - d.removeEntry(key, entry) + d.removeEntry(key, entry, FlowCloseTimeout) loaded = false } if loaded { @@ -171,7 +193,7 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for if packet.protocol == uint8(header.TCPProtocolNumber) { if packet.tcpFlags&header.TCPFlagRst != 0 { d.forwardToPort(flow, packet, raw) - flow.closed.Store(true) + flow.close(FlowCloseReset) d.tombstoneEntry(entry, now) return true } @@ -186,7 +208,7 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for case ActionAccept: entry.deadline = now + int64(entry.idle) if packet.protocol == uint8(header.TCPProtocolNumber) && packet.tcpFlags&header.TCPFlagRst != 0 { - d.removeEntry(key, entry) + d.removeEntry(key, entry, FlowCloseReset) } return false case ActionReject: @@ -200,7 +222,11 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for } func (d *ForwardDispatcher) judgeAndInstall(key flowKey, packet *forwardPacket, raw []byte, now int64) bool { - verdict := d.handler.JudgeFlow(packet.protocol, packet.source, packet.destination) + var firstPacket []byte + if packet.protocol == uint8(header.UDPProtocolNumber) { + firstPacket = header.UDP(packet.transport).Payload() + } + verdict := d.handler.JudgeFlow(packet.protocol, packet.source, packet.destination, firstPacket) switch verdict.Action { case ActionFlow: if verdict.Port != nil { @@ -249,6 +275,9 @@ func (d *ForwardDispatcher) idleTimeout(protocol uint8, established bool) time.D } func (d *ForwardDispatcher) flowIdle(flow *forwardFlow) time.Duration { + if flow.protocol == uint8(header.TCPProtocolNumber) && flow.finForward && flow.finReverse.Load() { + return tcpClosingTimeout + } established := flow.established.Load() && !flow.finForward && !flow.finReverse.Load() return d.idleTimeout(flow.protocol, established) } @@ -324,6 +353,12 @@ func (d *ForwardDispatcher) createFlow(packet *forwardPacket, verdict FlowVerdic flow.reverseRule.sourcePort = clientDestinationPort flow.reverseRule.rewriteSourcePort = true } + if verdict.NewTracker != nil { + flow.tracker = verdict.NewTracker() + if flow.tracker != nil { + flow.tracker.AttachFlow(flow) + } + } nat.insert(reverseKey, flow) return flow, true } @@ -354,6 +389,9 @@ func (d *ForwardDispatcher) natFor(port Port) *portNAT { 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) { + if flow.tracker != nil { + flow.tracker.CountForward(len(raw)) + } d.rewriteForward(flow, packet) d.resegmentTCP(flow, packet, raw) return @@ -361,6 +399,9 @@ func (d *ForwardDispatcher) forwardToPort(flow *forwardFlow, packet *forwardPack if packet.ipVersion == 4 { ipHdr := packet.network.(header.IPv4) if ipHdr.Flags()&header.IPv4FlagDontFragment == 0 { + if flow.tracker != nil { + flow.tracker.CountForward(len(raw)) + } d.rewriteForward(flow, packet) fragments, ok := fragmentIPv4Packet(ipHdr, flow.effectiveMTU) if ok { @@ -382,6 +423,9 @@ func (d *ForwardDispatcher) forwardToPort(flow *forwardFlow, packet *forwardPack } return } + if flow.tracker != nil { + flow.tracker.CountForward(len(raw)) + } d.rewriteForward(flow, packet) d.stagePort(flow.nat, raw) } @@ -460,10 +504,13 @@ func (d *ForwardDispatcher) tombstoneEntry(entry *flowEntry, now int64) { entry.deadline = now + int64(entry.idle) } -func (d *ForwardDispatcher) removeEntry(key flowKey, entry *flowEntry) { +func (d *ForwardDispatcher) removeEntry(key flowKey, entry *flowEntry, reason FlowCloseReason) { delete(d.table, key) if entry.flow != nil { - entry.flow.closed.Store(true) + if reason == FlowCloseTimeout && entry.flow.finForward && entry.flow.finReverse.Load() { + reason = FlowCloseFinished + } + entry.flow.close(reason) entry.flow.nat.delete(entry.flow.reverseKey) } } @@ -484,7 +531,7 @@ func (d *ForwardDispatcher) evictEntries(now int64) { ) for key, entry := range d.table { if d.entryExpired(entry, now) { - d.removeEntry(key, entry) + d.removeEntry(key, entry, FlowCloseTimeout) freed++ } else if oldest == nil || entry.deadline < oldest.deadline { oldestKey = key @@ -496,7 +543,7 @@ func (d *ForwardDispatcher) evictEntries(now int64) { } } if freed == 0 && oldest != nil { - d.removeEntry(oldestKey, oldest) + d.removeEntry(oldestKey, oldest, FlowCloseEvicted) } } @@ -510,7 +557,7 @@ func (d *ForwardDispatcher) maybeSweep(now int64) { if entry.action == ActionFlow && entry.flow.closed.Load() { d.tombstoneEntry(entry, now) } else if d.entryExpired(entry, now) { - d.removeEntry(key, entry) + d.removeEntry(key, entry, FlowCloseTimeout) } visited++ if visited >= flowSweepLimit { @@ -579,6 +626,9 @@ func (r *forwardReturn) ReturnPackets(packets [][]byte) [][]byte { if flow.closed.Load() { continue } + if flow.tracker != nil { + flow.tracker.CountReverse(len(raw) - headroom) + } flow.observeReverse(&parsed, now) if parsed.isTCPSyn() { applyRewriteRaw(&parsed, &flow.reverseRule) diff --git a/nfqueue_linux.go b/nfqueue_linux.go index 8513aa6..7046204 100644 --- a/nfqueue_linux.go +++ b/nfqueue_linux.go @@ -269,6 +269,7 @@ func (h *nfqueueHandler) handlePacket(attr nfqueue.Attribute) int { packet.protocol, packet.source, packet.destination, + packet.firstPacket, ) // Use NfRepeat for bypass/reset so the packet re-enters the chain diff --git a/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index 979f4ad..96466c1 100644 --- a/stack_gvisor_icmp.go +++ b/stack_gvisor_icmp.go @@ -28,7 +28,7 @@ type ICMPForwarder struct { returnPath icmpForwarderReturn flowAccess sync.Mutex - flows map[icmpFlowKey]time.Time + flows map[icmpFlowKey]*icmpFlow lastSweep time.Time attachedPorts map[Port]bool } @@ -40,12 +40,32 @@ type icmpFlowKey struct { identifier uint16 } +type icmpFlow struct { + port Port + tracker FlowTracker + deadline time.Time + closed atomic.Bool +} + +func (f *icmpFlow) close(reason FlowCloseReason) { + if !f.closed.CompareAndSwap(false, true) { + return + } + if f.tracker != nil { + f.tracker.CloseFlow(reason) + } +} + +func (f *icmpFlow) CloseFlow() { + f.close(FlowCloseInterrupted) +} + func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger) *ICMPForwarder { forwarder := &ICMPForwarder{ stack: stack, handler: handler, logger: logger, - flows: make(map[icmpFlowKey]time.Time), + flows: make(map[icmpFlowKey]*icmpFlow), attachedPorts: make(map[Port]bool), } forwarder.returnPath.forwarder = forwarder @@ -56,6 +76,10 @@ func (f *ICMPForwarder) Close() error { f.returnPath.closed.Store(true) f.flowAccess.Lock() defer f.flowAccess.Unlock() + for key, flow := range f.flows { + flow.close(FlowCloseShutdown) + delete(f.flows, key) + } for port := range f.attachedPorts { port.DetachReturn(&f.returnPath) delete(f.attachedPorts, port) @@ -71,16 +95,25 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa return false } identifier := icmpHdr.Ident() + key := icmpFlowKey{ + source: AddrFromAddress(ipHdr.SourceAddress()), + destination: AddrFromAddress(ipHdr.DestinationAddress()), + identifier: identifier, + } + if f.forwardCached(key, pkt) { + return true + } verdict := f.handler.JudgeFlow( uint8(header.ICMPv4ProtocolNumber), - netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier), - netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier), + netip.AddrPortFrom(key.source, identifier), + netip.AddrPortFrom(key.destination, identifier), + nil, ) switch verdict.Action { case ActionReject, ActionDrop: return true case ActionFlow: - if f.forwardFlow(verdict.Port, false, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) { + if f.installFlow(key, verdict, pkt) { return true } } @@ -117,16 +150,26 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa return false } identifier := icmpHdr.Ident() + key := icmpFlowKey{ + v6: true, + source: AddrFromAddress(ipHdr.SourceAddress()), + destination: AddrFromAddress(ipHdr.DestinationAddress()), + identifier: identifier, + } + if f.forwardCached(key, pkt) { + return true + } verdict := f.handler.JudgeFlow( uint8(header.ICMPv6ProtocolNumber), - netip.AddrPortFrom(AddrFromAddress(ipHdr.SourceAddress()), identifier), - netip.AddrPortFrom(AddrFromAddress(ipHdr.DestinationAddress()), identifier), + netip.AddrPortFrom(key.source, identifier), + netip.AddrPortFrom(key.destination, identifier), + nil, ) switch verdict.Action { case ActionReject, ActionDrop: return true case ActionFlow: - if f.forwardFlow(verdict.Port, true, AddrFromAddress(ipHdr.SourceAddress()), AddrFromAddress(ipHdr.DestinationAddress()), identifier, pkt) { + if f.installFlow(key, verdict, pkt) { return true } } @@ -163,13 +206,38 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa } } -func (f *ICMPForwarder) forwardFlow(port Port, v6 bool, source netip.Addr, destination netip.Addr, identifier uint16, pkt *stack.PacketBuffer) bool { +func (f *ICMPForwarder) forwardCached(key icmpFlowKey, pkt *stack.PacketBuffer) bool { + now := time.Now() + f.flowAccess.Lock() + flow, loaded := f.flows[key] + if loaded { + if flow.closed.Load() { + delete(f.flows, key) + loaded = false + } else if now.After(flow.deadline) { + delete(f.flows, key) + flow.close(FlowCloseTimeout) + loaded = false + } else { + flow.deadline = now.Add(defaultICMPTimeout) + } + } + f.flowAccess.Unlock() + if !loaded { + return false + } + f.writeToPort(flow, pkt) + return true +} + +func (f *ICMPForwarder) installFlow(key icmpFlowKey, verdict FlowVerdict, pkt *stack.PacketBuffer) bool { + port := verdict.Port if port == nil { return false } inet4Address, inet6Address := port.PortAddresses() portAddress := inet4Address - if v6 { + if key.v6 { portAddress = inet6Address } if !portAddress.IsValid() || !portAddress.IsUnspecified() { @@ -188,14 +256,29 @@ func (f *ICMPForwarder) forwardFlow(port Port, v6 bool, source netip.Addr, desti now := time.Now() if now.Sub(f.lastSweep) >= defaultICMPTimeout { f.lastSweep = now - for key, deadline := range f.flows { - if now.After(deadline) { - delete(f.flows, key) + for flowKey, cachedFlow := range f.flows { + if cachedFlow.closed.Load() { + delete(f.flows, flowKey) + } else if now.After(cachedFlow.deadline) { + delete(f.flows, flowKey) + cachedFlow.close(FlowCloseTimeout) } } } - f.flows[icmpFlowKey{v6: v6, source: source, destination: destination, identifier: identifier}] = now.Add(defaultICMPTimeout) + flow := &icmpFlow{port: port, deadline: now.Add(defaultICMPTimeout)} + if verdict.NewTracker != nil { + flow.tracker = verdict.NewTracker() + } + f.flows[key] = flow f.flowAccess.Unlock() + if flow.tracker != nil { + flow.tracker.AttachFlow(flow) + } + f.writeToPort(flow, pkt) + return true +} + +func (f *ICMPForwarder) writeToPort(flow *icmpFlow, pkt *stack.PacketBuffer) { networkSlice := pkt.NetworkHeader().Slice() transportSlice := pkt.TransportHeader().Slice() dataSlice := pkt.Data().AsRange().ToSlice() @@ -203,27 +286,34 @@ func (f *ICMPForwarder) forwardFlow(port Port, v6 bool, source netip.Addr, desti packetSlice = append(packetSlice, networkSlice...) packetSlice = append(packetSlice, transportSlice...) packetSlice = append(packetSlice, dataSlice...) - err := port.WritePackets([][]byte{packetSlice}) + if flow.tracker != nil { + flow.tracker.CountForward(len(packetSlice)) + } + err := flow.port.WritePackets([][]byte{packetSlice}) if err != nil { f.logger.Trace(E.Cause(err, "forward ICMP packet")) } - return true } -func (f *ICMPForwarder) lookupFlow(key icmpFlowKey) bool { +func (f *ICMPForwarder) lookupFlow(key icmpFlowKey) *icmpFlow { f.flowAccess.Lock() defer f.flowAccess.Unlock() - deadline, loaded := f.flows[key] + flow, loaded := f.flows[key] if !loaded { - return false + return nil + } + if flow.closed.Load() { + delete(f.flows, key) + return nil } now := time.Now() - if now.After(deadline) { + if now.After(flow.deadline) { delete(f.flows, key) - return false + flow.close(FlowCloseTimeout) + return nil } - f.flows[key] = now.Add(defaultICMPTimeout) - return true + flow.deadline = now.Add(defaultICMPTimeout) + return flow } type icmpForwarderReturn struct { @@ -289,9 +379,13 @@ func (f *ICMPForwarder) returnPacket(packet []byte) bool { default: return false } - if !f.lookupFlow(key) { + flow := f.lookupFlow(key) + if flow == nil { return false } + if flow.tracker != nil { + flow.tracker.CountReverse(len(packet)) + } return f.writeBack(packet, header.IPv4ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress()) case header.IPv6Version: ipHdr := header.IPv6(packet) @@ -308,9 +402,13 @@ func (f *ICMPForwarder) returnPacket(packet []byte) bool { destination: AddrFromAddress(ipHdr.SourceAddress()), identifier: icmpHdr.Ident(), } - if !f.lookupFlow(key) { + flow := f.lookupFlow(key) + if flow == nil { return false } + if flow.tracker != nil { + flow.tracker.CountReverse(len(packet)) + } return f.writeBack(packet, header.IPv6ProtocolNumber, ipHdr.SourceAddress(), ipHdr.DestinationAddress()) default: return false diff --git a/stack_gvisor_tcp.go b/stack_gvisor_tcp.go index 371c480..f432e28 100644 --- a/stack_gvisor_tcp.go +++ b/stack_gvisor_tcp.go @@ -77,7 +77,7 @@ func (f *TCPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pac func (f *TCPForwarder) Forward(r *tcp.ForwarderRequest) { source := M.SocksaddrFrom(AddrFromAddress(r.ID().RemoteAddress), r.ID().RemotePort) destination := M.SocksaddrFrom(AddrFromAddress(r.ID().LocalAddress), r.ID().LocalPort) - switch f.handler.JudgeFlow(uint8(header.TCPProtocolNumber), source.AddrPort(), destination.AddrPort()).Action { + switch f.handler.JudgeFlow(uint8(header.TCPProtocolNumber), source.AddrPort(), destination.AddrPort(), nil).Action { case ActionReject: r.Complete(true) return diff --git a/stack_gvisor_udp.go b/stack_gvisor_udp.go index fb31ba9..3e4afdd 100644 --- a/stack_gvisor_udp.go +++ b/stack_gvisor_udp.go @@ -57,7 +57,16 @@ func (f *UDPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pac func rangeIterate(r stack.Range, fn func(*buffer.View)) func (f *UDPForwarder) PreparePacketConnection(source M.Socksaddr, destination M.Socksaddr, userData any) (bool, context.Context, N.PacketWriter, N.CloseHandlerFunc) { - switch f.handler.JudgeFlow(uint8(header.UDPProtocolNumber), source.AddrPort(), destination.AddrPort()).Action { + firstPacketBuffer := userData.(*stack.PacketBuffer) + var firstPacket []byte + rangeIterate(firstPacketBuffer.Data().AsRange(), func(view *buffer.View) { + if firstPacket == nil { + firstPacket = view.AsSlice() + } else { + firstPacket = append(firstPacket[:len(firstPacket):len(firstPacket)], view.AsSlice()...) + } + }) + switch f.handler.JudgeFlow(uint8(header.UDPProtocolNumber), source.AddrPort(), destination.AddrPort(), firstPacket).Action { case ActionReject: gWriteUnreachable(f.stack, userData.(*stack.PacketBuffer)) return false, nil, nil, nil diff --git a/tun.go b/tun.go index 7d10c75..c6518f4 100644 --- a/tun.go +++ b/tun.go @@ -19,7 +19,7 @@ import ( ) type Handler interface { - JudgeFlow(network uint8, source netip.AddrPort, destination netip.AddrPort) FlowVerdict + JudgeFlow(network uint8, source netip.AddrPort, destination netip.AddrPort, firstPacket []byte) FlowVerdict N.TCPConnectionHandlerEx N.UDPConnectionHandlerEx }