diff --git a/flow_dispatch.go b/flow_dispatch.go index c5e6ea6..87f2c02 100644 --- a/flow_dispatch.go +++ b/flow_dispatch.go @@ -114,11 +114,12 @@ type ForwardDispatcher struct { udpTimeout time.Duration icmpTimeout time.Duration - table map[flowKey]*flowEntry - lastSweep int64 - ports map[Port]*portNAT - natList atomic.Pointer[[]*portNAT] - revNAT atomic.Pointer[map[netip.Addr]*portNAT] + table map[flowKey]*flowEntry + lastSweep int64 + resetPending atomic.Bool + ports map[Port]*portNAT + natList atomic.Pointer[[]*portNAT] + revNAT atomic.Pointer[map[netip.Addr]*portNAT] activeNATs []*portNAT writebackBatch [][]byte @@ -544,10 +545,22 @@ func (d *ForwardDispatcher) stageReject(packet *forwardPacket) { } } +func (d *ForwardDispatcher) ResetNetwork() { + if d == nil { + return + } + d.resetPending.Store(true) +} + func (d *ForwardDispatcher) Flush() { if d == nil { return } + if d.resetPending.Swap(false) { + for key, entry := range d.table { + d.removeEntry(key, entry, FlowCloseReset) + } + } for _, nat := range d.activeNATs { d.flushPort(nat) } diff --git a/stack.go b/stack.go index b2d9568..613e45d 100644 --- a/stack.go +++ b/stack.go @@ -14,6 +14,7 @@ import ( type Stack interface { Start() error + ResetNetwork() Close() error } diff --git a/stack_gvisor.go b/stack_gvisor.go index 8a02601..c226d05 100644 --- a/stack_gvisor.go +++ b/stack_gvisor.go @@ -134,6 +134,16 @@ func (t *GVisor) Start() error { return nil } +func (t *GVisor) ResetNetwork() { + if t.udpForwarder != nil { + t.udpForwarder.udpNat.Purge() + } + if t.icmpForwarder != nil { + t.icmpForwarder.Purge() + } + t.dispatcher.ResetNetwork() +} + func (t *GVisor) Close() error { t.dispatcher.Close() if t.icmpForwarder != nil { diff --git a/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index 11e82af..55cbbd5 100644 --- a/stack_gvisor_icmp.go +++ b/stack_gvisor_icmp.go @@ -72,6 +72,15 @@ func NewICMPForwarder(stack *stack.Stack, handler Handler, logger logger.Logger) return forwarder } +func (f *ICMPForwarder) Purge() { + f.flowAccess.Lock() + for key, flow := range f.flows { + flow.close(FlowCloseReset) + delete(f.flows, key) + } + f.flowAccess.Unlock() +} + func (f *ICMPForwarder) Close() error { f.returnPath.closed.Store(true) f.flowAccess.Lock() diff --git a/stack_mixed.go b/stack_mixed.go index a238622..69c8b27 100644 --- a/stack_mixed.go +++ b/stack_mixed.go @@ -62,6 +62,13 @@ func (m *Mixed) Start() error { return nil } +func (m *Mixed) ResetNetwork() { + m.System.ResetNetwork() + if m.udpForwarder != nil { + m.udpForwarder.udpNat.Purge() + } +} + func (m *Mixed) Close() error { if m.stack == nil { return nil diff --git a/stack_system.go b/stack_system.go index 148515a..dd561da 100644 --- a/stack_system.go +++ b/stack_system.go @@ -113,6 +113,16 @@ func NewSystem(options StackOptions) (Stack, error) { return stack, nil } +func (s *System) ResetNetwork() { + if s.tcpNat != nil { + s.tcpNat.Purge() + } + if s.udpNat != nil { + s.udpNat.Purge() + } + s.dispatcher.ResetNetwork() +} + func (s *System) Close() error { s.dispatcher.Close() if s.udpNat != nil { diff --git a/stack_system_nat.go b/stack_system_nat.go index 2fec29c..1dd5377 100644 --- a/stack_system_nat.go +++ b/stack_system_nat.go @@ -86,6 +86,15 @@ func (n *TCPNat) checkTimeout() { n.addrAccess.Unlock() } +func (n *TCPNat) Purge() { + n.addrAccess.Lock() + n.portAccess.Lock() + clear(n.addrMap) + clear(n.portMap) + n.portAccess.Unlock() + n.addrAccess.Unlock() +} + func (n *TCPNat) LookupBack(port uint16) *TCPSession { n.portAccess.RLock() session := n.portMap[port]