From 9064d4931ff6e33c263831c5c0ec47f0264a531f Mon Sep 17 00:00:00 2001 From: Cameron Gutman Date: Sat, 29 Aug 2020 18:54:11 -0700 Subject: [PATCH] Add createSocket() helper to reduce code duplication --- src/ConnectionTester.c | 24 ++-------- src/PlatformSockets.c | 103 +++++++++++++++++++++++------------------ src/PlatformSockets.h | 1 + 3 files changed, 63 insertions(+), 65 deletions(-) diff --git a/src/ConnectionTester.c b/src/ConnectionTester.c index 641e4a0..8664390 100644 --- a/src/ConnectionTester.c +++ b/src/ConnectionTester.c @@ -109,9 +109,10 @@ unsigned int LiTestClientConnectivity(const char* testServer, unsigned short ref for (i = 0; i < PORT_FLAGS_MAX_COUNT; i++) { if (testPortFlags & (1 << i)) { - sockets[i] = socket(address.ss_family, - LiGetProtocolFromPortFlagIndex(i) == IPPROTO_UDP ? SOCK_DGRAM : SOCK_STREAM, - LiGetProtocolFromPortFlagIndex(i)); + sockets[i] = createSocket(address.ss_family, + LiGetProtocolFromPortFlagIndex(i) == IPPROTO_UDP ? SOCK_DGRAM : SOCK_STREAM, + LiGetProtocolFromPortFlagIndex(i), + 1); if (sockets[i] == INVALID_SOCKET) { err = LastSocketFail(); Limelog("Failed to create socket: %d\n", err); @@ -119,25 +120,8 @@ unsigned int LiTestClientConnectivity(const char* testServer, unsigned short ref goto Exit; } - #ifdef LC_DARWIN - { - // Disable SIGPIPE on iOS - int val = 1; - setsockopt(sockets[i], SOL_SOCKET, SO_NOSIGPIPE, (char*)&val, sizeof(val)); - } - #endif - ((struct sockaddr_in6*)&address)->sin6_port = htons(LiGetPortFromPortFlagIndex(i)); if (LiGetProtocolFromPortFlagIndex(i) == IPPROTO_TCP) { - // Enable non-blocking I/O for connect timeout support - if (setSocketNonBlocking(sockets[i] , 1) != 0) { - // If non-blocking sockets are not available, TCP tests are not supported - err = LastSocketFail(); - Limelog("Failed to enable non-blocking I/O: %d\n", err); - failingPortFlags = ML_TEST_RESULT_INCONCLUSIVE; - goto Exit; - } - // Initiate an asynchronous connection err = connect(sockets[i], (struct sockaddr*)&address, address_length); if (err < 0) { diff --git a/src/PlatformSockets.c b/src/PlatformSockets.c index bb2ca0c..6faa55d 100644 --- a/src/PlatformSockets.c +++ b/src/PlatformSockets.c @@ -180,9 +180,8 @@ SOCKET bindUdpSocket(int addrfamily, int bufferSize) { LC_ASSERT(addrfamily == AF_INET || addrfamily == AF_INET6); - s = socket(addrfamily, SOCK_DGRAM, IPPROTO_UDP); + s = createSocket(addrfamily, SOCK_DGRAM, IPPROTO_UDP, 0); if (s == INVALID_SOCKET) { - Limelog("socket() failed: %d\n", (int)LastSocketError()); return INVALID_SOCKET; } @@ -251,25 +250,42 @@ int setSocketNonBlocking(SOCKET s, int val) { #endif } -SOCKET connectTcpSocket(struct sockaddr_storage* dstaddr, SOCKADDR_LEN addrlen, unsigned short port, int timeoutSec) { +SOCKET createSocket(int addressFamily, int socketType, int protocol, int nonBlocking) { SOCKET s; - struct sockaddr_in6 addr; - int err; - int nonBlocking; int val; - s = socket(dstaddr->ss_family, SOCK_STREAM, IPPROTO_TCP); + s = socket(addressFamily, socketType, protocol); if (s == INVALID_SOCKET) { Limelog("socket() failed: %d\n", (int)LastSocketError()); return INVALID_SOCKET; } - + #ifdef LC_DARWIN // Disable SIGPIPE on iOS val = 1; setsockopt(s, SOL_SOCKET, SO_NOSIGPIPE, (char*)&val, sizeof(val)); #endif + if (nonBlocking) { + setSocketNonBlocking(s, 1); + } + + return s; +} + +SOCKET connectTcpSocket(struct sockaddr_storage* dstaddr, SOCKADDR_LEN addrlen, unsigned short port, int timeoutSec) { + SOCKET s; + struct sockaddr_in6 addr; + struct pollfd pfd; + int err; + int val; + + // Create a non-blocking TCP socket + s = createSocket(dstaddr->ss_family, SOCK_STREAM, IPPROTO_TCP, 1); + if (s == INVALID_SOCKET) { + return INVALID_SOCKET; + } + // Some broken routers/firewalls (or routes with multiple broken routers) may result in TCP packets // being dropped without without us receiving an ICMP Fragmentation Needed packet. For example, // a router can elect to drop rather than fragment even without DF set. A misconfigured firewall @@ -311,9 +327,6 @@ SOCKET connectTcpSocket(struct sockaddr_storage* dstaddr, SOCKADDR_LEN addrlen, Limelog("setsockopt(TCP_MAXSEG, %d) failed: %d\n", val, (int)LastSocketError()); } #endif - - // Enable non-blocking I/O for connect timeout support - nonBlocking = setSocketNonBlocking(s, 1) == 0; // Start connection memcpy(&addr, dstaddr, addrlen); @@ -321,44 +334,44 @@ SOCKET connectTcpSocket(struct sockaddr_storage* dstaddr, SOCKADDR_LEN addrlen, err = connect(s, (struct sockaddr*) &addr, addrlen); if (err < 0) { err = (int)LastSocketError(); + if (err != EWOULDBLOCK && err != EAGAIN && err != EINPROGRESS) { + goto Exit; + } } - if (nonBlocking) { - struct pollfd pfd; + // Wait for the connection to complete or the timeout to elapse + pfd.fd = s; + pfd.events = POLLOUT; + err = pollSockets(&pfd, 1, timeoutSec * 1000); + if (err < 0) { + // pollSockets() failed + err = LastSocketError(); + Limelog("pollSockets() failed: %d\n", err); + closeSocket(s); + SetLastSocketError(err); + return INVALID_SOCKET; + } + else if (err == 0) { + // pollSockets() timed out + Limelog("Connection timed out after %d seconds (TCP port %u)\n", timeoutSec, port); + closeSocket(s); + SetLastSocketError(ETIMEDOUT); + return INVALID_SOCKET; + } + else { + // The socket was signalled + SOCKADDR_LEN len = sizeof(err); + getsockopt(s, SOL_SOCKET, SO_ERROR, (char*)&err, &len); + if (err != 0 || (pfd.revents & POLLERR)) { + // Get the error code + err = (err != 0) ? err : LastSocketFail(); + } + } - // Wait for the connection to complete or the timeout to elapse - pfd.fd = s; - pfd.events = POLLOUT; - err = pollSockets(&pfd, 1, timeoutSec * 1000); - if (err < 0) { - // pollSockets() failed - err = LastSocketError(); - Limelog("pollSockets() failed: %d\n", err); - closeSocket(s); - SetLastSocketError(err); - return INVALID_SOCKET; - } - else if (err == 0) { - // pollSockets() timed out - Limelog("Connection timed out after %d seconds (TCP port %u)\n", timeoutSec, port); - closeSocket(s); - SetLastSocketError(ETIMEDOUT); - return INVALID_SOCKET; - } - else { - // The socket was signalled - SOCKADDR_LEN len = sizeof(err); - getsockopt(s, SOL_SOCKET, SO_ERROR, (char*)&err, &len); - if (err != 0 || (pfd.revents & POLLERR)) { - // Get the error code - err = (err != 0) ? err : LastSocketFail(); - } - } - - // Disable non-blocking I/O now that the connection is established - setSocketNonBlocking(s, 0); - } + // Disable non-blocking I/O now that the connection is established + setSocketNonBlocking(s, 0); +Exit: if (err != 0) { Limelog("connect() failed: %d\n", err); closeSocket(s); diff --git a/src/PlatformSockets.h b/src/PlatformSockets.h index 5bb106c..4517319 100644 --- a/src/PlatformSockets.h +++ b/src/PlatformSockets.h @@ -51,6 +51,7 @@ typedef socklen_t SOCKADDR_LEN; #define URLSAFESTRING_LEN (INET6_ADDRSTRLEN+2) void addrToUrlSafeString(struct sockaddr_storage* addr, char* string); +SOCKET createSocket(int addressFamily, int socketType, int protocol, int nonBlocking); int resolveHostName(const char* host, int family, int tcpTestPort, struct sockaddr_storage* addr, SOCKADDR_LEN* addrLen); SOCKET connectTcpSocket(struct sockaddr_storage* dstaddr, SOCKADDR_LEN addrlen, unsigned short port, int timeoutSec); int sendMtuSafe(SOCKET s, char* buffer, int size);