From 9b22eb46047dcac0004097f25e73f3f113df8a2f Mon Sep 17 00:00:00 2001 From: Fred Nicolson Date: Fri, 16 Nov 2018 16:45:18 +0000 Subject: [PATCH] SocketSelector bug fixes. Don't disconnect socket on read/write error. --- src/SSLSocket.cpp | 9 +++------ src/SocketSelector.cpp | 30 ++++++++++++++++++++++-------- src/TcpSocket.cpp | 6 ++---- src/WebFrame.cpp | 4 +--- 4 files changed, 28 insertions(+), 21 deletions(-) diff --git a/src/SSLSocket.cpp b/src/SSLSocket.cpp index eea11e6..560dd4b 100644 --- a/src/SSLSocket.cpp +++ b/src/SSLSocket.cpp @@ -56,8 +56,7 @@ namespace fr } else if(response < 0) { - close_socket(); - return Socket::Status::Disconnected; + return Socket::Status::Error; } } @@ -77,8 +76,7 @@ namespace fr return Socket::Status::WouldBlock; } - close_socket(); - return Socket::Status::Disconnected; + return Socket::Status::Error; } } else @@ -97,8 +95,7 @@ namespace fr continue; //try again, interrupted before anything could be received } - close_socket(); - return Socket::Status::Disconnected; + return Socket::Status::Error; } break; } while(true); diff --git a/src/SocketSelector.cpp b/src/SocketSelector.cpp index bc764c8..6cab833 100644 --- a/src/SocketSelector.cpp +++ b/src/SocketSelector.cpp @@ -10,8 +10,9 @@ namespace fr { #ifndef _WIN32 + SocketSelector::SocketSelector() - : epoll_fd(-1) + : epoll_fd(-1) { epoll_fd = epoll_create1(O_CLOEXEC); if(epoll_fd < 0) @@ -32,7 +33,8 @@ namespace fr throw std::logic_error("Can't add disconnected socket"); } - auto add_iter = added_sockets.emplace(socket->get_socket_descriptor(), Opaque(socket->get_socket_descriptor(), socket, opaque)); + auto add_iter = added_sockets.emplace(socket->get_socket_descriptor(), + Opaque(socket->get_socket_descriptor(), socket, opaque)); if(!add_iter.second) { throw std::logic_error("Can't add duplicate socket: " + std::to_string(socket->get_socket_descriptor())); @@ -59,17 +61,25 @@ namespace fr throw std::runtime_error("epoll_wait returned: " + std::to_string(errno)); } - std::vector, void*>> ret; + std::vector, void *>> ret; for(int a = 0; a < event_count; ++a) { - auto *opaque = static_cast(events[a].data.ptr); + auto *opaque = static_cast(events[a].data.ptr); ret.emplace_back(opaque->socket, opaque->opaque); if(events[a].events & EPOLLERR || events[a].events & EPOLLHUP || events[a].events & EPOLLRDHUP) { auto iter = added_sockets.find(opaque->descriptor); if(iter != added_sockets.end()) { + epoll_event event = {0}; + auto remove_ret = epoll_ctl(epoll_fd, EPOLL_CTL_DEL, opaque->descriptor, &event); added_sockets.erase(opaque->descriptor); + if(remove_ret < 0) + { + throw std::runtime_error( + "Failed to remove socket: " + std::to_string(opaque->descriptor) + ". Errno: " + + std::to_string(errno)); + } } } } @@ -83,19 +93,23 @@ namespace fr { throw std::runtime_error("Can't remove disconnected socket"); } + auto iter = added_sockets.find(descriptor); - if(iter != added_sockets.end()) + if(iter == added_sockets.end()) { - added_sockets.erase(iter); + return; } + added_sockets.erase(iter); epoll_event event = {0}; if(epoll_ctl(epoll_fd, EPOLL_CTL_DEL, descriptor, &event) < 0) { - throw std::runtime_error("Failed to remove socket: " + std::to_string(errno)); + throw std::runtime_error( + "Failed to remove socket: " + std::to_string(descriptor) + ". Errno: " + std::to_string(errno)); } - delete static_cast(event.data.ptr); + delete static_cast(event.data.ptr); } + #endif } diff --git a/src/TcpSocket.cpp b/src/TcpSocket.cpp index a09c4e6..60e4340 100644 --- a/src/TcpSocket.cpp +++ b/src/TcpSocket.cpp @@ -35,8 +35,7 @@ namespace fr } else if(errno != EWOULDBLOCK && errno != EAGAIN) //Don't exit if the socket just couldn't block { - close_socket(); - return Socket::Status::Disconnected; + return Socket::Status::Error; } } return Socket::Status::Success; @@ -68,8 +67,7 @@ namespace fr continue; //try again, interrupted before anything could be received } - close_socket(); - return Socket::Status::Disconnected; + return Socket::Status::Error; } break; } while(true); diff --git a/src/WebFrame.cpp b/src/WebFrame.cpp index 4862420..9ae26af 100644 --- a/src/WebFrame.cpp +++ b/src/WebFrame.cpp @@ -104,8 +104,7 @@ namespace fr auto mask = static_cast((first_2bytes >> 7) & 0x1); if(mask == socket_->is_client()) { - socket->disconnect(); - return fr::Socket::Disconnected; + return fr::Socket::Error; } @@ -130,7 +129,6 @@ namespace fr //Verify that payload length isn't too large if(socket->get_max_receive_size() && payload_length > socket->get_max_receive_size()) { - socket->disconnect(); //We're forced to disconnect, otherwise we'll be out of sync with the server return Socket::MaxPacketSizeExceeded; }