Merge remote-tracking branch 'nccl/master' into HEAD
Este cometimento está contido em:
+431
-177
@@ -20,6 +20,52 @@
|
||||
|
||||
static std::vector<std::pair<int, std::unordered_set<std::string>>> clientPortPool;
|
||||
|
||||
static ncclResult_t socketProgressOpt(int op, struct ncclSocket* sock, void* ptr, int size, int* offset, int block, int* closed) {
|
||||
int bytes = 0;
|
||||
*closed = 0;
|
||||
char* data = (char*)ptr;
|
||||
char line[SOCKET_NAME_MAXLEN+1];
|
||||
do {
|
||||
if (op == NCCL_SOCKET_RECV) bytes = recv(sock->fd, data+(*offset), size-(*offset), block ? 0 : MSG_DONTWAIT);
|
||||
if (op == NCCL_SOCKET_SEND) bytes = send(sock->fd, data+(*offset), size-(*offset), block ? MSG_NOSIGNAL : MSG_DONTWAIT | MSG_NOSIGNAL);
|
||||
if (op == NCCL_SOCKET_RECV && bytes == 0) {
|
||||
*closed = 1;
|
||||
return ncclSuccess;
|
||||
}
|
||||
if (bytes == -1) {
|
||||
if (errno != EINTR && errno != EWOULDBLOCK && errno != EAGAIN) {
|
||||
WARN("socketProgressOpt: Call to recv from %s failed : %s", ncclSocketToString(&sock->addr, line), strerror(errno));
|
||||
return ncclRemoteError;
|
||||
} else {
|
||||
bytes = 0;
|
||||
}
|
||||
}
|
||||
(*offset) += bytes;
|
||||
if (sock->abortFlag && *sock->abortFlag != 0) {
|
||||
INFO(NCCL_NET, "socketProgressOpt: abort called");
|
||||
return ncclInternalError;
|
||||
}
|
||||
} while (bytes > 0 && (*offset) < size);
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t socketProgress(int op, struct ncclSocket* sock, void* ptr, int size, int* offset) {
|
||||
int closed;
|
||||
NCCLCHECK(socketProgressOpt(op, sock, ptr, size, offset, 0, &closed));
|
||||
if (closed) {
|
||||
char line[SOCKET_NAME_MAXLEN+1];
|
||||
WARN("socketProgress: Connection closed by remote peer %s", ncclSocketToString(&sock->addr, line, 0));
|
||||
return ncclRemoteError;
|
||||
}
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t socketWait(int op, struct ncclSocket* sock, void* ptr, int size, int* offset) {
|
||||
while (*offset < size)
|
||||
NCCLCHECK(socketProgress(op, sock, ptr, size, offset));
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
/* Format a string representation of a (union ncclSocketAddress *) socket address using getnameinfo()
|
||||
*
|
||||
* Output: "IPv4/IPv6 address<port>"
|
||||
@@ -202,7 +248,7 @@ int ncclFindInterfaceMatchSubnet(char* ifNames, union ncclSocketAddress* localAd
|
||||
return found;
|
||||
}
|
||||
|
||||
ncclResult_t ncclGetSocketAddrFromString(union ncclSocketAddress* ua, const char* ip_port_pair) {
|
||||
ncclResult_t ncclSocketGetAddrFromString(union ncclSocketAddress* ua, const char* ip_port_pair) {
|
||||
if (!(ip_port_pair && strlen(ip_port_pair) > 1)) {
|
||||
WARN("Net : string is null");
|
||||
return ncclInvalidArgument;
|
||||
@@ -304,7 +350,7 @@ int ncclFindInterfaces(char* ifNames, union ncclSocketAddress *ifAddrs, int ifNa
|
||||
INFO(NCCL_ENV, "NCCL_COMM_ID set by environment to %s", commId);
|
||||
// Try to find interface that is in the same subnet as the IP in comm id
|
||||
union ncclSocketAddress idAddr;
|
||||
ncclGetSocketAddrFromString(&idAddr, commId);
|
||||
ncclSocketGetAddrFromString(&idAddr, commId);
|
||||
nIfs = ncclFindInterfaceMatchSubnet(ifNames, ifAddrs, &idAddr, ifNameMaxSize, maxIfs);
|
||||
}
|
||||
}
|
||||
@@ -318,39 +364,31 @@ int ncclFindInterfaces(char* ifNames, union ncclSocketAddress *ifAddrs, int ifNa
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketListen(struct ncclSocket* sock) {
|
||||
/* IPv4/IPv6 support */
|
||||
int family = sock->addr.sa.sa_family;
|
||||
int salen = (family == AF_INET) ? sizeof(struct sockaddr_in) : sizeof(struct sockaddr_in6);
|
||||
int flags;
|
||||
|
||||
/* Create socket and bind it to a port */
|
||||
int fd = socket(family, SOCK_STREAM, 0);
|
||||
if (fd == -1) {
|
||||
WARN("Net : Socket creation failed : %s", strerror(errno));
|
||||
return ncclSystemError;
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketListen: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
if (sock->fd == -1) {
|
||||
WARN("ncclSocketListen: file descriptor is -1");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
|
||||
if (socketToPort(&sock->addr)) {
|
||||
// Port is forced by env. Make sure we get the port.
|
||||
int opt = 1;
|
||||
#if defined(SO_REUSEPORT)
|
||||
SYSCHECK(setsockopt(fd, SOL_SOCKET, SO_REUSEADDR | SO_REUSEPORT, &opt, sizeof(opt)), "setsockopt");
|
||||
SYSCHECK(setsockopt(sock->fd, SOL_SOCKET, SO_REUSEADDR | SO_REUSEPORT, &opt, sizeof(opt)), "setsockopt");
|
||||
#else
|
||||
SYSCHECK(setsockopt(fd, SOL_SOCKET, SO_REUSEADDR, &opt, sizeof(opt)), "setsockopt");
|
||||
SYSCHECK(setsockopt(sock->fd, SOL_SOCKET, SO_REUSEADDR, &opt, sizeof(opt)), "setsockopt");
|
||||
#endif
|
||||
}
|
||||
|
||||
/* The socket is set non-blocking for OS level, but asyncFlag is used to control
|
||||
* blocking and non-blocking behavior in user level. */
|
||||
EQCHECK(flags = fcntl(fd, F_GETFL), -1);
|
||||
SYSCHECK(fcntl(fd, F_SETFL, flags | O_NONBLOCK), "fcntl");
|
||||
|
||||
// addr port should be 0 (Any port)
|
||||
SYSCHECK(bind(fd, &sock->addr.sa, salen), "bind");
|
||||
SYSCHECK(bind(sock->fd, &sock->addr.sa, sock->salen), "bind");
|
||||
|
||||
/* Get the assigned Port */
|
||||
socklen_t size = salen;
|
||||
SYSCHECK(getsockname(fd, &sock->addr.sa, &size), "getsockname");
|
||||
socklen_t size = sock->salen;
|
||||
SYSCHECK(getsockname(sock->fd, &sock->addr.sa, &size), "getsockname");
|
||||
|
||||
#ifdef ENABLE_TRACE
|
||||
char line[SOCKET_NAME_MAXLEN+1];
|
||||
@@ -360,76 +398,226 @@ ncclResult_t ncclSocketListen(struct ncclSocket* sock) {
|
||||
/* Put the socket in listen mode
|
||||
* NB: The backlog will be silently truncated to the value in /proc/sys/net/core/somaxconn
|
||||
*/
|
||||
SYSCHECK(listen(fd, 16384), "listen");
|
||||
sock->fd = fd;
|
||||
SYSCHECK(listen(sock->fd, 16384), "listen");
|
||||
sock->state = ncclSocketStateReady;
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t getFdState(int fd, enum ncclSocketState* state) {
|
||||
struct pollfd pfd;
|
||||
int timeout = 1, ret;
|
||||
socklen_t rlen = sizeof(int);
|
||||
|
||||
memset(&pfd, 0, sizeof(struct pollfd));
|
||||
pfd.fd = fd;
|
||||
pfd.events = POLLOUT;
|
||||
SYSCHECK(ret = poll(&pfd, 1, timeout), "poll");
|
||||
if (ret == 0) {
|
||||
ret = EINPROGRESS;
|
||||
} else {
|
||||
/* check socket status */
|
||||
EQCHECK(ret == 1 && (pfd.revents & POLLOUT), 0);
|
||||
SYSCHECK(getsockopt(fd, SOL_SOCKET, SO_ERROR, (void*)&ret, &rlen), "getsockopt");
|
||||
}
|
||||
|
||||
if (ret == EINPROGRESS || ret == ECONNREFUSED)
|
||||
*state = ncclSocketConnecting;
|
||||
else if (ret == 0)
|
||||
*state = ncclSocketConnected;
|
||||
else
|
||||
*state = ncclSocketError;
|
||||
return ncclSuccess;
|
||||
ncclResult_t ncclSocketGetAddr(struct ncclSocket* sock, union ncclSocketAddress* addr) {
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketGetAddr: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
if (sock->state != ncclSocketStateReady) return ncclInternalError;
|
||||
memcpy(addr, &sock->addr, sizeof(union ncclSocketAddress));
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t ncclGetSocketState(struct ncclSocket* sock, enum ncclSocketState* state) {
|
||||
NCCLCHECK(getFdState(sock->fd, state));
|
||||
sock->state = *state;
|
||||
static ncclResult_t socketTryAccept(struct ncclSocket* sock) {
|
||||
socklen_t socklen = sizeof(union ncclSocketAddress);
|
||||
sock->fd = accept(sock->acceptFd, &sock->addr.sa, &socklen);
|
||||
if (sock->fd != -1) {
|
||||
sock->state = ncclSocketStateAccepted;
|
||||
} else if (errno != EAGAIN && errno != EWOULDBLOCK) {
|
||||
WARN("socketTryAccept: get errno %d that is not EAGAIN or EWOULDBLOCK", errno);
|
||||
return ncclSystemError;
|
||||
}
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t socketFinalizeAccept(struct ncclSocket* sock) {
|
||||
uint64_t magic;
|
||||
enum ncclSocketType type;
|
||||
int received = 0;
|
||||
NCCLCHECK(ncclSocketProgress(NCCL_SOCKET_RECV, sock, &magic, sizeof(magic), &received));
|
||||
if (received == 0) return ncclSuccess;
|
||||
NCCLCHECK(socketWait(NCCL_SOCKET_RECV, sock, &magic, sizeof(magic), &received));
|
||||
if (magic != sock->magic) {
|
||||
WARN("socketFinalizeAccept: wrong magic %lx != %lx", magic, sock->magic);
|
||||
close(sock->fd);
|
||||
sock->fd = -1;
|
||||
// Ignore spurious connection and accept again
|
||||
sock->state = ncclSocketStateAccepting;
|
||||
return ncclSuccess;
|
||||
} else {
|
||||
received = 0;
|
||||
NCCLCHECK(socketWait(NCCL_SOCKET_RECV, sock, &type, sizeof(type), &received));
|
||||
if (type != sock->type) {
|
||||
WARN("socketFinalizeAccept: wrong type %d != %d", type, sock->type);
|
||||
sock->state = ncclSocketStateError;
|
||||
close(sock->fd);
|
||||
sock->fd = -1;
|
||||
return ncclInternalError;
|
||||
} else {
|
||||
sock->state = ncclSocketStateReady;
|
||||
}
|
||||
}
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t socketStartConnect(struct ncclSocket* sock) {
|
||||
/* blocking/non-blocking connect() is determined by asyncFlag. */
|
||||
int ret = connect(sock->fd, &sock->addr.sa, sock->salen);
|
||||
|
||||
if (ret == 0) {
|
||||
sock->state = ncclSocketStateConnected;
|
||||
return ncclSuccess;
|
||||
} else if (errno == EINPROGRESS) {
|
||||
sock->state = ncclSocketStateConnectPolling;
|
||||
return ncclSuccess;
|
||||
} else if (errno == ECONNREFUSED) {
|
||||
if (++sock->refusedRetries == RETRY_REFUSED_TIMES) {
|
||||
sock->state = ncclSocketStateError;
|
||||
WARN("socketStartConnect: exceeded retries (%d)", sock->refusedRetries);
|
||||
return ncclRemoteError;
|
||||
}
|
||||
usleep(SLEEP_INT);
|
||||
if (sock->refusedRetries % 1000 == 0) INFO(NCCL_ALL, "Call to connect returned %s, retrying", strerror(errno));
|
||||
return ncclSuccess;
|
||||
} else if (errno == ETIMEDOUT) {
|
||||
if (++sock->timedOutRetries == RETRY_TIMEDOUT_TIMES) {
|
||||
sock->state = ncclSocketStateError;
|
||||
WARN("socketStartConnect: exceeded timeouts (%d)", sock->timedOutRetries);
|
||||
return ncclRemoteError;
|
||||
}
|
||||
usleep(SLEEP_INT);
|
||||
return ncclSuccess;
|
||||
} else {
|
||||
char line[SOCKET_NAME_MAXLEN+1];
|
||||
sock->state = ncclSocketStateError;
|
||||
WARN("socketStartConnect: Connect to %s failed : %s", ncclSocketToString(&sock->addr, line), strerror(errno));
|
||||
return ncclSystemError;
|
||||
}
|
||||
}
|
||||
|
||||
static ncclResult_t socketPollConnect(struct ncclSocket* sock) {
|
||||
struct pollfd pfd;
|
||||
int timeout = 1, ret;
|
||||
socklen_t rlen = sizeof(int);
|
||||
|
||||
memset(&pfd, 0, sizeof(struct pollfd));
|
||||
pfd.fd = sock->fd;
|
||||
pfd.events = POLLOUT;
|
||||
SYSCHECK(ret = poll(&pfd, 1, timeout), "poll");
|
||||
if (ret == 0) return ncclSuccess;
|
||||
|
||||
/* check socket status */
|
||||
EQCHECK(ret == 1 && (pfd.revents & POLLOUT), 0);
|
||||
SYSCHECK(getsockopt(sock->fd, SOL_SOCKET, SO_ERROR, (void*)&ret, &rlen), "getsockopt");
|
||||
|
||||
if (ret == 0) {
|
||||
sock->state = ncclSocketStateConnected;
|
||||
} else if (ret == ECONNREFUSED) {
|
||||
if (++sock->refusedRetries == RETRY_REFUSED_TIMES) {
|
||||
sock->state = ncclSocketStateError;
|
||||
WARN("socketPollConnect: exceeded retries (%d)", sock->refusedRetries);
|
||||
return ncclRemoteError;
|
||||
}
|
||||
if (sock->refusedRetries % 1000 == 0) INFO(NCCL_ALL, "Call to connect returned %s, retrying", strerror(errno));
|
||||
usleep(SLEEP_INT);
|
||||
sock->state = ncclSocketStateConnecting;
|
||||
} else if (ret == ETIMEDOUT) {
|
||||
if (++sock->timedOutRetries == RETRY_TIMEDOUT_TIMES) {
|
||||
sock->state = ncclSocketStateError;
|
||||
WARN("socketPollConnect: exceeded timeouts (%d)", sock->timedOutRetries);
|
||||
return ncclRemoteError;
|
||||
}
|
||||
usleep(SLEEP_INT);
|
||||
sock->state = ncclSocketStateConnecting;
|
||||
} else if (ret != EINPROGRESS) {
|
||||
sock->state = ncclSocketStateError;
|
||||
return ncclSystemError;
|
||||
}
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketPollConnect(struct ncclSocket* sock) {
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketPollConnect: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
NCCLCHECK(socketPollConnect(sock));
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t socketFinalizeConnect(struct ncclSocket* sock) {
|
||||
int sent = 0;
|
||||
NCCLCHECK(socketProgress(NCCL_SOCKET_SEND, sock, &sock->magic, sizeof(sock->magic), &sent));
|
||||
if (sent == 0) return ncclSuccess;
|
||||
NCCLCHECK(socketWait(NCCL_SOCKET_SEND, sock, &sock->magic, sizeof(sock->magic), &sent));
|
||||
sent = 0;
|
||||
NCCLCHECK(socketWait(NCCL_SOCKET_SEND, sock, &sock->type, sizeof(sock->type), &sent));
|
||||
sock->state = ncclSocketStateReady;
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t socketProgressState(struct ncclSocket* sock) {
|
||||
if (sock->state == ncclSocketStateAccepting) {
|
||||
NCCLCHECK(socketTryAccept(sock));
|
||||
}
|
||||
if (sock->state == ncclSocketStateAccepted) {
|
||||
NCCLCHECK(socketFinalizeAccept(sock));
|
||||
}
|
||||
if (sock->state == ncclSocketStateConnecting) {
|
||||
NCCLCHECK(socketStartConnect(sock));
|
||||
}
|
||||
if (sock->state == ncclSocketStateConnectPolling) {
|
||||
NCCLCHECK(socketPollConnect(sock));
|
||||
}
|
||||
if (sock->state == ncclSocketStateConnected) {
|
||||
NCCLCHECK(socketFinalizeConnect(sock));
|
||||
}
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketReady(struct ncclSocket* sock, int *running) {
|
||||
if (sock == NULL) {
|
||||
*running = 0;
|
||||
return ncclSuccess;
|
||||
}
|
||||
if (sock->state == ncclSocketStateError || sock->state == ncclSocketStateClosed) {
|
||||
WARN("ncclSocketReady: unexpected socket state %d", sock->state);
|
||||
return ncclRemoteError;
|
||||
}
|
||||
*running = (sock->state == ncclSocketStateReady) ? 1 : 0;
|
||||
if (*running == 0) {
|
||||
NCCLCHECK(socketProgressState(sock));
|
||||
*running = (sock->state == ncclSocketStateReady) ? 1 : 0;
|
||||
}
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketConnect(struct ncclSocket* sock, int portReuse) {
|
||||
char line[SOCKET_NAME_MAXLEN+1];
|
||||
/* IPv4/IPv6 support */
|
||||
int family = sock->addr.sa.sa_family;
|
||||
if (family != AF_INET && family != AF_INET6) {
|
||||
WARN("Net : connecting to address %s with family %d is neither AF_INET(%d) nor AF_INET6(%d)",
|
||||
ncclSocketToString(&sock->addr, line), family, AF_INET, AF_INET6);
|
||||
const int one = 1;
|
||||
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketConnect: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
if (sock->fd == -1) {
|
||||
WARN("ncclSocketConnect: file descriptor is -1");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
|
||||
if (sock->state != ncclSocketStateInitialized) {
|
||||
WARN("ncclSocketConnect: wrong socket state %d", sock->state);
|
||||
if (sock->state == ncclSocketStateError) return ncclRemoteError;
|
||||
return ncclInternalError;
|
||||
}
|
||||
int salen = (family == AF_INET) ? sizeof(struct sockaddr_in) : sizeof(struct sockaddr_in6);
|
||||
int flags;
|
||||
TRACE(NCCL_INIT|NCCL_NET,"Connecting to socket %s", ncclSocketToString(&sock->addr, line));
|
||||
|
||||
/* Connect to a hostname / port */
|
||||
int fd = socket(family, SOCK_STREAM, 0);
|
||||
if (fd == -1) {
|
||||
WARN("Net : Socket creation failed : %s", strerror(errno));
|
||||
return ncclSystemError;
|
||||
}
|
||||
|
||||
const int one = 1;
|
||||
SYSCHECK(setsockopt(fd, IPPROTO_TCP, TCP_NODELAY, (char*)&one, sizeof(int)), "setsockopt");
|
||||
|
||||
/* The socket is set non-blocking for OS level, but asyncFlag is used to control
|
||||
* blocking and non-blocking behavior in user level. */
|
||||
EQCHECK(flags = fcntl(fd, F_GETFL), -1);
|
||||
SYSCHECK(fcntl(fd, F_SETFL, flags | O_NONBLOCK), "fcntl");
|
||||
|
||||
/* const int bufsize = 128*1024;
|
||||
SYSCHECK(setsockopt(fd, SOL_SOCKET, SO_SNDBUF, (char*)&bufsize, sizeof(int)), "setsockopt");
|
||||
SYSCHECK(setsockopt(fd, SOL_SOCKET, SO_RCVBUF, (char*)&bufsize, sizeof(int)), "setsockopt");*/
|
||||
SYSCHECK(setsockopt(sock->fd, IPPROTO_TCP, TCP_NODELAY, (char*)&one, sizeof(int)), "setsockopt");
|
||||
|
||||
if (portReuse) {
|
||||
// pre-define ports according to tid, to avoid extra lock for race condition
|
||||
int family = sock->addr.sa.sa_family;
|
||||
if (family != AF_INET && family != AF_INET6) {
|
||||
WARN("Net : connecting to address %s with family %d is neither AF_INET(%d) nor AF_INET6(%d)",
|
||||
ncclSocketToString(&sock->addr, line), family, AF_INET, AF_INET6);
|
||||
return ncclInternalError;
|
||||
}
|
||||
int salen = (family == AF_INET) ? sizeof(struct sockaddr_in) : sizeof(struct sockaddr_in6); // pre-define ports according to tid, to avoid extra lock for race condition
|
||||
|
||||
if (clientPortPool.size() == 0) {
|
||||
for (int tid = syscall(SYS_gettid), i = 1; i < 5; i++) {
|
||||
clientPortPool.push_back(std::make_pair(60000 + i * 1000 + tid % 1000, std::unordered_set<std::string>()));
|
||||
@@ -448,161 +636,227 @@ ncclResult_t ncclSocketConnect(struct ncclSocket* sock, int portReuse) {
|
||||
// bind the port in fd for connect system call
|
||||
if (reused_port != -1) {
|
||||
int opt = 1;
|
||||
SYSCHECK(setsockopt(fd, SOL_SOCKET, SO_REUSEADDR | SO_REUSEPORT, &opt, sizeof(opt)), "setsockopt");
|
||||
SYSCHECK(setsockopt(sock->fd, SOL_SOCKET, SO_REUSEADDR | SO_REUSEPORT, &opt, sizeof(opt)), "setsockopt");
|
||||
struct sockaddr_in sin;
|
||||
sin.sin_family = family;
|
||||
sin.sin_addr.s_addr = htonl(INADDR_ANY);
|
||||
sin.sin_port = htons(reused_port);
|
||||
SYSCHECK(bind(fd, (struct sockaddr *)&sin, salen), "bind_client_port");
|
||||
SYSCHECK(bind(sock->fd, (struct sockaddr *)&sin, salen), "bind_client_port");
|
||||
}
|
||||
}
|
||||
|
||||
TRACE(NCCL_INIT|NCCL_NET,"Connecting to socket %s", ncclSocketToString(&sock->addr, line));
|
||||
sock->state = ncclSocketStateConnecting;
|
||||
do {
|
||||
NCCLCHECK(socketProgressState(sock));
|
||||
} while (sock->asyncFlag == 0 &&
|
||||
(sock->abortFlag == NULL || *sock->abortFlag == 0) &&
|
||||
(sock->state == ncclSocketStateConnecting ||
|
||||
sock->state == ncclSocketStateConnectPolling ||
|
||||
sock->state == ncclSocketStateConnected));
|
||||
|
||||
int ret;
|
||||
int timedout_retries = 0;
|
||||
int refused_retries = 0;
|
||||
retry:
|
||||
/* blocking/non-blocking connect() is determined by asyncFlag. */
|
||||
ret = connect(fd, &sock->addr.sa, salen);
|
||||
if (sock->abortFlag && *sock->abortFlag != 0) return ncclInternalError;
|
||||
|
||||
if (!sock->asyncFlag) {
|
||||
/* blocking socket, need retry if connect fails. */
|
||||
if (errno == EINPROGRESS || errno == EAGAIN || errno == EALREADY ||
|
||||
(errno == ECONNREFUSED && ++refused_retries < RETRY_REFUSED_TIMES) ||
|
||||
(errno == ETIMEDOUT && ++timedout_retries < RETRY_TIMEDOUT_TIMES)) {
|
||||
/* check abortFlag as long as we have chance to retry. */
|
||||
if (sock->abortFlag && *sock->abortFlag != 0) return ncclInternalError;
|
||||
if (errno == ECONNREFUSED && refused_retries % 1000 == 0) INFO(NCCL_ALL, "Call to connect returned %s, retrying", strerror(errno));
|
||||
usleep(SLEEP_INT);
|
||||
goto retry;
|
||||
}
|
||||
|
||||
/* If connect() fails with errno == EAGAIN/EINPROGRESS/ETIMEDOUT, we may want to try connect again.
|
||||
* However, it can return EISCONN instead of success which indicates connection is built up in
|
||||
* background already. No need to call connect() again. */
|
||||
if (ret == 0 || errno == EISCONN) {
|
||||
sock->fd = fd;
|
||||
switch (sock->state) {
|
||||
case ncclSocketStateConnecting:
|
||||
case ncclSocketStateConnectPolling:
|
||||
case ncclSocketStateConnected:
|
||||
case ncclSocketStateReady:
|
||||
return ncclSuccess;
|
||||
}
|
||||
} else {
|
||||
sock->fd = fd;
|
||||
return ncclSuccess;
|
||||
case ncclSocketStateError:
|
||||
return ncclSystemError;
|
||||
default:
|
||||
WARN("ncclSocketConnect: wrong socket state %d", sock->state);
|
||||
return ncclInternalError;
|
||||
}
|
||||
|
||||
WARN("Net : Connect to %s failed : %s", ncclSocketToString(&sock->addr, line), strerror(errno));
|
||||
return ncclRemoteError;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketAccept(struct ncclSocket* sock, struct ncclSocket* listenSocket) {
|
||||
socklen_t socklen = sizeof(union ncclSocketAddress);
|
||||
struct pollfd pollfd;
|
||||
int tmpFd = sock->fd = -1;
|
||||
int pollret;
|
||||
ncclResult_t ncclSocketAccept(struct ncclSocket* sock, struct ncclSocket* listenSock) {
|
||||
ncclResult_t ret = ncclSuccess;
|
||||
|
||||
pollfd.fd = listenSocket->fd;
|
||||
pollfd.events = POLLIN;
|
||||
retry:
|
||||
if ((pollret = poll(&pollfd, 1, listenSocket->asyncFlag ? 0 : 100)) < 0) {
|
||||
return ncclSystemError;
|
||||
} else {
|
||||
tmpFd = accept(listenSocket->fd, &sock->addr.sa, &socklen);
|
||||
if (listenSock == NULL || sock == NULL) {
|
||||
WARN("ncclSocketAccept: pass NULL socket");
|
||||
ret = ncclInvalidArgument;
|
||||
goto exit;
|
||||
}
|
||||
if (listenSock->state != ncclSocketStateReady) {
|
||||
WARN("ncclSocketAccept: wrong socket state %d", listenSock->state);
|
||||
if (listenSock->state == ncclSocketStateError)
|
||||
ret = ncclSystemError;
|
||||
else
|
||||
ret = ncclInternalError;
|
||||
goto exit;
|
||||
}
|
||||
|
||||
if (!listenSocket->asyncFlag) {
|
||||
/* blocking socket, if tmpFd is still -1, we need to retry */
|
||||
if (tmpFd == -1 && (errno == EAGAIN || errno == EWOULDBLOCK)) {
|
||||
if (listenSocket->abortFlag && *listenSocket->abortFlag != 0) return ncclInternalError;
|
||||
goto retry;
|
||||
}
|
||||
EQCHECK(tmpFd, -1);
|
||||
if (sock->acceptFd == -1) {
|
||||
memcpy(sock, listenSock, sizeof(struct ncclSocket));
|
||||
sock->acceptFd = listenSock->fd;
|
||||
sock->state = ncclSocketStateAccepting;
|
||||
}
|
||||
|
||||
sock->fd = tmpFd;
|
||||
return ncclSuccess;
|
||||
do {
|
||||
NCCLCHECKGOTO(socketProgressState(sock), ret, exit);
|
||||
} while (sock->asyncFlag == 0 &&
|
||||
(sock->abortFlag == NULL || *sock->abortFlag == 0) &&
|
||||
(sock->state == ncclSocketStateAccepting ||
|
||||
sock->state == ncclSocketStateAccepted));
|
||||
|
||||
if (sock->abortFlag && *sock->abortFlag != 0) return ncclInternalError;
|
||||
|
||||
switch (sock->state) {
|
||||
case ncclSocketStateAccepting:
|
||||
case ncclSocketStateAccepted:
|
||||
case ncclSocketStateReady:
|
||||
ret = ncclSuccess;
|
||||
break;
|
||||
case ncclSocketStateError:
|
||||
ret = ncclSystemError;
|
||||
break;
|
||||
default:
|
||||
WARN("ncclSocketAccept: wrong socket state %d", sock->state);
|
||||
ret = ncclInternalError;
|
||||
break;
|
||||
}
|
||||
|
||||
exit:
|
||||
return ret;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketInit(struct ncclSocket* sock, union ncclSocketAddress* addr, volatile uint32_t* abortFlag, int asyncFlag) {
|
||||
if (sock == NULL)
|
||||
return ncclSuccess;
|
||||
ncclResult_t ncclSocketInit(struct ncclSocket* sock, union ncclSocketAddress* addr, uint64_t magic, enum ncclSocketType type, volatile uint32_t* abortFlag, int asyncFlag) {
|
||||
ncclResult_t ret = ncclSuccess;
|
||||
|
||||
if (sock == NULL) goto exit;
|
||||
sock->timedOutRetries = 0;
|
||||
sock->refusedRetries = 0;
|
||||
sock->abortFlag = abortFlag;
|
||||
sock->asyncFlag = asyncFlag;
|
||||
sock->state = ncclSocketStateInitialized;
|
||||
sock->magic = magic;
|
||||
sock->type = type;
|
||||
sock->fd = -1;
|
||||
sock->acceptFd = -1;
|
||||
|
||||
if (addr) {
|
||||
/* IPv4/IPv6 support */
|
||||
int family;
|
||||
memcpy(&sock->addr, addr, sizeof(union ncclSocketAddress));
|
||||
family = sock->addr.sa.sa_family;
|
||||
if (family != AF_INET && family != AF_INET6) {
|
||||
char line[SOCKET_NAME_MAXLEN+1];
|
||||
WARN("ncclSocketInit: connecting to address %s with family %d is neither AF_INET(%d) nor AF_INET6(%d)",
|
||||
ncclSocketToString(&sock->addr, line), family, AF_INET, AF_INET6);
|
||||
ret = ncclInternalError;
|
||||
goto fail;
|
||||
}
|
||||
sock->salen = (family == AF_INET) ? sizeof(struct sockaddr_in) : sizeof(struct sockaddr_in6);
|
||||
|
||||
/* Connect to a hostname / port */
|
||||
sock->fd = socket(family, SOCK_STREAM, 0);
|
||||
if (sock->fd == -1) {
|
||||
WARN("ncclSocketInit: Socket creation failed : %s", strerror(errno));
|
||||
ret = ncclSystemError;
|
||||
goto fail;
|
||||
}
|
||||
} else {
|
||||
memset(&sock->addr, 0, sizeof(union ncclSocketAddress));
|
||||
}
|
||||
sock->abortFlag = abortFlag;
|
||||
sock->asyncFlag = asyncFlag;
|
||||
sock->state = ncclSocketStateNum;
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t ncclSocketProgressOpt(int op, struct ncclSocket* sock, void* ptr, int size, int* offset, int block, int* closed) {
|
||||
int bytes = 0;
|
||||
*closed = 0;
|
||||
char* data = (char*)ptr;
|
||||
char line[SOCKET_NAME_MAXLEN+1];
|
||||
do {
|
||||
if (op == NCCL_SOCKET_RECV) bytes = recv(sock->fd, data+(*offset), size-(*offset), block ? 0 : MSG_DONTWAIT);
|
||||
if (op == NCCL_SOCKET_SEND) bytes = send(sock->fd, data+(*offset), size-(*offset), block ? MSG_NOSIGNAL : MSG_DONTWAIT | MSG_NOSIGNAL);
|
||||
if (op == NCCL_SOCKET_RECV && bytes == 0) {
|
||||
*closed = 1;
|
||||
return ncclSuccess;
|
||||
}
|
||||
if (bytes == -1) {
|
||||
if (errno != EINTR && errno != EWOULDBLOCK && errno != EAGAIN) {
|
||||
WARN("Net : Call to recv from %s failed : %s", ncclSocketToString(&sock->addr, line), strerror(errno));
|
||||
return ncclRemoteError;
|
||||
} else {
|
||||
bytes = 0;
|
||||
}
|
||||
}
|
||||
(*offset) += bytes;
|
||||
if (sock->abortFlag && *sock->abortFlag != 0) {
|
||||
INFO(NCCL_NET, "Socket progress: abort called");
|
||||
return ncclInternalError;
|
||||
}
|
||||
} while (bytes > 0 && (*offset) < size);
|
||||
return ncclSuccess;
|
||||
/* Set socket as non-blocking if async or if we need to be able to abort */
|
||||
if ((sock->asyncFlag || sock->abortFlag) && sock->fd >= 0) {
|
||||
int flags;
|
||||
EQCHECKGOTO(flags = fcntl(sock->fd, F_GETFL), -1, ret, fail);
|
||||
SYSCHECKGOTO(fcntl(sock->fd, F_SETFL, flags | O_NONBLOCK), ret, fail);
|
||||
}
|
||||
|
||||
exit:
|
||||
return ret;
|
||||
fail:
|
||||
goto exit;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketProgress(int op, struct ncclSocket* sock, void* ptr, int size, int* offset) {
|
||||
int closed;
|
||||
NCCLCHECK(ncclSocketProgressOpt(op, sock, ptr, size, offset, 0, &closed));
|
||||
if (closed) {
|
||||
char line[SOCKET_NAME_MAXLEN+1];
|
||||
WARN("Net : Connection closed by remote peer %s", ncclSocketToString(&sock->addr, line, 0));
|
||||
return ncclRemoteError;
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketProgress: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
NCCLCHECK(socketProgress(op, sock, ptr, size, offset));
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketWait(int op, struct ncclSocket* sock, void* ptr, int size, int* offset) {
|
||||
while (*offset < size)
|
||||
NCCLCHECK(ncclSocketProgress(op, sock, ptr, size, offset));
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketWait: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
NCCLCHECK(socketWait(op, sock, ptr, size, offset));
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketSend(struct ncclSocket* sock, void* ptr, int size) {
|
||||
int offset = 0;
|
||||
NCCLCHECK(ncclSocketWait(NCCL_SOCKET_SEND, sock, ptr, size, &offset));
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketSend: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
if (sock->state != ncclSocketStateReady) {
|
||||
WARN("ncclSocketSend: socket state (%d) is not ready", sock->state);
|
||||
return ncclInternalError;
|
||||
}
|
||||
NCCLCHECK(socketWait(NCCL_SOCKET_SEND, sock, ptr, size, &offset));
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketRecv(struct ncclSocket* sock, void* ptr, int size) {
|
||||
int offset = 0;
|
||||
NCCLCHECK(ncclSocketWait(NCCL_SOCKET_RECV, sock, ptr, size, &offset));
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketRecv: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
if (sock->state != ncclSocketStateReady) {
|
||||
WARN("ncclSocketRecv: socket state (%d) is not ready", sock->state);
|
||||
return ncclInternalError;
|
||||
}
|
||||
NCCLCHECK(socketWait(NCCL_SOCKET_RECV, sock, ptr, size, &offset));
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
// Receive or detect connection closed
|
||||
ncclResult_t ncclSocketTryRecv(struct ncclSocket* sock, void* ptr, int size, int* closed) {
|
||||
int offset = 0;
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketTryRecv: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
*closed = 0;
|
||||
while (offset < size) {
|
||||
NCCLCHECK(ncclSocketProgressOpt(NCCL_SOCKET_RECV, sock, ptr, size, &offset, 0, closed));
|
||||
NCCLCHECK(socketProgressOpt(NCCL_SOCKET_RECV, sock, ptr, size, &offset, 0, closed));
|
||||
if (*closed) return ncclSuccess;
|
||||
}
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketClose(struct ncclSocket* sock) {
|
||||
if (sock != NULL) {
|
||||
if (sock->fd >= 0) close(sock->fd);
|
||||
sock->state = ncclSocketStateClosed;
|
||||
sock->fd = -1;
|
||||
}
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketGetFd(struct ncclSocket* sock, int* fd) {
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketGetFd: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
if (fd) *fd = sock->fd;
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t ncclSocketSetFd(int fd, struct ncclSocket* sock) {
|
||||
if (sock == NULL) {
|
||||
WARN("ncclSocketGetFd: pass NULL socket");
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
sock->fd = fd;
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
Criar uma nova questão referindo esta
Bloquear um utilizador