diff --git a/pkg/sentry/socket/hostinet/BUILD b/pkg/sentry/socket/hostinet/BUILD index 1276f9571..a2cef7974 100644 --- a/pkg/sentry/socket/hostinet/BUILD +++ b/pkg/sentry/socket/hostinet/BUILD @@ -21,6 +21,7 @@ go_library( visibility = ["//pkg/sentry:internal"], deps = [ "//pkg/abi/linux", + "//pkg/atomicbitops", "//pkg/binary", "//pkg/context", "//pkg/errors/linuxerr", diff --git a/pkg/sentry/socket/hostinet/socket.go b/pkg/sentry/socket/hostinet/socket.go index 80dbf2afa..084c681fe 100644 --- a/pkg/sentry/socket/hostinet/socket.go +++ b/pkg/sentry/socket/hostinet/socket.go @@ -19,6 +19,7 @@ import ( "golang.org/x/sys/unix" "gvisor.dev/gvisor/pkg/abi/linux" + "gvisor.dev/gvisor/pkg/atomicbitops" "gvisor.dev/gvisor/pkg/context" "gvisor.dev/gvisor/pkg/errors/linuxerr" "gvisor.dev/gvisor/pkg/fdnotifier" @@ -119,6 +120,10 @@ type Socket struct { // will return EWOULDBLOCK instead of blocking on the host. This allows us to // handle blocking behavior independently in the sentry. fd int + + // recvClosed indicates that the socket has been shutdown for reading + // (SHUT_RD or SHUT_RDWR). + recvClosed atomicbitops.Bool } var _ = socket.Socket(&Socket{}) @@ -358,7 +363,7 @@ func (s *Socket) Connect(t *kernel.Task, sockaddr []byte, blocking bool) *syserr // to state CONNECTED, which we can do by calling connect() a second // time ourselves. _, _, errno = unix.Syscall(unix.SYS_CONNECT, uintptr(s.fd), uintptr(firstBytePtr(sockaddr)), uintptr(len(sockaddr))) - if errno != 0 { + if errno != 0 && errno != unix.EALREADY { return syserr.FromError(translateIOSyscallError(errno)) } return nil @@ -390,7 +395,7 @@ func (s *Socket) Accept(t *kernel.Task, peerRequested bool, flags int, blocking } } else { var e waiter.Entry - e, ch = waiter.NewChannelEntry(waiter.ReadableEvents) + e, ch = waiter.NewChannelEntry(waiter.ReadableEvents | waiter.EventHUp | waiter.EventErr) s.EventRegister(&e) defer s.EventUnregister(&e) } @@ -445,7 +450,11 @@ func (s *Socket) Listen(_ *kernel.Task, backlog int) *syserr.Error { // Shutdown implements socket.Socket.Shutdown. func (s *Socket) Shutdown(_ *kernel.Task, how int) *syserr.Error { switch how { - case unix.SHUT_RD, unix.SHUT_WR, unix.SHUT_RDWR: + case unix.SHUT_RD, unix.SHUT_RDWR: + // Mark the socket as closed for reading. + s.recvClosed.Store(true) + fallthrough + case unix.SHUT_WR: return syserr.FromError(unix.Shutdown(s.fd, how)) default: return syserr.ErrInvalidArgument @@ -537,6 +546,10 @@ func (s *Socket) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags int, have if n != 0 { panic(fmt.Sprintf("CopyOutFrom: got (%d, %v), wanted (0, %v)", n, err, err)) } + // Are we closed for reading? No sense in trying to read if so. + if s.recvClosed.Load() { + break + } if ch != nil { if err = t.BlockWithDeadline(ch, haveDeadline, deadline); err != nil { if linuxerr.Equals(linuxerr.ETIMEDOUT, err) { @@ -546,7 +559,7 @@ func (s *Socket) RecvMsg(t *kernel.Task, dst usermem.IOSequence, flags int, have } } else { var e waiter.Entry - e, ch = waiter.NewChannelEntry(waiter.ReadableEvents) + e, ch = waiter.NewChannelEntry(waiter.ReadableEvents | waiter.EventRdHUp | waiter.EventHUp | waiter.EventErr) s.EventRegister(&e) defer s.EventUnregister(&e) } @@ -754,7 +767,7 @@ func (s *Socket) SendMsg(t *kernel.Task, src usermem.IOSequence, to []byte, flag } } else { var e waiter.Entry - e, ch = waiter.NewChannelEntry(waiter.WritableEvents) + e, ch = waiter.NewChannelEntry(waiter.WritableEvents | waiter.EventHUp | waiter.EventErr) s.EventRegister(&e) defer s.EventUnregister(&e) } diff --git a/pkg/waiter/waiter.go b/pkg/waiter/waiter.go index ff8495989..89a332c1c 100644 --- a/pkg/waiter/waiter.go +++ b/pkg/waiter/waiter.go @@ -75,8 +75,9 @@ const ( EventRdNorm EventMask = 0x0040 // POLLRDNORM EventWrNorm EventMask = 0x0100 // POLLWRNORM EventInternal EventMask = 0x1000 + EventRdHUp EventMask = 0x2000 // POLLRDHUP - allEvents EventMask = 0x1f | EventRdNorm | EventWrNorm + allEvents EventMask = 0x1f | EventRdNorm | EventWrNorm | EventRdHUp ReadableEvents EventMask = EventIn | EventRdNorm WritableEvents EventMask = EventOut | EventWrNorm ) diff --git a/test/syscalls/BUILD b/test/syscalls/BUILD index fe21696ca..004b8ac13 100644 --- a/test/syscalls/BUILD +++ b/test/syscalls/BUILD @@ -743,6 +743,7 @@ syscall_test( syscall_test( size = "large", + add_hostinet = True, shard_count = most_shards, test = "//test/syscalls/linux:socket_inet_loopback_test", ) diff --git a/test/syscalls/linux/socket_inet_loopback.cc b/test/syscalls/linux/socket_inet_loopback.cc index b81e8148a..68a16d84b 100644 --- a/test/syscalls/linux/socket_inet_loopback.cc +++ b/test/syscalls/linux/socket_inet_loopback.cc @@ -498,7 +498,8 @@ TEST_P(SocketInetLoopbackTest, TCPInfoState) { int n = poll(&pfd, 1, kTimeout); ASSERT_GE(n, 0) << strerror(errno); ASSERT_EQ(n, 1); - if (IsRunningOnGvisor() && GvisorPlatform() != Platform::kFuchsia) { + if (IsRunningOnGvisor() && !IsRunningWithHostinet() && + GvisorPlatform() != Platform::kFuchsia) { // TODO(gvisor.dev/issue/6015): Notify POLLRDHUP on incoming FIN. ASSERT_EQ(pfd.revents, POLLIN); } else { @@ -649,9 +650,17 @@ void TestListenHangupConnectingRead(const SocketInetTestParam& param, hangup(listen_fd); + int connecting_client_error = ECONNREFUSED; + if (IsRunningWithHostinet()) { + // TODO(b/267210840): For some reason the connecting client gets + // ECONNRESET on hostinet. Maybe the intervening poll() implementation + // changes the socket state somehow? + connecting_client_error = ECONNRESET; + } + std::array, 2> sockets = { std::make_pair(established_client.get(), ECONNRESET), - std::make_pair(connecting_client.get(), ECONNREFUSED), + std::make_pair(connecting_client.get(), connecting_client_error), }; for (size_t i = 0; i < sockets.size(); i++) { SCOPED_TRACE(absl::StrCat("i=", i)); @@ -741,7 +750,8 @@ TEST_P(SocketInetLoopbackTest, TCPNonBlockingConnectClose) { ASSERT_GE(n, 0) << strerror(errno); ASSERT_EQ(n, 1); - if (IsRunningOnGvisor() && GvisorPlatform() != Platform::kFuchsia) { + if (IsRunningOnGvisor() && !IsRunningWithHostinet() && + GvisorPlatform() != Platform::kFuchsia) { // TODO(gvisor.dev/issue/6015): Notify POLLRDHUP on incoming FIN. ASSERT_EQ(pfd.revents, POLLIN); } else { diff --git a/test/syscalls/linux/tcp_socket.cc b/test/syscalls/linux/tcp_socket.cc index b9e6d9eeb..a8c2ca36b 100644 --- a/test/syscalls/linux/tcp_socket.cc +++ b/test/syscalls/linux/tcp_socket.cc @@ -2549,18 +2549,16 @@ TEST_P(SimpleTcpSocketTest, SynRcvdOnListenerShutdown) { POLLOUT #else []() { - const int expected_revents = - POLLIN | POLLOUT | POLLHUP | POLLRDNORM | POLLWRNORM; + const int expected_revents = POLLIN | POLLOUT | POLLHUP | + POLLRDNORM | POLLWRNORM | + POLLRDHUP; // TODO(gvisor.dev/issue/6666): POLLERR is still present // after getsockopt(..., SO_ERROR, ...) call (unless // hostinet is used). - if (IsRunningOnGvisor()) { - if (IsRunningWithHostinet()) { - return expected_revents; - } + if (IsRunningOnGvisor() && !IsRunningWithHostinet()) { return expected_revents | POLLPRI | POLLERR; } - return expected_revents | POLLRDHUP; + return expected_revents; }() #endif );