Add flow tracking

This commit is contained in:
世界 2026-07-06 23:37:38 +08:00
parent ed63adda33
commit 14c8f75f7a
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
7 changed files with 240 additions and 39 deletions

43
flow.go
View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

2
tun.go
View file

@ -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
}