diff --git a/pkg/sentry/socket/netstack/save_restore.go b/pkg/sentry/socket/netstack/save_restore.go index 49428fd60..9a054f81c 100644 --- a/pkg/sentry/socket/netstack/save_restore.go +++ b/pkg/sentry/socket/netstack/save_restore.go @@ -21,8 +21,8 @@ import ( ) // afterLoad is invoked by stateify. -func (s *Stack) afterLoad(context.Context) { - s.Stack = stack.StackFromEnv // FIXME(b/36201077) +func (s *Stack) afterLoad(ctx context.Context) { + s.Stack = stack.RestoreStackFromContext(ctx) if s.Stack == nil { panic("can't restore without netstack/tcpip/stack.Stack") } diff --git a/pkg/tcpip/stack/BUILD b/pkg/tcpip/stack/BUILD index c54639da6..86ced0f41 100644 --- a/pkg/tcpip/stack/BUILD +++ b/pkg/tcpip/stack/BUILD @@ -253,7 +253,6 @@ go_library( "route_mutex.go", "route_stack_mutex.go", "stack.go", - "stack_global_state.go", "stack_mutex.go", "stack_options.go", "state_conn_mutex.go", diff --git a/pkg/tcpip/stack/stack_global_state.go b/pkg/tcpip/stack/stack_global_state.go deleted file mode 100644 index dfec4258a..000000000 --- a/pkg/tcpip/stack/stack_global_state.go +++ /dev/null @@ -1,19 +0,0 @@ -// Copyright 2018 The gVisor Authors. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -package stack - -// StackFromEnv is the global stack created in restore run. -// FIXME(b/36201077) -var StackFromEnv *Stack diff --git a/pkg/tcpip/stack/tcp.go b/pkg/tcpip/stack/tcp.go index 44b866db5..3393e6b29 100644 --- a/pkg/tcpip/stack/tcp.go +++ b/pkg/tcpip/stack/tcp.go @@ -15,6 +15,7 @@ package stack import ( + "context" "time" "gvisor.dev/gvisor/pkg/atomicbitops" @@ -24,6 +25,19 @@ import ( "gvisor.dev/gvisor/pkg/tcpip/seqnum" ) +// contextID is this package's type for context.Context.Value keys. +type contextID int + +const ( + // CtxRestoreStack is a Context.Value key for the stack to be used in restore. + CtxRestoreStack contextID = iota +) + +// RestoreStackFromContext returns the stack to be used during restore. +func RestoreStackFromContext(ctx context.Context) *Stack { + return ctx.Value(CtxRestoreStack).(*Stack) +} + // TCPProbeFunc is the expected function type for a TCP probe function to be // passed to stack.AddTCPProbe. type TCPProbeFunc func(s *TCPEndpointState) diff --git a/pkg/tcpip/transport/icmp/endpoint_state.go b/pkg/tcpip/transport/icmp/endpoint_state.go index 6bc6b0e69..aefe3e5f4 100644 --- a/pkg/tcpip/transport/icmp/endpoint_state.go +++ b/pkg/tcpip/transport/icmp/endpoint_state.go @@ -35,8 +35,8 @@ func (p *icmpPacket) loadReceivedAt(nsec int64) { } // afterLoad is invoked by stateify. -func (e *endpoint) afterLoad(context.Context) { - stack.StackFromEnv.RegisterRestoredEndpoint(e) +func (e *endpoint) afterLoad(ctx context.Context) { + stack.RestoreStackFromContext(ctx).RegisterRestoredEndpoint(e) } // beforeSave is invoked by stateify. diff --git a/pkg/tcpip/transport/packet/endpoint_state.go b/pkg/tcpip/transport/packet/endpoint_state.go index 7543662d0..228f09b7a 100644 --- a/pkg/tcpip/transport/packet/endpoint_state.go +++ b/pkg/tcpip/transport/packet/endpoint_state.go @@ -41,11 +41,11 @@ func (ep *endpoint) beforeSave() { } // afterLoad is invoked by stateify. -func (ep *endpoint) afterLoad(context.Context) { +func (ep *endpoint) afterLoad(ctx context.Context) { ep.mu.Lock() defer ep.mu.Unlock() - ep.stack = stack.StackFromEnv + ep.stack = stack.RestoreStackFromContext(ctx) ep.ops.InitHandler(ep, ep.stack, tcpip.GetStackSendBufferLimits, tcpip.GetStackReceiveBufferLimits) if err := ep.stack.RegisterPacketEndpoint(ep.boundNIC, ep.boundNetProto, ep); err != nil { diff --git a/pkg/tcpip/transport/raw/endpoint_state.go b/pkg/tcpip/transport/raw/endpoint_state.go index 237acf7cd..c7b3f0af5 100644 --- a/pkg/tcpip/transport/raw/endpoint_state.go +++ b/pkg/tcpip/transport/raw/endpoint_state.go @@ -34,8 +34,8 @@ func (p *rawPacket) loadReceivedAt(nsec int64) { } // afterLoad is invoked by stateify. -func (e *endpoint) afterLoad(context.Context) { - stack.StackFromEnv.RegisterRestoredEndpoint(e) +func (e *endpoint) afterLoad(ctx context.Context) { + stack.RestoreStackFromContext(ctx).RegisterRestoredEndpoint(e) } // beforeSave is invoked by stateify. diff --git a/pkg/tcpip/transport/tcp/endpoint_state.go b/pkg/tcpip/transport/tcp/endpoint_state.go index 281326b10..9fae26fbd 100644 --- a/pkg/tcpip/transport/tcp/endpoint_state.go +++ b/pkg/tcpip/transport/tcp/endpoint_state.go @@ -110,13 +110,13 @@ func (e *endpoint) loadState(epState EndpointState) { } // afterLoad is invoked by stateify. -func (e *endpoint) afterLoad(context.Context) { +func (e *endpoint) afterLoad(ctx context.Context) { // RacyLoad() can be used because we are initializing e. e.origEndpointState = e.state.RacyLoad() // Restore the endpoint to InitialState as it will be moved to // its origEndpointState during Resume. e.state = atomicbitops.FromUint32(uint32(StateInitial)) - stack.StackFromEnv.RegisterRestoredEndpoint(e) + stack.RestoreStackFromContext(ctx).RegisterRestoredEndpoint(e) } // Resume implements tcpip.ResumableEndpoint.Resume. diff --git a/pkg/tcpip/transport/udp/endpoint_state.go b/pkg/tcpip/transport/udp/endpoint_state.go index 9fe4eb2fa..d6893b30e 100644 --- a/pkg/tcpip/transport/udp/endpoint_state.go +++ b/pkg/tcpip/transport/udp/endpoint_state.go @@ -35,8 +35,8 @@ func (p *udpPacket) loadReceivedAt(nsec int64) { } // afterLoad is invoked by stateify. -func (e *endpoint) afterLoad(context.Context) { - stack.StackFromEnv.RegisterRestoredEndpoint(e) +func (e *endpoint) afterLoad(ctx context.Context) { + stack.RestoreStackFromContext(ctx).RegisterRestoredEndpoint(e) } // beforeSave is invoked by stateify. diff --git a/runsc/boot/restore.go b/runsc/boot/restore.go index 91c1354c0..da44402a5 100644 --- a/runsc/boot/restore.go +++ b/runsc/boot/restore.go @@ -18,6 +18,7 @@ import ( "fmt" "os" + "gvisor.dev/gvisor/pkg/context" "gvisor.dev/gvisor/pkg/sentry/inet" "gvisor.dev/gvisor/pkg/sentry/kernel" "gvisor.dev/gvisor/pkg/sentry/socket/hostinet" @@ -36,32 +37,31 @@ type restorer struct { deviceFile *os.File } -func createNetworkNamespaceForRestore(l *Loader) (*inet.Namespace, error) { +func createNetworkNamespaceForRestore(l *Loader) (*stack.Stack, *inet.Namespace, error) { creds := getRootCredentials(l.root.spec, l.root.conf, nil /* UserNamespace */) if creds == nil { - return nil, fmt.Errorf("getting root credentials") + return nil, nil, fmt.Errorf("getting root credentials") } // Save the current network stack to slap on top of the one that was restored. curNetwork := l.k.RootNetworkNamespace().Stack() eps, ok := curNetwork.(*netstack.Stack) if !ok { - return inet.NewRootNamespace(hostinet.NewStack(), nil, creds.UserNamespace), nil + return nil, inet.NewRootNamespace(hostinet.NewStack(), nil, creds.UserNamespace), nil } - stack.StackFromEnv = eps.Stack // FIXME(b/36201077) creator := &sandboxNetstackCreator{ clock: l.k.Timekeeper(), uniqueID: l.k, allowPacketEndpointWrite: l.root.conf.AllowPacketEndpointWrite, } - return inet.NewRootNamespace(curNetwork, creator, creds.UserNamespace), nil + return eps.Stack, inet.NewRootNamespace(curNetwork, creator, creds.UserNamespace), nil } func (r *restorer) restore(l *Loader) error { // Create a new root network namespace with the network stack of the - // old kernel to preserve the exisiting network configuration. - netns, err := createNetworkNamespaceForRestore(l) + // old kernel to preserve the existing network configuration. + oldStack, netns, err := createNetworkNamespaceForRestore(l) if err != nil { return fmt.Errorf("creating network: %w", err) } @@ -107,6 +107,10 @@ func (r *restorer) restore(l *Loader) error { // Set up the restore environment. ctx := l.k.SupervisorContext() + if oldStack != nil { + ctx = context.WithValue(ctx, stack.CtxRestoreStack, oldStack) + } + // TODO(b/298078576): Need to process hints here probably mntr := newContainerMounter(&l.root, l.k, l.mountHints, l.sharedMounts, l.productName, l.sandboxID) ctx, err = mntr.configureRestore(ctx)