diff --git a/pkg/sentry/inet/inet.go b/pkg/sentry/inet/inet.go index d3067f1bf..9582ddca2 100644 --- a/pkg/sentry/inet/inet.go +++ b/pkg/sentry/inet/inet.go @@ -84,6 +84,9 @@ type Stack interface { // Pause pauses the network stack before save. Pause() + // Resume resumes the network stack after save. + Resume() + // Restore restarts the network stack after restore. Restore() diff --git a/pkg/sentry/inet/test_stack.go b/pkg/sentry/inet/test_stack.go index b9d8ced32..d8683645f 100644 --- a/pkg/sentry/inet/test_stack.go +++ b/pkg/sentry/inet/test_stack.go @@ -155,6 +155,9 @@ func (s *TestStack) Pause() {} // Restore implements Stack. func (s *TestStack) Restore() {} +// Resume implements Stack. +func (s *TestStack) Resume() {} + // RegisteredEndpoints implements Stack. func (s *TestStack) RegisteredEndpoints() []stack.TransportEndpoint { return nil diff --git a/pkg/sentry/kernel/kernel.go b/pkg/sentry/kernel/kernel.go index d0109772c..86f5c2a4e 100644 --- a/pkg/sentry/kernel/kernel.go +++ b/pkg/sentry/kernel/kernel.go @@ -623,7 +623,7 @@ func (k *Kernel) SaveTo(ctx context.Context, w wire.Writer) error { netstackPauseStart := time.Now() log.Infof("Pausing root network namespace") k.rootNetworkNamespace.Stack().Pause() - defer k.rootNetworkNamespace.Stack().Restore() + defer k.rootNetworkNamespace.Stack().Resume() log.Infof("Pausing root network namespace took [%s].", time.Since(netstackPauseStart)) } diff --git a/pkg/sentry/socket/hostinet/stack.go b/pkg/sentry/socket/hostinet/stack.go index 6cdd8d198..f98150eca 100644 --- a/pkg/sentry/socket/hostinet/stack.go +++ b/pkg/sentry/socket/hostinet/stack.go @@ -329,6 +329,9 @@ func (*Stack) Pause() {} // Restore implements inet.Stack.Restore. func (*Stack) Restore() {} +// Resume implements inet.Stack.Resume. +func (*Stack) Resume() {} + // RegisteredEndpoints implements inet.Stack.RegisteredEndpoints. func (*Stack) RegisteredEndpoints() []stack.TransportEndpoint { return nil } diff --git a/pkg/sentry/socket/netstack/stack.go b/pkg/sentry/socket/netstack/stack.go index 76d3aeaad..a82b7b2a7 100644 --- a/pkg/sentry/socket/netstack/stack.go +++ b/pkg/sentry/socket/netstack/stack.go @@ -479,6 +479,11 @@ func (s *Stack) Restore() { s.Stack.Restore() } +// Resume implements inet.Stack.Resume. +func (s *Stack) Resume() { + s.Stack.Resume() +} + // RegisteredEndpoints implements inet.Stack.RegisteredEndpoints. func (s *Stack) RegisteredEndpoints() []stack.TransportEndpoint { return s.Stack.RegisteredEndpoints() diff --git a/pkg/tcpip/stack/stack.go b/pkg/tcpip/stack/stack.go index 269373927..b305c7adb 100644 --- a/pkg/tcpip/stack/stack.go +++ b/pkg/tcpip/stack/stack.go @@ -57,6 +57,12 @@ type RestoredEndpoint interface { Restore(*Stack) } +// ResumableEndpoint is an endpoint that needs to be resumed after save. +type ResumableEndpoint interface { + // Resume resumes an endpoint. + Resume() +} + // uniqueIDGenerator is a default unique ID generator. type uniqueIDGenerator atomicbitops.Uint64 @@ -119,6 +125,10 @@ type Stack struct { // stack is being restored. restoredEndpoints []RestoredEndpoint + // resumableEndpoints is a list of endpoints that need to be resumed + // after save. + resumableEndpoints []ResumableEndpoint + // icmpRateLimiter is a global rate limiter for all ICMP messages generated // by the stack. icmpRateLimiter *ICMPRateLimiter @@ -1715,14 +1725,24 @@ func (s *Stack) UnregisterRawTransportEndpoint(netProto tcpip.NetworkProtocolNum // this stack. func (s *Stack) RegisterRestoredEndpoint(e RestoredEndpoint) { s.mu.Lock() + defer s.mu.Unlock() + s.restoredEndpoints = append(s.restoredEndpoints, e) - s.mu.Unlock() +} + +// RegisterResumableEndpoint records e as an endpoint that has to be resumed. +func (s *Stack) RegisterResumableEndpoint(e ResumableEndpoint) { + s.mu.Lock() + defer s.mu.Unlock() + + s.resumableEndpoints = append(s.resumableEndpoints, e) } // RegisteredEndpoints returns all endpoints which are currently registered. func (s *Stack) RegisteredEndpoints() []TransportEndpoint { s.mu.Lock() defer s.mu.Unlock() + var es []TransportEndpoint for _, e := range s.demux.protocol { es = append(es, e.transportEndpoints()...) @@ -1733,11 +1753,12 @@ func (s *Stack) RegisteredEndpoints() []TransportEndpoint { // CleanupEndpoints returns endpoints currently in the cleanup state. func (s *Stack) CleanupEndpoints() []TransportEndpoint { s.cleanupEndpointsMu.Lock() + defer s.cleanupEndpointsMu.Unlock() + es := make([]TransportEndpoint, 0, len(s.cleanupEndpoints)) for e := range s.cleanupEndpoints { es = append(es, e) } - s.cleanupEndpointsMu.Unlock() return es } @@ -1745,10 +1766,11 @@ func (s *Stack) CleanupEndpoints() []TransportEndpoint { // for restoring a stack after a save. func (s *Stack) RestoreCleanupEndpoints(es []TransportEndpoint) { s.cleanupEndpointsMu.Lock() + defer s.cleanupEndpointsMu.Unlock() + for _, e := range es { s.cleanupEndpoints[e] = struct{}{} } - s.cleanupEndpointsMu.Unlock() } // Close closes all currently registered transport endpoints. @@ -1829,6 +1851,21 @@ func (s *Stack) Restore() { } } +// Resume resumes the stack after a save. +func (s *Stack) Resume() { + s.mu.Lock() + eps := s.resumableEndpoints + s.resumableEndpoints = nil + s.mu.Unlock() + for _, e := range eps { + e.Resume() + } + // Now resume any protocol level background workers. + for _, p := range s.transportProtocols { + p.proto.Resume() + } +} + // RegisterPacketEndpoint registers ep with the stack, causing it to receive // all traffic of the specified netProto on the given NIC. If nicID is 0, it // receives traffic from every NIC. diff --git a/pkg/tcpip/transport/icmp/endpoint_state.go b/pkg/tcpip/transport/icmp/endpoint_state.go index d7ce14230..134797e8b 100644 --- a/pkg/tcpip/transport/icmp/endpoint_state.go +++ b/pkg/tcpip/transport/icmp/endpoint_state.go @@ -42,6 +42,7 @@ func (e *endpoint) afterLoad(ctx context.Context) { // beforeSave is invoked by stateify. func (e *endpoint) beforeSave() { e.freeze() + e.stack.RegisterResumableEndpoint(e) } // Restore implements tcpip.RestoredEndpoint.Restore. @@ -68,3 +69,8 @@ func (e *endpoint) Restore(s *stack.Stack) { panic(fmt.Sprintf("unhandled state = %s", state)) } } + +// Resume implements tcpip.ResumableEndpoint.Resume. +func (e *endpoint) Resume() { + e.thaw() +} diff --git a/pkg/tcpip/transport/packet/endpoint_state.go b/pkg/tcpip/transport/packet/endpoint_state.go index 78e4a27d3..16be7d6b3 100644 --- a/pkg/tcpip/transport/packet/endpoint_state.go +++ b/pkg/tcpip/transport/packet/endpoint_state.go @@ -38,6 +38,7 @@ func (ep *endpoint) beforeSave() { ep.rcvMu.Lock() defer ep.rcvMu.Unlock() ep.rcvDisabled = true + ep.stack.RegisterResumableEndpoint(ep) } // afterLoad is invoked by stateify. @@ -56,3 +57,10 @@ func (ep *endpoint) afterLoad(ctx context.Context) { ep.rcvDisabled = false ep.rcvMu.Unlock() } + +// Resume implements tcpip.ResumableEndpoint.Resume. +func (ep *endpoint) Resume() { + ep.rcvMu.Lock() + defer ep.rcvMu.Unlock() + ep.rcvDisabled = false +} diff --git a/pkg/tcpip/transport/raw/endpoint_state.go b/pkg/tcpip/transport/raw/endpoint_state.go index 427a9d8e8..d915ade2e 100644 --- a/pkg/tcpip/transport/raw/endpoint_state.go +++ b/pkg/tcpip/transport/raw/endpoint_state.go @@ -41,6 +41,7 @@ func (e *endpoint) afterLoad(ctx context.Context) { // beforeSave is invoked by stateify. func (e *endpoint) beforeSave() { e.setReceiveDisabled(true) + e.stack.RegisterResumableEndpoint(e) } // Restore implements tcpip.RestoredEndpoint.Restore. @@ -58,3 +59,8 @@ func (e *endpoint) Restore(s *stack.Stack) { } } } + +// Resume implements tcpip.ResumableEndpoint.Resume. +func (e *endpoint) Resume() { + e.setReceiveDisabled(false) +} diff --git a/pkg/tcpip/transport/tcp/endpoint_state.go b/pkg/tcpip/transport/tcp/endpoint_state.go index 7a6bc8d7f..f23a22b03 100644 --- a/pkg/tcpip/transport/tcp/endpoint_state.go +++ b/pkg/tcpip/transport/tcp/endpoint_state.go @@ -57,6 +57,8 @@ func (e *endpoint) beforeSave() { default: panic(fmt.Sprintf("endpoint in unknown state %v", e.EndpointState())) } + + e.stack.RegisterResumableEndpoint(e) } // saveEndpoints is invoked by stateify. @@ -270,3 +272,8 @@ func (e *endpoint) Restore(s *stack.Stack) { tcpip.DeleteDanglingEndpoint(e) } } + +// Resume implements tcpip.ResumableEndpoint.Resume. +func (e *endpoint) Resume() { + e.segmentQueue.thaw() +} diff --git a/pkg/tcpip/transport/udp/endpoint_state.go b/pkg/tcpip/transport/udp/endpoint_state.go index 46b5afc86..488e46600 100644 --- a/pkg/tcpip/transport/udp/endpoint_state.go +++ b/pkg/tcpip/transport/udp/endpoint_state.go @@ -42,6 +42,7 @@ func (e *endpoint) afterLoad(ctx context.Context) { // beforeSave is invoked by stateify. func (e *endpoint) beforeSave() { e.freeze() + e.stack.RegisterResumableEndpoint(e) } // Restore implements tcpip.RestoredEndpoint.Restore. @@ -76,3 +77,8 @@ func (e *endpoint) Restore(s *stack.Stack) { panic(fmt.Sprintf("unhandled state = %s", state)) } } + +// Resume implements tcpip.ResumableEndpoint.Resume. +func (e *endpoint) Resume() { + e.thaw() +}