diff --git a/include/Cubed/gameplay/packet.hpp b/include/Cubed/gameplay/packet.hpp index 5e7b764..fe58b0c 100644 --- a/include/Cubed/gameplay/packet.hpp +++ b/include/Cubed/gameplay/packet.hpp @@ -1,17 +1,20 @@ #pragma once +#include "Cubed/tools/compression.hpp" #include "packet.pb.h" // IWYU pragma: keep -#include +#include #include #include #include #include +#include #include #include #include - namespace Cubed { -constexpr int HEADER_LEN = 12; +constexpr size_t HEADER_LEN = + sizeof(uint16_t) + sizeof(uint16_t) + sizeof(uint32_t) + sizeof(uint32_t); +constexpr size_t PACKET_COMPRESSION_THRESHOLD = 100; using Packet = std::shared_ptr>; enum class CompressType : uint16_t { NONE = 0, @@ -27,17 +30,17 @@ inline CompressType get_compress_type(uint16_t id) { case to(ZSTD): return ZSTD; } - throw std::runtime_error(std::format("Unkown CompressType {}", id)); + throw std::runtime_error(std::format("Unknown CompressType {}", id)); } struct PacketHeader { - uint16_t cmd; - CompressType compress_type; // 0=none 1=zlib - uint32_t uncompressed_size; - uint32_t compressed_size; + uint16_t cmd{}; + CompressType compress_type{}; // 0=none 1=zlib + uint32_t uncompressed_size{}; + uint32_t compressed_size{}; }; -enum class PacketEnum { +enum class PacketEnum : uint16_t { LOGIN_REQ = 1001, LOGIN_RSP = 1002, LOGOUT_REQ = 1003, @@ -91,20 +94,40 @@ template constexpr uint16_t get_packet_id() { } else if constexpr (is_same_v) { return to_num(PONG); } else { - static_assert(always_false::value, "Unkonw Type"); + static_assert(always_false::value, "Unknown Type"); } } -template Packet make_packet(const T& msg) { +template + requires std::derived_from +Packet make_packet(const T& msg) { PacketHeader header{}; header.cmd = get_packet_id(); - uint32_t size = static_cast(msg.ByteSizeLong()); - header.uncompressed_size = size; - header.compressed_size = size; - header.compress_type = CompressType::NONE; + uint32_t raw_size = static_cast(msg.ByteSizeLong()); + std::vector raw(raw_size); - auto packet = std::make_shared>( - HEADER_LEN + header.compressed_size); + if (!msg.SerializeToArray(raw.data(), raw_size)) { + return {}; + } + std::vector payload; + if (raw_size >= PACKET_COMPRESSION_THRESHOLD) { + std::vector compressed = compress_data(raw); + if (compressed.size() < raw.size()) { + payload = std::move(compressed); + header.compress_type = CompressType::ZSTD; + } else { + payload = std::move(raw); + header.compress_type = CompressType::NONE; + } + } else { + payload = std::move(raw); + header.compress_type = CompressType::NONE; + } + header.uncompressed_size = raw_size; + header.compressed_size = static_cast(payload.size()); + + auto packet = + std::make_shared>(HEADER_LEN + payload.size()); uint16_t cmd_net = htons(header.cmd); uint16_t compress_type_net = @@ -120,16 +143,14 @@ template Packet make_packet(const T& msg) { sizeof(uncompressed_size_net)); std::memcpy(packet->data() + 8, &compressed_size_net, sizeof(compressed_size_net)); - if (!msg.SerializeToArray(packet->data() + HEADER_LEN, - static_cast(size))) { - return {}; - } + std::memcpy(packet->data() + HEADER_LEN, payload.data(), payload.size()); return packet; } -inline PacketHeader -decode_packet_header(const std::array& header) { +inline PacketHeader decode_packet_header(std::span header) { + if (header.size() < HEADER_LEN) + throw std::runtime_error("Invalid header"); uint16_t cmd_net; uint16_t compress_type_net; uint32_t uncompressed_size_net; @@ -145,5 +166,31 @@ decode_packet_header(const std::array& header) { return {ntohs(cmd_net), get_compress_type(ntohs(compress_type_net)), ntohl(uncompressed_size_net), ntohl(compressed_size_net)}; } +template + requires std::derived_from +bool decode_packet(T& message, std::span data, + const PacketHeader& header) { + if (data.size() != header.compressed_size) { + return false; + } + + if (header.compress_type == CompressType::NONE && + header.uncompressed_size != header.compressed_size) { + return false; + } + + switch (header.compress_type) { + case CompressType::NONE: { + return message.ParseFromArray( + data.data(), static_cast(header.uncompressed_size)); + } + case CompressType::ZSTD: { + auto raw = decompress_data(data, header.uncompressed_size); + return message.ParseFromArray(raw.data(), static_cast(raw.size())); + } + default: + return false; + } +} } // namespace Cubed diff --git a/include/Cubed/tools/compression.hpp b/include/Cubed/tools/compression.hpp new file mode 100644 index 0000000..fe556c1 --- /dev/null +++ b/include/Cubed/tools/compression.hpp @@ -0,0 +1,40 @@ +#pragma once +#include +#include +#include +#include +#include +#include +namespace Cubed { +constexpr int DEFAULT_ZSTD_LEVEL = 3; +inline std::vector compress_data(std::span data) { + size_t max_size = ZSTD_compressBound(data.size()); + std::vector compressed_data(max_size); + size_t compressed_bytes = + ZSTD_compress(compressed_data.data(), max_size, data.data(), + data.size(), DEFAULT_ZSTD_LEVEL); + if (ZSTD_isError(compressed_bytes)) { + throw std::runtime_error(std::format( + "Compress Fail {}", ZSTD_getErrorName(compressed_bytes))); + } + compressed_data.resize(compressed_bytes); + return compressed_data; +} + +inline std::vector decompress_data(std::span data, + uint32_t original_size) { + std::vector decompressed_data(original_size); + size_t decompressed_bytes = ZSTD_decompress( + decompressed_data.data(), original_size, data.data(), data.size()); + if (ZSTD_isError(decompressed_bytes)) { + throw std::runtime_error(std::format( + "Decompress Fail {}", ZSTD_getErrorName(decompressed_bytes))); + } + if (decompressed_bytes != original_size) { + throw std::runtime_error("Unexpected decompressed size"); + } + + return decompressed_data; +} + +} // namespace Cubed diff --git a/src/gameplay/network_client.cpp b/src/gameplay/network_client.cpp index e5692e6..c050ed0 100644 --- a/src/gameplay/network_client.cpp +++ b/src/gameplay/network_client.cpp @@ -45,7 +45,7 @@ asio::awaitable NetworkClient::connect(std::string ip, int port) { asio::awaitable NetworkClient::read_loop() { try { while (true) { - std::array header_buffer; + std::array header_buffer; co_await asio::async_read(m_socket, asio::buffer(header_buffer), asio::use_awaitable); auto header = decode_packet_header(header_buffer); @@ -66,7 +66,7 @@ asio::awaitable NetworkClient::read_loop() { case to_num(PacketEnum::LOGIN_RSP): { LoginRsp rsp; Logger::info("Client: Receive Login rsp"); - if (rsp.ParseFromArray(body_data.data(), body_data.size())) { + if (decode_packet(rsp, body_data, header)) { if (rsp.success()) { m_world.start_client_thread(rsp.uuid()); } else { @@ -76,34 +76,34 @@ asio::awaitable NetworkClient::read_loop() { } break; case to_num(PacketEnum::CHUNK_DATA_RSP): { ChunkDataRsp rsp; - // Logger::info("Client: Receive Chunk Data rsp, size {}mb", - // body_data.size() / 1024.0f / 1024); - if (rsp.ParseFromArray(body_data.data(), body_data.size())) { + Logger::info("Client: Receive Chunk Data rsp, size {}mb", + body_data.size() / 1024.0f / 1024); + if (decode_packet(rsp, body_data, header)) { m_world.receive_chunk(rsp); } } break; case to_num(PacketEnum::BLOCK_CHANGE_RSP): { BlockChangeRsp rsp; Logger::info("Client: Receive Block Change rsp"); - if (rsp.ParseFromArray(body_data.data(), body_data.size())) { + if (decode_packet(rsp, body_data, header)) { m_world.receive_block_change(rsp); } } break; case to_num(PacketEnum::UPDATE_TIME): { UpdateTime rsp; - if (rsp.ParseFromArray(body_data.data(), body_data.size())) { + if (decode_packet(rsp, body_data, header)) { m_world.receive_time(rsp); } } break; case to_num(PacketEnum::PLAYER_INFO_RSP): { PlayerInfoRsp rsp; - if (rsp.ParseFromArray(body_data.data(), body_data.size())) { + if (decode_packet(rsp, body_data, header)) { m_world.receive_other_player(rsp); } } break; case to_num(PacketEnum::LOGOUT_RSP): { LogoutRsp rsp; - if (rsp.ParseFromArray(body_data.data(), body_data.size())) { + if (decode_packet(rsp, body_data, header)) { m_world.receive_player_logout(rsp); } } break; diff --git a/src/gameplay/session.cpp b/src/gameplay/session.cpp index 3c32810..cdf3a22 100644 --- a/src/gameplay/session.cpp +++ b/src/gameplay/session.cpp @@ -37,7 +37,7 @@ const std::string& Session::uuid() const { return m_uuid; } asio::awaitable Session::read_loop() { try { while (true) { - std::array header_buffer; + std::array header_buffer; co_await asio::async_read(m_socket, asio::buffer(header_buffer), asio::use_awaitable); @@ -58,14 +58,14 @@ asio::awaitable Session::read_loop() { if (cmd_id == to_num(PacketEnum::LOGIN_REQ)) { LoginReq req; Logger::info("Session: Receive Login req"); - if (req.ParseFromArray(body_data.data(), body_data.size())) { + if (decode_packet(req, body_data, header)) { m_server_world.handle_player_login(req.name(), shared_from_this()); } } if (cmd_id == to_num(PacketEnum::PLAYER_POS)) { PlayerPos pos; - if (pos.ParseFromArray(body_data.data(), body_data.size())) { + if (decode_packet(pos, body_data, header)) { m_server_world.sync_player_pos(pos.uuid(), pos.pos().x(), pos.pos().y(), pos.pos().z()); @@ -74,7 +74,7 @@ asio::awaitable Session::read_loop() { if (cmd_id == to_num(PacketEnum::CHUNK_DATA_REQ)) { ChunkDataReq req; // Logger::info("Session: Receive Chunk Data req"); - if (req.ParseFromArray(body_data.data(), body_data.size())) { + if (decode_packet(req, body_data, header)) { m_server_world.handle_chunk_req( req.uuid(), ChunkPos(req.pos().x(), req.pos().z())); } @@ -82,13 +82,13 @@ asio::awaitable Session::read_loop() { if (cmd_id == to_num(PacketEnum::BLOCK_CHANGE_REQ)) { BlockChangeReq req; Logger::info("Session: Receive Block Change req"); - if (req.ParseFromArray(body_data.data(), body_data.size())) { + if (decode_packet(req, body_data, header)) { m_server_world.handle_block_change(req); } } if (cmd_id == to_num(PacketEnum::LOGOUT_REQ)) { LogoutReq req; - if (req.ParseFromArray(body_data.data(), body_data.size())) { + if (decode_packet(req, body_data, header)) { m_server_world.handle_player_exit(req.uuid()); } }