Allow dual stack sockets to operate on AF_INET

Fixes #1490
Fixes #1495

PiperOrigin-RevId: 289523250
This commit is contained in:
Tamir Duberstein
2020-01-13 14:47:22 -08:00
committed by gVisor bot
parent fff0476951
commit debd213da6
10 changed files with 278 additions and 82 deletions
+51 -14
View File
@@ -324,22 +324,15 @@ func bytesToIPAddress(addr []byte) tcpip.Address {
// converts it to the FullAddress format. It supports AF_UNIX, AF_INET,
// AF_INET6, and AF_PACKET addresses.
//
// strict indicates whether addresses with the AF_UNSPEC family are accepted of not.
//
// AddressAndFamily returns an address and its family.
func AddressAndFamily(sfamily int, addr []byte, strict bool) (tcpip.FullAddress, uint16, *syserr.Error) {
func AddressAndFamily(addr []byte) (tcpip.FullAddress, uint16, *syserr.Error) {
// Make sure we have at least 2 bytes for the address family.
if len(addr) < 2 {
return tcpip.FullAddress{}, 0, syserr.ErrInvalidArgument
}
family := usermem.ByteOrder.Uint16(addr)
if family != uint16(sfamily) && (strict || family != linux.AF_UNSPEC) {
return tcpip.FullAddress{}, family, syserr.ErrAddressFamilyNotSupported
}
// Get the rest of the fields based on the address family.
switch family {
switch family := usermem.ByteOrder.Uint16(addr); family {
case linux.AF_UNIX:
path := addr[2:]
if len(path) > linux.UnixPathMax {
@@ -638,10 +631,40 @@ func (s *SocketOperations) Readiness(mask waiter.EventMask) waiter.EventMask {
return r
}
func (s *SocketOperations) checkFamily(family uint16, exact bool) *syserr.Error {
if family == uint16(s.family) {
return nil
}
if !exact && family == linux.AF_INET && s.family == linux.AF_INET6 {
v, err := s.Endpoint.GetSockOptBool(tcpip.V6OnlyOption)
if err != nil {
return syserr.TranslateNetstackError(err)
}
if !v {
return nil
}
}
return syserr.ErrInvalidArgument
}
// mapFamily maps the AF_INET ANY address to the IPv4-mapped IPv6 ANY if the
// receiver's family is AF_INET6.
//
// This is a hack to work around the fact that both IPv4 and IPv6 ANY are
// represented by the empty string.
//
// TODO(gvisor.dev/issues/1556): remove this function.
func (s *SocketOperations) mapFamily(addr tcpip.FullAddress, family uint16) tcpip.FullAddress {
if len(addr.Addr) == 0 && s.family == linux.AF_INET6 && family == linux.AF_INET {
addr.Addr = "\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\xff\xff\x00\x00\x00\x00"
}
return addr
}
// Connect implements the linux syscall connect(2) for sockets backed by
// tpcip.Endpoint.
func (s *SocketOperations) Connect(t *kernel.Task, sockaddr []byte, blocking bool) *syserr.Error {
addr, family, err := AddressAndFamily(s.family, sockaddr, false /* strict */)
addr, family, err := AddressAndFamily(sockaddr)
if err != nil {
return err
}
@@ -653,6 +676,12 @@ func (s *SocketOperations) Connect(t *kernel.Task, sockaddr []byte, blocking boo
}
return syserr.TranslateNetstackError(err)
}
if err := s.checkFamily(family, false /* exact */); err != nil {
return err
}
addr = s.mapFamily(addr, family)
// Always return right away in the non-blocking case.
if !blocking {
return syserr.TranslateNetstackError(s.Endpoint.Connect(addr))
@@ -681,10 +710,14 @@ func (s *SocketOperations) Connect(t *kernel.Task, sockaddr []byte, blocking boo
// Bind implements the linux syscall bind(2) for sockets backed by
// tcpip.Endpoint.
func (s *SocketOperations) Bind(t *kernel.Task, sockaddr []byte) *syserr.Error {
addr, _, err := AddressAndFamily(s.family, sockaddr, true /* strict */)
addr, family, err := AddressAndFamily(sockaddr)
if err != nil {
return err
}
if err := s.checkFamily(family, true /* exact */); err != nil {
return err
}
addr = s.mapFamily(addr, family)
// Issue the bind request to the endpoint.
return syserr.TranslateNetstackError(s.Endpoint.Bind(addr))
@@ -2080,8 +2113,8 @@ func ConvertAddress(family int, addr tcpip.FullAddress) (linux.SockAddr, uint32)
case linux.AF_INET6:
var out linux.SockAddrInet6
if len(addr.Addr) == 4 {
// Copy address is v4-mapped format.
if len(addr.Addr) == header.IPv4AddressSize {
// Copy address in v4-mapped format.
copy(out.Addr[12:], addr.Addr)
out.Addr[10] = 0xff
out.Addr[11] = 0xff
@@ -2395,10 +2428,14 @@ func (s *SocketOperations) SendMsg(t *kernel.Task, src usermem.IOSequence, to []
var addr *tcpip.FullAddress
if len(to) > 0 {
addrBuf, _, err := AddressAndFamily(s.family, to, true /* strict */)
addrBuf, family, err := AddressAndFamily(to)
if err != nil {
return 0, err
}
if err := s.checkFamily(family, false /* exact */); err != nil {
return 0, err
}
addrBuf = s.mapFamily(addrBuf, family)
addr = &addrBuf
}
+4 -1
View File
@@ -116,13 +116,16 @@ func (s *SocketOperations) Endpoint() transport.Endpoint {
// extractPath extracts and validates the address.
func extractPath(sockaddr []byte) (string, *syserr.Error) {
addr, _, err := netstack.AddressAndFamily(linux.AF_UNIX, sockaddr, true /* strict */)
addr, family, err := netstack.AddressAndFamily(sockaddr)
if err != nil {
if err == syserr.ErrAddressFamilyNotSupported {
err = syserr.ErrInvalidArgument
}
return "", err
}
if family != linux.AF_UNIX {
return "", syserr.ErrInvalidArgument
}
// The address is trimmed by GetAddress.
p := string(addr.Addr)
+1 -1
View File
@@ -341,7 +341,7 @@ func sockAddr(t *kernel.Task, addr usermem.Addr, length uint32) string {
switch family {
case linux.AF_INET, linux.AF_INET6, linux.AF_UNIX:
fa, _, err := netstack.AddressAndFamily(int(family), b, true /* strict */)
fa, _, err := netstack.AddressAndFamily(b)
if err != nil {
return fmt.Sprintf("%#x {Family: %s, error extracting address: %v}", addr, familyStr, err)
}
+43
View File
@@ -547,6 +547,49 @@ type TransportEndpointInfo struct {
RegisterNICID tcpip.NICID
}
// AddrNetProto unwraps the specified address if it is a V4-mapped V6 address
// and returns the network protocol number to be used to communicate with the
// specified address. It returns an error if the passed address is incompatible
// with the receiver.
func (e *TransportEndpointInfo) AddrNetProto(addr tcpip.FullAddress, v6only bool) (tcpip.FullAddress, tcpip.NetworkProtocolNumber, *tcpip.Error) {
netProto := e.NetProto
switch len(addr.Addr) {
case header.IPv4AddressSize:
netProto = header.IPv4ProtocolNumber
case header.IPv6AddressSize:
if header.IsV4MappedAddress(addr.Addr) {
netProto = header.IPv4ProtocolNumber
addr.Addr = addr.Addr[header.IPv6AddressSize-header.IPv4AddressSize:]
if addr.Addr == header.IPv4Any {
addr.Addr = ""
}
}
}
switch len(e.ID.LocalAddress) {
case header.IPv4AddressSize:
if len(addr.Addr) == header.IPv6AddressSize {
return tcpip.FullAddress{}, 0, tcpip.ErrInvalidEndpointState
}
case header.IPv6AddressSize:
if len(addr.Addr) == header.IPv4AddressSize {
return tcpip.FullAddress{}, 0, tcpip.ErrNetworkUnreachable
}
}
switch {
case netProto == e.NetProto:
case netProto == header.IPv4ProtocolNumber && e.NetProto == header.IPv6ProtocolNumber:
if v6only {
return tcpip.FullAddress{}, 0, tcpip.ErrNoRoute
}
default:
return tcpip.FullAddress{}, 0, tcpip.ErrInvalidEndpointState
}
return addr, netProto, nil
}
// IsEndpointInfo is an empty method to implement the tcpip.EndpointInfo
// marker interface.
func (*TransportEndpointInfo) IsEndpointInfo() {}
+8 -14
View File
@@ -288,7 +288,7 @@ func (e *endpoint) write(p tcpip.Payloader, opts tcpip.WriteOptions) (int64, <-c
toCopy := *to
to = &toCopy
netProto, err := e.checkV4Mapped(to, true)
netProto, err := e.checkV4Mapped(to)
if err != nil {
return 0, nil, err
}
@@ -475,18 +475,12 @@ func send6(r *stack.Route, ident uint16, data buffer.View, ttl uint8) *tcpip.Err
})
}
func (e *endpoint) checkV4Mapped(addr *tcpip.FullAddress, allowMismatch bool) (tcpip.NetworkProtocolNumber, *tcpip.Error) {
netProto := e.NetProto
if header.IsV4MappedAddress(addr.Addr) {
return 0, tcpip.ErrNoRoute
func (e *endpoint) checkV4Mapped(addr *tcpip.FullAddress) (tcpip.NetworkProtocolNumber, *tcpip.Error) {
unwrapped, netProto, err := e.TransportEndpointInfo.AddrNetProto(*addr, false /* v6only */)
if err != nil {
return 0, err
}
// Fail if we're bound to an address length different from the one we're
// checking.
if l := len(e.ID.LocalAddress); !allowMismatch && l != 0 && l != len(addr.Addr) {
return 0, tcpip.ErrInvalidEndpointState
}
*addr = unwrapped
return netProto, nil
}
@@ -518,7 +512,7 @@ func (e *endpoint) Connect(addr tcpip.FullAddress) *tcpip.Error {
return tcpip.ErrInvalidEndpointState
}
netProto, err := e.checkV4Mapped(&addr, false)
netProto, err := e.checkV4Mapped(&addr)
if err != nil {
return err
}
@@ -631,7 +625,7 @@ func (e *endpoint) bindLocked(addr tcpip.FullAddress) *tcpip.Error {
return tcpip.ErrInvalidEndpointState
}
netProto, err := e.checkV4Mapped(&addr, false)
netProto, err := e.checkV4Mapped(&addr)
if err != nil {
return err
}
+4 -19
View File
@@ -1691,26 +1691,11 @@ func (e *endpoint) GetSockOpt(opt interface{}) *tcpip.Error {
}
func (e *endpoint) checkV4Mapped(addr *tcpip.FullAddress) (tcpip.NetworkProtocolNumber, *tcpip.Error) {
netProto := e.NetProto
if header.IsV4MappedAddress(addr.Addr) {
// Fail if using a v4 mapped address on a v6only endpoint.
if e.v6only {
return 0, tcpip.ErrNoRoute
}
netProto = header.IPv4ProtocolNumber
addr.Addr = addr.Addr[header.IPv6AddressSize-header.IPv4AddressSize:]
if addr.Addr == header.IPv4Any {
addr.Addr = ""
}
unwrapped, netProto, err := e.TransportEndpointInfo.AddrNetProto(*addr, e.v6only)
if err != nil {
return 0, err
}
// Fail if we're bound to an address length different from the one we're
// checking.
if l := len(e.ID.LocalAddress); l != 0 && len(addr.Addr) != 0 && l != len(addr.Addr) {
return 0, tcpip.ErrInvalidEndpointState
}
*addr = unwrapped
return netProto, nil
}
+9 -32
View File
@@ -402,7 +402,7 @@ func (e *endpoint) write(p tcpip.Payloader, opts tcpip.WriteOptions) (int64, <-c
return 0, nil, tcpip.ErrBroadcastDisabled
}
netProto, err := e.checkV4Mapped(to, false)
netProto, err := e.checkV4Mapped(to)
if err != nil {
return 0, nil, err
}
@@ -501,7 +501,7 @@ func (e *endpoint) SetSockOpt(opt interface{}) *tcpip.Error {
defer e.mu.Unlock()
fa := tcpip.FullAddress{Addr: v.InterfaceAddr}
netProto, err := e.checkV4Mapped(&fa, false)
netProto, err := e.checkV4Mapped(&fa)
if err != nil {
return err
}
@@ -839,35 +839,12 @@ func sendUDP(r *stack.Route, data buffer.VectorisedView, localPort, remotePort u
return nil
}
func (e *endpoint) checkV4Mapped(addr *tcpip.FullAddress, allowMismatch bool) (tcpip.NetworkProtocolNumber, *tcpip.Error) {
netProto := e.NetProto
if len(addr.Addr) == 0 {
return netProto, nil
func (e *endpoint) checkV4Mapped(addr *tcpip.FullAddress) (tcpip.NetworkProtocolNumber, *tcpip.Error) {
unwrapped, netProto, err := e.TransportEndpointInfo.AddrNetProto(*addr, e.v6only)
if err != nil {
return 0, err
}
if header.IsV4MappedAddress(addr.Addr) {
// Fail if using a v4 mapped address on a v6only endpoint.
if e.v6only {
return 0, tcpip.ErrNoRoute
}
netProto = header.IPv4ProtocolNumber
addr.Addr = addr.Addr[header.IPv6AddressSize-header.IPv4AddressSize:]
if addr.Addr == header.IPv4Any {
addr.Addr = ""
}
// Fail if we are bound to an IPv6 address.
if !allowMismatch && len(e.ID.LocalAddress) == 16 {
return 0, tcpip.ErrNetworkUnreachable
}
}
// Fail if we're bound to an address length different from the one we're
// checking.
if l := len(e.ID.LocalAddress); l != 0 && l != len(addr.Addr) {
return 0, tcpip.ErrInvalidEndpointState
}
*addr = unwrapped
return netProto, nil
}
@@ -916,7 +893,7 @@ func (e *endpoint) Disconnect() *tcpip.Error {
// Connect connects the endpoint to its peer. Specifying a NIC is optional.
func (e *endpoint) Connect(addr tcpip.FullAddress) *tcpip.Error {
netProto, err := e.checkV4Mapped(&addr, false)
netProto, err := e.checkV4Mapped(&addr)
if err != nil {
return err
}
@@ -1074,7 +1051,7 @@ func (e *endpoint) bindLocked(addr tcpip.FullAddress) *tcpip.Error {
return tcpip.ErrInvalidEndpointState
}
netProto, err := e.checkV4Mapped(&addr, true)
netProto, err := e.checkV4Mapped(&addr)
if err != nil {
return err
}
+1 -1
View File
@@ -73,7 +73,7 @@ function install_runsc() {
sudo "${RUNSC_BIN}" install --experimental=true --runtime="${runtime}" -- --debug-log "${RUNSC_LOGS}" "$@"
# Clear old logs files that may exist.
sudo rm -f "${RUNSC_LOGS_DIR}"/*
sudo rm -f "${RUNSC_LOGS_DIR}"/'*'
# Restart docker to pick up the new runtime configuration.
sudo systemctl restart docker
+1
View File
@@ -2693,6 +2693,7 @@ cc_binary(
srcs = ["socket_inet_loopback.cc"],
linkstatic = 1,
deps = [
":ip_socket_test_util",
":socket_test_util",
"//test/util:file_descriptor",
"//test/util:posix_error",
+156
View File
@@ -32,6 +32,7 @@
#include "absl/strings/str_cat.h"
#include "absl/time/clock.h"
#include "absl/time/time.h"
#include "test/syscalls/linux/ip_socket_test_util.h"
#include "test/syscalls/linux/socket_test_util.h"
#include "test/util/file_descriptor.h"
#include "test/util/posix_error.h"
@@ -102,6 +103,161 @@ TEST(BadSocketPairArgs, ValidateErrForBadCallsToSocketPair) {
SyscallFailsWithErrno(EAFNOSUPPORT));
}
enum class Operation {
Bind,
Connect,
SendTo,
};
std::string OperationToString(Operation operation) {
switch (operation) {
case Operation::Bind:
return "Bind";
case Operation::Connect:
return "Connect";
case Operation::SendTo:
return "SendTo";
}
}
using OperationSequence = std::vector<Operation>;
using DualStackSocketTest =
::testing::TestWithParam<std::tuple<TestAddress, OperationSequence>>;
TEST_P(DualStackSocketTest, AddressOperations) {
const FileDescriptor fd =
ASSERT_NO_ERRNO_AND_VALUE(Socket(AF_INET6, SOCK_DGRAM, 0));
const TestAddress& addr = std::get<0>(GetParam());
const OperationSequence& operations = std::get<1>(GetParam());
auto addr_in = reinterpret_cast<const sockaddr*>(&addr.addr);
// sockets may only be bound once. Both `connect` and `sendto` cause a socket
// to be bound.
bool bound = false;
for (const Operation& operation : operations) {
bool sockname = false;
bool peername = false;
switch (operation) {
case Operation::Bind: {
ASSERT_NO_ERRNO(SetAddrPort(
addr.family(), const_cast<sockaddr_storage*>(&addr.addr), 0));
int bind_ret = bind(fd.get(), addr_in, addr.addr_len);
// Dual stack sockets may only be bound to AF_INET6.
if (!bound && addr.family() == AF_INET6) {
EXPECT_THAT(bind_ret, SyscallSucceeds());
bound = true;
sockname = true;
} else {
EXPECT_THAT(bind_ret, SyscallFailsWithErrno(EINVAL));
}
break;
}
case Operation::Connect: {
ASSERT_NO_ERRNO(SetAddrPort(
addr.family(), const_cast<sockaddr_storage*>(&addr.addr), 1337));
EXPECT_THAT(connect(fd.get(), addr_in, addr.addr_len),
SyscallSucceeds())
<< GetAddrStr(addr_in);
bound = true;
sockname = true;
peername = true;
break;
}
case Operation::SendTo: {
const char payload[] = "hello";
ASSERT_NO_ERRNO(SetAddrPort(
addr.family(), const_cast<sockaddr_storage*>(&addr.addr), 1337));
ssize_t sendto_ret = sendto(fd.get(), &payload, sizeof(payload), 0,
addr_in, addr.addr_len);
EXPECT_THAT(sendto_ret, SyscallSucceedsWithValue(sizeof(payload)));
sockname = !bound;
bound = true;
break;
}
}
if (sockname) {
sockaddr_storage sock_addr;
socklen_t addrlen = sizeof(sock_addr);
ASSERT_THAT(getsockname(fd.get(), reinterpret_cast<sockaddr*>(&sock_addr),
&addrlen),
SyscallSucceeds());
ASSERT_EQ(addrlen, sizeof(struct sockaddr_in6));
auto sock_addr_in6 = reinterpret_cast<const sockaddr_in6*>(&sock_addr);
if (operation == Operation::SendTo) {
EXPECT_EQ(sock_addr_in6->sin6_family, AF_INET6);
EXPECT_TRUE(IN6_IS_ADDR_UNSPECIFIED(sock_addr_in6->sin6_addr.s6_addr32))
<< OperationToString(operation) << " getsocknam="
<< GetAddrStr(reinterpret_cast<sockaddr*>(&sock_addr));
EXPECT_NE(sock_addr_in6->sin6_port, 0);
} else if (IN6_IS_ADDR_V4MAPPED(
reinterpret_cast<const sockaddr_in6*>(addr_in)
->sin6_addr.s6_addr32)) {
EXPECT_TRUE(IN6_IS_ADDR_V4MAPPED(sock_addr_in6->sin6_addr.s6_addr32))
<< OperationToString(operation) << " getsocknam="
<< GetAddrStr(reinterpret_cast<sockaddr*>(&sock_addr));
}
}
if (peername) {
sockaddr_storage peer_addr;
socklen_t addrlen = sizeof(peer_addr);
ASSERT_THAT(getpeername(fd.get(), reinterpret_cast<sockaddr*>(&peer_addr),
&addrlen),
SyscallSucceeds());
ASSERT_EQ(addrlen, sizeof(struct sockaddr_in6));
if (addr.family() == AF_INET ||
IN6_IS_ADDR_V4MAPPED(reinterpret_cast<const sockaddr_in6*>(addr_in)
->sin6_addr.s6_addr32)) {
EXPECT_TRUE(IN6_IS_ADDR_V4MAPPED(
reinterpret_cast<const sockaddr_in6*>(&peer_addr)
->sin6_addr.s6_addr32))
<< OperationToString(operation) << " getpeername="
<< GetAddrStr(reinterpret_cast<sockaddr*>(&peer_addr));
}
}
}
}
// TODO(gvisor.dev/issues/1556): uncomment V4MappedAny.
INSTANTIATE_TEST_SUITE_P(
All, DualStackSocketTest,
::testing::Combine(
::testing::Values(V4Any(), V4Loopback(), /*V4MappedAny(),*/
V4MappedLoopback(), V6Any(), V6Loopback()),
::testing::ValuesIn<OperationSequence>(
{{Operation::Bind, Operation::Connect, Operation::SendTo},
{Operation::Bind, Operation::SendTo, Operation::Connect},
{Operation::Connect, Operation::Bind, Operation::SendTo},
{Operation::Connect, Operation::SendTo, Operation::Bind},
{Operation::SendTo, Operation::Bind, Operation::Connect},
{Operation::SendTo, Operation::Connect, Operation::Bind}})),
[](::testing::TestParamInfo<
std::tuple<TestAddress, OperationSequence>> const& info) {
const TestAddress& addr = std::get<0>(info.param);
const OperationSequence& operations = std::get<1>(info.param);
std::string s = addr.description;
for (const Operation& operation : operations) {
absl::StrAppend(&s, OperationToString(operation));
}
return s;
});
void tcpSimpleConnectTest(TestAddress const& listener,
TestAddress const& connector, bool unbound) {
// Create the listening socket.