From abe7cee09692bb98ebb067539e0546b6926db3e6 Mon Sep 17 00:00:00 2001 From: Andrei Vagin Date: Tue, 1 Aug 2023 13:57:53 -0700 Subject: [PATCH] kernel: don't use atomic pointers for task.netns task.netns is always changed from a task goroutine under task.mu. It means that we can access it without any locks from a task goroutine we don't need to increment a reference counter in such cases. In all other cases, we need to take task.mu. PiperOrigin-RevId: 552913323 --- pkg/sentry/inet/BUILD | 12 ------------ pkg/sentry/kernel/task.go | 4 ++-- pkg/sentry/kernel/task_clone.go | 9 +++++---- pkg/sentry/kernel/task_exit.go | 3 ++- pkg/sentry/kernel/task_net.go | 8 ++++---- pkg/sentry/kernel/task_start.go | 2 +- 6 files changed, 14 insertions(+), 24 deletions(-) diff --git a/pkg/sentry/inet/BUILD b/pkg/sentry/inet/BUILD index ee5dfe6ca..11f9ece67 100644 --- a/pkg/sentry/inet/BUILD +++ b/pkg/sentry/inet/BUILD @@ -18,21 +18,9 @@ go_template_instance( }, ) -go_template_instance( - name = "atomicptr_netns", - out = "atomicptr_netns_unsafe.go", - package = "inet", - prefix = "Namespace", - template = "//pkg/sync/atomicptr:generic_atomicptr", - types = { - "Value": "Namespace", - }, -) - go_library( name = "inet", srcs = [ - "atomicptr_netns_unsafe.go", "context.go", "inet.go", "namespace.go", diff --git a/pkg/sentry/kernel/task.go b/pkg/sentry/kernel/task.go index 7cb91cbda..9fb943fd1 100644 --- a/pkg/sentry/kernel/task.go +++ b/pkg/sentry/kernel/task.go @@ -508,8 +508,8 @@ type Task struct { // netns is the task's network namespace. It has to be changed under mu // so that GetNetworkNamespace can take a reference before it is - // released. - netns inet.NamespaceAtomicPtr + // released. It is changed only from the task goroutine. + netns *inet.Namespace // If rseqPreempted is true, before the next call to p.Switch(), // interrupt rseq critical regions as defined by rseqAddr and diff --git a/pkg/sentry/kernel/task_clone.go b/pkg/sentry/kernel/task_clone.go index b89e8cd0e..a8324d004 100644 --- a/pkg/sentry/kernel/task_clone.go +++ b/pkg/sentry/kernel/task_clone.go @@ -126,7 +126,7 @@ func (t *Task) Clone(args *linux.CloneArgs) (ThreadID, *SyscallControl, error) { ipcns.DecRef(t) }) - netns := t.netns.Load() + netns := t.netns if args.Flags&linux.CLONE_NEWNET != 0 { netns = inet.NewNamespace(netns, userns) inode := nsfs.NewInode(t, t.k.nsfsMount, netns) @@ -444,7 +444,7 @@ func (t *Task) Setns(fd *vfs.FileDescription, flags int32) error { oldNS := t.NetworkNamespace() ns.IncRef() t.mu.Lock() - t.netns.Store(ns) + t.netns = ns t.mu.Unlock() oldNS.DecRef(t) return nil @@ -570,9 +570,10 @@ func (t *Task) Unshare(flags int32) error { netnsInode := nsfs.NewInode(t, t.k.nsfsMount, netns) netns.SetInode(netnsInode) t.mu.Lock() - netns = t.netns.Swap(netns) + oldNetns := t.netns + t.netns = netns t.mu.Unlock() - netns.DecRef(t) + oldNetns.DecRef(t) } cu := cleanup.Cleanup{} diff --git a/pkg/sentry/kernel/task_exit.go b/pkg/sentry/kernel/task_exit.go index 21f02b82e..99dd4362f 100644 --- a/pkg/sentry/kernel/task_exit.go +++ b/pkg/sentry/kernel/task_exit.go @@ -288,7 +288,8 @@ func (*runExitMain) execute(t *Task) taskRunState { mntns := t.mountNamespace t.mountNamespace = nil ipcns := t.ipcns - netns := t.netns.Swap(nil) + netns := t.netns + t.netns = nil t.mu.Unlock() if mntns != nil { mntns.DecRef(t) diff --git a/pkg/sentry/kernel/task_net.go b/pkg/sentry/kernel/task_net.go index c6698fb9f..e448bb18b 100644 --- a/pkg/sentry/kernel/task_net.go +++ b/pkg/sentry/kernel/task_net.go @@ -20,7 +20,7 @@ import ( // IsNetworkNamespaced returns true if t is in a non-root network namespace. func (t *Task) IsNetworkNamespaced() bool { - return !t.netns.Load().IsRoot() + return !t.netns.IsRoot() } // NetworkContext returns the network stack used by the task. NetworkContext @@ -29,12 +29,12 @@ func (t *Task) IsNetworkNamespaced() bool { // TODO(gvisor.dev/issue/1833): Migrate callers of this method to // NetworkNamespace(). func (t *Task) NetworkContext() inet.Stack { - return t.netns.Load().Stack() + return t.netns.Stack() } // NetworkNamespace returns the network namespace observed by the task. func (t *Task) NetworkNamespace() *inet.Namespace { - return t.netns.Load() + return t.netns } // GetNetworkNamespace takes a reference on the task network namespace and @@ -43,7 +43,7 @@ func (t *Task) GetNetworkNamespace() *inet.Namespace { // t.mu is required to be sure that the network namespace will not be // released. t.mu.Lock() - netns := t.netns.Load() + netns := t.netns if netns != nil { netns.IncRef() } diff --git a/pkg/sentry/kernel/task_start.go b/pkg/sentry/kernel/task_start.go index d9dbde37e..40ca0e316 100644 --- a/pkg/sentry/kernel/task_start.go +++ b/pkg/sentry/kernel/task_start.go @@ -171,7 +171,7 @@ func (ts *TaskSet) newTask(ctx context.Context, cfg *TaskConfig) (*Task, error) cgroups: make(map[Cgroup]struct{}), userCounters: cfg.UserCounters, } - t.netns.Store(cfg.NetworkNamespace) + t.netns = cfg.NetworkNamespace t.creds.Store(cfg.Credentials) t.endStopCond.L = &t.tg.signalHandlers.mu t.ptraceTracer.Store((*Task)(nil))