diff --git a/pkg/sentry/kernel/abstract_socket_namespace.go b/pkg/sentry/kernel/abstract_socket_namespace.go index 57c98e714..bc360d95d 100644 --- a/pkg/sentry/kernel/abstract_socket_namespace.go +++ b/pkg/sentry/kernel/abstract_socket_namespace.go @@ -16,12 +16,13 @@ package kernel import ( "fmt" + "math/rand" - "golang.org/x/sys/unix" "gvisor.dev/gvisor/pkg/context" "gvisor.dev/gvisor/pkg/refs" "gvisor.dev/gvisor/pkg/sentry/socket/unix/transport" "gvisor.dev/gvisor/pkg/sync" + "gvisor.dev/gvisor/pkg/syserr" ) // +stateify savable @@ -89,22 +90,44 @@ func (a *AbstractSocketNamespace) BoundEndpoint(name string) transport.BoundEndp // // When the last reference managed by socket is dropped, ep may be removed from the // namespace. -func (a *AbstractSocketNamespace) Bind(ctx context.Context, name string, ep transport.BoundEndpoint, socket refs.TryRefCounter) error { +func (a *AbstractSocketNamespace) Bind(ctx context.Context, path string, ep transport.BoundEndpoint, socket refs.TryRefCounter) (string, *syserr.Error) { a.mu.Lock() defer a.mu.Unlock() - // Check if there is already a socket (which has not yet been destroyed) bound at name. - if ep, ok := a.endpoints[name]; ok { - if ep.socket.TryIncRef() { - ep.socket.DecRef(ctx) - return unix.EADDRINUSE + name := "" + if path == "" { + // Autobind feature. + mask := uint32(0xFFFFF) + r := rand.Uint32() + for i := uint32(0); i <= mask; i++ { + p := fmt.Sprintf("X%05x", (r+i)&mask) + if _, ok := a.endpoints[p[1:]]; ok { + continue + } + b := ([]byte)(p) + b[0] = 0 + path = string(b) + break + } + if path == "" { + return "", syserr.ErrNoSpace + } + name = path[1:] + } else { + name = path[1:] + // Check if there is already a socket (which has not yet been destroyed) bound at name. + if ep, ok := a.endpoints[name]; ok { + if ep.socket.TryIncRef() { + ep.socket.DecRef(ctx) + return name, syserr.ErrPortInUse + } } } ae := abstractEndpoint{ep: ep, name: name, ns: a} ae.socket = socket a.endpoints[name] = ae - return nil + return path, nil } // Remove removes the specified socket at name from the abstract socket diff --git a/pkg/sentry/socket/unix/unix.go b/pkg/sentry/socket/unix/unix.go index db1d13ecf..45896d123 100644 --- a/pkg/sentry/socket/unix/unix.go +++ b/pkg/sentry/socket/unix/unix.go @@ -211,17 +211,18 @@ func (s *Socket) Bind(t *kernel.Task, sockaddr []byte) *syserr.Error { return syserr.ErrInvalidArgument } - if p[0] == 0 { + // If path is empty, the socket is autobound to an abstract address. + if len(p) == 0 || 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 + p, err := asn.Bind(t, p, bep, s) + if err != nil { + return err } + name := p[1:] if err := s.ep.Bind(transport.Address{Addr: p}); err != nil { asn.Remove(name, s) return err @@ -449,11 +450,7 @@ func extractPath(sockaddr []byte) (string, *syserr.Error) { // The address is trimmed by GetAddress. p := addr.Addr - if p == "" { - // Not allowed. - return "", syserr.ErrInvalidArgument - } - if p[len(p)-1] == '/' { + if len(p) > 0 && p[len(p)-1] == '/' { // Weird, they tried to bind '/a/b/c/'? return "", syserr.ErrIsDir } @@ -526,6 +523,10 @@ func extractEndpoint(t *kernel.Task, sockaddr []byte) (transport.BoundEndpoint, if err != nil { return nil, err } + if path == "" { + // Not allowed. + return nil, syserr.ErrInvalidArgument + } // Is it abstract? if path[0] == 0 { diff --git a/test/syscalls/linux/socket_unix_unbound_abstract.cc b/test/syscalls/linux/socket_unix_unbound_abstract.cc index 3a91e97ac..04d263d0c 100644 --- a/test/syscalls/linux/socket_unix_unbound_abstract.cc +++ b/test/syscalls/linux/socket_unix_unbound_abstract.cc @@ -12,6 +12,8 @@ // See the License for the specific language governing permissions and // limitations under the License. +#include +#include #include #include @@ -72,6 +74,50 @@ TEST_P(UnboundAbstractUnixSocketPairTest, BindNothing) { SyscallSucceeds()); } +TEST_P(UnboundAbstractUnixSocketPairTest, AutoBindSuccess) { + auto sockets = ASSERT_NO_ERRNO_AND_VALUE(NewSocketPair()); + struct sockaddr_un addr = {.sun_family = AF_UNIX}; + ASSERT_THAT( + bind(sockets->first_fd(), reinterpret_cast(&addr), + sizeof(sa_family_t)), + SyscallSucceeds()); + socklen_t addr_len = sizeof(addr); + ASSERT_THAT(getsockname(sockets->first_fd(), + reinterpret_cast(&addr), &addr_len), + SyscallSucceeds()); + // The address consists of a null byte followed by 5 bytes in the character + // set [0-9a-f]. + EXPECT_EQ(offsetof(struct sockaddr_un, sun_path) + 6, addr_len); + EXPECT_EQ(addr.sun_path[0], 0); + for (int i = 1; i < 6; i++) { + char c = addr.sun_path[i]; + EXPECT_TRUE((c >= '0' && c <= '9') || (c >= 'a' && c <= 'f')); + } + if ((GetParam().type & SOCK_DGRAM) == 0) { + ASSERT_THAT(listen(sockets->first_fd(), 0 /* backlog */), + SyscallSucceeds()); + } + ASSERT_THAT(connect(sockets->second_fd(), + reinterpret_cast(&addr), addr_len), + SyscallSucceeds()); +} + +TEST_P(UnboundAbstractUnixSocketPairTest, AutoBindAddrInUse) { + auto sockets = ASSERT_NO_ERRNO_AND_VALUE(NewSocketPair()); + struct sockaddr_un addr = {.sun_family = AF_UNIX}; + ASSERT_THAT( + bind(sockets->first_fd(), reinterpret_cast(&addr), + sizeof(sa_family_t)), + SyscallSucceeds()); + socklen_t addr_len = sizeof(addr); + ASSERT_THAT(getsockname(sockets->first_fd(), + reinterpret_cast(&addr), &addr_len), + SyscallSucceeds()); + ASSERT_THAT(bind(sockets->second_fd(), + reinterpret_cast(&addr), addr_len), + SyscallFailsWithErrno(EADDRINUSE)); +} + TEST_P(UnboundAbstractUnixSocketPairTest, ListenZeroBacklog) { SKIP_IF((GetParam().type & SOCK_DGRAM) != 0); auto sockets = ASSERT_NO_ERRNO_AND_VALUE(NewSocketPair());