Enable socket_inet_loopback test on hostinet.

A few minor fixes. The biggest change is that the blocking implementation needs
to wait on POLLHUP and POLLERR events, in addition to readable/writable events.
We also need to track shutdown state in the socket.

PiperOrigin-RevId: 529816115
This commit is contained in:
Nicolas Lacasse
2023-05-05 14:35:57 -07:00
committed by gVisor bot
parent 6763252ef0
commit 4b35f1242d
6 changed files with 40 additions and 16 deletions
+1
View File
@@ -21,6 +21,7 @@ go_library(
visibility = ["//pkg/sentry:internal"],
deps = [
"//pkg/abi/linux",
"//pkg/atomicbitops",
"//pkg/binary",
"//pkg/context",
"//pkg/errors/linuxerr",
+18 -5
View File
@@ -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)
}
+2 -1
View File
@@ -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
)
+1
View File
@@ -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",
)
+13 -3
View File
@@ -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<std::pair<int, int>, 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 {
+5 -7
View File
@@ -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
);