diff --git a/tpu_sync/transport/block_transport.cc b/tpu_sync/transport/block_transport.cc index 78d603ad..6260e801 100644 --- a/tpu_sync/transport/block_transport.cc +++ b/tpu_sync/transport/block_transport.cc @@ -22,6 +22,7 @@ #include #include +#include #include #include #include @@ -378,15 +379,15 @@ absl::Status BlockTransport::HandleIncomingPush( if (header.op == 1) { ASSIGN_OR_RETURN(allocated_ids, block_delegate_->AllocateBlocks( header.count_or_size, header.uuid)); - RETURN_IF_ERROR(WriteExact(client_fd, allocated_ids.data(), - header.count_or_size * sizeof(int))); + const std::vector s_ids = lib::SerializeBlockIds(allocated_ids); + RETURN_IF_ERROR(WriteExact(client_fd, s_ids.data(), s_ids.size())); } else { - allocated_ids.resize(header.count_or_size, 0); - RETURN_IF_ERROR(ReadExact(client_fd, allocated_ids.data(), - header.count_or_size * sizeof(int))); - src_block_ids.resize(header.count_or_size, 0); - RETURN_IF_ERROR(ReadExact(client_fd, src_block_ids.data(), - header.count_or_size * sizeof(int))); + std::vector ids_buf(header.count_or_size * sizeof(uint32_t)); + RETURN_IF_ERROR(ReadExact(client_fd, ids_buf.data(), ids_buf.size())); + allocated_ids = lib::DeserializeBlockIds(ids_buf); + + RETURN_IF_ERROR(ReadExact(client_fd, ids_buf.data(), ids_buf.size())); + src_block_ids = lib::DeserializeBlockIds(ids_buf); uint8_t ack = 1; RETURN_IF_ERROR(WriteExact(client_fd, &ack, 1)); } @@ -397,9 +398,9 @@ absl::Status BlockTransport::HandleIncomingPush( header.count_or_size, [&](size_t l, size_t sh, size_t k) -> absl::Status { ABSL_DCHECK_LT(k, allocated_ids.size()); const int dst_id = allocated_ids[k]; - uint32_t sender_size = 0; - RETURN_IF_ERROR( - ReadExact(client_fd, &sender_size, sizeof(sender_size))); + uint8_t size_buf[lib::kChunkSizeFieldSize]; + RETURN_IF_ERROR(ReadExact(client_fd, size_buf, sizeof(size_buf))); + const uint32_t sender_size = lib::DeserializeChunkSize(size_buf); const int64_t block_id_val = dst_id; int64_t src_bid = -1; @@ -631,7 +632,9 @@ void BlockTransport::TriggerNextSendStep( } uint32_t total_size = GetChunksTotalSize(chunks); - s = WriteExact(state->client_fd, &total_size, sizeof(total_size)); + const std::array s_size = + lib::SerializeChunkSize(total_size); + s = WriteExact(state->client_fd, s_size.data(), s_size.size()); if (!s.ok()) { LOG(ERROR) << "Write size failed: " << s.ToString(); shutdown(state->client_fd, SHUT_RDWR); @@ -1041,10 +1044,12 @@ absl::Status BlockTransport::ProcessSocketPush( if (socket_opcode == 6) { ABSL_DCHECK_LE(block_offset + block_count, dst_block_ids.size()); - RETURN_IF_ERROR(WriteExact(fd, &dst_block_ids[block_offset], - block_count * sizeof(int))); - RETURN_IF_ERROR(WriteExact(fd, &src_block_ids[block_offset], - block_count * sizeof(int))); + const auto s_dst_ids = lib::SerializeBlockIds( + {dst_block_ids.data() + block_offset, block_count}); + RETURN_IF_ERROR(WriteExact(fd, s_dst_ids.data(), s_dst_ids.size())); + const auto s_src_ids = lib::SerializeBlockIds( + {src_block_ids.data() + block_offset, block_count}); + RETURN_IF_ERROR(WriteExact(fd, s_src_ids.data(), s_src_ids.size())); uint8_t ack = 0; s = ReadExact(fd, &ack, 1); if (!s.ok() || ack != 1) { @@ -1054,9 +1059,10 @@ absl::Status BlockTransport::ProcessSocketPush( allocated_ids[block_offset + k] = dst_block_ids[block_offset + k]; } } else { - std::vector stream_allocated_ids(block_count, 0); - RETURN_IF_ERROR( - ReadExact(fd, stream_allocated_ids.data(), block_count * sizeof(int))); + std::vector ids_buf(block_count * sizeof(uint32_t)); + RETURN_IF_ERROR(ReadExact(fd, ids_buf.data(), ids_buf.size())); + const std::vector stream_allocated_ids = + lib::DeserializeBlockIds(ids_buf); for (size_t k = 0; k < block_count; ++k) { ABSL_DCHECK_LT(block_offset + k, allocated_ids.size()); @@ -1080,7 +1086,9 @@ absl::Status BlockTransport::ProcessSocketPush( ++j; } - RETURN_IF_ERROR(WriteExact(fd, &total_size, sizeof(total_size))); + const std::array s_size = + lib::SerializeChunkSize(total_size); + RETURN_IF_ERROR(WriteExact(fd, s_size.data(), s_size.size())); if (total_size > 0) { RETURN_IF_ERROR(WriteVExact(fd, absl::MakeSpan(iov))); stream_bytes_sent += total_size; @@ -1309,8 +1317,9 @@ void BlockTransport::H2hReadWorker( expected_size += chunk.size; } - uint32_t sender_size = 0; - RETURN_IF_ERROR(ReadExact(fd, &sender_size, sizeof(sender_size))); + uint8_t size_buf[lib::kChunkSizeFieldSize]; + RETURN_IF_ERROR(ReadExact(fd, size_buf, sizeof(size_buf))); + const uint32_t sender_size = lib::DeserializeChunkSize(size_buf); if (sender_size != expected_size) { return absl::InternalError(absl::StrCat( diff --git a/tpu_sync/transport/lib/BUILD b/tpu_sync/transport/lib/BUILD index d2860d41..c93a744a 100644 --- a/tpu_sync/transport/lib/BUILD +++ b/tpu_sync/transport/lib/BUILD @@ -55,6 +55,7 @@ cc_library( deps = [ ":chunk", ":chunk_cc_fbs", + "//third_party/flatbuffers", "@com_google_absl//absl/container:inlined_vector", "@com_google_absl//absl/log:check", "@com_google_absl//absl/status", diff --git a/tpu_sync/transport/lib/chunk_serializer.cc b/tpu_sync/transport/lib/chunk_serializer.cc index 481a03fc..4b01ae7c 100644 --- a/tpu_sync/transport/lib/chunk_serializer.cc +++ b/tpu_sync/transport/lib/chunk_serializer.cc @@ -14,8 +14,10 @@ #include "tpu_sync/transport/lib/chunk_serializer.h" +#include #include #include +#include #include "absl/container/inlined_vector.h" #include "absl/log/check.h" @@ -23,6 +25,7 @@ #include "absl/status/statusor.h" #include "absl/strings/str_cat.h" #include "absl/types/span.h" +#include "flatbuffers/include/flatbuffers/base.h" #include "tpu_sync/transport/lib/chunk.h" #include "tpu_sync/transport/lib/chunk_generated.h" @@ -139,4 +142,41 @@ absl::StatusOr DeserializeChunkMetadata(absl::Span s, return metadata; } +std::vector SerializeBlockIds(absl::Span ids) { + std::vector buf(ids.size() * sizeof(uint32_t)); + for (size_t i = 0; i < ids.size(); ++i) { + uint32_t val = flatbuffers::EndianScalar(static_cast(ids[i])); + std::memcpy(buf.data() + i * sizeof(uint32_t), &val, sizeof(uint32_t)); + } + return buf; +} + +std::vector DeserializeBlockIds(absl::Span bytes) { + DCHECK_EQ(bytes.size() % sizeof(uint32_t), 0); + const size_t count = bytes.size() / sizeof(uint32_t); + std::vector ids; + ids.reserve(count); + for (size_t i = 0; i < count; ++i) { + uint32_t val = 0; + std::memcpy(&val, bytes.data() + i * sizeof(uint32_t), sizeof(uint32_t)); + ids.push_back(static_cast(flatbuffers::EndianScalar(val))); + } + return ids; +} + +std::array SerializeChunkSize( + uint32_t size_bytes) { + std::array buf; + uint32_t val = flatbuffers::EndianScalar(size_bytes); + std::memcpy(buf.data(), &val, sizeof(uint32_t)); + return buf; +} + +uint32_t DeserializeChunkSize(absl::Span bytes) { + DCHECK_EQ(bytes.size(), kChunkSizeFieldSize); + uint32_t val = 0; + std::memcpy(&val, bytes.data(), sizeof(uint32_t)); + return flatbuffers::EndianScalar(val); +} + } // namespace tpu_raiden::transport::lib diff --git a/tpu_sync/transport/lib/chunk_serializer.h b/tpu_sync/transport/lib/chunk_serializer.h index 1c200585..5bf779f8 100644 --- a/tpu_sync/transport/lib/chunk_serializer.h +++ b/tpu_sync/transport/lib/chunk_serializer.h @@ -15,8 +15,10 @@ #ifndef THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_LIB_CHUNK_SERIALIZER_H_ #define THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_LIB_CHUNK_SERIALIZER_H_ +#include #include #include +#include #include "absl/container/inlined_vector.h" #include "absl/status/status.h" @@ -29,6 +31,7 @@ namespace tpu_raiden::transport::lib { inline constexpr size_t kChunkHeaderSize = 64; inline constexpr size_t kMaxMetadataSize = 24; +inline constexpr size_t kChunkSizeFieldSize = sizeof(uint32_t); inline constexpr uint16_t kRaidenMagic = static_cast(flatbuf::Constant_MAGIC); @@ -62,6 +65,19 @@ absl::InlinedVector SerializeChunkMetadata( absl::StatusOr DeserializeChunkMetadata( absl::Span bytes, uint16_t ver); +// Serializes a span of integer block IDs to a byte vector. +std::vector SerializeBlockIds(absl::Span ids); + +// Parses a vector of integer block IDs from its serialized binary bytes. +std::vector DeserializeBlockIds(absl::Span bytes); + +// Serializes a 32-bit chunk size to a 4-byte array. +std::array SerializeChunkSize( + uint32_t size_bytes); + +// Parses a 32-bit chunk size from its serialized binary bytes. +uint32_t DeserializeChunkSize(absl::Span bytes); + } // namespace tpu_raiden::transport::lib #endif // THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_LIB_CHUNK_SERIALIZER_H_ diff --git a/tpu_sync/transport/lib/chunk_serializer_test.cc b/tpu_sync/transport/lib/chunk_serializer_test.cc index b1758362..d6c70296 100644 --- a/tpu_sync/transport/lib/chunk_serializer_test.cc +++ b/tpu_sync/transport/lib/chunk_serializer_test.cc @@ -15,6 +15,7 @@ #include "tpu_sync/transport/lib/chunk_serializer.h" #include +#include #include #include @@ -145,5 +146,38 @@ TEST(ChunkMetadataSerializerTest, DeserializeLittleEndian) { IsOkAndHolds(MakeSampleMetadataV1())); } +TEST(BlockIdsSerializerTest, SerializeAndDeserialize) { + const std::vector original = {0, 1, -1, 42, 0x12345678, -12345678}; + const auto bytes = SerializeBlockIds(original); + + EXPECT_EQ(DeserializeBlockIds(bytes), original); +} + +TEST(BlockIdsSerializerTest, SerializeToLittleEndian) { + const std::vector original = {0x12345678, 0x01020304}; + const auto bytes = SerializeBlockIds(original); + ASSERT_EQ(bytes.size(), 8); + + const uint8_t expected_wire[8] = { + 0x78, 0x56, 0x34, 0x12, 0x04, 0x03, 0x02, 0x01, + }; + EXPECT_THAT(bytes, ElementsAreArray(expected_wire)); +} + +TEST(ChunkSizeSerializerTest, SerializeAndDeserialize) { + for (uint32_t original : {0u, 1u, 1024u, 0x12345678u, 0xFFFFFFFFu}) { + const auto bytes = SerializeChunkSize(original); + EXPECT_EQ(DeserializeChunkSize(bytes), original); + } +} + +TEST(ChunkSizeSerializerTest, SerializeToLittleEndian) { + const auto bytes = SerializeChunkSize(0x12345678); + ASSERT_EQ(bytes.size(), 4); + + const uint8_t expected_wire[4] = {0x78, 0x56, 0x34, 0x12}; + EXPECT_THAT(bytes, ElementsAreArray(expected_wire)); +} + } // namespace } // namespace tpu_raiden::transport::lib