Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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;
}
```
54 changes: 28 additions & 26 deletions Tests/socket_tests.cpp
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
#include <gtest/gtest.h>

#include "../src/socket.h"
#include "../src/exceptions.h"


constexpr auto TEST_IP_ADDR = "127.0.0.1";
#ifdef _WIN32
Expand All @@ -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<char> 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());
Expand All @@ -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);
}
Expand All @@ -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<uint8_t> 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());
Expand Down
30 changes: 30 additions & 0 deletions src/exceptions.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
#pragma once

#include <exception>
#include <string>

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")
68 changes: 34 additions & 34 deletions src/socket.cpp
Original file line number Diff line number Diff line change
@@ -1,103 +1,106 @@
#include "socket.h"

#include <cstring>
#include <stdexcept>
#include <system_error>
#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<sockaddr *>(&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<sockaddr *>(&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<const char *>(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<char *>(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
Expand All @@ -107,8 +110,5 @@ void Socket::close()

Socket::~Socket()
{
close();
#ifdef _WIN32
WSACleanup();
#endif
Close();
}
29 changes: 15 additions & 14 deletions src/socket.h
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
Loading