diff --git a/flow.go b/flow.go index 0954601..1b1f34d 100644 --- a/flow.go +++ b/flow.go @@ -41,27 +41,16 @@ 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" + return "connection reset" } } diff --git a/flow_dispatch.go b/flow_dispatch.go index 8911bf1..3f71751 100644 --- a/flow_dispatch.go +++ b/flow_dispatch.go @@ -59,15 +59,16 @@ type forwardFlow struct { dnatAddress bool dnatPort bool - finForward bool + finForward atomic.Bool established atomic.Bool finReverse atomic.Bool + reported atomic.Bool closed atomic.Bool lastReverse atomic.Int64 } -func (f *forwardFlow) close(reason FlowCloseReason) { - if !f.closed.CompareAndSwap(false, true) { +func (f *forwardFlow) report(reason FlowCloseReason) { + if !f.reported.CompareAndSwap(false, true) { return } if f.tracker != nil { @@ -75,8 +76,15 @@ func (f *forwardFlow) close(reason FlowCloseReason) { } } +func (f *forwardFlow) close(reason FlowCloseReason) { + if !f.closed.CompareAndSwap(false, true) { + return + } + f.report(reason) +} + func (f *forwardFlow) CloseFlow() { - f.close(FlowCloseInterrupted) + f.close(FlowCloseReset) } func (f *forwardFlow) observeReverse(packet *forwardPacket, now int64) { @@ -93,6 +101,8 @@ func (f *forwardFlow) observeReverse(packet *forwardPacket, now int64) { } if packet.tcpFlags&header.TCPFlagFin != 0 { f.finReverse.Store(true) + } else if f.finReverse.Load() && f.finForward.Load() { + f.report(FlowCloseFinished) } } @@ -159,7 +169,7 @@ func (d *ForwardDispatcher) Close() { d.returnPath.closed.Store(true) for _, entry := range d.table { if entry.flow != nil { - entry.flow.close(FlowCloseShutdown) + entry.flow.close(FlowCloseReset) } } for port, nat := range d.ports { @@ -202,6 +212,7 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for d.tombstoneEntry(entry, now) return true } + var flowFinished bool if packet.protocol == uint8(header.TCPProtocolNumber) { if packet.tcpFlags&header.TCPFlagRst != 0 { d.forwardToPort(flow, packet, raw) @@ -210,12 +221,17 @@ func (d *ForwardDispatcher) handleHit(key flowKey, entry *flowEntry, packet *for return true } if packet.tcpFlags&header.TCPFlagFin != 0 { - flow.finForward = true + flow.finForward.Store(true) + } else if flow.finForward.Load() && flow.finReverse.Load() { + flowFinished = true } } entry.idle = d.flowIdle(flow) entry.deadline = now + int64(entry.idle) d.forwardToPort(flow, packet, raw) + if flowFinished { + flow.report(FlowCloseFinished) + } return true case ActionAccept: if packet.protocol == uint8(header.TCPProtocolNumber) { @@ -302,13 +318,13 @@ 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() { + if flow.protocol == uint8(header.TCPProtocolNumber) && flow.finForward.Load() && flow.finReverse.Load() { return tcpClosingTimeout } if flow.udpTimeout > 0 { return flow.udpTimeout } - established := flow.established.Load() && !flow.finForward && !flow.finReverse.Load() + established := flow.established.Load() && !flow.finForward.Load() && !flow.finReverse.Load() return d.idleTimeout(flow.protocol, established) } @@ -569,7 +585,7 @@ func (d *ForwardDispatcher) tombstoneEntry(entry *flowEntry, now int64) { func (d *ForwardDispatcher) removeEntry(key flowKey, entry *flowEntry, reason FlowCloseReason) { delete(d.table, key) if entry.flow != nil { - if reason == FlowCloseTimeout && entry.flow.finForward && entry.flow.finReverse.Load() { + if reason == FlowCloseTimeout && entry.flow.finForward.Load() && entry.flow.finReverse.Load() { reason = FlowCloseFinished } entry.flow.close(reason) @@ -605,7 +621,7 @@ func (d *ForwardDispatcher) evictEntries(now int64) { } } if freed == 0 && oldest != nil { - d.removeEntry(oldestKey, oldest, FlowCloseEvicted) + d.removeEntry(oldestKey, oldest, FlowCloseReset) } } diff --git a/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index 9b742e5..11e82af 100644 --- a/stack_gvisor_icmp.go +++ b/stack_gvisor_icmp.go @@ -57,7 +57,7 @@ func (f *icmpFlow) close(reason FlowCloseReason) { } func (f *icmpFlow) CloseFlow() { - f.close(FlowCloseInterrupted) + f.close(FlowCloseReset) } func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger) *ICMPForwarder { @@ -77,7 +77,7 @@ func (f *ICMPForwarder) Close() error { f.flowAccess.Lock() defer f.flowAccess.Unlock() for key, flow := range f.flows { - flow.close(FlowCloseShutdown) + flow.close(FlowCloseReset) delete(f.flows, key) } for port := range f.attachedPorts {