From 794482cd918418be83ea8ac5341b5b007754a23e Mon Sep 17 00:00:00 2001 From: Dylan Hollemaert Date: Thu, 3 Sep 2026 19:54:59 +0200 Subject: [PATCH 1/3] add: Networking system * Supports TCP and UDP Signed-off-by: Dylan Hollemaert --- Source/Engine/CMakeLists.txt | 1 + Source/Engine/Network/CMakeLists.txt | 27 ++ .../Network/Common/Private/NetworkAddress.cpp | 21 + .../Network/Common/Public/NetworkAddress.cppm | 35 ++ .../Network/Common/Public/NetworkAddress.inl | 15 + .../Network/Common/Public/NetworkError.cppm | 41 ++ .../Manager/Private/NetworkManager.cpp | 74 +++ .../Manager/Public/NetworkManager.cppm | 54 +++ .../Network/Platform/Linux/LinuxSocket.cpp | 298 ++++++++++++ .../Platform/Linux/LinuxSocket_Internal.cpp | 228 +++++++++ .../Platform/Linux/LinuxSocket_Internal.hpp | 66 +++ .../Network/Platform/Public/Socket.cppm | 68 +++ .../Windows/WindowsScoket_Internal.hpp | 68 +++ .../Platform/Windows/WindowsSocket.cpp | 328 +++++++++++++ .../Windows/WindowsSocket_Internal.cpp | 246 ++++++++++ .../Protocol/Private/PacketDispatcher.cpp | 70 +++ .../Protocol/Public/PacketDispatcher.cppm | 35 ++ .../Network/Serialization/Private/Packet.cpp | 37 ++ .../Serialization/Private/PacketWriter.cpp | 58 +++ .../Network/Serialization/Public/Packet.cppm | 79 +++ .../Network/Serialization/Public/Packet.inl | 69 +++ .../Serialization/Public/PacketWriter.cppm | 44 ++ .../Serialization/Public/PacketWriter.inl | 18 + .../Network/TCP/Private/StreamReassembler.cpp | 101 ++++ .../Engine/Network/TCP/Private/TCPClient.cpp | 210 ++++++++ .../Engine/Network/TCP/Private/TCPServer.cpp | 96 ++++ .../Network/TCP/Public/StreamReassembler.cppm | 37 ++ .../Engine/Network/TCP/Public/TCPClient.cppm | 62 +++ .../Engine/Network/TCP/Public/TCPServer.cppm | 42 ++ .../Engine/Network/UDP/Private/UDPSocket.cpp | 65 +++ .../Engine/Network/UDP/Public/UDPSocket.cppm | 48 ++ Source/Engine/NexusEngine.cppm | 12 + Tests/Engine/Network/NetworkAddressTests.cpp | 78 +++ Tests/Engine/Network/NetworkErrorTests.cpp | 58 +++ Tests/Engine/Network/NetworkManagerTests.cpp | 226 +++++++++ .../Engine/Network/PacketDispatcherTests.cpp | 169 +++++++ Tests/Engine/Network/PacketTests.cpp | 242 ++++++++++ Tests/Engine/Network/PacketWriterTests.cpp | 155 ++++++ Tests/Engine/Network/SocketTests.cpp | 452 ++++++++++++++++++ .../Engine/Network/StreamReassemblerTests.cpp | 169 +++++++ Tests/Engine/Network/TCPClientTests.cpp | 311 ++++++++++++ Tests/Engine/Network/TCPServerTests.cpp | 311 ++++++++++++ Tests/Engine/Network/UDPSocketTests.cpp | 167 +++++++ 43 files changed, 4991 insertions(+) create mode 100644 Source/Engine/Network/CMakeLists.txt create mode 100644 Source/Engine/Network/Common/Private/NetworkAddress.cpp create mode 100644 Source/Engine/Network/Common/Public/NetworkAddress.cppm create mode 100644 Source/Engine/Network/Common/Public/NetworkAddress.inl create mode 100644 Source/Engine/Network/Common/Public/NetworkError.cppm create mode 100644 Source/Engine/Network/Manager/Private/NetworkManager.cpp create mode 100644 Source/Engine/Network/Manager/Public/NetworkManager.cppm create mode 100644 Source/Engine/Network/Platform/Linux/LinuxSocket.cpp create mode 100644 Source/Engine/Network/Platform/Linux/LinuxSocket_Internal.cpp create mode 100644 Source/Engine/Network/Platform/Linux/LinuxSocket_Internal.hpp create mode 100644 Source/Engine/Network/Platform/Public/Socket.cppm create mode 100644 Source/Engine/Network/Platform/Windows/WindowsScoket_Internal.hpp create mode 100644 Source/Engine/Network/Platform/Windows/WindowsSocket.cpp create mode 100644 Source/Engine/Network/Platform/Windows/WindowsSocket_Internal.cpp create mode 100644 Source/Engine/Network/Protocol/Private/PacketDispatcher.cpp create mode 100644 Source/Engine/Network/Protocol/Public/PacketDispatcher.cppm create mode 100644 Source/Engine/Network/Serialization/Private/Packet.cpp create mode 100644 Source/Engine/Network/Serialization/Private/PacketWriter.cpp create mode 100644 Source/Engine/Network/Serialization/Public/Packet.cppm create mode 100644 Source/Engine/Network/Serialization/Public/Packet.inl create mode 100644 Source/Engine/Network/Serialization/Public/PacketWriter.cppm create mode 100644 Source/Engine/Network/Serialization/Public/PacketWriter.inl create mode 100644 Source/Engine/Network/TCP/Private/StreamReassembler.cpp create mode 100644 Source/Engine/Network/TCP/Private/TCPClient.cpp create mode 100644 Source/Engine/Network/TCP/Private/TCPServer.cpp create mode 100644 Source/Engine/Network/TCP/Public/StreamReassembler.cppm create mode 100644 Source/Engine/Network/TCP/Public/TCPClient.cppm create mode 100644 Source/Engine/Network/TCP/Public/TCPServer.cppm create mode 100644 Source/Engine/Network/UDP/Private/UDPSocket.cpp create mode 100644 Source/Engine/Network/UDP/Public/UDPSocket.cppm create mode 100644 Tests/Engine/Network/NetworkAddressTests.cpp create mode 100644 Tests/Engine/Network/NetworkErrorTests.cpp create mode 100644 Tests/Engine/Network/NetworkManagerTests.cpp create mode 100644 Tests/Engine/Network/PacketDispatcherTests.cpp create mode 100644 Tests/Engine/Network/PacketTests.cpp create mode 100644 Tests/Engine/Network/PacketWriterTests.cpp create mode 100644 Tests/Engine/Network/SocketTests.cpp create mode 100644 Tests/Engine/Network/StreamReassemblerTests.cpp create mode 100644 Tests/Engine/Network/TCPClientTests.cpp create mode 100644 Tests/Engine/Network/TCPServerTests.cpp create mode 100644 Tests/Engine/Network/UDPSocketTests.cpp diff --git a/Source/Engine/CMakeLists.txt b/Source/Engine/CMakeLists.txt index e70aad5..40edae3 100644 --- a/Source/Engine/CMakeLists.txt +++ b/Source/Engine/CMakeLists.txt @@ -19,6 +19,7 @@ add_subdirectory(Asset) add_subdirectory(Audio) add_subdirectory(Core) add_subdirectory(Input) +add_subdirectory(Network) add_subdirectory(Physics) add_subdirectory(Platform) add_subdirectory(Renderer) diff --git a/Source/Engine/Network/CMakeLists.txt b/Source/Engine/Network/CMakeLists.txt new file mode 100644 index 0000000..49f4440 --- /dev/null +++ b/Source/Engine/Network/CMakeLists.txt @@ -0,0 +1,27 @@ +# SPDX-License-Identifier: MIT + +file(GLOB_RECURSE PRIVATE_SOURCES ${CMAKE_CURRENT_SOURCE_DIR}/*/Private/*.cpp) +file(GLOB_RECURSE PUBLIC_MODULES ${CMAKE_CURRENT_SOURCE_DIR}/*/Public/*.cppm) +file(GLOB_RECURSE PRIVATE_MODULES ${CMAKE_CURRENT_SOURCE_DIR}/*/Private/*.cppm) + +if(WIN32) + file(GLOB_RECURSE WINDOWS_SOURCES ${CMAKE_CURRENT_SOURCE_DIR}/Platform/Windows/*.cpp) + list(APPEND PRIVATE_SOURCES ${WINDOWS_SOURCES}) +elseif(LINUX) + file(GLOB_RECURSE LINUX_SOURCES ${CMAKE_CURRENT_SOURCE_DIR}/Platform/Linux/*.cpp) + list(APPEND PRIVATE_SOURCES ${LINUX_SOURCES}) +endif() + +target_sources(NexusEngine PRIVATE ${PRIVATE_SOURCES}) +target_sources(NexusEngine PUBLIC + FILE_SET engine_network_modules TYPE CXX_MODULES + FILES ${PUBLIC_MODULES} +) +target_sources(NexusEngine PUBLIC + FILE_SET engine_network_modules_prv TYPE CXX_MODULES + FILES ${PRIVATE_MODULES} +) + +if(WIN32) + target_link_libraries(NexusEngine PRIVATE ws2_32) +endif() diff --git a/Source/Engine/Network/Common/Private/NetworkAddress.cpp b/Source/Engine/Network/Common/Private/NetworkAddress.cpp new file mode 100644 index 0000000..86f6e37 --- /dev/null +++ b/Source/Engine/Network/Common/Private/NetworkAddress.cpp @@ -0,0 +1,21 @@ +// SPDX-License-Identifier: MIT + +module NE.Engine.Network.Common.NetworkAddress; + +import NE.Engine.Core.Types; + +import std; + +namespace Nexus::Network { + [[nodiscard]] const std::string& NetworkAddress::Host() const noexcept { + return m_host; + } + + [[nodiscard]] uint16 NetworkAddress::Port() const noexcept { + return m_port; + } + + [[nodiscard]] bool NetworkAddress::operator==(const NetworkAddress& other) const noexcept { + return m_port == other.m_port && m_host == other.m_host; + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Common/Public/NetworkAddress.cppm b/Source/Engine/Network/Common/Public/NetworkAddress.cppm new file mode 100644 index 0000000..7038538 --- /dev/null +++ b/Source/Engine/Network/Common/Public/NetworkAddress.cppm @@ -0,0 +1,35 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.Common.NetworkAddress; + +import NE.Engine.Core.Types; + +import std; + +export namespace Nexus::Network { + class NEXUS_API NetworkAddress { + public: + NetworkAddress() = default; + + NetworkAddress(std::string host, uint16 port) : m_host(std::move(host)), m_port(port) {} + + [[nodiscard]] const std::string& Host() const noexcept; + + [[nodiscard]] uint16 Port() const noexcept; + + [[nodiscard]] bool operator==(const NetworkAddress& other) const noexcept; + + private: + std::string m_host; + uint16 m_port = 0; + }; +} // namespace Nexus::Network + +template <> +struct NEXUS_API std::hash; + +#include "NetworkAddress.inl" diff --git a/Source/Engine/Network/Common/Public/NetworkAddress.inl b/Source/Engine/Network/Common/Public/NetworkAddress.inl new file mode 100644 index 0000000..5f8b8e8 --- /dev/null +++ b/Source/Engine/Network/Common/Public/NetworkAddress.inl @@ -0,0 +1,15 @@ +// SPDX-License-Identifier: MIT + +#pragma once + +template <> +struct NEXUS_API std::hash { + Nexus::usize operator()(const Nexus::Network::NetworkAddress& address) const noexcept { + const Nexus::usize h1 = std::hash{}(address.Host()); + const Nexus::usize h2 = std::hash{}(address.Port()); + + // 64-bit mix inspired by the standard hash-combine pattern. + static constexpr auto PHI = 0x9e3779b97f4a7c15ULL; + return h1 ^ (h2 + static_cast(PHI) + (h1 << 6) + (h1 >> 2)); + } +}; diff --git a/Source/Engine/Network/Common/Public/NetworkError.cppm b/Source/Engine/Network/Common/Public/NetworkError.cppm new file mode 100644 index 0000000..62b7446 --- /dev/null +++ b/Source/Engine/Network/Common/Public/NetworkError.cppm @@ -0,0 +1,41 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.Common.NetworkError; + +import std; + +export namespace Nexus::Network { + enum class NetworkErrorCode { + None, + WouldBlock, + ConnectionInProgress, + ConnectionClosed, + ConnectionReset, + ConnectionAborted, + ConnectionRefused, + NotConnected, + AddressInUse, + HostUnreachable, + NetworkUnreachable, + Timeout, + MessageTooLarge, + InvalidAddress, + PermissionDenied, + InvalidOperation, + Unknown, + }; + + struct NEXUS_API NetworkError { + NetworkErrorCode code = NetworkErrorCode::None; + std::string message; + int platformErrno = 0; + + [[nodiscard]] explicit operator bool() const noexcept { + return code != NetworkErrorCode::None; + } + }; +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Manager/Private/NetworkManager.cpp b/Source/Engine/Network/Manager/Private/NetworkManager.cpp new file mode 100644 index 0000000..ceb7f6a --- /dev/null +++ b/Source/Engine/Network/Manager/Private/NetworkManager.cpp @@ -0,0 +1,74 @@ +// SPDX-License-Identifier: MIT + +module NE.Engine.Network.Manager; + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; +import NE.Engine.Network.Protocol.PacketDispatcher; +import NE.Engine.Network.Packet; +import NE.Engine.Network.PacketWriter; +import NE.Engine.Network.TCP.TCPClient; +import NE.Engine.Network.UDP.UDPSocket; + +import NE.Engine.Core.Types; + +import std; + +namespace Nexus::Network { + bool NetworkManager::StartClient(const NetworkAddress& server) { + m_serverAddress = server; + + const bool tcpStarted = m_tcp.Connect(server); + const bool udpStarted = m_udp.Bind(server.Port()); // binding to 0 for an ephemeral port also works + + m_lastConnectAttempt = std::chrono::steady_clock::now(); + + return tcpStarted || udpStarted; + } + + void NetworkManager::RegisterHandler(PacketType type, PacketDispatcher::Handler handler) { + m_dispatcher.Register(type, std::move(handler)); + } + + bool NetworkManager::SendReliable(std::vector framedPacket) { + return m_tcp.Send(std::move(framedPacket)); + } + + bool NetworkManager::SendUnreliable(std::vector framedPacket) { + return m_udp.Send(m_serverAddress, framedPacket); + } + + PacketWriter NetworkManager::MakeWriter(PacketType type) { + return PacketWriter(type, m_nextSequence++); + } + + bool NetworkManager::IsConnected() const noexcept { + return m_tcp.IsConnected(); + } + + bool NetworkManager::IsConnecting() const noexcept { + return m_tcp.IsConnecting(); + } + + void NetworkManager::Update() { + // If we're neither connected nor connecting, try to reconnect periodically. + const auto now = std::chrono::steady_clock::now(); + if (!IsConnected() && !IsConnecting()) { + if (now - m_lastConnectAttempt >= m_connectRetryInterval) { + m_tcp.Connect(m_serverAddress); + m_lastConnectAttempt = now; + } + } + + for (auto& frame : m_tcp.Poll()) + DispatchIncoming(frame); + + std::optional udpError; + while (auto message = m_udp.Receive(udpError)) + DispatchIncoming(message->data); + } + + void NetworkManager::DispatchIncoming(std::span frame) const { + m_dispatcher.Dispatch(frame); + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Manager/Public/NetworkManager.cppm b/Source/Engine/Network/Manager/Public/NetworkManager.cppm new file mode 100644 index 0000000..4f15746 --- /dev/null +++ b/Source/Engine/Network/Manager/Public/NetworkManager.cppm @@ -0,0 +1,54 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.Manager; + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Protocol.PacketDispatcher; +import NE.Engine.Network.Packet; +import NE.Engine.Network.PacketWriter; +import NE.Engine.Network.TCP.TCPClient; +import NE.Engine.Network.UDP.UDPSocket; + +import NE.Engine.Core.Types; + +import std; + +export namespace Nexus::Network { + class NEXUS_API NetworkManager { + public: + bool StartClient(const NetworkAddress& server); + + void RegisterHandler(PacketType type, PacketDispatcher::Handler handler); + + [[nodiscard]] bool SendReliable(std::vector framedPacket); + + [[nodiscard]] bool SendUnreliable(std::vector framedPacket); + + [[nodiscard]] PacketWriter MakeWriter(PacketType type); + + [[nodiscard]] bool IsConnected() const noexcept; + + [[nodiscard]] bool IsConnecting() const noexcept; + + /// Call once per frame from the main loop. + void Update(); + + private: + void DispatchIncoming(std::span frame) const; + + private: + TCPClient m_tcp; + UDPSocket m_udp; + NetworkAddress m_serverAddress; + PacketDispatcher m_dispatcher; + uint32 m_nextSequence = 0; + + // Reconnection state + std::chrono::steady_clock::time_point m_lastConnectAttempt = std::chrono::steady_clock::now(); + std::chrono::milliseconds m_connectRetryInterval = std::chrono::milliseconds(1000); + }; +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Platform/Linux/LinuxSocket.cpp b/Source/Engine/Network/Platform/Linux/LinuxSocket.cpp new file mode 100644 index 0000000..e7fe076 --- /dev/null +++ b/Source/Engine/Network/Platform/Linux/LinuxSocket.cpp @@ -0,0 +1,298 @@ +// SPDX-License-Identifier: MIT + +module; + +#include "LinuxSocket_Internal.hpp" + +// clang-format off +#include + +#include +// clang-format on + +module NE.Engine.Network.Platform.Socket; + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; + +import std; + +namespace Nexus::Network { + namespace { + using namespace Internal; + + NetworkErrorCode ToNetworkErrorCode(RawErrorKind kind) { + switch (kind) { + case RawErrorKind::WouldBlock: + return NetworkErrorCode::WouldBlock; + case RawErrorKind::ConnectionInProgress: + return NetworkErrorCode::ConnectionInProgress; + case RawErrorKind::ConnectionReset: + return NetworkErrorCode::ConnectionReset; + case RawErrorKind::ConnectionAborted: + return NetworkErrorCode::ConnectionAborted; + case RawErrorKind::ConnectionRefused: + return NetworkErrorCode::ConnectionRefused; + case RawErrorKind::NotConnected: + return NetworkErrorCode::NotConnected; + case RawErrorKind::AddressInUse: + return NetworkErrorCode::AddressInUse; + case RawErrorKind::HostUnreachable: + return NetworkErrorCode::HostUnreachable; + case RawErrorKind::NetworkUnreachable: + return NetworkErrorCode::NetworkUnreachable; + case RawErrorKind::Timeout: + return NetworkErrorCode::Timeout; + case RawErrorKind::MessageTooLarge: + return NetworkErrorCode::MessageTooLarge; + case RawErrorKind::PermissionDenied: + return NetworkErrorCode::PermissionDenied; + case RawErrorKind::None: + case RawErrorKind::Unknown: + default: + return NetworkErrorCode::Unknown; + } + } + + RawAddress ToRawAddress(const NetworkAddress& address) { + RawAddress raw{}; + static_cast(RawParseIPv4(address.Host().c_str(), address.Port(), raw)); + return raw; + } + + NetworkAddress FromRawAddress(const RawAddress& raw) { + return NetworkAddress(std::string(raw.host), raw.port); + } + + void SetError(std::optional& outError, NetworkErrorCode code, std::string message, + int platformErrno = 0) { + outError = NetworkError{code, std::move(message), platformErrno}; + } + + void ClearError(std::optional& outError) { + outError.reset(); + } + + void SetErrorFromRaw(std::optional& outError, const RawError& raw) { + outError = NetworkError{ToNetworkErrorCode(raw.kind), std::strerror(raw.platformErrno), raw.platformErrno}; + } + } // namespace + + NetworkError TranslateLastSocketError() { + const RawError raw = RawTranslateErrno(errno); + NetworkError result; + result.code = ToNetworkErrorCode(raw.kind); + result.platformErrno = raw.platformErrno; + result.message = std::strerror(raw.platformErrno); + return result; + } + + Socket::Socket(SocketType type) : m_type(type) { + const RawHandle handle = RawCreateSocket(type == SocketType::TCP); + if (handle >= 0) { + m_handle = static_cast(handle); + } + } + + Socket::~Socket() { + Close(); + } + + Socket::Socket(Socket&& other) noexcept + : m_handle(std::exchange(other.m_handle, INVALID_SOCKET)), + m_type(other.m_type) {} + + Socket& Socket::operator=(Socket&& other) noexcept { + if (this != &other) { + Close(); + m_handle = std::exchange(other.m_handle, INVALID_SOCKET); + m_type = other.m_type; + } + return *this; + } + + bool Socket::SetNonBlocking(bool enabled) { + if (!IsValid()) { + return false; + } + return RawSetNonBlocking(m_handle, enabled); + } + + bool Socket::Bind(uint16 port) { + if (!IsValid()) { + return false; + } + + static_cast(RawSetReuseAddr(m_handle)); + return RawBind(m_handle, port); + } + + bool Socket::Connect(const NetworkAddress& address, std::optional& outError) { + if (!IsValid() || m_type != SocketType::TCP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "Connect requires a valid TCP socket"); + return false; + } + + const RawAddress raw = ToRawAddress(address); + if (raw.host[0] == '\0') { + SetError(outError, NetworkErrorCode::InvalidAddress, "invalid IPv4 address"); + return false; + } + + const RawResult result = RawConnect(m_handle, raw); + + if (result.ok) { + ClearError(outError); + return true; + } + + const RawError translated = RawTranslateErrno(result.platformErrno); + if (translated.kind == RawErrorKind::ConnectionInProgress) { + SetErrorFromRaw(outError, translated); + return true; + } + + SetErrorFromRaw(outError, translated); + return false; + } + + bool Socket::CompleteConnect(std::optional& outError) { + if (!IsValid() || m_type != SocketType::TCP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "CompleteConnect requires a valid TCP socket"); + return false; + } + + const RawResult result = RawCompleteConnect(m_handle); + + if (!result.ok) { + SetErrorFromRaw(outError, RawTranslateErrno(result.platformErrno)); + return false; + } + + ClearError(outError); + return true; + } + + bool Socket::Listen(int backlog) { + if (!IsValid() || m_type != SocketType::TCP) { + return false; + } + return RawListen(m_handle, backlog); + } + + std::optional Socket::Accept() { + if (!IsValid() || m_type != SocketType::TCP) { + return std::nullopt; + } + + const RawHandle clientHandle = RawAccept(m_handle); + if (clientHandle < 0) { + return std::nullopt; + } + + Socket client; + client.m_handle = static_cast(clientHandle); + client.m_type = SocketType::TCP; + return client; + } + + std::optional Socket::Send(std::span data, std::optional& outError) { + if (!IsValid() || m_type != SocketType::TCP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "Send requires a valid TCP socket"); + return std::nullopt; + } + + if (data.empty()) { + ClearError(outError); + return usize{0}; + } + + const RawResult result = RawSend(m_handle, data.data(), static_cast(data.size())); + if (!result.ok) { + SetErrorFromRaw(outError, RawTranslateErrno(result.platformErrno)); + return std::nullopt; + } + + ClearError(outError); + return static_cast(result.value); + } + + std::optional Socket::ReceiveInto(std::span buffer, std::optional& outError) { + if (!IsValid() || m_type != SocketType::TCP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "Receive requires a valid TCP socket"); + return std::nullopt; + } + + if (buffer.empty()) { + ClearError(outError); + return usize{0}; + } + + const RawResult result = RawReceive(m_handle, buffer.data(), static_cast(buffer.size())); + + if (result.ok && result.value == 0) { + SetError(outError, NetworkErrorCode::ConnectionClosed, "peer closed"); + return std::nullopt; + } + + if (!result.ok) { + SetErrorFromRaw(outError, RawTranslateErrno(result.platformErrno)); + return std::nullopt; + } + + ClearError(outError); + return static_cast(result.value); + } + + std::optional Socket::SendTo(const NetworkAddress& dest, std::span data, + std::optional& outError) { + if (!IsValid() || m_type != SocketType::UDP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "SendTo requires a valid UDP socket"); + return std::nullopt; + } + + const RawAddress raw = ToRawAddress(dest); + if (raw.host[0] == '\0') { + SetError(outError, NetworkErrorCode::InvalidAddress, "invalid IPv4 address"); + return std::nullopt; + } + + const RawResult result = RawSendTo(m_handle, raw, data.data(), static_cast(data.size())); + if (!result.ok) { + SetErrorFromRaw(outError, RawTranslateErrno(result.platformErrno)); + return std::nullopt; + } + + ClearError(outError); + return static_cast(result.value); + } + + std::optional Socket::ReceiveFrom(std::span buffer, NetworkAddress& outSender, + std::optional& outError) { + if (!IsValid() || m_type != SocketType::UDP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "ReceiveFrom requires a valid UDP socket"); + return std::nullopt; + } + + const RawReceiveResult result = + RawReceiveFrom(m_handle, buffer.data(), static_cast(buffer.size())); + + if (!result.result.ok) { + SetErrorFromRaw(outError, RawTranslateErrno(result.result.platformErrno)); + return std::nullopt; + } + + outSender = FromRawAddress(result.sender); + ClearError(outError); + return static_cast(result.result.value); + } + + void Socket::Close() { + if (!IsValid()) { + return; + } + + RawCloseSocket(m_handle); + m_handle = INVALID_SOCKET; + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Platform/Linux/LinuxSocket_Internal.cpp b/Source/Engine/Network/Platform/Linux/LinuxSocket_Internal.cpp new file mode 100644 index 0000000..fba56f9 --- /dev/null +++ b/Source/Engine/Network/Platform/Linux/LinuxSocket_Internal.cpp @@ -0,0 +1,228 @@ +// SPDX-License-Identifier: MIT + +#include "LinuxSocket_Internal.hpp" + +// clang-format off +#include +#include +#include +#include +#include +#include + +#include +#include +// clang-format on + +namespace Nexus::Network::Internal { + namespace { + void ToSockAddr(const RawAddress& addr, sockaddr_in& out) { + out = {}; + out.sin_family = AF_INET; + out.sin_port = htons(addr.port); + ::inet_pton(AF_INET, addr.host, &out.sin_addr); + } + + void FromSockAddr(const sockaddr_in& sa, RawAddress& out) { + out = {}; + ::inet_ntop(AF_INET, &sa.sin_addr, out.host, sizeof(out.host)); + out.port = ntohs(sa.sin_port); + } + + RawResult MakeResult(RawHandle value) { + RawResult result; + result.value = value; + result.ok = value >= 0; + result.platformErrno = result.ok ? 0 : errno; + return result; + } + } // namespace + + RawError RawTranslateErrno(int err) { + RawError result; + result.platformErrno = err; + + switch (err) { + case EWOULDBLOCK: // == EAGAIN + result.kind = RawErrorKind::WouldBlock; + break; + case EINPROGRESS: + case EALREADY: + result.kind = RawErrorKind::ConnectionInProgress; + break; + case ECONNRESET: + result.kind = RawErrorKind::ConnectionReset; + break; + case ECONNABORTED: + result.kind = RawErrorKind::ConnectionAborted; + break; + case ECONNREFUSED: + result.kind = RawErrorKind::ConnectionRefused; + break; + case ENOTCONN: + result.kind = RawErrorKind::NotConnected; + break; + case EADDRINUSE: + result.kind = RawErrorKind::AddressInUse; + break; + case EHOSTUNREACH: + result.kind = RawErrorKind::HostUnreachable; + break; + case ENETUNREACH: + result.kind = RawErrorKind::NetworkUnreachable; + break; + case ETIMEDOUT: + result.kind = RawErrorKind::Timeout; + break; + case EMSGSIZE: + result.kind = RawErrorKind::MessageTooLarge; + break; + case EACCES: + result.kind = RawErrorKind::PermissionDenied; + break; + default: + result.kind = RawErrorKind::Unknown; + break; + } + return result; + } + + RawHandle RawCreateSocket(bool isTcp) { + const int stype = isTcp ? SOCK_STREAM : SOCK_DGRAM; + return ::socket(AF_INET, stype, 0); + } + + void RawCloseSocket(RawHandle handle) { + ::close(static_cast(handle)); + } + + bool RawSetNonBlocking(RawHandle handle, bool enabled) { + const int fd = static_cast(handle); + const int flags = ::fcntl(fd, F_GETFL, 0); + if (flags == -1) { + return false; + } + + const int newFlags = enabled ? (flags | O_NONBLOCK) : (flags & ~O_NONBLOCK); + return ::fcntl(fd, F_SETFL, newFlags) != -1; + } + + bool RawSetReuseAddr(RawHandle handle) { + const int reuse = 1; + return ::setsockopt(static_cast(handle), SOL_SOCKET, SO_REUSEADDR, &reuse, sizeof(reuse)) == 0; + } + + bool RawBind(RawHandle handle, unsigned short port) { + sockaddr_in sa{}; + sa.sin_family = AF_INET; + sa.sin_addr.s_addr = htonl(INADDR_ANY); + sa.sin_port = htons(port); + + return ::bind(static_cast(handle), reinterpret_cast(&sa), sizeof(sa)) == 0; + } + + bool RawListen(RawHandle handle, int backlog) { + return ::listen(static_cast(handle), backlog) == 0; + } + + RawHandle RawAccept(RawHandle handle) { + sockaddr_in sa{}; + socklen_t len = sizeof(sa); + return ::accept(static_cast(handle), reinterpret_cast(&sa), &len); + } + + bool RawParseIPv4(const char* host, unsigned short port, RawAddress& out) { + sockaddr_in sa{}; + sa.sin_family = AF_INET; + sa.sin_port = htons(port); + if (::inet_pton(AF_INET, host, &sa.sin_addr) != 1) { + return false; + } + + out = {}; + std::strncpy(out.host, host, sizeof(out.host) - 1); + out.port = port; + return true; + } + + RawResult RawConnect(RawHandle handle, const RawAddress& address) { + sockaddr_in sa{}; + ToSockAddr(address, sa); + + const int result = ::connect(static_cast(handle), reinterpret_cast(&sa), sizeof(sa)); + return MakeResult(result); + } + + RawResult RawCompleteConnect(RawHandle handle) { + const int fd = static_cast(handle); + + fd_set writeSet; + fd_set errorSet; + FD_ZERO(&writeSet); + FD_ZERO(&errorSet); + FD_SET(fd, &writeSet); + FD_SET(fd, &errorSet); + + timeval timeout{}; + const int ready = ::select(fd + 1, nullptr, &writeSet, &errorSet, &timeout); + if (ready < 0) { + return MakeResult(-1); + } + + if (ready == 0) { + RawResult pending; + pending.value = 0; + pending.ok = false; + pending.platformErrno = EINPROGRESS; + return pending; + } + + int socketError = 0; + socklen_t errorSize = sizeof(socketError); + if (::getsockopt(fd, SOL_SOCKET, SO_ERROR, &socketError, &errorSize) != 0) { + return MakeResult(-1); + } + + if (socketError != 0) { + RawResult failed; + failed.value = -1; + failed.ok = false; + failed.platformErrno = socketError; + return failed; + } + + return MakeResult(0); + } + + RawResult RawSend(RawHandle handle, const void* data, unsigned long size) { + const ssize_t n = ::send(static_cast(handle), data, size, 0); + return MakeResult(static_cast(n)); + } + + RawResult RawReceive(RawHandle handle, void* buffer, unsigned long size) { + const ssize_t n = ::recv(static_cast(handle), buffer, size, 0); + return MakeResult(static_cast(n)); + } + + RawResult RawSendTo(RawHandle handle, const RawAddress& dest, const void* data, unsigned long size) { + sockaddr_in sa{}; + ToSockAddr(dest, sa); + + const ssize_t n = + ::sendto(static_cast(handle), data, size, 0, reinterpret_cast(&sa), sizeof(sa)); + return MakeResult(static_cast(n)); + } + + RawReceiveResult RawReceiveFrom(RawHandle handle, void* buffer, unsigned long size) { + sockaddr_in sa{}; + socklen_t len = sizeof(sa); + const ssize_t n = ::recvfrom(static_cast(handle), buffer, size, 0, reinterpret_cast(&sa), &len); + + RawReceiveResult result; + result.result = MakeResult(static_cast(n)); + if (result.result.ok) { + FromSockAddr(sa, result.sender); + } + return result; + } +} // namespace Nexus::Network::Internal diff --git a/Source/Engine/Network/Platform/Linux/LinuxSocket_Internal.hpp b/Source/Engine/Network/Platform/Linux/LinuxSocket_Internal.hpp new file mode 100644 index 0000000..8bd9e49 --- /dev/null +++ b/Source/Engine/Network/Platform/Linux/LinuxSocket_Internal.hpp @@ -0,0 +1,66 @@ +// SPDX-License-Identifier: MIT + +#pragma once + +namespace Nexus::Network::Internal { + using RawHandle = long long; + + struct RawResult { + RawHandle value = -1; + int platformErrno = 0; + bool ok = false; + }; + + struct RawAddress { + char host[64] = {}; // IPv4 string, null-terminated + unsigned short port = 0; + }; + + struct RawReceiveResult { + RawResult result; + RawAddress sender; + }; + + enum class RawErrorKind { + None, + WouldBlock, + ConnectionInProgress, + ConnectionReset, + ConnectionAborted, + ConnectionRefused, + NotConnected, + AddressInUse, + HostUnreachable, + NetworkUnreachable, + Timeout, + MessageTooLarge, + PermissionDenied, + Unknown, + }; + + struct RawError { + RawErrorKind kind = RawErrorKind::None; + int platformErrno = 0; + }; + + [[nodiscard]] RawError RawTranslateErrno(int err); + + [[nodiscard]] RawHandle RawCreateSocket(bool isTcp); + void RawCloseSocket(RawHandle handle); + + [[nodiscard]] bool RawSetNonBlocking(RawHandle handle, bool enabled); + [[nodiscard]] bool RawSetReuseAddr(RawHandle handle); + + [[nodiscard]] bool RawBind(RawHandle handle, unsigned short port); + [[nodiscard]] bool RawListen(RawHandle handle, int backlog); + [[nodiscard]] RawHandle RawAccept(RawHandle handle); + + [[nodiscard]] bool RawParseIPv4(const char* host, unsigned short port, RawAddress& out); + [[nodiscard]] RawResult RawConnect(RawHandle handle, const RawAddress& address); + [[nodiscard]] RawResult RawCompleteConnect(RawHandle handle); + + [[nodiscard]] RawResult RawSend(RawHandle handle, const void* data, unsigned long size); + [[nodiscard]] RawResult RawReceive(RawHandle handle, void* buffer, unsigned long size); + [[nodiscard]] RawResult RawSendTo(RawHandle handle, const RawAddress& dest, const void* data, unsigned long size); + [[nodiscard]] RawReceiveResult RawReceiveFrom(RawHandle handle, void* buffer, unsigned long size); +} // namespace Nexus::Network::Internal diff --git a/Source/Engine/Network/Platform/Public/Socket.cppm b/Source/Engine/Network/Platform/Public/Socket.cppm new file mode 100644 index 0000000..d731849 --- /dev/null +++ b/Source/Engine/Network/Platform/Public/Socket.cppm @@ -0,0 +1,68 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.Platform.Socket; + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; + +import NE.Engine.Core.Types; + +import std; + +export namespace Nexus::Network { + enum class SocketType { + TCP, + UDP, + }; + + using NativeSocketHandle = uptr; + inline constexpr NativeSocketHandle INVALID_SOCKET = ~0; // not -1 because unsigned + + [[nodiscard]] NEXUS_API NetworkError TranslateLastSocketError(); + + class NEXUS_API Socket { + public: + Socket() = default; + explicit Socket(SocketType type); + ~Socket(); + + Socket(const Socket&) = delete; + Socket& operator=(const Socket&) = delete; + + Socket(Socket&& other) noexcept; + Socket& operator=(Socket&& other) noexcept; + + [[nodiscard]] bool IsValid() const noexcept { + return m_handle != INVALID_SOCKET; + } + + bool SetNonBlocking(bool enabled); + bool Bind(uint16 port); + + bool Connect(const NetworkAddress& address, std::optional& outError); + [[nodiscard]] bool CompleteConnect(std::optional& outError); + bool Listen(int backlog = 16); + [[nodiscard]] std::optional Accept(); + + std::optional Send(std::span data, std::optional& outError); + + std::optional ReceiveInto(std::span buffer, std::optional& outError); + + std::optional SendTo(const NetworkAddress& dest, std::span data, + std::optional& outError); + + std::optional ReceiveFrom(std::span buffer, NetworkAddress& outSender, + std::optional& outError); + + void Close(); + + private: + NativeSocketHandle m_handle = INVALID_SOCKET; + SocketType m_type = SocketType::TCP; + }; + +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Platform/Windows/WindowsScoket_Internal.hpp b/Source/Engine/Network/Platform/Windows/WindowsScoket_Internal.hpp new file mode 100644 index 0000000..7b4957c --- /dev/null +++ b/Source/Engine/Network/Platform/Windows/WindowsScoket_Internal.hpp @@ -0,0 +1,68 @@ +// SPDX-License-Identifier: MIT + +#pragma once + +namespace Nexus::Network::Internal { + using RawHandle = long long; + + struct RawResult { + RawHandle value = -1; + int platformErrno = 0; + bool ok = false; + }; + + struct RawAddress { + char host[64] = {}; // IPv4 string, null-terminated + unsigned short port = 0; + }; + + struct RawReceiveResult { + RawResult result; + RawAddress sender; + }; + + enum class RawErrorKind { + None, + WouldBlock, + ConnectionInProgress, + ConnectionReset, + ConnectionAborted, + ConnectionRefused, + NotConnected, + AddressInUse, + HostUnreachable, + NetworkUnreachable, + Timeout, + MessageTooLarge, + PermissionDenied, + Unknown, + }; + + struct RawError { + RawErrorKind kind = RawErrorKind::None; + int platformErrno = 0; + }; + + [[nodiscard]] RawError RawTranslateLastError(); + [[nodiscard]] RawError RawTranslateErrorCode(int platformErrno); + + void RawEnsureInitialized(); + [[nodiscard]] RawHandle RawCreateSocket(bool isTcp); + void RawCloseSocket(RawHandle handle); + + [[nodiscard]] bool RawSetNonBlocking(RawHandle handle, bool enabled); + [[nodiscard]] bool RawSetReuseAddr(RawHandle handle); + + [[nodiscard]] bool RawBind(RawHandle handle, unsigned short port); + [[nodiscard]] bool RawListen(RawHandle handle, int backlog); + [[nodiscard]] RawHandle RawAccept(RawHandle handle); + + [[nodiscard]] bool RawParseIPv4(const char* host, unsigned short port, RawAddress& out); + [[nodiscard]] RawResult RawConnect(RawHandle handle, const RawAddress& address); + [[nodiscard]] RawResult RawCompleteConnect(RawHandle handle); + + [[nodiscard]] RawResult RawSend(RawHandle handle, const void* data, unsigned long size); + [[nodiscard]] RawResult RawReceive(RawHandle handle, void* buffer, unsigned long size); + [[nodiscard]] RawResult RawSendTo(RawHandle handle, const RawAddress& dest, const void* data, unsigned long size); + [[nodiscard]] RawReceiveResult RawReceiveFrom(RawHandle handle, void* buffer, unsigned long size); +} // namespace Nexus::Network::Internal diff --git a/Source/Engine/Network/Platform/Windows/WindowsSocket.cpp b/Source/Engine/Network/Platform/Windows/WindowsSocket.cpp new file mode 100644 index 0000000..3d044db --- /dev/null +++ b/Source/Engine/Network/Platform/Windows/WindowsSocket.cpp @@ -0,0 +1,328 @@ +// SPDX-License-Identifier: MIT + +module; + +#include "WindowsScoket_Internal.hpp" + +#include + +module NE.Engine.Network.Platform.Socket; + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; + +import std; + +namespace Nexus::Network { + namespace { + using namespace Internal; + + NetworkErrorCode ToNetworkErrorCode(RawErrorKind kind) { + switch (kind) { + case RawErrorKind::WouldBlock: + return NetworkErrorCode::WouldBlock; + case RawErrorKind::ConnectionInProgress: + return NetworkErrorCode::ConnectionInProgress; + case RawErrorKind::ConnectionReset: + return NetworkErrorCode::ConnectionReset; + case RawErrorKind::ConnectionAborted: + return NetworkErrorCode::ConnectionAborted; + case RawErrorKind::ConnectionRefused: + return NetworkErrorCode::ConnectionRefused; + case RawErrorKind::NotConnected: + return NetworkErrorCode::NotConnected; + case RawErrorKind::AddressInUse: + return NetworkErrorCode::AddressInUse; + case RawErrorKind::HostUnreachable: + return NetworkErrorCode::HostUnreachable; + case RawErrorKind::NetworkUnreachable: + return NetworkErrorCode::NetworkUnreachable; + case RawErrorKind::Timeout: + return NetworkErrorCode::Timeout; + case RawErrorKind::MessageTooLarge: + return NetworkErrorCode::MessageTooLarge; + case RawErrorKind::PermissionDenied: + return NetworkErrorCode::PermissionDenied; + case RawErrorKind::None: + case RawErrorKind::Unknown: + default: + return NetworkErrorCode::Unknown; + } + } + + RawAddress ToRawAddress(const NetworkAddress& address) { + RawAddress raw{}; + static_cast(RawParseIPv4(address.Host().c_str(), address.Port(), raw)); + return raw; + } + + NetworkAddress FromRawAddress(const RawAddress& raw) { + return NetworkAddress(std::string(raw.host), raw.port); + } + + void SetError(std::optional& outError, NetworkErrorCode code, std::string message, + int platformError = 0) { + outError = NetworkError{code, std::move(message), platformError}; + } + + void ClearError(std::optional& outError) { + outError.reset(); + } + + std::string MessageForKind(RawErrorKind kind, int platformErrno) { + switch (kind) { + case RawErrorKind::WouldBlock: + return "would block"; + case RawErrorKind::ConnectionInProgress: + return "connection in progress"; + case RawErrorKind::ConnectionReset: + return "connection reset"; + case RawErrorKind::ConnectionAborted: + return "connection aborted"; + case RawErrorKind::ConnectionRefused: + return "connection refused"; + case RawErrorKind::NotConnected: + return "not connected"; + case RawErrorKind::AddressInUse: + return "address in use"; + case RawErrorKind::HostUnreachable: + return "host unreachable"; + case RawErrorKind::NetworkUnreachable: + return "network unreachable"; + case RawErrorKind::Timeout: + return "timed out"; + case RawErrorKind::MessageTooLarge: + return "message too large"; + case RawErrorKind::PermissionDenied: + return "permission denied"; + case RawErrorKind::None: + case RawErrorKind::Unknown: + default: + return "winsock error " + std::to_string(platformErrno); + } + } + + void SetErrorFromRaw(std::optional& outError, const RawError& raw) { + outError = NetworkError{ToNetworkErrorCode(raw.kind), MessageForKind(raw.kind, raw.platformErrno), + raw.platformErrno}; + } + } // namespace + + [[nodiscard]] NEXUS_API NetworkError TranslateLastSocketError() { + const RawError raw = RawTranslateLastError(); + NetworkError result; + result.code = ToNetworkErrorCode(raw.kind); + result.platformErrno = raw.platformErrno; + result.message = MessageForKind(raw.kind, raw.platformErrno); + return result; + } + + Socket::Socket(SocketType type) : m_type(type) { + const RawHandle handle = RawCreateSocket(type == SocketType::TCP); + if (handle >= 0) { + m_handle = static_cast(handle); + } + } + + Socket::~Socket() { + Close(); + } + + Socket::Socket(Socket&& other) noexcept + : m_handle(std::exchange(other.m_handle, INVALID_SOCKET)), + m_type(other.m_type) {} + + Socket& Socket::operator=(Socket&& other) noexcept { + if (this != &other) { + Close(); + m_handle = std::exchange(other.m_handle, INVALID_SOCKET); + m_type = other.m_type; + } + return *this; + } + + bool Socket::SetNonBlocking(bool enabled) { + if (!IsValid()) { + return false; + } + return RawSetNonBlocking(m_handle, enabled); + } + + bool Socket::Bind(uint16 port) { + if (!IsValid()) { + return false; + } + + static_cast(RawSetReuseAddr(m_handle)); + return RawBind(m_handle, port); + } + + bool Socket::Connect(const NetworkAddress& address, std::optional& outError) { + if (!IsValid() || m_type != SocketType::TCP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "Connect requires a valid TCP socket"); + return false; + } + + const RawAddress raw = ToRawAddress(address); + if (raw.host[0] == '\0') { + SetError(outError, NetworkErrorCode::InvalidAddress, "invalid IPv4 address"); + return false; + } + + const RawResult result = RawConnect(m_handle, raw); + + if (result.ok) { + ClearError(outError); + return true; + } + + const RawError translated = RawTranslateErrorCode(result.platformErrno); + if (translated.kind == RawErrorKind::WouldBlock || translated.kind == RawErrorKind::ConnectionInProgress) { + SetErrorFromRaw(outError, RawError{RawErrorKind::ConnectionInProgress, translated.platformErrno}); + return true; + } + + SetErrorFromRaw(outError, translated); + return false; + } + + bool Socket::CompleteConnect(std::optional& outError) { + if (!IsValid() || m_type != SocketType::TCP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "CompleteConnect requires a valid TCP socket"); + return false; + } + + const RawResult result = RawCompleteConnect(m_handle); + + if (!result.ok) { + SetErrorFromRaw(outError, RawTranslateErrorCode(result.platformErrno)); + return false; + } + + ClearError(outError); + return true; + } + + bool Socket::Listen(int backlog) { + if (!IsValid() || m_type != SocketType::TCP) { + return false; + } + return RawListen(m_handle, backlog); + } + + std::optional Socket::Accept() { + if (!IsValid() || m_type != SocketType::TCP) { + return std::nullopt; + } + + const RawHandle clientHandle = RawAccept(m_handle); + if (clientHandle < 0) { + return std::nullopt; + } + + Socket client; + client.m_handle = static_cast(clientHandle); + client.m_type = SocketType::TCP; + return client; + } + + std::optional Socket::Send(std::span data, std::optional& outError) { + if (!IsValid() || m_type != SocketType::TCP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "Send requires a valid TCP socket"); + return std::nullopt; + } + + if (data.empty()) { + ClearError(outError); + return usize{0}; + } + + const RawResult result = RawSend(m_handle, data.data(), static_cast(data.size())); + if (!result.ok) { + SetErrorFromRaw(outError, RawTranslateErrorCode(result.platformErrno)); + return std::nullopt; + } + + ClearError(outError); + return static_cast(result.value); + } + + std::optional Socket::ReceiveInto(std::span buffer, std::optional& outError) { + if (!IsValid() || m_type != SocketType::TCP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "Receive requires a valid TCP socket"); + return std::nullopt; + } + + if (buffer.empty()) { + ClearError(outError); + return usize{0}; + } + + const RawResult result = RawReceive(m_handle, buffer.data(), static_cast(buffer.size())); + + if (result.ok && result.value == 0) { + SetError(outError, NetworkErrorCode::ConnectionClosed, "peer closed"); + return std::nullopt; + } + + if (!result.ok) { + SetErrorFromRaw(outError, RawTranslateErrorCode(result.platformErrno)); + return std::nullopt; + } + + ClearError(outError); + return static_cast(result.value); + } + + std::optional Socket::SendTo(const NetworkAddress& dest, std::span data, + std::optional& outError) { + if (!IsValid() || m_type != SocketType::UDP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "SendTo requires a valid UDP socket"); + return std::nullopt; + } + + const RawAddress raw = ToRawAddress(dest); + if (raw.host[0] == '\0') { + SetError(outError, NetworkErrorCode::InvalidAddress, "invalid IPv4 address"); + return std::nullopt; + } + + const RawResult result = RawSendTo(m_handle, raw, data.data(), static_cast(data.size())); + if (!result.ok) { + SetErrorFromRaw(outError, RawTranslateErrorCode(result.platformErrno)); + return std::nullopt; + } + + ClearError(outError); + return static_cast(result.value); + } + + std::optional Socket::ReceiveFrom(std::span buffer, NetworkAddress& outSender, + std::optional& outError) { + if (!IsValid() || m_type != SocketType::UDP) { + SetError(outError, NetworkErrorCode::InvalidOperation, "ReceiveFrom requires a valid UDP socket"); + return std::nullopt; + } + + const RawReceiveResult result = + RawReceiveFrom(m_handle, buffer.data(), static_cast(buffer.size())); + + if (!result.result.ok) { + SetErrorFromRaw(outError, RawTranslateErrorCode(result.result.platformErrno)); + return std::nullopt; + } + + outSender = FromRawAddress(result.sender); + ClearError(outError); + return static_cast(result.result.value); + } + + void Socket::Close() { + if (!IsValid()) { + return; + } + + RawCloseSocket(m_handle); + m_handle = INVALID_SOCKET; + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Platform/Windows/WindowsSocket_Internal.cpp b/Source/Engine/Network/Platform/Windows/WindowsSocket_Internal.cpp new file mode 100644 index 0000000..74f3691 --- /dev/null +++ b/Source/Engine/Network/Platform/Windows/WindowsSocket_Internal.cpp @@ -0,0 +1,246 @@ +// SPDX-License-Identifier: MIT + +#include "WindowsScoket_Internal.hpp" + +// clang-format off +#if !defined(NOMINMAX) +# define NOMINMAX +#endif +#if !defined(WIN32_LEAN_AND_MEAN) +# define WIN32_LEAN_AND_MEAN +#endif +#include +#include + +#include +// clang-format on + +namespace Nexus::Network::Internal { + namespace { + void ToSockAddr(const RawAddress& addr, sockaddr_in& out) { + out = {}; + out.sin_family = AF_INET; + out.sin_port = htons(addr.port); + InetPtonA(AF_INET, addr.host, &out.sin_addr); + } + + void FromSockAddr(const sockaddr_in& sa, RawAddress& out) { + out = {}; + InetNtopA(AF_INET, const_cast(&sa.sin_addr), out.host, sizeof(out.host)); + out.port = ntohs(sa.sin_port); + } + + RawResult MakeResult(RawHandle value) { + RawResult result; + result.value = value; + result.ok = value != SOCKET_ERROR; + result.platformErrno = result.ok ? 0 : WSAGetLastError(); + return result; + } + } // namespace + + void RawEnsureInitialized() { + static const bool initialized = [] { + WSADATA data{}; + return WSAStartup(MAKEWORD(2, 2), &data) == 0; + }(); + static_cast(initialized); + } + + RawError RawTranslateErrorCode(int err) { + RawError result; + result.platformErrno = err; + + switch (err) { + case WSAEWOULDBLOCK: + result.kind = RawErrorKind::WouldBlock; + break; + case WSAEINPROGRESS: + case WSAEALREADY: + result.kind = RawErrorKind::ConnectionInProgress; + break; + case WSAECONNRESET: + result.kind = RawErrorKind::ConnectionReset; + break; + case WSAECONNABORTED: + result.kind = RawErrorKind::ConnectionAborted; + break; + case WSAECONNREFUSED: + result.kind = RawErrorKind::ConnectionRefused; + break; + case WSAENOTCONN: + result.kind = RawErrorKind::NotConnected; + break; + case WSAEADDRINUSE: + result.kind = RawErrorKind::AddressInUse; + break; + case WSAEHOSTUNREACH: + result.kind = RawErrorKind::HostUnreachable; + break; + case WSAENETUNREACH: + result.kind = RawErrorKind::NetworkUnreachable; + break; + case WSAETIMEDOUT: + result.kind = RawErrorKind::Timeout; + break; + case WSAEMSGSIZE: + result.kind = RawErrorKind::MessageTooLarge; + break; + case WSAEACCES: + result.kind = RawErrorKind::PermissionDenied; + break; + default: + result.kind = RawErrorKind::Unknown; + break; + } + return result; + } + + RawError RawTranslateLastError() { + return RawTranslateErrorCode(WSAGetLastError()); + } + + RawHandle RawCreateSocket(bool isTcp) { + RawEnsureInitialized(); + + const int socketType = isTcp ? SOCK_STREAM : SOCK_DGRAM; + const int protocol = isTcp ? IPPROTO_TCP : IPPROTO_UDP; + const SOCKET handle = ::socket(AF_INET, socketType, protocol); + return handle == INVALID_SOCKET ? -1 : static_cast(handle); + } + + void RawCloseSocket(RawHandle handle) { + ::closesocket(static_cast(handle)); + } + + bool RawSetNonBlocking(RawHandle handle, bool enabled) { + u_long mode = enabled ? 1UL : 0UL; + return ioctlsocket(static_cast(handle), FIONBIO, &mode) == 0; + } + + bool RawSetReuseAddr(RawHandle handle) { + const BOOL reuse = TRUE; + return ::setsockopt(static_cast(handle), SOL_SOCKET, SO_REUSEADDR, + reinterpret_cast(&reuse), sizeof(reuse)) == 0; + } + + bool RawBind(RawHandle handle, unsigned short port) { + sockaddr_in sa{}; + sa.sin_family = AF_INET; + sa.sin_addr.s_addr = htonl(INADDR_ANY); + sa.sin_port = htons(port); + + return ::bind(static_cast(handle), reinterpret_cast(&sa), sizeof(sa)) == 0; + } + + bool RawListen(RawHandle handle, int backlog) { + return ::listen(static_cast(handle), backlog) == 0; + } + + RawHandle RawAccept(RawHandle handle) { + sockaddr_in sa{}; + int len = sizeof(sa); + const SOCKET clientSocket = ::accept(static_cast(handle), reinterpret_cast(&sa), &len); + return clientSocket == INVALID_SOCKET ? -1 : static_cast(clientSocket); + } + + bool RawParseIPv4(const char* host, unsigned short port, RawAddress& out) { + sockaddr_in sa{}; + if (InetPtonA(AF_INET, host, &sa.sin_addr) != 1) { + return false; + } + + out = {}; + std::strncpy(out.host, host, sizeof(out.host) - 1); + out.port = port; + return true; + } + + RawResult RawConnect(RawHandle handle, const RawAddress& address) { + sockaddr_in sa{}; + ToSockAddr(address, sa); + + const int result = ::connect(static_cast(handle), reinterpret_cast(&sa), sizeof(sa)); + return MakeResult(result == 0 ? 0 : SOCKET_ERROR); + } + + RawResult RawCompleteConnect(RawHandle handle) { + fd_set writeSet{}; + fd_set errorSet{}; + FD_ZERO(&writeSet); + FD_ZERO(&errorSet); + + const SOCKET socketHandle = static_cast(handle); + FD_SET(socketHandle, &writeSet); + FD_SET(socketHandle, &errorSet); + + timeval timeout{}; + const int ready = ::select(0, nullptr, &writeSet, &errorSet, &timeout); + if (ready == SOCKET_ERROR) { + return MakeResult(SOCKET_ERROR); + } + + if (ready == 0) { + RawResult pending; + pending.value = 0; + pending.ok = false; + pending.platformErrno = WSAEINPROGRESS; + return pending; + } + + int socketError = 0; + int errorSize = sizeof(socketError); + if (getsockopt(socketHandle, SOL_SOCKET, SO_ERROR, reinterpret_cast(&socketError), &errorSize) != 0) { + return MakeResult(SOCKET_ERROR); + } + + if (socketError != 0) { + WSASetLastError(socketError); + return MakeResult(SOCKET_ERROR); + } + + return MakeResult(0); + } + + RawResult RawSend(RawHandle handle, const void* data, unsigned long size) { + static constexpr auto MAXCONN = static_cast(SOMAXCONN); + const int clamped = static_cast(size > MAXCONN ? MAXCONN : size); + const int n = ::send(static_cast(handle), static_cast(data), clamped, 0); + return MakeResult(n); + } + + RawResult RawReceive(RawHandle handle, void* buffer, unsigned long size) { + static constexpr auto MAXCONN = static_cast(SOMAXCONN); + const int clamped = static_cast(size > MAXCONN ? MAXCONN : size); + const int n = ::recv(static_cast(handle), static_cast(buffer), clamped, 0); + return MakeResult(n); + } + + RawResult RawSendTo(RawHandle handle, const RawAddress& dest, const void* data, unsigned long size) { + sockaddr_in sa{}; + ToSockAddr(dest, sa); + + static constexpr auto MAXCONN = static_cast(SOMAXCONN); + const int clamped = static_cast(size > MAXCONN ? MAXCONN : size); + const int n = ::sendto(static_cast(handle), static_cast(data), clamped, 0, + reinterpret_cast(&sa), sizeof(sa)); + return MakeResult(n); + } + + RawReceiveResult RawReceiveFrom(RawHandle handle, void* buffer, unsigned long size) { + sockaddr_in sa{}; + int len = sizeof(sa); + + static constexpr auto MAXCONN = static_cast(SOMAXCONN); + const int clamped = static_cast(size > MAXCONN ? MAXCONN : size); + const int n = ::recvfrom(static_cast(handle), static_cast(buffer), clamped, 0, + reinterpret_cast(&sa), &len); + + RawReceiveResult result; + result.result = MakeResult(n); + if (result.result.ok) { + FromSockAddr(sa, result.sender); + } + return result; + } +} // namespace Nexus::Network::Internal diff --git a/Source/Engine/Network/Protocol/Private/PacketDispatcher.cpp b/Source/Engine/Network/Protocol/Private/PacketDispatcher.cpp new file mode 100644 index 0000000..1fe6ef5 --- /dev/null +++ b/Source/Engine/Network/Protocol/Private/PacketDispatcher.cpp @@ -0,0 +1,70 @@ +// SPDX-License-Identifier: MIT + +module NE.Engine.Network.Protocol.PacketDispatcher; + +import NE.Engine.Network.Packet; + +import NE.Engine.Core.Types; + +import std; + +namespace Nexus::Network { + void PacketDispatcher::Register(PacketType type, Handler handler) { + if (!handler) { + m_handlers.erase(type); + return; + } + m_handlers.insert_or_assign(type, std::move(handler)); + } + + void PacketDispatcher::Unregister(PacketType type) { + m_handlers.erase(type); + } + + void PacketDispatcher::Clear() { + m_handlers.clear(); + } + + bool PacketDispatcher::HasHandler(PacketType type) const { + return m_handlers.contains(type); + } + + void PacketDispatcher::Dispatch(std::span frame) const { + if (frame.size() < HEADER_SIZE) { + return; + } + + usize offset = 0; + uint32 payloadSize = 0; + uint16 typeValue = 0; + uint32 sequence = 0; + + if (!Wire::Read(frame, offset, payloadSize) || !Wire::Read(frame, offset, typeValue) || + !Wire::Read(frame, offset, sequence)) { + return; + } + + if (payloadSize > MAX_PAYLOAD_SIZE || static_cast(payloadSize) != frame.size() - HEADER_SIZE) { + return; + } + + if (typeValue > static_cast(PacketType::ChatMessage)) { + return; + } + + const PacketType type = static_cast(typeValue); + const auto it = m_handlers.find(type); + if (it == m_handlers.end()) { + return; + } + + std::vector payload(frame.begin() + static_cast(HEADER_SIZE), frame.end()); + + Packet packet(payload); + it->second(packet); + } + + void PacketDispatcher::Dispatch(const std::vector& frame) const { + Dispatch(std::span(frame)); + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Protocol/Public/PacketDispatcher.cppm b/Source/Engine/Network/Protocol/Public/PacketDispatcher.cppm new file mode 100644 index 0000000..b6c2df9 --- /dev/null +++ b/Source/Engine/Network/Protocol/Public/PacketDispatcher.cppm @@ -0,0 +1,35 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.Protocol.PacketDispatcher; + +import NE.Engine.Network.Packet; + +import NE.Engine.Core.Types; + +import std; + +export namespace Nexus::Network { + class NEXUS_API PacketDispatcher { + public: + using Handler = std::function; + + void Register(PacketType type, Handler handler); + + void Unregister(PacketType type); + + void Clear(); + + [[nodiscard]] bool HasHandler(PacketType type) const; + + void Dispatch(std::span frame) const; + + void Dispatch(const std::vector& frame) const; + + private: + std::unordered_map m_handlers; + }; +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Serialization/Private/Packet.cpp b/Source/Engine/Network/Serialization/Private/Packet.cpp new file mode 100644 index 0000000..31f70eb --- /dev/null +++ b/Source/Engine/Network/Serialization/Private/Packet.cpp @@ -0,0 +1,37 @@ +// SPDX-License-Identifier: MIT + +module NE.Engine.Network.Packet; + +import NE.Engine.Core.Types; + +import std; + +namespace Nexus::Network { + bool Packet::ReadString(std::string& out) { + uint32 length = 0; + if (!Read(length)) { + return false; + } + + if (static_cast(length) > m_storage.size() - m_readPos) { + return false; + } + + const auto* chars = reinterpret_cast(m_storage.data() + m_readPos); + out.assign(chars, length); + m_readPos += length; + return true; + } + + std::span Packet::Data() const noexcept { + return m_storage; + } + + usize Packet::Remaining() const noexcept { + return m_storage.size() - std::min(m_readPos, m_storage.size()); + } + + void Packet::Reset() noexcept { + m_readPos = 0; + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Serialization/Private/PacketWriter.cpp b/Source/Engine/Network/Serialization/Private/PacketWriter.cpp new file mode 100644 index 0000000..033ca15 --- /dev/null +++ b/Source/Engine/Network/Serialization/Private/PacketWriter.cpp @@ -0,0 +1,58 @@ +// SPDX-License-Identifier: MIT + +module NE.Engine.Network.PacketWriter; + +import NE.Engine.Network.Packet; + +import NE.Engine.Core.Types; + +import std; + +namespace Nexus::Network { + PacketWriter::PacketWriter(PacketType type, uint32 sequence, usize reserveBytes) + : m_type(type), + m_sequence(sequence) { + m_buffer.reserve(std::max(reserveBytes, HEADER_SIZE)); + m_buffer.resize(HEADER_SIZE); + } + + bool PacketWriter::WriteString(std::string_view text) { + if (text.size() > std::numeric_limits::max()) { + return false; + } + + Write(static_cast(text.size())); + const auto* bytes = reinterpret_cast(text.data()); + m_buffer.insert(m_buffer.end(), bytes, bytes + text.size()); + return true; + } + + std::vector PacketWriter::Build() { + if (m_built) { + return {}; + } + + const usize payloadSize = m_buffer.size() - HEADER_SIZE; + if (payloadSize > MAX_PAYLOAD_SIZE) { + return {}; + } + + usize offset = 0; + const auto WriteHeaderField = [this, &offset](auto value) { + const auto wireValue = Wire::HostToBigEndian(value); + std::memcpy(m_buffer.data() + offset, &wireValue, sizeof(wireValue)); + offset += sizeof(wireValue); + }; + + WriteHeaderField(static_cast(payloadSize)); + WriteHeaderField(static_cast(m_type)); + WriteHeaderField(m_sequence); + + m_built = true; + return std::move(m_buffer); + } + + std::span PacketWriter::Data() const noexcept { + return m_buffer; + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Serialization/Public/Packet.cppm b/Source/Engine/Network/Serialization/Public/Packet.cppm new file mode 100644 index 0000000..3258a27 --- /dev/null +++ b/Source/Engine/Network/Serialization/Public/Packet.cppm @@ -0,0 +1,79 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.Packet; + +import NE.Engine.Core.Types; + +import std; + +export namespace Nexus::Network { + // Placeholder types, changes guaranteed at some point + enum class PacketType : uint16 { + ClientHello, + ServerHello, + PlayerInput, + PlayerState, + SpawnEntity, + DestroyEntity, + ChatMessage, + }; + + struct NEXUS_API PacketHeader { + uint32 size = 0; + PacketType type = {}; + uint32 sequence = 0; + }; + + inline constexpr usize HEADER_SIZE = sizeof(uint32) + sizeof(PacketType) + sizeof(uint32); + + inline constexpr uint32 MAX_PAYLOAD_SIZE = 1u << 20; // 1 MiB + inline constexpr usize MAX_FRAME_SIZE = HEADER_SIZE + static_cast(MAX_PAYLOAD_SIZE); + + static_assert(HEADER_SIZE == 10); + + /// Wire values are always sent big-endian. + namespace Wire { + template + [[nodiscard]] constexpr T HostToBigEndian(T value) noexcept; + + template + [[nodiscard]] constexpr T BigEndianToHost(T value) noexcept; + + template + void Append(std::vector& buffer, T value); + + template + [[nodiscard]] bool Read(std::span data, usize& offset, T& out); + } // namespace Wire + + /// A read/write cursor over a byte buffer. Integral types go through Wire + /// (big-endian on the wire); other trivially copyable types are copied raw. + class NEXUS_API Packet { + public: + explicit Packet(std::vector& storage) noexcept : m_storage(storage) {} + + template + void Write(const T& value); + + template + [[nodiscard]] bool Read(T& out); + + [[nodiscard]] bool ReadString(std::string& out); + + [[nodiscard]] std::span Data() const noexcept; + + [[nodiscard]] usize Remaining() const noexcept; + + void Reset() noexcept; + + private: + std::vector& m_storage; + usize m_readPos = 0; + }; +} // namespace Nexus::Network + +#include "Packet.inl" diff --git a/Source/Engine/Network/Serialization/Public/Packet.inl b/Source/Engine/Network/Serialization/Public/Packet.inl new file mode 100644 index 0000000..ee74a44 --- /dev/null +++ b/Source/Engine/Network/Serialization/Public/Packet.inl @@ -0,0 +1,69 @@ +// SPDX-License-Identifier: MIT + +#pragma once + +namespace Nexus::Network { + namespace Wire { + template + [[nodiscard]] constexpr T HostToBigEndian(T value) noexcept { + if constexpr (std::endian::native == std::endian::little && sizeof(T) > 1) { + return std::byteswap(value); + } else { + return value; + } + } + + template + [[nodiscard]] constexpr T BigEndianToHost(T value) noexcept { + return HostToBigEndian(value); + } + + template + void Append(std::vector& buffer, T value) { + const T wireValue = HostToBigEndian(value); + const auto* bytes = reinterpret_cast(&wireValue); + buffer.insert(buffer.end(), bytes, bytes + sizeof(T)); + } + + template + [[nodiscard]] bool Read(std::span data, usize& offset, T& out) { + if (offset > data.size() || sizeof(T) > data.size() - offset) { + return false; + } + + T wireValue{}; + std::memcpy(&wireValue, data.data() + offset, sizeof(T)); + offset += sizeof(T); + out = BigEndianToHost(wireValue); + return true; + } + } // namespace Wire + + template + void Packet::Write(const T& value) { + static_assert(std::is_trivially_copyable_v); + + if constexpr (std::integral) { + Wire::Append(m_storage, value); + } else { + const auto* bytes = reinterpret_cast(&value); + m_storage.insert(m_storage.end(), bytes, bytes + sizeof(T)); + } + } + + template + [[nodiscard]] bool Packet::Read(T& out) { + static_assert(std::is_trivially_copyable_v); + + if constexpr (std::integral) { + return Wire::Read(m_storage, m_readPos, out); + } else { + if (m_readPos > m_storage.size() || sizeof(T) > m_storage.size() - m_readPos) { + return false; + } + std::memcpy(&out, m_storage.data() + m_readPos, sizeof(T)); + m_readPos += sizeof(T); + return true; + } + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/Serialization/Public/PacketWriter.cppm b/Source/Engine/Network/Serialization/Public/PacketWriter.cppm new file mode 100644 index 0000000..5b60eb9 --- /dev/null +++ b/Source/Engine/Network/Serialization/Public/PacketWriter.cppm @@ -0,0 +1,44 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.PacketWriter; + +import NE.Engine.Network.Packet; + +import NE.Engine.Core.Types; + +import std; + +export namespace Nexus::Network { + class NEXUS_API PacketWriter { + public: + explicit PacketWriter(PacketType type, uint32 sequence, usize reserveBytes = 64); + + PacketWriter(const PacketWriter&) = delete; + PacketWriter& operator=(const PacketWriter&) = delete; + PacketWriter(PacketWriter&&) = default; + PacketWriter& operator=(PacketWriter&&) = default; + + template + void Write(const T& value); + + bool WriteString(std::string_view text); + + /// Fills in the 10-byte header (payload size, type, sequence) and returns + /// the completed frame. Can only be called once. + [[nodiscard]] std::vector Build(); + + [[nodiscard]] std::span Data() const noexcept; + + private: + std::vector m_buffer; + PacketType m_type; + uint32 m_sequence; + bool m_built = false; + }; +} // namespace Nexus::Network + +#include "PacketWriter.inl" diff --git a/Source/Engine/Network/Serialization/Public/PacketWriter.inl b/Source/Engine/Network/Serialization/Public/PacketWriter.inl new file mode 100644 index 0000000..3131fbf --- /dev/null +++ b/Source/Engine/Network/Serialization/Public/PacketWriter.inl @@ -0,0 +1,18 @@ +// SPDX-License-Identifier: MIT + +#pragma once + +namespace Nexus::Network { + template + void PacketWriter::Write(const T& value) { + static_assert(!std::is_pointer_v); + static_assert(std::is_trivially_copyable_v); + + if constexpr (std::integral) { + Wire::Append(m_buffer, value); + } else { + const auto* bytes = reinterpret_cast(&value); + m_buffer.insert(m_buffer.end(), bytes, bytes + sizeof(T)); + } + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/TCP/Private/StreamReassembler.cpp b/Source/Engine/Network/TCP/Private/StreamReassembler.cpp new file mode 100644 index 0000000..ed878bb --- /dev/null +++ b/Source/Engine/Network/TCP/Private/StreamReassembler.cpp @@ -0,0 +1,101 @@ +// SPDX-License-Identifier: MIT + +module NE.Engine.Network.TCP.StreamReassembler; + +import NE.Engine.Network.Packet; + +import NE.Engine.Core.Types; + +import std; + +namespace Nexus::Network { + void StreamReassembler::Feed(std::span rawBytes) { + if (m_invalid || rawBytes.empty()) { + return; + } + + if (rawBytes.size() > MAX_BUFFERED_BYTES - BufferedSize()) { + m_buffer.clear(); + m_readPos = 0; + m_invalid = true; + return; + } + + m_buffer.insert(m_buffer.end(), rawBytes.begin(), rawBytes.end()); + m_invalid = false; + } + + std::optional> StreamReassembler::TryExtract() { + if (m_invalid) { + return std::nullopt; + } + + const std::span pending = std::span(m_buffer).subspan(m_readPos); + + if (pending.size() < HEADER_SIZE) { + Compact(); + return std::nullopt; + } + + usize offset = 0; + uint32 payloadSize = 0; + uint16 typeValue = 0; + uint32 sequence = 0; + + if (!Wire::Read(pending, offset, payloadSize) || !Wire::Read(pending, offset, typeValue) || + !Wire::Read(pending, offset, sequence)) { + return std::nullopt; + } + + static_cast(sequence); + + if (payloadSize > MAX_PAYLOAD_SIZE) { + m_buffer.clear(); + m_readPos = 0; + m_invalid = true; + return std::nullopt; + } + + const usize totalFrameSize = HEADER_SIZE + static_cast(payloadSize); + if (pending.size() < totalFrameSize) { + return std::nullopt; + } + + std::vector frame(pending.begin(), pending.begin() + static_cast(totalFrameSize)); + + m_readPos += totalFrameSize; + Compact(); + return frame; + } + + bool StreamReassembler::HasInvalidData() const noexcept { + return m_invalid; + } + + void StreamReassembler::Reset() noexcept { + m_buffer.clear(); + m_readPos = 0; + m_invalid = false; + } + + usize StreamReassembler::BufferedSize() const noexcept { + return m_buffer.size() - std::min(m_readPos, m_buffer.size()); + } + + void StreamReassembler::Compact() { + if (m_readPos == 0) { + return; + } + + if (m_readPos == m_buffer.size()) { + m_buffer.clear(); + m_readPos = 0; + return; + } + + if (m_readPos >= 4096 || m_readPos * 2 >= m_buffer.size()) { + m_buffer.erase(m_buffer.begin(), m_buffer.begin() + static_cast(m_readPos)); + m_readPos = 0; + } + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/TCP/Private/TCPClient.cpp b/Source/Engine/Network/TCP/Private/TCPClient.cpp new file mode 100644 index 0000000..2c2928c --- /dev/null +++ b/Source/Engine/Network/TCP/Private/TCPClient.cpp @@ -0,0 +1,210 @@ +// SPDX-License-Identifier: MIT + +module NE.Engine.Network.TCP.TCPClient; + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; +import NE.Engine.Network.Platform.Socket; +import NE.Engine.Network.TCP.StreamReassembler; + +import NE.Engine.Core.Types; + +import std; + +namespace Nexus::Network { + bool TCPClient::Connect(const NetworkAddress& address) { + Disconnect(); + + m_socket = Socket(SocketType::TCP); + if (!m_socket.IsValid()) { + return false; + } + + if (!m_socket.SetNonBlocking(true)) { + m_socket.Close(); + return false; + } + + std::optional error; + if (!m_socket.Connect(address, error)) { + m_socket.Close(); + return false; + } + + if (error && error->code == NetworkErrorCode::ConnectionInProgress) { + m_connecting = true; + m_connected = false; + } else { + m_connecting = false; + m_connected = true; + } + + return true; + } + + TCPClient::TCPClient(Socket&& acceptedSocket) + : m_socket(std::move(acceptedSocket)), + m_connected(m_socket.IsValid()) { + m_connecting = false; + + if (m_connected && !m_socket.SetNonBlocking(true)) { + Disconnect(); + } + } + + void TCPClient::Disconnect() noexcept { + m_socket.Close(); + m_connecting = false; + m_connected = false; + m_outgoing.clear(); + m_queuedSendBytes = 0; + m_reassembler.Reset(); + } + + bool TCPClient::IsConnected() const noexcept { + return m_connected; + } + + bool TCPClient::IsConnecting() const noexcept { + return m_connecting; + } + + bool TCPClient::Send(std::vector&& framedData) { + if ((!m_connected && !m_connecting) || framedData.empty()) { + return false; + } + + if (framedData.size() > MAX_QUEUED_SEND_BYTES - m_queuedSendBytes) { + return false; + } + + m_queuedSendBytes += framedData.size(); + m_outgoing.emplace_back(std::move(framedData), 0); + + if (m_connected) { + FlushOutgoing(); + } + + return m_connected || m_connecting; + } + + bool TCPClient::Send(std::span framedData) { + if (framedData.empty()) { + return false; + } + + std::vector copy(framedData.begin(), framedData.end()); + return Send(std::move(copy)); + } + + std::vector> TCPClient::Poll() { + std::vector> frames; + + if (!m_socket.IsValid()) { + return frames; + } + + if (m_connecting) { + std::optional error; + + if (!m_socket.CompleteConnect(error)) { + if (error && error->code == NetworkErrorCode::ConnectionInProgress) { + return frames; + } + + Disconnect(); + return frames; + } + + m_connecting = false; + m_connected = true; + } + + if (!m_connected) { + return frames; + } + + FlushOutgoing(); + if (!m_connected) { + return frames; + } + + std::array scratch{}; + + while (true) { + std::optional error; + const auto received = m_socket.ReceiveInto(scratch, error); + + if (!received) { + if (error && error->code == NetworkErrorCode::WouldBlock) { + break; + } + + Disconnect(); + break; + } + + if (*received == 0) { + Disconnect(); + break; + } + + m_reassembler.Feed(std::span(scratch.data(), *received)); + + if (m_reassembler.HasInvalidData()) { + Disconnect(); + break; + } + } + + while (auto frame = m_reassembler.TryExtract()) { + frames.push_back(std::move(*frame)); + } + + return frames; + } + + bool TCPClient::FlushOutgoing() { + if (!m_connected || !m_socket.IsValid()) { + return true; + } + + while (!m_outgoing.empty()) { + PendingSend& pending = m_outgoing.front(); + + if (pending.offset >= pending.data.size()) { + m_outgoing.pop_front(); + continue; + } + + const std::span remaining(pending.data.data() + pending.offset, + pending.data.size() - pending.offset); + + std::optional error; + const auto sent = m_socket.Send(remaining, error); + + if (!sent) { + if (error && error->code == NetworkErrorCode::WouldBlock) { + return true; + } + + Disconnect(); + return false; + } + + if (*sent == 0) { + Disconnect(); + return false; + } + + pending.offset += *sent; + m_queuedSendBytes -= *sent; + + if (pending.offset == pending.data.size()) { + m_outgoing.pop_front(); + } + } + + return true; + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/TCP/Private/TCPServer.cpp b/Source/Engine/Network/TCP/Private/TCPServer.cpp new file mode 100644 index 0000000..cc8a421 --- /dev/null +++ b/Source/Engine/Network/TCP/Private/TCPServer.cpp @@ -0,0 +1,96 @@ +// SPDX-License-Identifier: MIT + +module NE.Engine.Network.TCP.TCPServer; + +import NE.Engine.Network.TCP.TCPClient; +import NE.Engine.Network.Platform.Socket; + +import NE.Engine.Core.Types; + +import std; + +namespace Nexus::Network { + bool TCPServer::Listen(uint16 port) { + m_clients.clear(); + m_listenSocket.Close(); + + m_listenSocket = Socket(SocketType::TCP); + if (!m_listenSocket.IsValid()) { + return false; + } + + if (!m_listenSocket.SetNonBlocking(true)) { + m_listenSocket.Close(); + return false; + } + + if (!m_listenSocket.Bind(port) || !m_listenSocket.Listen()) { + m_listenSocket.Close(); + return false; + } + + m_nextClientId = 1; + return true; + } + + void TCPServer::Stop() { + m_listenSocket.Close(); + m_clients.clear(); + } + + std::vector TCPServer::AcceptPending() { + std::vector newClients; + if (!m_listenSocket.IsValid()) { + return newClients; + } + + while (auto accepted = m_listenSocket.Accept()) { + const ClientId id = m_nextClientId++; + m_clients.emplace(id, TCPClient(std::move(*accepted))); + newClients.push_back(id); + } + + return newClients; + } + + std::vector>> TCPServer::PollAll() { + std::vector>> allFrames; + std::vector toRemove; + + for (auto& [id, client] : m_clients) { + for (auto& frame : client.Poll()) { + allFrames.emplace_back(id, std::move(frame)); + } + + if (!client.IsConnected()) { + toRemove.push_back(id); + } + } + + for (const ClientId id : toRemove) { + m_clients.erase(id); + } + + return allFrames; + } + + bool TCPServer::SendTo(ClientId id, std::span framedData) { + const auto it = m_clients.find(id); + return it != m_clients.end() && it->second.Send(framedData); + } + + bool TCPServer::SendTo(ClientId id, std::vector&& framedData) { + const auto it = m_clients.find(id); + return it != m_clients.end() && it->second.Send(std::move(framedData)); + } + + void TCPServer::Broadcast(std::span framedData) { + for (TCPClient& client : m_clients | std::views::values) { + static_cast(client.Send(framedData)); + } + } + + usize TCPServer::ClientCount() const noexcept { + return m_clients.size(); + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/TCP/Public/StreamReassembler.cppm b/Source/Engine/Network/TCP/Public/StreamReassembler.cppm new file mode 100644 index 0000000..4e52dbd --- /dev/null +++ b/Source/Engine/Network/TCP/Public/StreamReassembler.cppm @@ -0,0 +1,37 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.TCP.StreamReassembler; + +import NE.Engine.Network.Packet; + +import NE.Engine.Core.Types; + +import std; + +export namespace Nexus::Network { + class NEXUS_API StreamReassembler { + public: + static constexpr usize MAX_BUFFERED_BYTES = MAX_FRAME_SIZE * 2; + + void Feed(std::span rawBytes); + + [[nodiscard]] std::optional> TryExtract(); + + [[nodiscard]] bool HasInvalidData() const noexcept; + + void Reset() noexcept; + + private: + [[nodiscard]] usize BufferedSize() const noexcept; + + void Compact(); + + std::vector m_buffer; + usize m_readPos = 0; + bool m_invalid = false; + }; +} // namespace Nexus::Network diff --git a/Source/Engine/Network/TCP/Public/TCPClient.cppm b/Source/Engine/Network/TCP/Public/TCPClient.cppm new file mode 100644 index 0000000..a94e922 --- /dev/null +++ b/Source/Engine/Network/TCP/Public/TCPClient.cppm @@ -0,0 +1,62 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.TCP.TCPClient; + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; +import NE.Engine.Network.Platform.Socket; +import NE.Engine.Network.TCP.StreamReassembler; + +import NE.Engine.Core.Types; + +import std; + +export namespace Nexus::Network { + class NEXUS_API TCPClient { + public: + static constexpr usize MAX_QUEUED_SEND_BYTES = 4u << 20; // 4 MiB + + TCPClient() = default; + + TCPClient(const TCPClient&) = delete; + TCPClient& operator=(const TCPClient&) = delete; + + TCPClient(TCPClient&& other) noexcept = default; + TCPClient& operator=(TCPClient&& other) noexcept = default; + + bool Connect(const NetworkAddress& address); + + explicit TCPClient(Socket&& acceptedSocket); + + void Disconnect() noexcept; + + [[nodiscard]] bool IsConnected() const noexcept; + + [[nodiscard]] bool IsConnecting() const noexcept; + + bool Send(std::vector&& framedData); + + bool Send(std::span framedData); + + std::vector> Poll(); + + private: + struct PendingSend { + std::vector data; + usize offset = 0; + }; + + bool FlushOutgoing(); + + Socket m_socket; + StreamReassembler m_reassembler; + std::deque m_outgoing; + usize m_queuedSendBytes = 0; + bool m_connected = false; + bool m_connecting = false; + }; +} // namespace Nexus::Network diff --git a/Source/Engine/Network/TCP/Public/TCPServer.cppm b/Source/Engine/Network/TCP/Public/TCPServer.cppm new file mode 100644 index 0000000..2bb85ac --- /dev/null +++ b/Source/Engine/Network/TCP/Public/TCPServer.cppm @@ -0,0 +1,42 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.TCP.TCPServer; + +import NE.Engine.Network.TCP.TCPClient; +import NE.Engine.Network.Platform.Socket; + +import NE.Engine.Core.Types; + +import std; + +export namespace Nexus::Network { + using ClientId = uint32; + + class NEXUS_API TCPServer { + public: + bool Listen(uint16 port); + + void Stop(); + + std::vector AcceptPending(); + + std::vector>> PollAll(); + + bool SendTo(ClientId id, std::span framedData); + + bool SendTo(ClientId id, std::vector&& framedData); + + void Broadcast(std::span framedData); + + [[nodiscard]] usize ClientCount() const noexcept; + + private: + Socket m_listenSocket; + std::unordered_map m_clients; + ClientId m_nextClientId = 1; + }; +} // namespace Nexus::Network diff --git a/Source/Engine/Network/UDP/Private/UDPSocket.cpp b/Source/Engine/Network/UDP/Private/UDPSocket.cpp new file mode 100644 index 0000000..fd1ae9a --- /dev/null +++ b/Source/Engine/Network/UDP/Private/UDPSocket.cpp @@ -0,0 +1,65 @@ +// SPDX-License-Identifier: MIT + +module NE.Engine.Network.UDP.UDPSocket; + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; +import NE.Engine.Network.Platform.Socket; + +import NE.Engine.Core.Types; + +import std; + +namespace Nexus::Network { + bool UDPSocket::Bind(uint16 port) { + m_socket.Close(); + m_socket = Socket(SocketType::UDP); + if (!m_socket.IsValid()) { + return false; + } + + if (!m_socket.SetNonBlocking(true) || !m_socket.Bind(port)) { + m_socket.Close(); + return false; + } + return true; + } + + void UDPSocket::Close() { + m_socket.Close(); + } + + bool UDPSocket::IsOpen() const noexcept { + return m_socket.IsValid(); + } + + bool UDPSocket::Send(const NetworkAddress& destination, std::span data) { + if (!m_socket.IsValid() || data.size() > MAX_UDP_PAYLOAD) { + return false; + } + + std::optional error; + const auto sent = m_socket.SendTo(destination, data, error); + return sent.has_value() && *sent == data.size(); + } + + std::optional UDPSocket::Receive(std::optional& outError) { + if (!m_socket.IsValid()) { + outError = NetworkError{NetworkErrorCode::InvalidOperation, "UDP socket is not open", 0}; + return std::nullopt; + } + + std::array scratch{}; + NetworkAddress sender; + const auto received = m_socket.ReceiveFrom(scratch, sender, outError); + + if (!received) { + return std::nullopt; + } + + UDPMessage message; + message.sender = std::move(sender); + message.data.assign(scratch.begin(), scratch.begin() + static_cast(*received)); + return message; + } +} // namespace Nexus::Network diff --git a/Source/Engine/Network/UDP/Public/UDPSocket.cppm b/Source/Engine/Network/UDP/Public/UDPSocket.cppm new file mode 100644 index 0000000..a8c94be --- /dev/null +++ b/Source/Engine/Network/UDP/Public/UDPSocket.cppm @@ -0,0 +1,48 @@ +// SPDX-License-Identifier: MIT + +module; + +#include + +export module NE.Engine.Network.UDP.UDPSocket; + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; +import NE.Engine.Network.Platform.Socket; + +import NE.Engine.Core.Types; + +import std; + +export namespace Nexus::Network { + inline constexpr usize MAX_UDP_PAYLOAD = 65507; + + struct NEXUS_API UDPMessage { + NetworkAddress sender; + std::vector data; + }; + + class NEXUS_API UDPSocket { + public: + UDPSocket() = default; + + UDPSocket(const UDPSocket&) = delete; + UDPSocket& operator=(const UDPSocket&) = delete; + UDPSocket(UDPSocket&&) noexcept = default; + UDPSocket& operator=(UDPSocket&&) noexcept = default; + + bool Bind(uint16 port); + + void Close(); + + [[nodiscard]] bool IsOpen() const noexcept; + + bool Send(const NetworkAddress& destination, std::span data); + + /// Returns one datagram. Call repeatedly until it returns nullopt. + [[nodiscard]] std::optional Receive(std::optional& outError); + + private: + Socket m_socket; + }; +} // namespace Nexus::Network diff --git a/Source/Engine/NexusEngine.cppm b/Source/Engine/NexusEngine.cppm index 93b1514..685f163 100644 --- a/Source/Engine/NexusEngine.cppm +++ b/Source/Engine/NexusEngine.cppm @@ -16,3 +16,15 @@ export import NE.Engine.Core.Window; export import NE.Engine.Math.Mat; export import NE.Engine.Math.Quaternion; export import NE.Engine.Math.Vec; + +// Network +export import NE.Engine.Network.Common.NetworkAddress; +export import NE.Engine.Network.Common.NetworkError; +export import NE.Engine.Network.Manager; +export import NE.Engine.Network.Packet; +export import NE.Engine.Network.PacketWriter; +export import NE.Engine.Network.Protocol.PacketDispatcher; +export import NE.Engine.Network.TCP.StreamReassembler; +export import NE.Engine.Network.TCP.TCPClient; +export import NE.Engine.Network.TCP.TCPServer; +export import NE.Engine.Network.UDP.UDPSocket; diff --git a/Tests/Engine/Network/NetworkAddressTests.cpp b/Tests/Engine/Network/NetworkAddressTests.cpp new file mode 100644 index 0000000..bbed0e3 --- /dev/null +++ b/Tests/Engine/Network/NetworkAddressTests.cpp @@ -0,0 +1,78 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Common.NetworkAddress; + +import NE.Engine.Core.Types; + +import std; + +using Nexus::Network::NetworkAddress; + +TEST(NetworkAddressTest, DefaultConstruction) { + const NetworkAddress address; + + EXPECT_TRUE(address.Host().empty()); + EXPECT_EQ(address.Port(), 0); +} + +TEST(NetworkAddressTest, ConstructionWithHostAndPort) { + const NetworkAddress address("127.0.0.1", 8080); + + EXPECT_EQ(address.Host(), "127.0.0.1"); + EXPECT_EQ(address.Port(), 8080); +} + +TEST(NetworkAddressTest, EqualityWhenHostAndPortMatch) { + const NetworkAddress first("localhost", 27015); + const NetworkAddress second("localhost", 27015); + + EXPECT_TRUE(first == second); +} + +TEST(NetworkAddressTest, InequalityWhenHostDiffers) { + const NetworkAddress first("localhost", 27015); + const NetworkAddress second("example.com", 27015); + + EXPECT_FALSE(first == second); +} + +TEST(NetworkAddressTest, InequalityWhenPortDiffers) { + const NetworkAddress first("localhost", 27015); + const NetworkAddress second("localhost", 27016); + + EXPECT_FALSE(first == second); +} + +TEST(NetworkAddressTest, HashIsConsistentForEqualAddresses) { + const NetworkAddress first("localhost", 27015); + const NetworkAddress second("localhost", 27015); + + const std::hash hasher; + EXPECT_EQ(hasher(first), hasher(second)); +} + +TEST(NetworkAddressTest, HashDiffersForDifferentPorts) { + const NetworkAddress first("localhost", 27015); + const NetworkAddress second("localhost", 27016); + + const std::hash hasher; + + EXPECT_NE(hasher(first), hasher(second)); +} + +TEST(NetworkAddressTest, UsableAsUnorderedMapKey) { + std::unordered_map map; + map[NetworkAddress("127.0.0.1", 1000)] = 1; + map[NetworkAddress("127.0.0.1", 2000)] = 2; + + EXPECT_EQ(map.at(NetworkAddress("127.0.0.1", 1000)), 1); + EXPECT_EQ(map.at(NetworkAddress("127.0.0.1", 2000)), 2); + EXPECT_EQ(map.size(), 2u); +} + +TEST(NetworkAddressTest, PortAtMaxUint16) { + const NetworkAddress address("host", 65535); + EXPECT_EQ(address.Port(), 65535); +} diff --git a/Tests/Engine/Network/NetworkErrorTests.cpp b/Tests/Engine/Network/NetworkErrorTests.cpp new file mode 100644 index 0000000..cd34df0 --- /dev/null +++ b/Tests/Engine/Network/NetworkErrorTests.cpp @@ -0,0 +1,58 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Common.NetworkError; + +import std; + +using Nexus::Network::NetworkError; +using Nexus::Network::NetworkErrorCode; + +TEST(NetworkErrorTest, DefaultConstructedHasNoneCode) { + const NetworkError error; + + EXPECT_EQ(error.code, NetworkErrorCode::None); + EXPECT_TRUE(error.message.empty()); + EXPECT_EQ(error.platformErrno, 0); +} + +TEST(NetworkErrorTest, DefaultConstructedIsFalsy) { + const NetworkError error; + + EXPECT_FALSE(static_cast(error)); +} + +TEST(NetworkErrorTest, NonNoneCodeIsTruthy) { + NetworkError error; + error.code = NetworkErrorCode::Timeout; + + EXPECT_TRUE(static_cast(error)); +} + +TEST(NetworkErrorTest, StoresMessageAndPlatformErrno) { + const NetworkError error{NetworkErrorCode::ConnectionReset, "connection reset", 104}; + + EXPECT_EQ(error.code, NetworkErrorCode::ConnectionReset); + EXPECT_EQ(error.message, "connection reset"); + EXPECT_EQ(error.platformErrno, 104); +} + +TEST(NetworkErrorTest, EveryNonNoneCodeIsTruthy) { + constexpr NetworkErrorCode codes[] = { + NetworkErrorCode::WouldBlock, NetworkErrorCode::ConnectionInProgress, + NetworkErrorCode::ConnectionClosed, NetworkErrorCode::ConnectionReset, + NetworkErrorCode::ConnectionAborted, NetworkErrorCode::ConnectionRefused, + NetworkErrorCode::NotConnected, NetworkErrorCode::AddressInUse, + NetworkErrorCode::HostUnreachable, NetworkErrorCode::NetworkUnreachable, + NetworkErrorCode::Timeout, NetworkErrorCode::MessageTooLarge, + NetworkErrorCode::InvalidAddress, NetworkErrorCode::PermissionDenied, + NetworkErrorCode::InvalidOperation, NetworkErrorCode::Unknown, + }; + + for (const auto code : codes) { + NetworkError error; + error.code = code; + EXPECT_TRUE(static_cast(error)) << "code index " << static_cast(code); + } +} diff --git a/Tests/Engine/Network/NetworkManagerTests.cpp b/Tests/Engine/Network/NetworkManagerTests.cpp new file mode 100644 index 0000000..bb58798 --- /dev/null +++ b/Tests/Engine/Network/NetworkManagerTests.cpp @@ -0,0 +1,226 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; +import NE.Engine.Network.Manager; +import NE.Engine.Network.Packet; +import NE.Engine.Network.PacketWriter; +import NE.Engine.Network.Platform.Socket; + +import NE.Engine.Core.Types; + +import std; + +using namespace Nexus::Network; + +namespace { + template + bool WaitUntil(Predicate predicate, std::chrono::milliseconds timeout = std::chrono::milliseconds(2000)) { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (std::chrono::steady_clock::now() < deadline) { + if (predicate()) { + return true; + } + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + return predicate(); + } + + Nexus::uint16 NextTestPort() { + static std::atomic counter{55000}; + return counter.fetch_add(1); + } + + // A raw-Socket "server" used as the peer for NetworkManager's TCP side, + // so NetworkManager is tested in isolation rather than against TCPServer. + struct RawServer { + Socket listener; + Nexus::uint16 port = 0; + std::optional accepted; + + static std::optional Create() { + RawServer result; + result.port = NextTestPort(); + result.listener = Socket(SocketType::TCP); + if (!result.listener.IsValid() || !result.listener.SetNonBlocking(true) || + !result.listener.Bind(result.port) || !result.listener.Listen()) { + return std::nullopt; + } + return result; + } + + bool WaitForAccept() { + return WaitUntil([&] { + if (!accepted) { + accepted = listener.Accept(); + if (accepted) { + static_cast(accepted->SetNonBlocking(true)); + } + } + return accepted.has_value(); + }); + } + }; + + class NetworkManagerTest : public ::testing::Test {}; +} // namespace + +TEST_F(NetworkManagerTest, DefaultConstructedIsNotConnected) { + const NetworkManager manager; + EXPECT_FALSE(manager.IsConnected()); + EXPECT_FALSE(manager.IsConnecting()); +} + +TEST_F(NetworkManagerTest, StartClientBeginsConnecting) { + auto server = RawServer::Create(); + ASSERT_TRUE(server.has_value()); + + NetworkManager manager; + EXPECT_TRUE(manager.StartClient(NetworkAddress("127.0.0.1", server->port))); +} + +TEST_F(NetworkManagerTest, UpdateCompletesConnectionAfterAccept) { + auto server = RawServer::Create(); + ASSERT_TRUE(server.has_value()); + + NetworkManager manager; + ASSERT_TRUE(manager.StartClient(NetworkAddress("127.0.0.1", server->port))); + ASSERT_TRUE(server->WaitForAccept()); + + const bool connected = WaitUntil([&] { + manager.Update(); + return manager.IsConnected(); + }); + + EXPECT_TRUE(connected); + EXPECT_FALSE(manager.IsConnecting()); +} + +TEST_F(NetworkManagerTest, RegisterHandlerReceivesDispatchedFrameFromPeer) { + auto server = RawServer::Create(); + ASSERT_TRUE(server.has_value()); + + NetworkManager manager; + + std::string receivedText; + manager.RegisterHandler(PacketType::ChatMessage, + [&receivedText](Packet& packet) { static_cast(packet.ReadString(receivedText)); }); + + ASSERT_TRUE(manager.StartClient(NetworkAddress("127.0.0.1", server->port))); + ASSERT_TRUE(server->WaitForAccept()); + + WaitUntil([&] { + manager.Update(); + return manager.IsConnected(); + }); + ASSERT_TRUE(manager.IsConnected()); + + PacketWriter writer(PacketType::ChatMessage, 1); + writer.WriteString("greetings from the server"); + const auto frame = writer.Build(); + + std::optional sendError; + ASSERT_TRUE(server->accepted->Send(frame, sendError).has_value()); + + const bool got = WaitUntil([&] { + manager.Update(); + return !receivedText.empty(); + }); + + ASSERT_TRUE(got); + EXPECT_EQ(receivedText, "greetings from the server"); +} + +TEST_F(NetworkManagerTest, SendReliableDeliversFrameToPeer) { + auto server = RawServer::Create(); + ASSERT_TRUE(server.has_value()); + + NetworkManager manager; + ASSERT_TRUE(manager.StartClient(NetworkAddress("127.0.0.1", server->port))); + ASSERT_TRUE(server->WaitForAccept()); + + WaitUntil([&] { + manager.Update(); + return manager.IsConnected(); + }); + ASSERT_TRUE(manager.IsConnected()); + + PacketWriter writer = manager.MakeWriter(PacketType::ChatMessage); + writer.WriteString("via SendReliable"); + const auto frame = writer.Build(); + + EXPECT_TRUE(manager.SendReliable(frame)); + + std::array buffer{}; + std::optional received; + const bool got = WaitUntil([&] { + std::optional error; + received = server->accepted->ReceiveInto(buffer, error); + return received.has_value(); + }); + + ASSERT_TRUE(got); + ASSERT_EQ(*received, frame.size()); + EXPECT_TRUE(std::equal(frame.begin(), frame.end(), buffer.begin())); +} + +TEST_F(NetworkManagerTest, MakeWriterAssignsIncrementingSequenceNumbers) { + NetworkManager manager; + + PacketWriter first = manager.MakeWriter(PacketType::PlayerInput); + const auto firstFrame = first.Build(); + + PacketWriter second = manager.MakeWriter(PacketType::PlayerInput); + const auto secondFrame = second.Build(); + + Nexus::usize firstOffset = sizeof(Nexus::uint32) + sizeof(Nexus::uint16); + Nexus::usize secondOffset = firstOffset; + Nexus::uint32 firstSequence = 0; + Nexus::uint32 secondSequence = 0; + + ASSERT_TRUE(Wire::Read(firstFrame, firstOffset, firstSequence)); + ASSERT_TRUE(Wire::Read(secondFrame, secondOffset, secondSequence)); + + EXPECT_EQ(secondSequence, firstSequence + 1); +} + +TEST_F(NetworkManagerTest, SendReliableBeforeConnectingFails) { + NetworkManager manager; + + PacketWriter writer = manager.MakeWriter(PacketType::ChatMessage); + writer.WriteString("nobody there yet"); + + EXPECT_FALSE(manager.SendReliable(writer.Build())); +} + +TEST_F(NetworkManagerTest, UnregisteredPacketTypeIsSilentlyIgnored) { + auto server = RawServer::Create(); + ASSERT_TRUE(server.has_value()); + + NetworkManager manager; + // No handlers registered at all. + ASSERT_TRUE(manager.StartClient(NetworkAddress("127.0.0.1", server->port))); + ASSERT_TRUE(server->WaitForAccept()); + + WaitUntil([&] { + manager.Update(); + return manager.IsConnected(); + }); + ASSERT_TRUE(manager.IsConnected()); + + PacketWriter writer(PacketType::SpawnEntity, 1); + writer.Write(Nexus::uint32{1}); + const auto frame = writer.Build(); + + std::optional sendError; + ASSERT_TRUE(server->accepted->Send(frame, sendError).has_value()); + + EXPECT_NO_THROW({ + for (int i = 0; i < 20; ++i) { + manager.Update(); + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } + }); +} diff --git a/Tests/Engine/Network/PacketDispatcherTests.cpp b/Tests/Engine/Network/PacketDispatcherTests.cpp new file mode 100644 index 0000000..202b47e --- /dev/null +++ b/Tests/Engine/Network/PacketDispatcherTests.cpp @@ -0,0 +1,169 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Packet; +import NE.Engine.Network.PacketWriter; +import NE.Engine.Network.Protocol.PacketDispatcher; + +import NE.Engine.Core.Types; + +import std; + +using namespace Nexus::Network; + +namespace { + class PacketDispatcherTest : public ::testing::Test { + protected: + PacketDispatcher m_dispatcher; + }; + + std::vector BuildChatFrame(Nexus::uint32 sequence, std::string_view text) { + PacketWriter writer(PacketType::ChatMessage, sequence); + writer.WriteString(text); + return writer.Build(); + } +} // namespace + +TEST_F(PacketDispatcherTest, HasHandlerFalseWhenNoneRegistered) { + EXPECT_FALSE(m_dispatcher.HasHandler(PacketType::ChatMessage)); +} + +TEST_F(PacketDispatcherTest, RegisterMakesHasHandlerTrue) { + m_dispatcher.Register(PacketType::ChatMessage, [](Packet&) {}); + EXPECT_TRUE(m_dispatcher.HasHandler(PacketType::ChatMessage)); +} + +TEST_F(PacketDispatcherTest, RegisterWithEmptyHandlerUnregisters) { + m_dispatcher.Register(PacketType::ChatMessage, [](Packet&) {}); + ASSERT_TRUE(m_dispatcher.HasHandler(PacketType::ChatMessage)); + + m_dispatcher.Register(PacketType::ChatMessage, PacketDispatcher::Handler{}); + EXPECT_FALSE(m_dispatcher.HasHandler(PacketType::ChatMessage)); +} + +TEST_F(PacketDispatcherTest, UnregisterRemovesHandler) { + m_dispatcher.Register(PacketType::ChatMessage, [](Packet&) {}); + m_dispatcher.Unregister(PacketType::ChatMessage); + + EXPECT_FALSE(m_dispatcher.HasHandler(PacketType::ChatMessage)); +} + +TEST_F(PacketDispatcherTest, ClearRemovesAllHandlers) { + m_dispatcher.Register(PacketType::ChatMessage, [](Packet&) {}); + m_dispatcher.Register(PacketType::PlayerInput, [](Packet&) {}); + + m_dispatcher.Clear(); + + EXPECT_FALSE(m_dispatcher.HasHandler(PacketType::ChatMessage)); + EXPECT_FALSE(m_dispatcher.HasHandler(PacketType::PlayerInput)); +} + +TEST_F(PacketDispatcherTest, DispatchInvokesRegisteredHandler) { + bool invoked = false; + m_dispatcher.Register(PacketType::ChatMessage, [&invoked](Packet&) { invoked = true; }); + + const auto frame = BuildChatFrame(1, "hi"); + m_dispatcher.Dispatch(frame); + + EXPECT_TRUE(invoked); +} + +TEST_F(PacketDispatcherTest, DispatchPassesReadablePayloadToHandler) { + std::string received; + m_dispatcher.Register(PacketType::ChatMessage, + [&received](Packet& packet) { static_cast(packet.ReadString(received)); }); + + const auto frame = BuildChatFrame(1, "hello world"); + m_dispatcher.Dispatch(frame); + + EXPECT_EQ(received, "hello world"); +} + +TEST_F(PacketDispatcherTest, DispatchDoesNotInvokeHandlerForDifferentType) { + bool chatInvoked = false; + bool inputInvoked = false; + m_dispatcher.Register(PacketType::ChatMessage, [&chatInvoked](Packet&) { chatInvoked = true; }); + m_dispatcher.Register(PacketType::PlayerInput, [&inputInvoked](Packet&) { inputInvoked = true; }); + + const auto frame = BuildChatFrame(1, "hi"); + m_dispatcher.Dispatch(frame); + + EXPECT_TRUE(chatInvoked); + EXPECT_FALSE(inputInvoked); +} + +TEST_F(PacketDispatcherTest, DispatchIgnoresFrameWithNoRegisteredHandler) { + const auto frame = BuildChatFrame(1, "hi"); + EXPECT_NO_THROW(m_dispatcher.Dispatch(frame)); +} + +TEST_F(PacketDispatcherTest, DispatchIgnoresFrameShorterThanHeader) { + bool invoked = false; + m_dispatcher.Register(PacketType::ChatMessage, [&invoked](Packet&) { invoked = true; }); + + const std::vector tooShort(HEADER_SIZE - 1, Nexus::byte{0}); + m_dispatcher.Dispatch(tooShort); + + EXPECT_FALSE(invoked); +} + +TEST_F(PacketDispatcherTest, DispatchIgnoresFrameWithMismatchedPayloadSize) { + bool invoked = false; + m_dispatcher.Register(PacketType::ChatMessage, [&invoked](Packet&) { invoked = true; }); + + auto frame = BuildChatFrame(1, "hello world"); + Nexus::usize offset = 0; + const auto corruptedSize = Wire::HostToBigEndian(Nexus::uint32{9999}); + std::memcpy(frame.data() + offset, &corruptedSize, sizeof(corruptedSize)); + + m_dispatcher.Dispatch(frame); + + EXPECT_FALSE(invoked); +} + +TEST_F(PacketDispatcherTest, DispatchIgnoresFrameWithUnknownType) { + bool invoked = false; + m_dispatcher.Register(PacketType::ChatMessage, [&invoked](Packet&) { invoked = true; }); + + auto frame = BuildChatFrame(1, "hi"); + + Nexus::usize offset = sizeof(Nexus::uint32); + const auto bogusType = Wire::HostToBigEndian(Nexus::uint16{9999}); + std::memcpy(frame.data() + offset, &bogusType, sizeof(bogusType)); + + EXPECT_NO_THROW(m_dispatcher.Dispatch(frame)); + EXPECT_FALSE(invoked); +} + +TEST_F(PacketDispatcherTest, DispatchAcceptsSpanOverload) { + bool invoked = false; + m_dispatcher.Register(PacketType::ChatMessage, [&invoked](Packet&) { invoked = true; }); + + const auto frame = BuildChatFrame(1, "hi"); + m_dispatcher.Dispatch(std::span(frame)); + + EXPECT_TRUE(invoked); +} + +TEST_F(PacketDispatcherTest, DispatchWithEmptyPayloadInvokesHandler) { + bool invoked = false; + m_dispatcher.Register(PacketType::ClientHello, [&invoked](Packet&) { invoked = true; }); + + PacketWriter writer(PacketType::ClientHello, 1); + const auto frame = writer.Build(); + m_dispatcher.Dispatch(frame); + + EXPECT_TRUE(invoked); +} + +TEST_F(PacketDispatcherTest, ReRegisteringReplacesPreviousHandler) { + int callCount = 0; + m_dispatcher.Register(PacketType::ChatMessage, [&callCount](Packet&) { callCount += 1; }); + m_dispatcher.Register(PacketType::ChatMessage, [&callCount](Packet&) { callCount += 100; }); + + const auto frame = BuildChatFrame(1, "hi"); + m_dispatcher.Dispatch(frame); + + EXPECT_EQ(callCount, 100); +} diff --git a/Tests/Engine/Network/PacketTests.cpp b/Tests/Engine/Network/PacketTests.cpp new file mode 100644 index 0000000..3e43664 --- /dev/null +++ b/Tests/Engine/Network/PacketTests.cpp @@ -0,0 +1,242 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Packet; + +import NE.Engine.Core.Types; + +import std; + +using namespace Nexus::Network; + +namespace { + class PacketTest : public ::testing::Test { + protected: + std::vector m_storage; + }; +} // namespace + +// ---- Wire ------------------------------------------------------------- + +TEST(WireTest, HostToBigEndianRoundTripsThroughBigEndianToHost) { + constexpr Nexus::uint32 value = 0x12345678u; + const Nexus::uint32 wire = Wire::HostToBigEndian(value); + const Nexus::uint32 back = Wire::BigEndianToHost(wire); + + EXPECT_EQ(back, value); +} + +TEST(WireTest, HostToBigEndianProducesBigEndianByteOrder) { + constexpr Nexus::uint32 value = 0x12345678u; + const Nexus::uint32 wire = Wire::HostToBigEndian(value); + + std::array bytes{}; + std::memcpy(bytes.data(), &wire, sizeof(wire)); + + // Big-endian: most significant Nexus::byte first, regardless of host endianness. + EXPECT_EQ(static_cast(bytes[0]), 0x12); + EXPECT_EQ(static_cast(bytes[1]), 0x34); + EXPECT_EQ(static_cast(bytes[2]), 0x56); + EXPECT_EQ(static_cast(bytes[3]), 0x78); +} + +TEST(WireTest, SingleByteTypeIsUnaffectedByByteOrder) { + constexpr Nexus::uint8 value = 0xAB; + EXPECT_EQ(Wire::HostToBigEndian(value), value); +} + +TEST(WireTest, AppendThenReadRoundTrips) { + std::vector buffer; + Wire::Append(buffer, Nexus::uint32{0xDEADBEEF}); + + Nexus::usize offset = 0; + Nexus::uint32 out = 0; + ASSERT_TRUE(Wire::Read(buffer, offset, out)); + EXPECT_EQ(out, 0xDEADBEEFu); + EXPECT_EQ(offset, sizeof(Nexus::uint32)); +} + +TEST(WireTest, ReadFailsWhenBufferTooShort) { + const std::vector buffer(2, Nexus::byte{0}); + Nexus::usize offset = 0; + Nexus::uint32 out = 0; + + EXPECT_FALSE(Wire::Read(buffer, offset, out)); +} + +TEST(WireTest, ReadFailsWhenOffsetPastEnd) { + const std::vector buffer(4, Nexus::byte{0}); + Nexus::usize offset = 10; + Nexus::uint32 out = 0; + + EXPECT_FALSE(Wire::Read(buffer, offset, out)); +} + +TEST(WireTest, MultipleAppendsAreReadBackInOrder) { + std::vector buffer; + Wire::Append(buffer, Nexus::uint16{111}); + Wire::Append(buffer, Nexus::uint32{222}); + Wire::Append(buffer, Nexus::uint8{33}); + + Nexus::usize offset = 0; + Nexus::uint16 a = 0; + Nexus::uint32 b = 0; + Nexus::uint8 c = 0; + + ASSERT_TRUE(Wire::Read(buffer, offset, a)); + ASSERT_TRUE(Wire::Read(buffer, offset, b)); + ASSERT_TRUE(Wire::Read(buffer, offset, c)); + + EXPECT_EQ(a, 111); + EXPECT_EQ(b, 222u); + EXPECT_EQ(c, 33); +} + +// ---- Packet ------------------------------------------------------------- + +TEST_F(PacketTest, WriteThenReadIntegralRoundTrips) { + Packet writer(m_storage); + writer.Write(Nexus::uint32{424242}); + + Packet reader(m_storage); + Nexus::uint32 out = 0; + ASSERT_TRUE(reader.Read(out)); + EXPECT_EQ(out, 424242u); +} + +TEST_F(PacketTest, WriteThenReadBoolRoundTrips) { + Packet writer(m_storage); + writer.Write(true); + writer.Write(false); + + Packet reader(m_storage); + bool first = false; + bool second = true; + ASSERT_TRUE(reader.Read(first)); + ASSERT_TRUE(reader.Read(second)); + + EXPECT_TRUE(first); + EXPECT_FALSE(second); +} + +TEST_F(PacketTest, WriteThenReadFloatRoundTrips) { + Packet writer(m_storage); + writer.Write(3.14159f); + + Packet reader(m_storage); + Nexus::float32 out = 0.0f; + ASSERT_TRUE(reader.Read(out)); + EXPECT_FLOAT_EQ(out, 3.14159f); +} + +TEST_F(PacketTest, WriteThenReadStringRoundTrips) { + Packet writer(m_storage); + std::string text = "hello, packet"; + + // WriteString helper that replicates the exact wire format ReadString expects + // Nexus::uint32 length followed by the raw Nexus::bytes. + writer.Write(static_cast(text.size())); + m_storage.insert(m_storage.end(), reinterpret_cast(text.data()), + reinterpret_cast(text.data()) + text.size()); + + Packet reader(m_storage); + std::string out; + ASSERT_TRUE(reader.ReadString(out)); + EXPECT_EQ(out, text); +} + +TEST_F(PacketTest, ReadStringFailsWhenLengthExceedsRemainingBytes) { + Packet writer(m_storage); + writer.Write(Nexus::uint32{1000}); // claims 1000 Nexus::bytes follow, but none do + + Packet reader(m_storage); + std::string out; + EXPECT_FALSE(reader.ReadString(out)); +} + +TEST_F(PacketTest, ReadFailsPastEndOfBuffer) { + Packet writer(m_storage); + writer.Write(Nexus::uint8{1}); + + Packet reader(m_storage); + Nexus::uint8 first = 0; + Nexus::uint8 second = 0; + ASSERT_TRUE(reader.Read(first)); + EXPECT_FALSE(reader.Read(second)); +} + +TEST_F(PacketTest, MultipleFieldsRoundTripInOrder) { + Packet writer(m_storage); + writer.Write(Nexus::uint32{1}); + writer.Write(Nexus::uint16{2}); + writer.Write(true); + writer.Write(4.5f); + + Packet reader(m_storage); + Nexus::uint32 a = 0; + Nexus::uint16 b = 0; + bool c = false; + Nexus::float32 d = 0.0f; + + ASSERT_TRUE(reader.Read(a)); + ASSERT_TRUE(reader.Read(b)); + ASSERT_TRUE(reader.Read(c)); + ASSERT_TRUE(reader.Read(d)); + + EXPECT_EQ(a, 1u); + EXPECT_EQ(b, 2); + EXPECT_TRUE(c); + EXPECT_FLOAT_EQ(d, 4.5f); +} + +TEST_F(PacketTest, RemainingReflectsUnreadBytes) { + Packet writer(m_storage); + writer.Write(Nexus::uint32{1}); + writer.Write(Nexus::uint32{2}); + + Packet reader(m_storage); + EXPECT_EQ(reader.Remaining(), 2 * sizeof(Nexus::uint32)); + + Nexus::uint32 value = 0; + ASSERT_TRUE(reader.Read(value)); + EXPECT_EQ(reader.Remaining(), sizeof(Nexus::uint32)); + + ASSERT_TRUE(reader.Read(value)); + EXPECT_EQ(reader.Remaining(), 0u); +} + +TEST_F(PacketTest, ResetRewindsReadCursor) { + Packet writer(m_storage); + writer.Write(Nexus::uint32{99}); + + Packet reader(m_storage); + Nexus::uint32 first = 0; + ASSERT_TRUE(reader.Read(first)); + EXPECT_EQ(reader.Remaining(), 0u); + + reader.Reset(); + EXPECT_EQ(reader.Remaining(), sizeof(Nexus::uint32)); + + Nexus::uint32 second = 0; + ASSERT_TRUE(reader.Read(second)); + EXPECT_EQ(second, 99u); +} + +TEST_F(PacketTest, DataReflectsUnderlyingStorage) { + Packet writer(m_storage); + writer.Write(Nexus::uint8{7}); + + EXPECT_EQ(writer.Data().size(), 1u); + EXPECT_EQ(static_cast(writer.Data()[0]), 7); +} + +TEST_F(PacketTest, EmptyStorageHasZeroRemaining) { + Packet reader(m_storage); + EXPECT_EQ(reader.Remaining(), 0u); +} + +TEST_F(PacketTest, HeaderSizeConstantMatchesWireLayout) { + // Nexus::uint32 (payload size) + Nexus::uint16 (type) + Nexus::uint32 (sequence) = 10 Nexus::bytes. + EXPECT_EQ(HEADER_SIZE, 10u); +} diff --git a/Tests/Engine/Network/PacketWriterTests.cpp b/Tests/Engine/Network/PacketWriterTests.cpp new file mode 100644 index 0000000..3bff3d4 --- /dev/null +++ b/Tests/Engine/Network/PacketWriterTests.cpp @@ -0,0 +1,155 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Packet; +import NE.Engine.Network.PacketWriter; + +import NE.Engine.Core.Types; + +import std; + +using namespace Nexus::Network; + +TEST(PacketWriterTest, BuildProducesFrameWithCorrectHeader) { + PacketWriter writer(PacketType::ChatMessage, 7); + writer.Write(Nexus::uint32{123}); + + const std::vector frame = writer.Build(); + ASSERT_GE(frame.size(), HEADER_SIZE); + + Nexus::usize offset = 0; + Nexus::uint32 payloadSize = 0; + Nexus::uint16 type = 0; + Nexus::uint32 sequence = 0; + + ASSERT_TRUE(Wire::Read(frame, offset, payloadSize)); + ASSERT_TRUE(Wire::Read(frame, offset, type)); + ASSERT_TRUE(Wire::Read(frame, offset, sequence)); + + EXPECT_EQ(payloadSize, sizeof(Nexus::uint32)); + EXPECT_EQ(type, static_cast(PacketType::ChatMessage)); + EXPECT_EQ(sequence, 7u); +} + +TEST(PacketWriterTest, BuildWithNoPayloadHasZeroPayloadSize) { + PacketWriter writer(PacketType::ClientHello, 1); + const std::vector frame = writer.Build(); + + ASSERT_EQ(frame.size(), HEADER_SIZE); + + Nexus::usize offset = 0; + Nexus::uint32 payloadSize = 0; + ASSERT_TRUE(Wire::Read(frame, offset, payloadSize)); + EXPECT_EQ(payloadSize, 0u); +} + +TEST(PacketWriterTest, FrameSizeIsHeaderPlusPayload) { + PacketWriter writer(PacketType::PlayerState, 1); + writer.Write(Nexus::uint32{1}); + writer.Write(Nexus::uint32{2}); + writer.Write(Nexus::uint32{3}); + + const std::vector frame = writer.Build(); + EXPECT_EQ(frame.size(), HEADER_SIZE + 3 * sizeof(Nexus::uint32)); +} + +TEST(PacketWriterTest, PayloadBytesFollowHeaderInFrame) { + PacketWriter writer(PacketType::PlayerInput, 1); + writer.Write(Nexus::uint8{0xAB}); + + const std::vector frame = writer.Build(); + ASSERT_EQ(frame.size(), HEADER_SIZE + 1); + EXPECT_EQ(static_cast(frame[HEADER_SIZE]), 0xAB); +} + +TEST(PacketWriterTest, WriteStringEncodesLengthPrefixAndBytes) { + PacketWriter writer(PacketType::ChatMessage, 1); + ASSERT_TRUE(writer.WriteString("hi")); + + const std::vector frame = writer.Build(); + ASSERT_EQ(frame.size(), HEADER_SIZE + sizeof(Nexus::uint32) + 2); + + Nexus::usize offset = HEADER_SIZE; + Nexus::uint32 length = 0; + ASSERT_TRUE(Wire::Read(frame, offset, length)); + EXPECT_EQ(length, 2u); + + EXPECT_EQ(static_cast(frame[offset]), 'h'); + EXPECT_EQ(static_cast(frame[offset + 1]), 'i'); +} + +TEST(PacketWriterTest, WriteStringHandlesEmptyString) { + PacketWriter writer(PacketType::ChatMessage, 1); + ASSERT_TRUE(writer.WriteString("")); + + const std::vector frame = writer.Build(); + EXPECT_EQ(frame.size(), HEADER_SIZE + sizeof(Nexus::uint32)); +} + +TEST(PacketWriterTest, MultipleWriteStringCallsAppendSequentially) { + PacketWriter writer(PacketType::ChatMessage, 1); + ASSERT_TRUE(writer.WriteString("abc")); + ASSERT_TRUE(writer.WriteString("de")); + + const std::vector frame = writer.Build(); + + Nexus::usize offset = HEADER_SIZE; + Nexus::uint32 firstLength = 0; + ASSERT_TRUE(Wire::Read(frame, offset, firstLength)); + EXPECT_EQ(firstLength, 3u); + offset += firstLength; + + Nexus::uint32 secondLength = 0; + ASSERT_TRUE(Wire::Read(frame, offset, secondLength)); + EXPECT_EQ(secondLength, 2u); +} + +TEST(PacketWriterTest, BuildCanOnlyBeCalledOnce) { + PacketWriter writer(PacketType::ChatMessage, 1); + writer.Write(Nexus::uint32{1}); + + const std::vector first = writer.Build(); + EXPECT_FALSE(first.empty()); + + const std::vector second = writer.Build(); + EXPECT_TRUE(second.empty()); +} + +TEST(PacketWriterTest, DataReflectsBufferBeforeBuild) { + PacketWriter writer(PacketType::ChatMessage, 1); + writer.Write(Nexus::uint32{42}); + + // Before Build(), Data() includes the still-unfilled 10-Nexus::byte header plus + // the payload written so far. + EXPECT_EQ(writer.Data().size(), HEADER_SIZE + sizeof(Nexus::uint32)); +} + +TEST(PacketWriterTest, IsMoveOnly) { + static_assert(!std::is_copy_constructible_v); + static_assert(!std::is_copy_assignable_v); + static_assert(std::is_move_constructible_v); + static_assert(std::is_move_assignable_v); + + SUCCEED(); +} + +TEST(PacketWriterTest, DifferentPacketTypesEncodeDifferentTypeValues) { + PacketWriter helloWriter(PacketType::ClientHello, 1); + PacketWriter chatWriter(PacketType::ChatMessage, 1); + + const std::vector helloFrame = helloWriter.Build(); + const std::vector chatFrame = chatWriter.Build(); + + Nexus::usize helloOffset = sizeof(Nexus::uint32); + Nexus::usize chatOffset = sizeof(Nexus::uint32); + Nexus::uint16 helloType = 0; + Nexus::uint16 chatType = 0; + + ASSERT_TRUE(Wire::Read(helloFrame, helloOffset, helloType)); + ASSERT_TRUE(Wire::Read(chatFrame, chatOffset, chatType)); + + EXPECT_NE(helloType, chatType); + EXPECT_EQ(helloType, static_cast(PacketType::ClientHello)); + EXPECT_EQ(chatType, static_cast(PacketType::ChatMessage)); +} diff --git a/Tests/Engine/Network/SocketTests.cpp b/Tests/Engine/Network/SocketTests.cpp new file mode 100644 index 0000000..fe9d80c --- /dev/null +++ b/Tests/Engine/Network/SocketTests.cpp @@ -0,0 +1,452 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; +import NE.Engine.Network.Platform.Socket; + +import NE.Engine.Core.Types; + +import std; + +using namespace Nexus::Network; + +namespace { + template + bool WaitUntil(Predicate predicate, std::chrono::milliseconds timeout = std::chrono::milliseconds(2000)) { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (std::chrono::steady_clock::now() < deadline) { + if (predicate()) { + return true; + } + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + return predicate(); + } + + Nexus::uint16 NextTestPort() { + // avoids collisions on ports that are still active/waiting to be closed + static std::atomic counter{51000}; + return counter.fetch_add(1); + } + + class SocketTest : public ::testing::Test {}; +} // namespace + +// ---- Construction / validity ------------------------------------------- + +TEST_F(SocketTest, DefaultConstructedIsInvalid) { + const Socket socket; + EXPECT_FALSE(socket.IsValid()); +} + +TEST_F(SocketTest, ConstructedWithTypeIsValid) { + const Socket tcp(SocketType::TCP); + EXPECT_TRUE(tcp.IsValid()); + + const Socket udp(SocketType::UDP); + EXPECT_TRUE(udp.IsValid()); +} + +TEST_F(SocketTest, CloseInvalidatesSocket) { + Socket socket(SocketType::TCP); + ASSERT_TRUE(socket.IsValid()); + + socket.Close(); + EXPECT_FALSE(socket.IsValid()); +} + +TEST_F(SocketTest, CloseOnAlreadyInvalidSocketIsSafe) { + Socket socket; + EXPECT_NO_THROW(socket.Close()); +} + +TEST_F(SocketTest, MoveConstructionTransfersValidity) { + Socket original(SocketType::TCP); + ASSERT_TRUE(original.IsValid()); + + Socket moved(std::move(original)); + EXPECT_TRUE(moved.IsValid()); + EXPECT_FALSE(original.IsValid()); +} + +TEST_F(SocketTest, MoveAssignmentTransfersValidityAndClosesTarget) { + Socket target(SocketType::TCP); + Socket source(SocketType::UDP); + ASSERT_TRUE(target.IsValid()); + ASSERT_TRUE(source.IsValid()); + + target = std::move(source); + EXPECT_TRUE(target.IsValid()); + EXPECT_FALSE(source.IsValid()); +} + +// ---- Bind / Listen on an invalid socket --------------------------------- + +TEST_F(SocketTest, BindOnInvalidSocketFails) { + Socket socket; + EXPECT_FALSE(socket.Bind(0)); +} + +TEST_F(SocketTest, ListenOnInvalidSocketFails) { + Socket socket; + EXPECT_FALSE(socket.Listen()); +} + +TEST_F(SocketTest, SetNonBlockingOnInvalidSocketFails) { + Socket socket; + EXPECT_FALSE(socket.SetNonBlocking(true)); +} + +// ---- Operation/type guards ------------------------------------------------ + +TEST_F(SocketTest, ConnectOnUdpSocketFailsWithInvalidOperation) { + Socket udp(SocketType::UDP); + std::optional error; + + EXPECT_FALSE(udp.Connect(NetworkAddress("127.0.0.1", 12345), error)); + ASSERT_TRUE(error.has_value()); + EXPECT_EQ(error->code, NetworkErrorCode::InvalidOperation); +} + +TEST_F(SocketTest, SendToOnTcpSocketFailsWithInvalidOperation) { + Socket tcp(SocketType::TCP); + std::optional error; + const std::array data{}; + + const auto result = tcp.SendTo(NetworkAddress("127.0.0.1", 12345), data, error); + EXPECT_FALSE(result.has_value()); + ASSERT_TRUE(error.has_value()); + EXPECT_EQ(error->code, NetworkErrorCode::InvalidOperation); +} + +TEST_F(SocketTest, SendOnUdpSocketFailsWithInvalidOperation) { + Socket udp(SocketType::UDP); + std::optional error; + const std::array data{}; + + const auto result = udp.Send(data, error); + EXPECT_FALSE(result.has_value()); + ASSERT_TRUE(error.has_value()); + EXPECT_EQ(error->code, NetworkErrorCode::InvalidOperation); +} + +TEST_F(SocketTest, ListenOnUdpSocketFails) { + Socket udp(SocketType::UDP); + EXPECT_FALSE(udp.Listen()); +} + +TEST_F(SocketTest, AcceptOnUdpSocketReturnsNullopt) { + Socket udp(SocketType::UDP); + EXPECT_FALSE(udp.Accept().has_value()); +} + +// ---- TCP: bind / listen / connect / accept ------------------------------- + +TEST_F(SocketTest, BindThenListenSucceeds) { + Socket listener(SocketType::TCP); + ASSERT_TRUE(listener.IsValid()); + ASSERT_TRUE(listener.SetNonBlocking(true)); + + const Nexus::uint16 port = NextTestPort(); + ASSERT_TRUE(listener.Bind(port)); + EXPECT_TRUE(listener.Listen()); +} + +TEST_F(SocketTest, AcceptOnListenerWithNoPendingConnectionReturnsNullopt) { + Socket listener(SocketType::TCP); + ASSERT_TRUE(listener.SetNonBlocking(true)); + const Nexus::uint16 port = NextTestPort(); + ASSERT_TRUE(listener.Bind(port)); + ASSERT_TRUE(listener.Listen()); + + EXPECT_FALSE(listener.Accept().has_value()); +} + +TEST_F(SocketTest, ClientCanConnectAndServerCanAccept) { + const Nexus::uint16 port = NextTestPort(); + + Socket listener(SocketType::TCP); + ASSERT_TRUE(listener.SetNonBlocking(true)); + ASSERT_TRUE(listener.Bind(port)); + ASSERT_TRUE(listener.Listen()); + + Socket client(SocketType::TCP); + ASSERT_TRUE(client.SetNonBlocking(true)); + + std::optional connectError; + EXPECT_TRUE(client.Connect(NetworkAddress("127.0.0.1", port), connectError)); + + std::optional accepted; + const bool got = WaitUntil([&] { + if (!accepted) { + accepted = listener.Accept(); + } + return accepted.has_value(); + }); + + ASSERT_TRUE(got); + EXPECT_TRUE(accepted->IsValid()); +} + +TEST_F(SocketTest, ConnectToNonBlockingSocketReportsConnectionInProgressOrSucceeds) { + const Nexus::uint16 port = NextTestPort(); + + Socket listener(SocketType::TCP); + ASSERT_TRUE(listener.SetNonBlocking(true)); + ASSERT_TRUE(listener.Bind(port)); + ASSERT_TRUE(listener.Listen()); + + Socket client(SocketType::TCP); + ASSERT_TRUE(client.SetNonBlocking(true)); + + std::optional error; + const bool started = client.Connect(NetworkAddress("127.0.0.1", port), error); + + EXPECT_TRUE(started); + if (error.has_value()) { + EXPECT_EQ(error->code, NetworkErrorCode::ConnectionInProgress); + } +} + +TEST_F(SocketTest, CompleteConnectEventuallySucceeds) { + const Nexus::uint16 port = NextTestPort(); + + Socket listener(SocketType::TCP); + ASSERT_TRUE(listener.SetNonBlocking(true)); + ASSERT_TRUE(listener.Bind(port)); + ASSERT_TRUE(listener.Listen()); + + Socket client(SocketType::TCP); + ASSERT_TRUE(client.SetNonBlocking(true)); + std::optional connectError; + ASSERT_TRUE(client.Connect(NetworkAddress("127.0.0.1", port), connectError)); + + const bool completed = WaitUntil([&] { + static_cast(listener.Accept()); + std::optional completeError; + return client.CompleteConnect(completeError); + }); + + EXPECT_TRUE(completed); +} + +TEST_F(SocketTest, ConnectToClosedPortFails) { + Socket client(SocketType::TCP); + ASSERT_TRUE(client.SetNonBlocking(true)); + + // Nothing is listening on this port. + const Nexus::uint16 port = NextTestPort(); + std::optional connectError; + const bool started = client.Connect(NetworkAddress("127.0.0.1", port), connectError); + + if (!started) { + ASSERT_TRUE(connectError.has_value()); + return; + } + + bool sawSpuriousSuccess = false; + WaitUntil([&] { + std::optional completeError; + if (client.CompleteConnect(completeError)) { + sawSpuriousSuccess = true; + return true; + } + return completeError && completeError->code != NetworkErrorCode::ConnectionInProgress; + }); + + EXPECT_FALSE(sawSpuriousSuccess); +} + +// ---- TCP: send / receive -------------------------------------------------- + +namespace { + struct ConnectedPair { + Socket client; + Socket server; + }; + + std::optional MakeConnectedPair(Nexus::uint16 port) { + Socket listener(SocketType::TCP); + if (!listener.SetNonBlocking(true) || !listener.Bind(port) || !listener.Listen()) { + return std::nullopt; + } + + Socket client(SocketType::TCP); + if (!client.SetNonBlocking(true)) { + return std::nullopt; + } + + std::optional connectError; + if (!client.Connect(NetworkAddress("127.0.0.1", port), connectError)) { + return std::nullopt; + } + + std::optional server; + const bool connected = WaitUntil([&] { + if (!server) { + server = listener.Accept(); + } + std::optional completeError; + const bool clientReady = client.CompleteConnect(completeError); + return server.has_value() && clientReady; + }); + + if (!connected || !server) { + return std::nullopt; + } + + return ConnectedPair{std::move(client), std::move(*server)}; + } +} // namespace + +TEST_F(SocketTest, SendThenReceiveDeliversData) { + auto pair = MakeConnectedPair(NextTestPort()); + ASSERT_TRUE(pair.has_value()); + + const std::array payload{Nexus::byte{'h'}, Nexus::byte{'e'}, Nexus::byte{'l'}, Nexus::byte{'l'}, + Nexus::byte{'o'}}; + std::optional sendError; + const auto sent = pair->client.Send(payload, sendError); + ASSERT_TRUE(sent.has_value()); + EXPECT_EQ(*sent, payload.size()); + + std::array buffer{}; + std::optional received; + const bool got = WaitUntil([&] { + std::optional receiveError; + received = pair->server.ReceiveInto(buffer, receiveError); + return received.has_value(); + }); + + ASSERT_TRUE(got); + ASSERT_EQ(*received, payload.size()); + EXPECT_TRUE(std::equal(payload.begin(), payload.end(), buffer.begin())); +} + +TEST_F(SocketTest, ReceiveOnSocketWithNoDataReturnsWouldBlock) { + auto pair = MakeConnectedPair(NextTestPort()); + ASSERT_TRUE(pair.has_value()); + + std::array buffer{}; + std::optional error; + const auto received = pair->server.ReceiveInto(buffer, error); + + EXPECT_FALSE(received.has_value()); + ASSERT_TRUE(error.has_value()); + EXPECT_EQ(error->code, NetworkErrorCode::WouldBlock); +} + +TEST_F(SocketTest, ReceiveAfterPeerClosesReportsConnectionClosed) { + auto pair = MakeConnectedPair(NextTestPort()); + ASSERT_TRUE(pair.has_value()); + + pair->client.Close(); + + std::array buffer{}; + std::optional received; + std::optional error; + const bool sawClose = WaitUntil([&] { + received = pair->server.ReceiveInto(buffer, error); + return !received.has_value() && error.has_value() && error->code == NetworkErrorCode::ConnectionClosed; + }); + + EXPECT_TRUE(sawClose); +} + +TEST_F(SocketTest, SendEmptySpanSucceedsWithZeroBytes) { + auto pair = MakeConnectedPair(NextTestPort()); + ASSERT_TRUE(pair.has_value()); + + std::optional error; + const auto sent = pair->client.Send(std::span{}, error); + + ASSERT_TRUE(sent.has_value()); + EXPECT_EQ(*sent, 0u); + EXPECT_FALSE(error.has_value()); +} + +// ---- UDP: bind / send / receive ------------------------------------------- + +TEST_F(SocketTest, UdpBindSucceeds) { + Socket udp(SocketType::UDP); + ASSERT_TRUE(udp.SetNonBlocking(true)); + EXPECT_TRUE(udp.Bind(NextTestPort())); +} + +TEST_F(SocketTest, UdpSendToThenReceiveFromDeliversDatagram) { + const Nexus::uint16 serverPort = NextTestPort(); + + Socket server(SocketType::UDP); + ASSERT_TRUE(server.SetNonBlocking(true)); + ASSERT_TRUE(server.Bind(serverPort)); + + Socket client(SocketType::UDP); + ASSERT_TRUE(client.SetNonBlocking(true)); + ASSERT_TRUE(client.Bind(0)); + + const std::array payload{Nexus::byte{1}, Nexus::byte{2}, Nexus::byte{3}}; + std::optional sendError; + const auto sent = client.SendTo(NetworkAddress("127.0.0.1", serverPort), payload, sendError); + ASSERT_TRUE(sent.has_value()); + EXPECT_EQ(*sent, payload.size()); + + std::array buffer{}; + NetworkAddress sender; + std::optional received; + const bool got = WaitUntil([&] { + std::optional receiveError; + received = server.ReceiveFrom(buffer, sender, receiveError); + return received.has_value(); + }); + + ASSERT_TRUE(got); + ASSERT_EQ(*received, payload.size()); + EXPECT_TRUE(std::equal(payload.begin(), payload.end(), buffer.begin())); + EXPECT_EQ(sender.Host(), "127.0.0.1"); +} + +TEST_F(SocketTest, UdpReceiveWithNoDataReturnsWouldBlock) { + Socket udp(SocketType::UDP); + ASSERT_TRUE(udp.SetNonBlocking(true)); + ASSERT_TRUE(udp.Bind(NextTestPort())); + + std::array buffer{}; + NetworkAddress sender; + std::optional error; + const auto received = udp.ReceiveFrom(buffer, sender, error); + + EXPECT_FALSE(received.has_value()); + ASSERT_TRUE(error.has_value()); + EXPECT_EQ(error->code, NetworkErrorCode::WouldBlock); +} + +TEST_F(SocketTest, UdpSenderAddressMatchesClientBoundPort) { + const Nexus::uint16 serverPort = NextTestPort(); + const Nexus::uint16 clientPort = NextTestPort(); + + Socket server(SocketType::UDP); + ASSERT_TRUE(server.SetNonBlocking(true)); + ASSERT_TRUE(server.Bind(serverPort)); + + Socket client(SocketType::UDP); + ASSERT_TRUE(client.SetNonBlocking(true)); + ASSERT_TRUE(client.Bind(clientPort)); + + const std::array payload{Nexus::byte{7}}; + std::optional sendError; + ASSERT_TRUE(client.SendTo(NetworkAddress("127.0.0.1", serverPort), payload, sendError).has_value()); + + std::array buffer{}; + NetworkAddress sender; + std::optional received; + const bool got = WaitUntil([&] { + std::optional receiveError; + received = server.ReceiveFrom(buffer, sender, receiveError); + return received.has_value(); + }); + + ASSERT_TRUE(got); + EXPECT_EQ(sender.Port(), clientPort); +} diff --git a/Tests/Engine/Network/StreamReassemblerTests.cpp b/Tests/Engine/Network/StreamReassemblerTests.cpp new file mode 100644 index 0000000..3b4199c --- /dev/null +++ b/Tests/Engine/Network/StreamReassemblerTests.cpp @@ -0,0 +1,169 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Packet; +import NE.Engine.Network.PacketWriter; +import NE.Engine.Network.TCP.StreamReassembler; + +import NE.Engine.Core.Types; + +import std; + +using namespace Nexus::Network; + +namespace { + class StreamReassemblerTest : public ::testing::Test { + protected: + StreamReassembler m_reassembler; + }; + + std::vector BuildFrame(PacketType type, Nexus::uint32 sequence, Nexus::uint32 payloadValue) { + PacketWriter writer(type, sequence); + writer.Write(payloadValue); + return writer.Build(); + } +} // namespace + +TEST_F(StreamReassemblerTest, TryExtractReturnsNulloptWhenEmpty) { + EXPECT_FALSE(m_reassembler.TryExtract().has_value()); +} + +TEST_F(StreamReassemblerTest, TryExtractReturnsNulloptForPartialHeader) { + const auto frame = BuildFrame(PacketType::ChatMessage, 1, 42); + m_reassembler.Feed(std::span(frame).first(3)); // fewer than HEADER_SIZE Bytes + + EXPECT_FALSE(m_reassembler.TryExtract().has_value()); +} + +TEST_F(StreamReassemblerTest, TryExtractReturnsNulloptForPartialPayload) { + const auto frame = BuildFrame(PacketType::ChatMessage, 1, 42); + // Feed the full header but only part of the payload. + m_reassembler.Feed(std::span(frame).first(HEADER_SIZE + 1)); + + EXPECT_FALSE(m_reassembler.TryExtract().has_value()); +} + +TEST_F(StreamReassemblerTest, ExtractsExactSingleFrame) { + const auto frame = BuildFrame(PacketType::ChatMessage, 1, 42); + m_reassembler.Feed(frame); + + const auto extracted = m_reassembler.TryExtract(); + ASSERT_TRUE(extracted.has_value()); + EXPECT_EQ(*extracted, frame); + + EXPECT_FALSE(m_reassembler.TryExtract().has_value()); +} + +TEST_F(StreamReassemblerTest, ExtractsMultipleFramesFedTogether) { + const auto first = BuildFrame(PacketType::ChatMessage, 1, 1); + const auto second = BuildFrame(PacketType::PlayerInput, 2, 2); + + std::vector combined; + combined.insert(combined.end(), first.begin(), first.end()); + combined.insert(combined.end(), second.begin(), second.end()); + m_reassembler.Feed(combined); + + const auto firstOut = m_reassembler.TryExtract(); + ASSERT_TRUE(firstOut.has_value()); + EXPECT_EQ(*firstOut, first); + + const auto secondOut = m_reassembler.TryExtract(); + ASSERT_TRUE(secondOut.has_value()); + EXPECT_EQ(*secondOut, second); + + EXPECT_FALSE(m_reassembler.TryExtract().has_value()); +} + +TEST_F(StreamReassemblerTest, ExtractsFrameSplitAcrossMultipleFeeds) { + const auto frame = BuildFrame(PacketType::ChatMessage, 1, 42); + + // Feed one Nexus::byte at a time to simulate a slow/fragmented TCP stream. + for (const Nexus::byte b : frame) { + m_reassembler.Feed(std::span(&b, 1)); + } + + const auto extracted = m_reassembler.TryExtract(); + ASSERT_TRUE(extracted.has_value()); + EXPECT_EQ(*extracted, frame); +} + +TEST_F(StreamReassemblerTest, FeedIgnoresEmptySpan) { + m_reassembler.Feed(std::span{}); + EXPECT_FALSE(m_reassembler.HasInvalidData()); + EXPECT_FALSE(m_reassembler.TryExtract().has_value()); +} + +TEST_F(StreamReassemblerTest, HasInvalidDataFalseInitially) { + EXPECT_FALSE(m_reassembler.HasInvalidData()); +} + +TEST_F(StreamReassemblerTest, OversizedPayloadMarksInvalid) { + // Hand-build a frame header claiming a payload larger than MAX_PAYLOAD_SIZE. + std::vector frame(HEADER_SIZE, Nexus::byte{0}); + Nexus::usize offset = 0; + const auto oversizedPayload = Wire::HostToBigEndian(MAX_PAYLOAD_SIZE + 1); + std::memcpy(frame.data() + offset, &oversizedPayload, sizeof(oversizedPayload)); + + m_reassembler.Feed(frame); + EXPECT_FALSE(m_reassembler.TryExtract().has_value()); + EXPECT_TRUE(m_reassembler.HasInvalidData()); +} + +TEST_F(StreamReassemblerTest, FeedingMoreThanCapacityMarksInvalid) { + const std::vector tooMuch(StreamReassembler::MAX_BUFFERED_BYTES + 1, Nexus::byte{0}); + m_reassembler.Feed(tooMuch); + + EXPECT_TRUE(m_reassembler.HasInvalidData()); +} + +TEST_F(StreamReassemblerTest, FeedAfterInvalidIsNoOp) { + const std::vector tooMuch(StreamReassembler::MAX_BUFFERED_BYTES + 1, Nexus::byte{0}); + m_reassembler.Feed(tooMuch); + ASSERT_TRUE(m_reassembler.HasInvalidData()); + + const auto frame = BuildFrame(PacketType::ChatMessage, 1, 42); + m_reassembler.Feed(frame); + + EXPECT_TRUE(m_reassembler.HasInvalidData()); + EXPECT_FALSE(m_reassembler.TryExtract().has_value()); +} + +TEST_F(StreamReassemblerTest, ResetClearsInvalidState) { + const std::vector tooMuch(StreamReassembler::MAX_BUFFERED_BYTES + 1, Nexus::byte{0}); + m_reassembler.Feed(tooMuch); + ASSERT_TRUE(m_reassembler.HasInvalidData()); + + m_reassembler.Reset(); + EXPECT_FALSE(m_reassembler.HasInvalidData()); + + const auto frame = BuildFrame(PacketType::ChatMessage, 1, 42); + m_reassembler.Feed(frame); + const auto extracted = m_reassembler.TryExtract(); + ASSERT_TRUE(extracted.has_value()); + EXPECT_EQ(*extracted, frame); +} + +TEST_F(StreamReassemblerTest, ResetClearsBufferedPartialFrame) { + const auto frame = BuildFrame(PacketType::ChatMessage, 1, 42); + m_reassembler.Feed(std::span(frame).first(HEADER_SIZE)); // header only, no payload yet + + m_reassembler.Reset(); + + m_reassembler.Feed(std::span(frame).subspan(HEADER_SIZE)); + EXPECT_FALSE(m_reassembler.TryExtract().has_value()); +} + +TEST_F(StreamReassemblerTest, PayloadContentIsPreservedThroughExtraction) { + const auto frame = BuildFrame(PacketType::PlayerState, 5, 0xCAFEBABE); + m_reassembler.Feed(frame); + + const auto extracted = m_reassembler.TryExtract(); + ASSERT_TRUE(extracted.has_value()); + + std::vector payload(extracted->begin() + static_cast(HEADER_SIZE), extracted->end()); + Packet reader(payload); + Nexus::uint32 value = 0; + ASSERT_TRUE(reader.Read(value)); + EXPECT_EQ(value, 0xCAFEBABEu); +} diff --git a/Tests/Engine/Network/TCPClientTests.cpp b/Tests/Engine/Network/TCPClientTests.cpp new file mode 100644 index 0000000..e967751 --- /dev/null +++ b/Tests/Engine/Network/TCPClientTests.cpp @@ -0,0 +1,311 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; +import NE.Engine.Network.Packet; +import NE.Engine.Network.PacketWriter; +import NE.Engine.Network.Platform.Socket; +import NE.Engine.Network.TCP.TCPClient; + +import NE.Engine.Core.Types; + +import std; + +using namespace Nexus::Network; + +namespace { + template + bool WaitUntil(Predicate predicate, std::chrono::milliseconds timeout = std::chrono::milliseconds(2000)) { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (std::chrono::steady_clock::now() < deadline) { + if (predicate()) { + return true; + } + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + return predicate(); + } + + Nexus::uint16 NextTestPort() { + static std::atomic counter{52000}; + return counter.fetch_add(1); + } + + struct RawListener { + Socket listener; + Nexus::uint16 port = 0; + + static std::optional Create() { + RawListener result; + result.port = NextTestPort(); + result.listener = Socket(SocketType::TCP); + if (!result.listener.IsValid() || !result.listener.SetNonBlocking(true) || + !result.listener.Bind(result.port) || !result.listener.Listen()) { + return std::nullopt; + } + return result; + } + + std::optional WaitForAccept() { + std::optional accepted; + WaitUntil([&] { + accepted = listener.Accept(); + return accepted.has_value(); + }); + return accepted; + } + }; + + std::vector BuildFrame(std::string_view text) { + PacketWriter writer(PacketType::ChatMessage, 1); + writer.WriteString(text); + return writer.Build(); + } + + class TCPClientTest : public ::testing::Test {}; +} // namespace + +TEST_F(TCPClientTest, DefaultConstructedIsNotConnected) { + const TCPClient client; + EXPECT_FALSE(client.IsConnected()); + EXPECT_FALSE(client.IsConnecting()); +} + +TEST_F(TCPClientTest, ConnectToClosedPortFailsOrEndsUpDisconnected) { + TCPClient client; + const Nexus::uint16 port = NextTestPort(); // nothing listening here + + const bool started = client.Connect(NetworkAddress("127.0.0.1", port)); + if (!started) { + EXPECT_FALSE(client.IsConnected()); + return; + } + + WaitUntil([&] { + client.Poll(); + return !client.IsConnecting(); + }); + + EXPECT_FALSE(client.IsConnected()); +} + +TEST_F(TCPClientTest, ConnectThenPollBecomesConnected) { + auto listener = RawListener::Create(); + ASSERT_TRUE(listener.has_value()); + + TCPClient client; + ASSERT_TRUE(client.Connect(NetworkAddress("127.0.0.1", listener->port))); + + auto accepted = listener->WaitForAccept(); + ASSERT_TRUE(accepted.has_value()); + + const bool connected = WaitUntil([&] { + client.Poll(); + return client.IsConnected(); + }); + + EXPECT_TRUE(connected); + EXPECT_FALSE(client.IsConnecting()); +} + +TEST_F(TCPClientTest, SendBeforeFullyConnectedIsQueuedAndFlushedOnPoll) { + auto listener = RawListener::Create(); + ASSERT_TRUE(listener.has_value()); + + TCPClient client; + ASSERT_TRUE(client.Connect(NetworkAddress("127.0.0.1", listener->port))); + + const auto frame = BuildFrame("queued while connecting"); + EXPECT_TRUE(client.Send(frame)); + + auto accepted = listener->WaitForAccept(); + ASSERT_TRUE(accepted.has_value()); + + WaitUntil([&] { + client.Poll(); + return client.IsConnected(); + }); + + std::array buffer{}; + std::optional received; + const bool got = WaitUntil([&] { + std::optional error; + received = accepted->ReceiveInto(buffer, error); + return received.has_value(); + }); + + ASSERT_TRUE(got); + ASSERT_EQ(*received, frame.size()); + EXPECT_TRUE(std::equal(frame.begin(), frame.end(), buffer.begin())); +} + +TEST_F(TCPClientTest, PollDeliversFramesSentByPeer) { + auto listener = RawListener::Create(); + ASSERT_TRUE(listener.has_value()); + + TCPClient client; + ASSERT_TRUE(client.Connect(NetworkAddress("127.0.0.1", listener->port))); + + auto accepted = listener->WaitForAccept(); + ASSERT_TRUE(accepted.has_value()); + ASSERT_TRUE(accepted->SetNonBlocking(true)); + + const auto frame = BuildFrame("hello from peer"); + std::optional sendError; + ASSERT_TRUE(accepted->Send(frame, sendError).has_value()); + + std::vector> received; + const bool got = WaitUntil([&] { + auto frames = client.Poll(); + received.insert(received.end(), std::make_move_iterator(frames.begin()), std::make_move_iterator(frames.end())); + return !received.empty(); + }); + + ASSERT_TRUE(got); + ASSERT_EQ(received.size(), 1u); + EXPECT_EQ(received.front(), frame); +} + +TEST_F(TCPClientTest, PollDeliversMultipleQueuedFrames) { + auto listener = RawListener::Create(); + ASSERT_TRUE(listener.has_value()); + + TCPClient client; + ASSERT_TRUE(client.Connect(NetworkAddress("127.0.0.1", listener->port))); + + auto accepted = listener->WaitForAccept(); + ASSERT_TRUE(accepted.has_value()); + ASSERT_TRUE(accepted->SetNonBlocking(true)); + + const auto first = BuildFrame("first"); + const auto second = BuildFrame("second"); + + std::vector combined; + combined.insert(combined.end(), first.begin(), first.end()); + combined.insert(combined.end(), second.begin(), second.end()); + + std::optional sendError; + ASSERT_TRUE(accepted->Send(combined, sendError).has_value()); + + std::vector> received; + WaitUntil([&] { + auto frames = client.Poll(); + received.insert(received.end(), std::make_move_iterator(frames.begin()), std::make_move_iterator(frames.end())); + return received.size() >= 2; + }); + + ASSERT_EQ(received.size(), 2u); + EXPECT_EQ(received[0], first); + EXPECT_EQ(received[1], second); +} + +TEST_F(TCPClientTest, DisconnectsWhenPeerCloses) { + auto listener = RawListener::Create(); + ASSERT_TRUE(listener.has_value()); + + TCPClient client; + ASSERT_TRUE(client.Connect(NetworkAddress("127.0.0.1", listener->port))); + + auto accepted = listener->WaitForAccept(); + ASSERT_TRUE(accepted.has_value()); + + WaitUntil([&] { + client.Poll(); + return client.IsConnected(); + }); + ASSERT_TRUE(client.IsConnected()); + + accepted->Close(); + + const bool disconnected = WaitUntil([&] { + client.Poll(); + return !client.IsConnected(); + }); + + EXPECT_TRUE(disconnected); +} + +TEST_F(TCPClientTest, SendOnDisconnectedClientFails) { + TCPClient client; + const auto frame = BuildFrame("nobody's listening"); + + EXPECT_FALSE(client.Send(frame)); +} + +TEST_F(TCPClientTest, SendEmptyFrameFails) { + auto listener = RawListener::Create(); + ASSERT_TRUE(listener.has_value()); + + TCPClient client; + ASSERT_TRUE(client.Connect(NetworkAddress("127.0.0.1", listener->port))); + + EXPECT_FALSE(client.Send(std::vector{})); +} + +TEST_F(TCPClientTest, DisconnectResetsState) { + auto listener = RawListener::Create(); + ASSERT_TRUE(listener.has_value()); + + TCPClient client; + ASSERT_TRUE(client.Connect(NetworkAddress("127.0.0.1", listener->port))); + + auto accepted = listener->WaitForAccept(); + ASSERT_TRUE(accepted.has_value()); + WaitUntil([&] { + client.Poll(); + return client.IsConnected(); + }); + ASSERT_TRUE(client.IsConnected()); + + client.Disconnect(); + + EXPECT_FALSE(client.IsConnected()); + EXPECT_FALSE(client.IsConnecting()); + EXPECT_TRUE(client.Poll().empty()); +} + +TEST_F(TCPClientTest, ConstructedFromAcceptedSocketIsImmediatelyConnected) { + auto listener = RawListener::Create(); + ASSERT_TRUE(listener.has_value()); + + Socket rawClient(SocketType::TCP); + ASSERT_TRUE(rawClient.SetNonBlocking(true)); + std::optional connectError; + ASSERT_TRUE(rawClient.Connect(NetworkAddress("127.0.0.1", listener->port), connectError)); + + auto accepted = listener->WaitForAccept(); + ASSERT_TRUE(accepted.has_value()); + + TCPClient serverSideClient(std::move(*accepted)); + EXPECT_TRUE(serverSideClient.IsConnected()); + EXPECT_FALSE(serverSideClient.IsConnecting()); +} + +TEST_F(TCPClientTest, ReconnectAfterDisconnectStartsFresh) { + auto listener = RawListener::Create(); + ASSERT_TRUE(listener.has_value()); + + TCPClient client; + ASSERT_TRUE(client.Connect(NetworkAddress("127.0.0.1", listener->port))); + + auto firstAccepted = listener->WaitForAccept(); + ASSERT_TRUE(firstAccepted.has_value()); + WaitUntil([&] { + client.Poll(); + return client.IsConnected(); + }); + client.Disconnect(); + + // Reconnect to the same listener. + ASSERT_TRUE(client.Connect(NetworkAddress("127.0.0.1", listener->port))); + auto secondAccepted = listener->WaitForAccept(); + ASSERT_TRUE(secondAccepted.has_value()); + + const bool reconnected = WaitUntil([&] { + client.Poll(); + return client.IsConnected(); + }); + EXPECT_TRUE(reconnected); +} diff --git a/Tests/Engine/Network/TCPServerTests.cpp b/Tests/Engine/Network/TCPServerTests.cpp new file mode 100644 index 0000000..506ab10 --- /dev/null +++ b/Tests/Engine/Network/TCPServerTests.cpp @@ -0,0 +1,311 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; +import NE.Engine.Network.Packet; +import NE.Engine.Network.PacketWriter; +import NE.Engine.Network.Platform.Socket; +import NE.Engine.Network.TCP.TCPServer; + +import NE.Engine.Core.Types; + +import std; + +using namespace Nexus::Network; + +namespace { + template + bool WaitUntil(Predicate predicate, std::chrono::milliseconds timeout = std::chrono::milliseconds(2000)) { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (std::chrono::steady_clock::now() < deadline) { + if (predicate()) { + return true; + } + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + return predicate(); + } + + Nexus::uint16 NextTestPort() { + static std::atomic counter{53000}; + return counter.fetch_add(1); + } + + std::optional ConnectRawClient(Nexus::uint16 port) { + Socket client(SocketType::TCP); + if (!client.IsValid() || !client.SetNonBlocking(true)) { + return std::nullopt; + } + std::optional error; + if (!client.Connect(NetworkAddress("127.0.0.1", port), error)) { + return std::nullopt; + } + + const bool ready = WaitUntil([&] { + std::optional completeError; + return client.CompleteConnect(completeError); + }); + if (!ready) { + return std::nullopt; + } + return client; + } + + std::vector BuildFrame(std::string_view text) { + PacketWriter writer(PacketType::ChatMessage, 1); + writer.WriteString(text); + return writer.Build(); + } + + class TCPServerTest : public ::testing::Test {}; +} // namespace + +TEST_F(TCPServerTest, ListenSucceeds) { + TCPServer server; + EXPECT_TRUE(server.Listen(NextTestPort())); +} + +TEST_F(TCPServerTest, ClientCountStartsAtZero) { + TCPServer server; + ASSERT_TRUE(server.Listen(NextTestPort())); + EXPECT_EQ(server.ClientCount(), 0u); +} + +TEST_F(TCPServerTest, AcceptPendingWithNoConnectionsReturnsEmpty) { + TCPServer server; + ASSERT_TRUE(server.Listen(NextTestPort())); + EXPECT_TRUE(server.AcceptPending().empty()); +} + +TEST_F(TCPServerTest, AcceptsIncomingConnection) { + TCPServer server; + const Nexus::uint16 port = NextTestPort(); + ASSERT_TRUE(server.Listen(port)); + + auto client = ConnectRawClient(port); + ASSERT_TRUE(client.has_value()); + + std::vector accepted; + const bool got = WaitUntil([&] { + auto ids = server.AcceptPending(); + accepted.insert(accepted.end(), ids.begin(), ids.end()); + return !accepted.empty(); + }); + + ASSERT_TRUE(got); + EXPECT_EQ(accepted.size(), 1u); + EXPECT_EQ(server.ClientCount(), 1u); +} + +TEST_F(TCPServerTest, ClientIdsAreUniqueAndIncrementing) { + TCPServer server; + const Nexus::uint16 port = NextTestPort(); + ASSERT_TRUE(server.Listen(port)); + + auto firstClient = ConnectRawClient(port); + ASSERT_TRUE(firstClient.has_value()); + std::vector firstIds; + WaitUntil([&] { + auto ids = server.AcceptPending(); + firstIds.insert(firstIds.end(), ids.begin(), ids.end()); + return !firstIds.empty(); + }); + ASSERT_EQ(firstIds.size(), 1u); + + auto secondClient = ConnectRawClient(port); + ASSERT_TRUE(secondClient.has_value()); + std::vector secondIds; + WaitUntil([&] { + auto ids = server.AcceptPending(); + secondIds.insert(secondIds.end(), ids.begin(), ids.end()); + return !secondIds.empty(); + }); + ASSERT_EQ(secondIds.size(), 1u); + + EXPECT_NE(firstIds.front(), secondIds.front()); + EXPECT_GT(secondIds.front(), firstIds.front()); +} + +TEST_F(TCPServerTest, PollAllReturnsFramesSentByClient) { + TCPServer server; + const Nexus::uint16 port = NextTestPort(); + ASSERT_TRUE(server.Listen(port)); + + auto client = ConnectRawClient(port); + ASSERT_TRUE(client.has_value()); + + ClientId id = 0; + WaitUntil([&] { + auto ids = server.AcceptPending(); + if (!ids.empty()) { + id = ids.front(); + } + return id != 0; + }); + ASSERT_NE(id, 0u); + + const auto frame = BuildFrame("hello server"); + std::optional sendError; + ASSERT_TRUE(client->Send(frame, sendError).has_value()); + + std::vector>> received; + const bool got = WaitUntil([&] { + auto frames = server.PollAll(); + received.insert(received.end(), std::make_move_iterator(frames.begin()), std::make_move_iterator(frames.end())); + return !received.empty(); + }); + + ASSERT_TRUE(got); + EXPECT_EQ(received.front().first, id); + EXPECT_EQ(received.front().second, frame); +} + +TEST_F(TCPServerTest, SendToDeliversFrameToCorrectClient) { + TCPServer server; + const Nexus::uint16 port = NextTestPort(); + ASSERT_TRUE(server.Listen(port)); + + auto client = ConnectRawClient(port); + ASSERT_TRUE(client.has_value()); + ASSERT_TRUE(client->SetNonBlocking(true)); + + ClientId id = 0; + WaitUntil([&] { + auto ids = server.AcceptPending(); + if (!ids.empty()) { + id = ids.front(); + } + return id != 0; + }); + ASSERT_NE(id, 0u); + + const auto frame = BuildFrame("hello client"); + EXPECT_TRUE(server.SendTo(id, frame)); + + std::array buffer{}; + std::optional received; + const bool got = WaitUntil([&] { + std::optional error; + received = client->ReceiveInto(buffer, error); + return received.has_value(); + }); + + ASSERT_TRUE(got); + ASSERT_EQ(*received, frame.size()); + EXPECT_TRUE(std::equal(frame.begin(), frame.end(), buffer.begin())); +} + +TEST_F(TCPServerTest, SendToUnknownClientIdFails) { + TCPServer server; + ASSERT_TRUE(server.Listen(NextTestPort())); + + const auto frame = BuildFrame("nobody home"); + EXPECT_FALSE(server.SendTo(999, frame)); +} + +TEST_F(TCPServerTest, BroadcastReachesAllConnectedClients) { + TCPServer server; + const Nexus::uint16 port = NextTestPort(); + ASSERT_TRUE(server.Listen(port)); + + auto clientA = ConnectRawClient(port); + auto clientB = ConnectRawClient(port); + ASSERT_TRUE(clientA.has_value()); + ASSERT_TRUE(clientB.has_value()); + ASSERT_TRUE(clientA->SetNonBlocking(true)); + ASSERT_TRUE(clientB->SetNonBlocking(true)); + + WaitUntil([&] { + server.AcceptPending(); + return server.ClientCount() == 2; + }); + ASSERT_EQ(server.ClientCount(), 2u); + + const auto frame = BuildFrame("hello everyone"); + server.Broadcast(frame); + + std::array bufferA{}; + std::array bufferB{}; + std::optional receivedA; + std::optional receivedB; + + const bool gotBoth = WaitUntil([&] { + std::optional errorA; + std::optional errorB; + if (!receivedA) { + receivedA = clientA->ReceiveInto(bufferA, errorA); + } + if (!receivedB) { + receivedB = clientB->ReceiveInto(bufferB, errorB); + } + return receivedA.has_value() && receivedB.has_value(); + }); + + ASSERT_TRUE(gotBoth); + EXPECT_TRUE(std::equal(frame.begin(), frame.end(), bufferA.begin())); + EXPECT_TRUE(std::equal(frame.begin(), frame.end(), bufferB.begin())); +} + +TEST_F(TCPServerTest, DisconnectedClientIsRemovedOnPollAll) { + TCPServer server; + const Nexus::uint16 port = NextTestPort(); + ASSERT_TRUE(server.Listen(port)); + + auto client = ConnectRawClient(port); + ASSERT_TRUE(client.has_value()); + + WaitUntil([&] { + server.AcceptPending(); + return server.ClientCount() == 1; + }); + ASSERT_EQ(server.ClientCount(), 1u); + + client->Close(); + + const bool removed = WaitUntil([&] { + server.PollAll(); + return server.ClientCount() == 0; + }); + + EXPECT_TRUE(removed); +} + +TEST_F(TCPServerTest, StopClearsClientsAndListenSocket) { + TCPServer server; + const Nexus::uint16 port = NextTestPort(); + ASSERT_TRUE(server.Listen(port)); + + auto client = ConnectRawClient(port); + ASSERT_TRUE(client.has_value()); + WaitUntil([&] { + server.AcceptPending(); + return server.ClientCount() == 1; + }); + ASSERT_EQ(server.ClientCount(), 1u); + + server.Stop(); + + EXPECT_EQ(server.ClientCount(), 0u); + EXPECT_TRUE(server.AcceptPending().empty()); +} + +TEST_F(TCPServerTest, ListenAgainResetsExistingClients) { + TCPServer server; + const Nexus::uint16 firstPort = NextTestPort(); + ASSERT_TRUE(server.Listen(firstPort)); + + auto client = ConnectRawClient(firstPort); + ASSERT_TRUE(client.has_value()); + WaitUntil([&] { + server.AcceptPending(); + return server.ClientCount() == 1; + }); + ASSERT_EQ(server.ClientCount(), 1u); + + const Nexus::uint16 secondPort = NextTestPort(); + ASSERT_TRUE(server.Listen(secondPort)); + + EXPECT_EQ(server.ClientCount(), 0u); +} diff --git a/Tests/Engine/Network/UDPSocketTests.cpp b/Tests/Engine/Network/UDPSocketTests.cpp new file mode 100644 index 0000000..f81f11a --- /dev/null +++ b/Tests/Engine/Network/UDPSocketTests.cpp @@ -0,0 +1,167 @@ +// SPDX-License-Identifier: MIT + +#include + +import NE.Engine.Network.Common.NetworkAddress; +import NE.Engine.Network.Common.NetworkError; +import NE.Engine.Network.UDP.UDPSocket; + +import NE.Engine.Core.Types; + +import std; + +using namespace Nexus::Network; + +namespace { + template + bool WaitUntil(Predicate predicate, std::chrono::milliseconds timeout = std::chrono::milliseconds(2000)) { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (std::chrono::steady_clock::now() < deadline) { + if (predicate()) { + return true; + } + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + return predicate(); + } + + Nexus::uint16 NextTestPort() { + static std::atomic counter{54000}; + return counter.fetch_add(1); + } + + class UDPSocketTest : public ::testing::Test {}; +} // namespace + +TEST_F(UDPSocketTest, DefaultConstructedIsNotOpen) { + const UDPSocket socket; + EXPECT_FALSE(socket.IsOpen()); +} + +TEST_F(UDPSocketTest, BindSucceedsAndOpensSocket) { + UDPSocket socket; + EXPECT_TRUE(socket.Bind(NextTestPort())); + EXPECT_TRUE(socket.IsOpen()); +} + +TEST_F(UDPSocketTest, BindToEphemeralPortSucceeds) { + UDPSocket socket; + EXPECT_TRUE(socket.Bind(0)); + EXPECT_TRUE(socket.IsOpen()); +} + +TEST_F(UDPSocketTest, CloseInvalidatesSocket) { + UDPSocket socket; + ASSERT_TRUE(socket.Bind(NextTestPort())); + + socket.Close(); + EXPECT_FALSE(socket.IsOpen()); +} + +TEST_F(UDPSocketTest, SendOnUnboundSocketFails) { + UDPSocket socket; + const std::array data{}; + + EXPECT_FALSE(socket.Send(NetworkAddress("127.0.0.1", 1234), data)); +} + +TEST_F(UDPSocketTest, ReceiveOnUnboundSocketFailsWithInvalidOperation) { + UDPSocket socket; + std::optional error; + + const auto message = socket.Receive(error); + + EXPECT_FALSE(message.has_value()); + ASSERT_TRUE(error.has_value()); + EXPECT_EQ(error->code, NetworkErrorCode::InvalidOperation); +} + +TEST_F(UDPSocketTest, SendThenReceiveRoundTripsPayload) { + const Nexus::uint16 serverPort = NextTestPort(); + + UDPSocket server; + ASSERT_TRUE(server.Bind(serverPort)); + + UDPSocket client; + ASSERT_TRUE(client.Bind(0)); + + const std::array payload{Nexus::byte{'p'}, Nexus::byte{'i'}, Nexus::byte{'n'}, Nexus::byte{'g'}}; + EXPECT_TRUE(client.Send(NetworkAddress("127.0.0.1", serverPort), payload)); + + std::optional message; + const bool got = WaitUntil([&] { + std::optional error; + message = server.Receive(error); + return message.has_value(); + }); + + ASSERT_TRUE(got); + ASSERT_EQ(message->data.size(), payload.size()); + EXPECT_TRUE(std::equal(payload.begin(), payload.end(), message->data.begin())); + EXPECT_EQ(message->sender.Host(), "127.0.0.1"); +} + +TEST_F(UDPSocketTest, ReceiveWithNoDatagramReturnsNulloptWithoutError) { + UDPSocket socket; + ASSERT_TRUE(socket.Bind(NextTestPort())); + + std::optional error; + const auto message = socket.Receive(error); + + EXPECT_FALSE(message.has_value()); + if (error.has_value()) { + EXPECT_EQ(error->code, NetworkErrorCode::WouldBlock); + } +} + +TEST_F(UDPSocketTest, SendPayloadLargerThanMaxUdpPayloadFails) { + UDPSocket socket; + ASSERT_TRUE(socket.Bind(NextTestPort())); + + const std::vector oversized(MAX_UDP_PAYLOAD + 1, Nexus::byte{0}); + EXPECT_FALSE(socket.Send(NetworkAddress("127.0.0.1", NextTestPort()), oversized)); +} + +TEST_F(UDPSocketTest, MultipleDatagramsAreReceivedInSeparateCalls) { + const Nexus::uint16 serverPort = NextTestPort(); + + UDPSocket server; + ASSERT_TRUE(server.Bind(serverPort)); + + UDPSocket client; + ASSERT_TRUE(client.Bind(0)); + + const std::array first{Nexus::byte{1}}; + const std::array second{Nexus::byte{2}}; + EXPECT_TRUE(client.Send(NetworkAddress("127.0.0.1", serverPort), first)); + EXPECT_TRUE(client.Send(NetworkAddress("127.0.0.1", serverPort), second)); + + std::vector receivedValues; + WaitUntil([&] { + std::optional error; + if (auto message = server.Receive(error)) { + receivedValues.insert(receivedValues.end(), message->data.begin(), message->data.end()); + } + return receivedValues.size() >= 2; + }); + + ASSERT_EQ(receivedValues.size(), 2u); + EXPECT_EQ(static_cast(receivedValues[0]), 1); + EXPECT_EQ(static_cast(receivedValues[1]), 2); +} + +TEST_F(UDPSocketTest, IsMoveOnly) { + EXPECT_FALSE(std::is_copy_constructible_v); + EXPECT_FALSE(std::is_copy_assignable_v); + EXPECT_TRUE(std::is_move_constructible_v); + EXPECT_TRUE(std::is_move_assignable_v); +} + +TEST_F(UDPSocketTest, RebindClosesPreviousSocket) { + UDPSocket socket; + ASSERT_TRUE(socket.Bind(NextTestPort())); + ASSERT_TRUE(socket.IsOpen()); + + ASSERT_TRUE(socket.Bind(NextTestPort())); + EXPECT_TRUE(socket.IsOpen()); +} From 0523a799e8372e7e3c75448ebc0925d2e9c19743 Mon Sep 17 00:00:00 2001 From: Dylan Hollemaert Date: Thu, 3 Sep 2026 20:30:22 +0200 Subject: [PATCH 2/3] Fix: Set servers non blocking. Signed-off-by: Dylan Hollemaert --- Tests/Engine/Network/NetworkManagerTests.cpp | 5 +++-- Tests/Engine/Network/SocketTests.cpp | 4 +++- Tests/Engine/Network/TCPClientTests.cpp | 1 + 3 files changed, 7 insertions(+), 3 deletions(-) diff --git a/Tests/Engine/Network/NetworkManagerTests.cpp b/Tests/Engine/Network/NetworkManagerTests.cpp index bb58798..1901fd6 100644 --- a/Tests/Engine/Network/NetworkManagerTests.cpp +++ b/Tests/Engine/Network/NetworkManagerTests.cpp @@ -53,13 +53,14 @@ namespace { bool WaitForAccept() { return WaitUntil([&] { + bool nonBlocking = false; if (!accepted) { accepted = listener.Accept(); if (accepted) { - static_cast(accepted->SetNonBlocking(true)); + nonBlocking = accepted->SetNonBlocking(true); } } - return accepted.has_value(); + return accepted.has_value() && nonBlocking; }); } }; diff --git a/Tests/Engine/Network/SocketTests.cpp b/Tests/Engine/Network/SocketTests.cpp index fe9d80c..e543834 100644 --- a/Tests/Engine/Network/SocketTests.cpp +++ b/Tests/Engine/Network/SocketTests.cpp @@ -285,12 +285,14 @@ namespace { std::optional server; const bool connected = WaitUntil([&] { + bool nonBlocking = false; if (!server) { server = listener.Accept(); + nonBlocking = server->SetNonBlocking(true); } std::optional completeError; const bool clientReady = client.CompleteConnect(completeError); - return server.has_value() && clientReady; + return server.has_value() && clientReady && nonBlocking; }); if (!connected || !server) { diff --git a/Tests/Engine/Network/TCPClientTests.cpp b/Tests/Engine/Network/TCPClientTests.cpp index e967751..3de54b1 100644 --- a/Tests/Engine/Network/TCPClientTests.cpp +++ b/Tests/Engine/Network/TCPClientTests.cpp @@ -54,6 +54,7 @@ namespace { accepted = listener.Accept(); return accepted.has_value(); }); + const bool nonBlocking = accepted->SetNonBlocking(true); return accepted; } }; From ef987ba01dcffa9269e461bc81409bbbdf07de90 Mon Sep 17 00:00:00 2001 From: Dylan Hollemaert <92250732+dark-dylan-dev@users.noreply.github.com> Date: Thu, 3 Sep 2026 21:07:34 +0200 Subject: [PATCH 3/3] Make the deploy job only run on master --- .github/workflows/ci.yml | 1 + 1 file changed, 1 insertion(+) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 91ea320..96ee40e 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -166,6 +166,7 @@ jobs: path: Docs/Generated_HTML deploy: + if: github.event_name == 'push' && github.ref == 'refs/heads/master' needs: docs environment: name: github-pages