diff --git a/pkg/sentry/socket/netlink/socket.go b/pkg/sentry/socket/netlink/socket.go index 955881cf4..9442dfa62 100644 --- a/pkg/sentry/socket/netlink/socket.go +++ b/pkg/sentry/socket/netlink/socket.go @@ -138,7 +138,7 @@ func NewSocket(t *kernel.Task, skType linux.SockType, protocol Protocol) (*Socke // Bind the endpoint for good measure so we can connect to it. The // bound address will never be exposed. - if err := ep.Bind(tcpip.FullAddress{Addr: "dummy"}, nil); err != nil { + if err := ep.Bind(tcpip.FullAddress{Addr: "dummy"}); err != nil { ep.Close(t) return nil, err } diff --git a/pkg/sentry/socket/netlink/socket_vfs2.go b/pkg/sentry/socket/netlink/socket_vfs2.go index c1fcfbe98..b746337fe 100644 --- a/pkg/sentry/socket/netlink/socket_vfs2.go +++ b/pkg/sentry/socket/netlink/socket_vfs2.go @@ -57,7 +57,7 @@ func NewVFS2(t *kernel.Task, skType linux.SockType, protocol Protocol) (*SocketV // Bind the endpoint for good measure so we can connect to it. The // bound address will never be exposed. - if err := ep.Bind(tcpip.FullAddress{Addr: "dummy"}, nil); err != nil { + if err := ep.Bind(tcpip.FullAddress{Addr: "dummy"}); err != nil { ep.Close(t) return nil, err } diff --git a/pkg/sentry/socket/unix/transport/connectioned.go b/pkg/sentry/socket/unix/transport/connectioned.go index 0d299dbb8..dfeb6a5c4 100644 --- a/pkg/sentry/socket/unix/transport/connectioned.go +++ b/pkg/sentry/socket/unix/transport/connectioned.go @@ -449,7 +449,7 @@ func (e *connectionedEndpoint) Accept(ctx context.Context, peerAddr *tcpip.FullA // // Bind will fail only if the socket is connected, bound or the passed address // is invalid (the empty string). -func (e *connectionedEndpoint) Bind(addr tcpip.FullAddress, commit func() *syserr.Error) *syserr.Error { +func (e *connectionedEndpoint) Bind(addr tcpip.FullAddress) *syserr.Error { e.Lock() defer e.Unlock() if e.isBound() || e.ListeningLocked() { @@ -459,11 +459,6 @@ func (e *connectionedEndpoint) Bind(addr tcpip.FullAddress, commit func() *syser // The empty string is not permitted. return syserr.ErrBadLocalAddress } - if commit != nil { - if err := commit(); err != nil { - return err - } - } // Save the bound address. e.path = string(addr.Addr) diff --git a/pkg/sentry/socket/unix/transport/connectionless.go b/pkg/sentry/socket/unix/transport/connectionless.go index 09c9b334b..9d5d63fbe 100644 --- a/pkg/sentry/socket/unix/transport/connectionless.go +++ b/pkg/sentry/socket/unix/transport/connectionless.go @@ -164,7 +164,7 @@ func (*connectionlessEndpoint) Accept(context.Context, *tcpip.FullAddress) (Endp // // Bind will fail only if the socket is connected, bound or the passed address // is invalid (the empty string). -func (e *connectionlessEndpoint) Bind(addr tcpip.FullAddress, commit func() *syserr.Error) *syserr.Error { +func (e *connectionlessEndpoint) Bind(addr tcpip.FullAddress) *syserr.Error { e.Lock() defer e.Unlock() if e.isBound() { @@ -174,11 +174,6 @@ func (e *connectionlessEndpoint) Bind(addr tcpip.FullAddress, commit func() *sys // The empty string is not permitted. return syserr.ErrBadLocalAddress } - if commit != nil { - if err := commit(); err != nil { - return err - } - } // Save the bound address. e.path = string(addr.Addr) diff --git a/pkg/sentry/socket/unix/transport/unix.go b/pkg/sentry/socket/unix/transport/unix.go index 361514a78..4829c4f95 100644 --- a/pkg/sentry/socket/unix/transport/unix.go +++ b/pkg/sentry/socket/unix/transport/unix.go @@ -166,11 +166,7 @@ type Endpoint interface { // Bind binds the endpoint to a specific local address and port. // Specifying a NIC is optional. - // - // An optional commit function will be executed atomically with respect - // to binding the endpoint. If this returns an error, the bind will not - // occur and the error will be propagated back to the caller. - Bind(address tcpip.FullAddress, commit func() *syserr.Error) *syserr.Error + Bind(address tcpip.FullAddress) *syserr.Error // Type return the socket type, typically either SockStream, SockDgram // or SockSeqpacket. diff --git a/pkg/sentry/socket/unix/unix.go b/pkg/sentry/socket/unix/unix.go index 038526fb4..ff726e62d 100644 --- a/pkg/sentry/socket/unix/unix.go +++ b/pkg/sentry/socket/unix/unix.go @@ -26,6 +26,7 @@ import ( "gvisor.dev/gvisor/pkg/errors/linuxerr" "gvisor.dev/gvisor/pkg/fspath" "gvisor.dev/gvisor/pkg/hostarch" + "gvisor.dev/gvisor/pkg/log" "gvisor.dev/gvisor/pkg/marshal" "gvisor.dev/gvisor/pkg/sentry/arch" "gvisor.dev/gvisor/pkg/sentry/fs" @@ -284,69 +285,80 @@ func (s *SocketOperations) Bind(t *kernel.Task, sockaddr []byte) *syserr.Error { return syserr.ErrInvalidArgument } - return s.ep.Bind(tcpip.FullAddress{Addr: tcpip.Address(p)}, func() *syserr.Error { - // Is it abstract? - if p[0] == 0 { - if t.IsNetworkNamespaced() { - return syserr.ErrInvalidEndpointState - } - asn := t.AbstractSockets() - name := p[1:] - if err := asn.Bind(t, name, bep, s); err != nil { - // syserr.ErrPortInUse corresponds to EADDRINUSE. - return syserr.ErrPortInUse - } - s.abstractName = name - s.abstractNamespace = asn - } else { - // The parent and name. - var d *fs.Dirent - var name string - - cwd := t.FSContext().WorkingDirectory() - defer cwd.DecRef(t) - - // Is there no slash at all? - if !strings.Contains(p, "/") { - d = cwd - name = p - } else { - root := t.FSContext().RootDirectory() - defer root.DecRef(t) - // Find the last path component, we know that something follows - // that final slash, otherwise extractPath() would have failed. - lastSlash := strings.LastIndex(p, "/") - subPath := p[:lastSlash] - if subPath == "" { - // Fix up subpath in case file is in root. - subPath = "/" - } - var err error - remainingTraversals := uint(fs.DefaultTraversalLimit) - d, err = t.MountNamespace().FindInode(t, root, cwd, subPath, &remainingTraversals) - if err != nil { - // No path available. - return syserr.ErrNoSuchFile - } - defer d.DecRef(t) - name = p[lastSlash+1:] - } - - // Create the socket. - // - // Note that the file permissions here are not set correctly (see - // gvisor.dev/issue/2324). There is no convenient way to get permissions - // on the socket referred to by s, so we will leave this discrepancy - // unresolved until VFS2 replaces this code. - childDir, err := d.Bind(t, t.FSContext().RootDirectory(), name, bep, fs.FilePermissions{User: fs.PermMask{Read: true}}) - if err != nil { - return syserr.ErrPortInUse - } - childDir.DecRef(t) + if p[0] == 0 { + // Abstract socket. See net/unix/af_unix.c:unix_bind_abstract(). + if t.IsNetworkNamespaced() { + return syserr.ErrInvalidEndpointState } - + asn := t.AbstractSockets() + name := p[1:] + if err := asn.Bind(t, name, bep, s); err != nil { + // syserr.ErrPortInUse corresponds to EADDRINUSE. + return syserr.ErrPortInUse + } + if err := s.ep.Bind(tcpip.FullAddress{Addr: tcpip.Address(p)}); err != nil { + asn.Remove(name, s) + return err + } + // The socket has been successfully bound. We can update the following. + s.abstractName = name + s.abstractNamespace = asn return nil - }) + } + + // See net/unix/af_unix.c:unix_bind_bsd(). + + // The parent and name. + var d *fs.Dirent + var name string + + cwd := t.FSContext().WorkingDirectory() + defer cwd.DecRef(t) + + // Is there no slash at all? + if !strings.Contains(p, "/") { + d = cwd + name = p + } else { + root := t.FSContext().RootDirectory() + defer root.DecRef(t) + // Find the last path component, we know that something follows + // that final slash, otherwise extractPath() would have failed. + lastSlash := strings.LastIndex(p, "/") + subPath := p[:lastSlash] + if subPath == "" { + // Fix up subpath in case file is in root. + subPath = "/" + } + var err error + remainingTraversals := uint(fs.DefaultTraversalLimit) + d, err = t.MountNamespace().FindInode(t, root, cwd, subPath, &remainingTraversals) + if err != nil { + // No path available. + return syserr.ErrNoSuchFile + } + defer d.DecRef(t) + name = p[lastSlash+1:] + } + + // Create the socket. + // + // Note that the file permissions here are not set correctly (see + // gvisor.dev/issue/2324). There is no convenient way to get permissions + // on the socket referred to by s, so we will leave this discrepancy + // unresolved until VFS2 replaces this code. + childDir, err := d.Bind(t, t.FSContext().RootDirectory(), name, bep, fs.FilePermissions{User: fs.PermMask{Read: true}}) + if err != nil { + return syserr.ErrPortInUse + } + childDir.DecRef(t) + if err := s.ep.Bind(tcpip.FullAddress{Addr: tcpip.Address(p)}); err != nil { + if removeErr := d.Remove(t, t.FSContext().RootDirectory(), name, false /* dirPath */); removeErr != nil { + log.Warningf("failed to remove socket file created for bind(%q): %v", p, removeErr) + } + return err + } + return nil } // extractEndpoint retrieves the transport.BoundEndpoint associated with a Unix diff --git a/pkg/sentry/socket/unix/unix_vfs2.go b/pkg/sentry/socket/unix/unix_vfs2.go index 0b3ad8de3..66fbad03b 100644 --- a/pkg/sentry/socket/unix/unix_vfs2.go +++ b/pkg/sentry/socket/unix/unix_vfs2.go @@ -20,6 +20,7 @@ import ( "gvisor.dev/gvisor/pkg/errors/linuxerr" "gvisor.dev/gvisor/pkg/fspath" "gvisor.dev/gvisor/pkg/hostarch" + "gvisor.dev/gvisor/pkg/log" "gvisor.dev/gvisor/pkg/marshal" "gvisor.dev/gvisor/pkg/sentry/arch" "gvisor.dev/gvisor/pkg/sentry/fsimpl/sockfs" @@ -198,52 +199,63 @@ func (s *SocketVFS2) Bind(t *kernel.Task, sockaddr []byte) *syserr.Error { return syserr.ErrInvalidArgument } - return s.ep.Bind(tcpip.FullAddress{Addr: tcpip.Address(p)}, func() *syserr.Error { - // Is it abstract? - if p[0] == 0 { - if t.IsNetworkNamespaced() { - return syserr.ErrInvalidEndpointState - } - asn := t.AbstractSockets() - name := p[1:] - if err := asn.Bind(t, name, bep, s); err != nil { - // syserr.ErrPortInUse corresponds to EADDRINUSE. - return syserr.ErrPortInUse - } - s.abstractName = name - s.abstractNamespace = asn - } else { - path := fspath.Parse(p) - root := t.FSContext().RootDirectoryVFS2() - defer root.DecRef(t) - start := root - relPath := !path.Absolute - if relPath { - start = t.FSContext().WorkingDirectoryVFS2() - defer start.DecRef(t) - } - pop := vfs.PathOperation{ - Root: root, - Start: start, - Path: path, - } - stat, err := s.vfsfd.Stat(t, vfs.StatOptions{Mask: linux.STATX_MODE}) - if err != nil { - return syserr.FromError(err) - } - err = t.Kernel().VFS().MknodAt(t, t.Credentials(), &pop, &vfs.MknodOptions{ - // File permissions correspond to net/unix/af_unix.c:unix_bind. - Mode: linux.FileMode(linux.S_IFSOCK | uint(stat.Mode)&^t.FSContext().Umask()), - Endpoint: bep, - }) - if linuxerr.Equals(linuxerr.EEXIST, err) { - return syserr.ErrAddressInUse - } - return syserr.FromError(err) + if p[0] == 0 { + // Abstract socket. See net/unix/af_unix.c:unix_bind_abstract(). + if t.IsNetworkNamespaced() { + return syserr.ErrInvalidEndpointState } - + asn := t.AbstractSockets() + name := p[1:] + if err := asn.Bind(t, name, bep, s); err != nil { + // syserr.ErrPortInUse corresponds to EADDRINUSE. + return syserr.ErrPortInUse + } + if err := s.ep.Bind(tcpip.FullAddress{Addr: tcpip.Address(p)}); err != nil { + asn.Remove(name, s) + return err + } + // The socket has been successfully bound. We can update the following. + s.abstractName = name + s.abstractNamespace = asn return nil + } + + // See net/unix/af_unix.c:unix_bind_bsd(). + path := fspath.Parse(p) + root := t.FSContext().RootDirectoryVFS2() + defer root.DecRef(t) + start := root + relPath := !path.Absolute + if relPath { + start = t.FSContext().WorkingDirectoryVFS2() + defer start.DecRef(t) + } + pop := vfs.PathOperation{ + Root: root, + Start: start, + Path: path, + } + stat, err := s.vfsfd.Stat(t, vfs.StatOptions{Mask: linux.STATX_MODE}) + if err != nil { + return syserr.FromError(err) + } + err = t.Kernel().VFS().MknodAt(t, t.Credentials(), &pop, &vfs.MknodOptions{ + Mode: linux.FileMode(linux.S_IFSOCK | uint(stat.Mode)&^t.FSContext().Umask()), + Endpoint: bep, }) + if linuxerr.Equals(linuxerr.EEXIST, err) { + return syserr.ErrAddressInUse + } + if err != nil { + return syserr.FromError(err) + } + if err := s.ep.Bind(tcpip.FullAddress{Addr: tcpip.Address(p)}); err != nil { + if unlinkErr := t.Kernel().VFS().UnlinkAt(t, t.Credentials(), &pop); unlinkErr != nil { + log.Warningf("failed to unlink socket file created for bind(%q): %v", p, unlinkErr) + } + return err + } + return nil } // Ioctl implements vfs.FileDescriptionImpl.