diff --git a/pkg/sentry/socket/netlink/socket.go b/pkg/sentry/socket/netlink/socket.go index d25614770..b4d64b179 100644 --- a/pkg/sentry/socket/netlink/socket.go +++ b/pkg/sentry/socket/netlink/socket.go @@ -18,6 +18,7 @@ package netlink import ( "io" "math" + "time" "gvisor.dev/gvisor/pkg/abi/linux" "gvisor.dev/gvisor/pkg/abi/linux/errno" @@ -387,6 +388,20 @@ func (s *Socket) GetSockOpt(t *kernel.Task, level int, name int, outPtr hostarch passcred = 1 } return &passcred, nil + + case linux.SO_SNDTIMEO: + if outLen < linux.SizeOfTimeval { + return nil, syserr.ErrInvalidArgument + } + sendTimeout := linux.NsecToTimeval(s.SendTimeout()) + return &sendTimeout, nil + + case linux.SO_RCVTIMEO: + if outLen < linux.SizeOfTimeval { + return nil, syserr.ErrInvalidArgument + } + recvTimeout := linux.NsecToTimeval(s.RecvTimeout()) + return &recvTimeout, nil } case linux.SOL_NETLINK: switch name { @@ -472,6 +487,32 @@ func (s *Socket) SetSockOpt(t *kernel.Task, level int, name int, opt []byte) *sy } return nil + + case linux.SO_SNDTIMEO: + if len(opt) < linux.SizeOfTimeval { + return syserr.ErrInvalidArgument + } + + var v linux.Timeval + v.UnmarshalBytes(opt) + if v.Usec < 0 || v.Usec >= int64(time.Second/time.Microsecond) { + return syserr.ErrDomain + } + s.SetSendTimeout(v.ToNsecCapped()) + return nil + + case linux.SO_RCVTIMEO: + if len(opt) < linux.SizeOfTimeval { + return syserr.ErrInvalidArgument + } + + var v linux.Timeval + v.UnmarshalBytes(opt) + if v.Usec < 0 || v.Usec >= int64(time.Second/time.Microsecond) { + return syserr.ErrDomain + } + s.SetRecvTimeout(v.ToNsecCapped()) + return nil } case linux.SOL_NETLINK: switch name { diff --git a/test/syscalls/linux/BUILD b/test/syscalls/linux/BUILD index ef71995ea..8049a214d 100644 --- a/test/syscalls/linux/BUILD +++ b/test/syscalls/linux/BUILD @@ -3357,6 +3357,7 @@ cc_binary( "//test/util:file_descriptor", "//test/util:socket_util", gtest, + "//test/util:posix_error", "//test/util:test_main", "//test/util:test_util", ], diff --git a/test/syscalls/linux/socket_netlink.cc b/test/syscalls/linux/socket_netlink.cc index c78529a14..f5b6aad4f 100644 --- a/test/syscalls/linux/socket_netlink.cc +++ b/test/syscalls/linux/socket_netlink.cc @@ -17,8 +17,10 @@ #include #include +#include "gmock/gmock.h" #include "gtest/gtest.h" #include "test/util/file_descriptor.h" +#include "test/util/posix_error.h" #include "test/util/socket_util.h" #include "test/util/test_util.h" @@ -143,6 +145,52 @@ TEST_P(NetlinkTest, GetPeerName) { EXPECT_EQ(addr.nl_pid, 0); } +TEST_P(NetlinkTest, GetSendTimeout) { + const int protocol = GetParam(); + + FileDescriptor fd = + ASSERT_NO_ERRNO_AND_VALUE(Socket(AF_NETLINK, SOCK_RAW, protocol)); + + // tv_usec should be a multiple of 4000 to work on most systems. + struct timeval tv_to_set { + .tv_sec = 1, .tv_usec = 40000 + }; + EXPECT_THAT(setsockopt(fd.get(), SOL_SOCKET, SO_SNDTIMEO, &tv_to_set, + sizeof(tv_to_set)), + SyscallSucceeds()); + struct timeval tv { + .tv_sec = -1, .tv_usec = -1 + }; + socklen_t len = sizeof(tv); + EXPECT_THAT(getsockopt(fd.get(), SOL_SOCKET, SO_SNDTIMEO, &tv, &len), + SyscallSucceeds()); + EXPECT_EQ(tv.tv_sec, tv_to_set.tv_sec); + EXPECT_EQ(tv.tv_usec, tv_to_set.tv_usec); +} + +TEST_P(NetlinkTest, GetReceiveTimeout) { + const int protocol = GetParam(); + + FileDescriptor fd = + ASSERT_NO_ERRNO_AND_VALUE(Socket(AF_NETLINK, SOCK_RAW, protocol)); + + // tv_usec should be a multiple of 4000 to work on most systems. + struct timeval tv_to_set { + .tv_sec = 1, .tv_usec = 8000 + }; + EXPECT_THAT(setsockopt(fd.get(), SOL_SOCKET, SO_RCVTIMEO, &tv_to_set, + sizeof(tv_to_set)), + SyscallSucceeds()); + struct timeval tv { + .tv_sec = -1, .tv_usec = -1 + }; + socklen_t len = sizeof(tv); + EXPECT_THAT(getsockopt(fd.get(), SOL_SOCKET, SO_RCVTIMEO, &tv, &len), + SyscallSucceeds()); + EXPECT_EQ(tv.tv_sec, tv_to_set.tv_sec); + EXPECT_EQ(tv.tv_usec, tv_to_set.tv_usec); +} + INSTANTIATE_TEST_SUITE_P(ProtocolTest, NetlinkTest, ::testing::Values(NETLINK_ROUTE, NETLINK_KOBJECT_UEVENT));