diff --git a/pkg/sentry/socket/hostinet/socket.go b/pkg/sentry/socket/hostinet/socket.go index c457cd427..e2fb93f3f 100644 --- a/pkg/sentry/socket/hostinet/socket.go +++ b/pkg/sentry/socket/hostinet/socket.go @@ -384,7 +384,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) } @@ -540,11 +540,21 @@ 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) } n, err = copyToDst() + + // If we got an RDHUP event, there's no use trying to + // read again since the socket has been shutdown. + if s.Readiness(waiter.EventRdHUp)&waiter.EventRdHUp > 0 { + // Don't return the EWOULDBLOCK. + if linuxerr.Equals(linuxerr.ErrWouldBlock, err) { + err = nil + } + break + } } } if err != nil { @@ -743,7 +753,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 2e048e04d..cbd4fcf7f 100644 --- a/test/syscalls/BUILD +++ b/test/syscalls/BUILD @@ -742,6 +742,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 37b7add32..f35b2b73e 100644 --- a/test/syscalls/linux/tcp_socket.cc +++ b/test/syscalls/linux/tcp_socket.cc @@ -2488,18 +2488,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; - } - return expected_revents | POLLPRI | POLLERR; + if (IsRunningWithHostinet()) { + return expected_revents; } - return expected_revents | POLLRDHUP; + return expected_revents | POLLPRI | POLLERR; }() #endif );