diff --git a/README.md b/README.md index e92d650..933c633 100644 --- a/README.md +++ b/README.md @@ -13,13 +13,13 @@ Copy the the files from `src` to your project (`socket.cpp` and `socket.h`), and //Create a socket Socket sock = Socket(); - sock.connect("127.0.0.1", 8000); + sock.Connect("127.0.0.1", 8000); //Send data std::string data_to_send = "Hello world :S"; - sock.send(data_to_send.data(),data_to_send.size()); + sock.Send(data_to_send.data(),data_to_send.size()); - sock.close(); + sock.Close(); return 0; } ``` diff --git a/Tests/socket_tests.cpp b/Tests/socket_tests.cpp index 7de163b..e89ec21 100644 --- a/Tests/socket_tests.cpp +++ b/Tests/socket_tests.cpp @@ -1,6 +1,8 @@ #include #include "../src/socket.h" +#include "../src/exceptions.h" + constexpr auto TEST_IP_ADDR = "127.0.0.1"; #ifdef _WIN32 @@ -21,26 +23,26 @@ TEST(create_tcp_sock, BasicAssertions) { // Create a listening socket Socket sock1 = Socket(); - sock1.bind(TEST_PORT_1); - sock1.listen(1); + sock1.Bind(TEST_PORT_1); + sock1.Listen(1); // Connect to listening socket Socket sock2 = Socket(); - sock2.connect(TEST_IP_ADDR, TEST_PORT_1); + sock2.Connect(TEST_IP_ADDR, TEST_PORT_1); // Accept conneciton - auto sock3 = sock1.accept(); + auto sock3 = sock1.Accept(); std::string demo_str = "hello"; // Try to send data - int snd = sock3.send(demo_str.data(), demo_str.size()); + int snd = sock3.Send(demo_str.data(), demo_str.size()); // Try to receive data std::vector recv_buffer(demo_str.size()); - sock2.recv(recv_buffer.data(), demo_str.size()); + sock2.Receive(recv_buffer.data(), demo_str.size()); recv_buffer.resize(demo_str.size()); std::string buffer_string; buffer_string.assign(recv_buffer.begin(), recv_buffer.end()); // Cleanup - sock3.close(); - sock2.close(); - sock1.close(); + sock3.Close(); + sock2.Close(); + sock1.Close(); // Expect equality. EXPECT_STREQ(demo_str.data(), buffer_string.data()); EXPECT_EQ(buffer_string.size(), demo_str.size()); @@ -50,25 +52,25 @@ TEST(try_bind_twice, HasCertainMessage) { // Bind a socket Socket sock1 = Socket(); - sock1.bind(TEST_PORT_2); + sock1.Bind(TEST_PORT_2); // Bind socket with same port Socket sock2 = Socket(); EXPECT_THROW( { - sock2.bind(TEST_PORT_2); + sock2.Bind(TEST_PORT_2); }, - std::runtime_error); + SocketBindException); - sock2.close(); - sock1.close(); + sock2.Close(); + sock1.Close(); } TEST(try_connect_to_closed_peer, BasicAssertions) { // Create a listening socket Socket sock1 = Socket(); - auto connect_result = sock1.connect(TEST_IP_ADDR, TEST_PORT_3); - sock1.close(); + auto connect_result = sock1.Connect(TEST_IP_ADDR, TEST_PORT_3); + sock1.Close(); // Expect connection error. EXPECT_EQ(connect_result, -1); } @@ -77,26 +79,26 @@ TEST(partial_recv, BasicAssertions) { // Create a listening socket Socket sock1 = Socket(); - sock1.bind(TEST_PORT_4); - sock1.listen(1); + sock1.Bind(TEST_PORT_4); + sock1.Listen(1); // Connect to listening socket Socket sock2 = Socket(); - sock2.connect(TEST_IP_ADDR, TEST_PORT_4); + sock2.Connect(TEST_IP_ADDR, TEST_PORT_4); // Accept conneciton - auto sock3 = sock1.accept(); + auto sock3 = sock1.Accept(); std::string demo_str = "abC1deFg%H"; // Try to send data - int snd = sock3.send(demo_str.c_str(), demo_str.size()); + int snd = sock3.Send(demo_str.c_str(), demo_str.size()); // Try to receive data std::vector recv_buffer(demo_str.size()); size_t total_read = 0; - total_read += sock2.recv(recv_buffer.data(), 5); - total_read += sock2.recv(&recv_buffer.at(total_read), demo_str.size()); + total_read += sock2.Receive(recv_buffer.data(), 5); + total_read += sock2.Receive(&recv_buffer.at(total_read), demo_str.size()); recv_buffer.resize(demo_str.size()); // Cleanup - sock3.close(); - sock2.close(); - sock1.close(); + sock3.Close(); + sock2.Close(); + sock1.Close(); // Expect equality. std::string buffer_string; buffer_string.assign(recv_buffer.begin(), recv_buffer.end()); diff --git a/src/exceptions.h b/src/exceptions.h new file mode 100644 index 0000000..55aad33 --- /dev/null +++ b/src/exceptions.h @@ -0,0 +1,30 @@ +#pragma once + +#include +#include + +class SocketException: public std::exception { +public: + SocketException(std::string&& value) : m_value(std::move(value)) { + } + + const char* what() const noexcept override { + return m_value.c_str(); + } +private: + std::string m_value; +}; + +#define DECLARE_EXCEPTION(exception_name, exception_value) \ + class exception_name final: public SocketException { \ + public: \ + exception_name(std::string&& value): SocketException(std::string(exception_value) + ": " + value) { \ + \ + } \ + }; + + +DECLARE_EXCEPTION(SocketCreateException, "Error creating socket") +DECLARE_EXCEPTION(SocketBindException, "Socket bind error") +DECLARE_EXCEPTION(SocketListenException, "Socket listen error") +DECLARE_EXCEPTION(SocketAcceptException, "Socket accept error") \ No newline at end of file diff --git a/src/socket.cpp b/src/socket.cpp index f554b21..aacd624 100644 --- a/src/socket.cpp +++ b/src/socket.cpp @@ -1,103 +1,106 @@ #include "socket.h" #include -#include -#include +#include "exceptions.h" constexpr auto SOCKET_ERROR_CODE = -1; -Socket::Socket() : -#ifdef _WIN32 - m_socket(INVALID_SOCKET) -#else - m_socket(0) -#endif +Socket::Socket() { #ifdef _WIN32 if (WSAStartup(MAKEWORD(2, 2), &m_wsaData) != 0) { - throw std::runtime_error(std::strerror(errno)); + throw SocketCreateException(std::strerror(errno)); } #endif - m_socket = socket(AF_INET, SOCK_STREAM, 0); + m_socket = socket(AF_INET, SOCK_STREAM, IPPROTO_TCP); if (m_socket == SOCKET_ERROR_CODE) { #ifdef _WIN32 WSACleanup(); #endif - throw std::runtime_error(std::strerror(errno)); + throw SocketCreateException(std::strerror(errno)); } } -Socket::Socket(SOCKET socket) : m_socket(std::move(socket)) +Socket::Socket(SOCKET socket) : m_socket(socket) {} + +Socket::Socket(Socket&& other) noexcept : m_socket(other.m_socket) { + other.m_socket = SOCKET_ERROR_CODE; + +} + +Socket& Socket::operator=(Socket&& other) noexcept { + m_socket = other.m_socket; + other.m_socket = SOCKET_ERROR_CODE; #ifdef _WIN32 - if (WSAStartup(MAKEWORD(2, 2), &m_wsaData) != 0) - { - throw std::runtime_error(std::strerror(errno)); - } + m_wsaData = other.m_wsaData; + ZeroMemory(&other.m_wsaData, sizeof(other.m_wsaData)); #endif + return *this; } -int Socket::connect(std::string address, int port) +int Socket::Connect(std::string address, int port) { - sockaddr_in addr; + sockaddr_in addr = {0}; addr.sin_family = AF_INET; addr.sin_port = htons(port); addr.sin_addr.s_addr = inet_addr(address.c_str()); - return ::connect(m_socket, (sockaddr *)&addr, sizeof(addr)); + return ::connect(m_socket, reinterpret_cast(&addr), sizeof(addr)); } -void Socket::bind(int port) +void Socket::Bind(int port) { - sockaddr_in addr; + sockaddr_in addr = {0}; addr.sin_family = AF_INET; addr.sin_port = htons(port); addr.sin_addr.s_addr = htonl(INADDR_ANY); - if (::bind(m_socket, (sockaddr *)&addr, sizeof(addr)) == SOCKET_ERROR_CODE) + if (::bind(m_socket, reinterpret_cast(&addr), sizeof(addr)) == SOCKET_ERROR_CODE) { - throw std::runtime_error(std::strerror(errno)); + throw SocketBindException(std::strerror(errno)); } } -void Socket::listen(int backlog) +void Socket::Listen(int backlog) { if (::listen(m_socket, backlog) == SOCKET_ERROR_CODE) { - throw std::runtime_error(std::strerror(errno)); + throw SocketListenException(std::strerror(errno)); } } -Socket Socket::accept() +Socket Socket::Accept() { sockaddr_in client_addr; socklen_t client_addr_size = sizeof(client_addr); SOCKET client_socket = ::accept(m_socket, (sockaddr *)&client_addr, &client_addr_size); if (client_socket == SOCKET_ERROR_CODE) { - throw std::runtime_error(std::strerror(errno)); + throw SocketAcceptException(std::strerror(errno)); } return Socket(std::move(client_socket)); } -int Socket::send(const void *data, size_t length) +int Socket::Send(const void *data, size_t length) { return ::send(m_socket, reinterpret_cast(data), length, 0); } -int Socket::recv(void *buffer, size_t length) +int Socket::Receive(void *buffer, size_t length) { return ::recv(m_socket, reinterpret_cast(buffer), length, 0); } -void Socket::close() +void Socket::Close() noexcept { if (m_socket != SOCKET_ERROR_CODE) { int result = 0; #ifdef _WIN32 result = ::closesocket(m_socket); + WSACleanup(); #else result = ::close(m_socket); #endif @@ -107,8 +110,5 @@ void Socket::close() Socket::~Socket() { - close(); -#ifdef _WIN32 - WSACleanup(); -#endif + Close(); } \ No newline at end of file diff --git a/src/socket.h b/src/socket.h index a673f61..c9c6f43 100644 --- a/src/socket.h +++ b/src/socket.h @@ -16,48 +16,49 @@ typedef int SOCKET; class Socket { public: - - /// @brief Default constructor of `Socket`. Socket(); - - ~Socket(); + explicit Socket(SOCKET socket); + Socket(const Socket&) = delete; + Socket(Socket&& other) noexcept; + Socket& operator=(Socket&& other) noexcept; + Socket& operator=(const Socket&) = delete; + + // Destructor, Close the socket + virtual ~Socket(); /// @brief Connect to endpoint. /// @param address (std::string). /// @param port (int). /// @return operation result (int) - `0` if connected, `-1` if there was an /// error. - int connect(std::string address, int port); + int Connect(std::string address, int port); /// @brief Bind operation. /// @param port - port to bind. - void bind(int port); + void Bind(int port); /// @brief Allows to listen for connections. /// @param backlog (int) - number of pending connections the socket will hold. - void listen(int backlog); + void Listen(int backlog); /// @brief Accept new connections. /// @return A new `Socket` class. - Socket accept(); + Socket Accept(); /// @brief Close the socket. - void close(); + void Close() noexcept; /// @brief Send data to endpoint. /// @param data (const void *) - pointer to data buffer. /// @param length (size_t) - length of buffer to send. /// @return (int) The number of bytes accepted by the kernel for sending. - int send(const void *data, size_t length); + int Send(const void *data, size_t length); /// @brief Receive data from endpoint. /// @param buffer (void *) - buffer to write received data. /// @param length (size_t) - length to read from socket. /// @return (int) The number of received bytes. - int recv(void *buffer, size_t length); - - protected: - Socket(SOCKET socket); + int Receive(void *buffer, size_t length); private: int m_socket;