diff --git a/tpu_sync/transport/lib/chunk.fbs b/tpu_sync/transport/lib/chunk.fbs index f4144c7cf..4619b27a6 100644 --- a/tpu_sync/transport/lib/chunk.fbs +++ b/tpu_sync/transport/lib/chunk.fbs @@ -1,10 +1,29 @@ namespace tpu_raiden.transport.lib.flatbuf; -// Do not change. enum Constant:uint16 { MAGIC = 0x4452, // 'RD' in little endian. } +// v1 header layout for backwards compatibility (ver = 1). +struct ChunkHeaderV1 { + magic:uint16; + ver:uint16; + op:uint8; + flags:uint8; + buffer_id:uint16; + reserved:uint16; + metadata_size:uint16; + remote_id:uint32; + local_id:uint32; + count_or_size:uint32; + uuid:uint64; + padding:uint64; + padding1:uint64; + padding2:uint64; + padding3:uint64; +} + +// Latest header layout corresponding to the latest version (ver = 2). struct ChunkHeader { // Two fixed fields. Do not change. magic:uint16; @@ -18,14 +37,14 @@ struct ChunkHeader { buffer_id:uint16; reserved:uint16; metadata_size:uint16; - remote_id:uint32; + padding0:uint32; // Explicit padding for 8-byte alignment of remote_id. + remote_id:uint64; local_id:uint32; - count_or_size:uint32; + padding1:uint32; // Explicit padding for 8-byte alignment of count_or_size. + count_or_size:uint64; uuid:uint64; // Paddings to keep the header size always 64 bytes. - padding0:uint64; - padding1:uint64; padding2:uint64; padding3:uint64; // LINT.ThenChange(chunk.h) diff --git a/tpu_sync/transport/lib/chunk.h b/tpu_sync/transport/lib/chunk.h index cea60246c..f4ced6b98 100644 --- a/tpu_sync/transport/lib/chunk.h +++ b/tpu_sync/transport/lib/chunk.h @@ -12,8 +12,8 @@ // See the License for the specific language governing permissions and // limitations under the License. -#ifndef THIRD_PARTY_TPU_RAIDEN_TRANSPORT_LIB_CHUNK_H_ -#define THIRD_PARTY_TPU_RAIDEN_TRANSPORT_LIB_CHUNK_H_ +#ifndef THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_LIB_CHUNK_H_ +#define THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_LIB_CHUNK_H_ #include @@ -23,7 +23,8 @@ inline constexpr uint8_t kOpBufferPull = 3; inline constexpr uint8_t kOpBufferPush = 5; inline constexpr uint8_t kOpBufferPushBatched = 7; -// Compact 32-byte binary chunk header layout. +// Current and latest chunk header layout being used. +// Compact 48-byte binary chunk header layout. struct alignas(8) ChunkHeader { // LINT.IfChange uint16_t version; // Header version @@ -33,15 +34,17 @@ struct alignas(8) ChunkHeader { uint16_t reserved; // Holds parallelism/expected chunks count uint16_t metadata_size; // Size of metadata item in bytes for batch push uint16_t padding; // Unused padding to align fields - uint32_t remote_id; // Remote block ID or linear memory offset + uint32_t padding2; // Explicit padding for 8-byte alignment of remote_id + uint64_t remote_id; // Remote block ID or linear memory offset uint32_t local_id; // Local block ID or target shard index - uint32_t count_or_size; // Number of blocks or continuous payload bytes + uint32_t padding3; // Explicit padding for 8-byte alignment of count_or_size + uint64_t count_or_size; // Number of blocks or continuous payload bytes uint64_t uuid; // Globally unique transaction routing ID // LINT.ThenChange(chunk.fbs) bool operator==(const ChunkHeader&) const = default; }; -static_assert(sizeof(ChunkHeader) == 32); +static_assert(sizeof(ChunkHeader) == 48); struct ChunkMetadata { // LINT.IfChange @@ -53,7 +56,8 @@ struct ChunkMetadata { bool operator==(const ChunkMetadata&) const = default; }; +static_assert(sizeof(ChunkMetadata) == 24); } // namespace tpu_raiden::transport::lib -#endif // THIRD_PARTY_TPU_RAIDEN_TRANSPORT_LIB_CHUNK_H_ +#endif // THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_LIB_CHUNK_H_ diff --git a/tpu_sync/transport/lib/chunk_serializer.cc b/tpu_sync/transport/lib/chunk_serializer.cc index 846ee459f..0aa314a84 100644 --- a/tpu_sync/transport/lib/chunk_serializer.cc +++ b/tpu_sync/transport/lib/chunk_serializer.cc @@ -33,13 +33,14 @@ namespace tpu_raiden::transport::lib { namespace { -absl::InlinedVector SerializeHeaderV1( + +absl::InlinedVector SerializeHeaderV2( const ChunkHeader& header) { - constexpr uint16_t kVer = 1; + constexpr uint16_t kVer = 2; const flatbuf::ChunkHeader h( kRaidenMagic, kVer, header.op, header.flags, header.buffer_id, - header.reserved, header.metadata_size, header.remote_id, header.local_id, - header.count_or_size, header.uuid, /*padding0=*/0, /*padding1=*/0, + header.reserved, header.metadata_size, /*padding0=*/0, header.remote_id, + header.local_id, /*padding1=*/0, header.count_or_size, header.uuid, /*padding2=*/0, /*padding3=*/0); absl::InlinedVector bytes(sizeof(h)); @@ -47,7 +48,7 @@ absl::InlinedVector SerializeHeaderV1( return bytes; } -void DeserializeHeaderV1(const flatbuf::ChunkHeader& h, ChunkHeader& header) { +void DeserializeHeaderV1(const flatbuf::ChunkHeaderV1& h, ChunkHeader& header) { DCHECK_EQ(h.ver(), 1); header.version = h.ver(); header.op = h.op(); @@ -61,7 +62,21 @@ void DeserializeHeaderV1(const flatbuf::ChunkHeader& h, ChunkHeader& header) { header.uuid = h.uuid(); } -absl::InlinedVector SerializeMetadataV1( +void DeserializeHeaderV2(const flatbuf::ChunkHeader& h, ChunkHeader& header) { + DCHECK_EQ(h.ver(), 2); + header.version = h.ver(); + header.op = h.op(); + header.flags = h.flags(); + header.buffer_id = h.buffer_id(); + header.reserved = h.reserved(); + header.metadata_size = h.metadata_size(); + header.remote_id = h.remote_id(); + header.local_id = h.local_id(); + header.count_or_size = h.count_or_size(); + header.uuid = h.uuid(); +} + +absl::InlinedVector SerializeMetadata( const ChunkMetadata& meta) { const flatbuf::ChunkMetadata m(meta.layer_idx, meta.dst_shard_idx, meta.dst_offset_bytes, meta.size_bytes); @@ -70,8 +85,8 @@ absl::InlinedVector SerializeMetadataV1( return bytes; } -void DeserializeMetadataV1(const flatbuf::ChunkMetadata& m, - ChunkMetadata& meta) { +void DeserializeMetadata(const flatbuf::ChunkMetadata& m, + ChunkMetadata& meta) { meta.layer_idx = m.layer_idx(); meta.dst_shard_idx = m.dst_shard_idx(); meta.dst_offset_bytes = m.dst_offset_bytes(); @@ -82,18 +97,17 @@ void DeserializeMetadataV1(const flatbuf::ChunkMetadata& m, absl::InlinedVector SerializeChunkHeader( const ChunkHeader& header) { - const auto bytes = SerializeHeaderV1(header); - DCHECK_EQ(bytes.size(), kChunkHeaderSize); - return bytes; + DCHECK_EQ(header.version, kChunkHeaderCurrentVersion); + return SerializeHeaderV2(header); } absl::StatusOr DeserializeChunkHeader(absl::Span s) { - flatbuf::ChunkHeader h; - DCHECK_EQ(sizeof(h), kChunkHeaderSize); if (s.size() != kChunkHeaderSize) { return absl::InvalidArgumentError("Invalid chunk header size"); } + flatbuf::ChunkHeader h; + DCHECK_EQ(sizeof(h), kChunkHeaderSize); std::memcpy(&h, s.data(), sizeof(h)); if (h.magic() != kRaidenMagic) { @@ -103,25 +117,40 @@ absl::StatusOr DeserializeChunkHeader(absl::Span s) { } const uint16_t ver = h.ver(); + if (auto status = ValidateChunkHeaderVersion(ver); !status.ok()) { + return status; + } + + // Support at most 1 version backwards compatible; fail if version diff by 2. switch (ver) { case 1: { + flatbuf::ChunkHeaderV1 h1; + std::memcpy(&h1, s.data(), sizeof(h1)); + ChunkHeader header = {}; + DeserializeHeaderV1(h1, header); + return header; + } + case 2: { ChunkHeader header = {}; - DeserializeHeaderV1(h, header); + DeserializeHeaderV2(h, header); return header; } default: return absl::FailedPreconditionError( - absl::StrCat("Unsupported chunk header flatbuf version: ", ver)); + absl::StrCat("Unhandled supported chunk header version: ", ver)); } } absl::InlinedVector SerializeChunkMetadata( const ChunkMetadata& meta) { - return SerializeMetadataV1(meta); + return SerializeMetadata(meta); } absl::StatusOr DeserializeChunkMetadata(absl::Span s, uint16_t ver) { + if (auto status = ValidateChunkHeaderVersion(ver); !status.ok()) { + return status; + } const size_t meta_size = GetChunkMetadataSize(ver); if (s.size() != meta_size) { return absl::InvalidArgumentError("Invalid chunk metadata size"); @@ -129,15 +158,16 @@ absl::StatusOr DeserializeChunkMetadata(absl::Span s, ChunkMetadata metadata = {}; switch (ver) { - case 1: { + case 1: + case 2: { flatbuf::ChunkMetadata m = {}; std::memcpy(&m, s.data(), meta_size); - DeserializeMetadataV1(m, metadata); + DeserializeMetadata(m, metadata); break; } default: return absl::FailedPreconditionError( - absl::StrCat("Unsupported chunk metadata flatbuf version: ", ver)); + absl::StrCat("Unhandled supported chunk metadata version: ", ver)); } return metadata; } diff --git a/tpu_sync/transport/lib/chunk_serializer.h b/tpu_sync/transport/lib/chunk_serializer.h index 5bf779f8c..6707e835d 100644 --- a/tpu_sync/transport/lib/chunk_serializer.h +++ b/tpu_sync/transport/lib/chunk_serializer.h @@ -23,6 +23,7 @@ #include "absl/container/inlined_vector.h" #include "absl/status/status.h" #include "absl/status/statusor.h" +#include "absl/strings/str_cat.h" #include "absl/types/span.h" #include "tpu_sync/transport/lib/chunk.h" #include "tpu_sync/transport/lib/chunk_generated.h" @@ -37,12 +38,40 @@ inline constexpr uint16_t kRaidenMagic = static_cast(flatbuf::Constant_MAGIC); static_assert(kRaidenMagic == 0x4452); +inline constexpr uint16_t kChunkHeaderCurrentVersion = 2; +inline constexpr uint16_t kChunkHeaderMinSupportedVersion = 1; +inline constexpr uint16_t kChunkHeaderMaxSupportedVersion = 2; + +// Returns true if `ver` is within the supported version range +// [kChunkHeaderMinSupportedVersion, kChunkHeaderMaxSupportedVersion]. +inline constexpr bool IsSupportedChunkHeaderVersion(uint16_t ver) { + return ver >= kChunkHeaderMinSupportedVersion && + ver <= kChunkHeaderMaxSupportedVersion; +} + +// Validates the incoming chunk header version against the supported range. +// Returns absl::FailedPreconditionError if the version is unsupported (e.g. +// diff >= 2). +inline absl::Status ValidateChunkHeaderVersion(uint16_t ver) { + if (ver < kChunkHeaderMinSupportedVersion || + ver > kChunkHeaderMaxSupportedVersion) { + return absl::FailedPreconditionError( + absl::StrCat("Unsupported chunk header flatbuf version: ", ver, + " (current version: ", kChunkHeaderCurrentVersion, + ", supported: ", kChunkHeaderMinSupportedVersion, " to ", + kChunkHeaderMaxSupportedVersion, ")")); + } + return absl::OkStatus(); +} + +static_assert(sizeof(flatbuf::ChunkHeaderV1) == kChunkHeaderSize); static_assert(sizeof(flatbuf::ChunkHeader) == kChunkHeaderSize); // Returns the size of a chunk metadata for the given version. constexpr size_t GetChunkMetadataSize(uint16_t ver) { switch (ver) { case 1: + case 2: return 24; default: return 0; diff --git a/tpu_sync/transport/lib/chunk_serializer_test.cc b/tpu_sync/transport/lib/chunk_serializer_test.cc index d6c702965..910fb4ae6 100644 --- a/tpu_sync/transport/lib/chunk_serializer_test.cc +++ b/tpu_sync/transport/lib/chunk_serializer_test.cc @@ -46,6 +46,21 @@ ChunkHeader MakeSampleHeaderV1() { }; } +ChunkHeader MakeSampleHeaderV2() { + return ChunkHeader{ + .version = 2, + .op = 0xAB, + .flags = 0xCD, + .buffer_id = 0x1234, + .reserved = 0x5678, + .metadata_size = 24, + .remote_id = 0x0123456789ABCDEFULL, + .local_id = 0x9ABCDEF0, + .count_or_size = 0xFEDCBA9876543210ULL, + .uuid = 0x1122334455667788ULL, + }; +} + ChunkMetadata MakeSampleMetadataV1() { return ChunkMetadata{ .layer_idx = 0x12345678, @@ -55,38 +70,45 @@ ChunkMetadata MakeSampleMetadataV1() { }; } -TEST(ChunkHeaderSerializerTest, SerializeAndDeserialize) { - const ChunkHeader original = MakeSampleHeaderV1(); +TEST(ChunkHeaderSerializerTest, SerializeAndDeserializeV2) { + const ChunkHeader original = MakeSampleHeaderV2(); const auto bytes = SerializeChunkHeader(original); EXPECT_THAT(DeserializeChunkHeader(bytes), IsOkAndHolds(original)); } -TEST(ChunkHeaderSerializerTest, SerializeToLittleEndian) { - const auto wire = SerializeChunkHeader(MakeSampleHeaderV1()); + +TEST(ChunkHeaderSerializerTest, SerializeToLittleEndianV2) { + const auto wire = SerializeChunkHeader(MakeSampleHeaderV2()); ASSERT_EQ(wire.size(), kChunkHeaderSize); alignas(8) const uint8_t expected_wire[64] = { - 0x52, 0x44, 0x01, 0x00, 0xAB, 0xCD, 0x34, 0x12, 0x78, 0x56, 0x18, - 0x00, 0x78, 0x56, 0x34, 0x12, 0xF0, 0xDE, 0xBC, 0x9A, 0x44, 0x33, - 0x22, 0x11, 0xEF, 0xCD, 0xAB, 0x89, 0x67, 0x45, 0x23, 0x01, + 0x52, 0x44, 0x02, 0x00, 0xAB, 0xCD, 0x34, 0x12, 0x78, 0x56, 0x18, + 0x00, 0x00, 0x00, 0x00, 0x00, 0xEF, 0xCD, 0xAB, 0x89, 0x67, 0x45, + 0x23, 0x01, 0xF0, 0xDE, 0xBC, 0x9A, 0x00, 0x00, 0x00, 0x00, 0x10, + 0x32, 0x54, 0x76, 0x98, 0xBA, 0xDC, 0xFE, 0x88, 0x77, 0x66, 0x55, + 0x44, 0x33, 0x22, 0x11, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, }; EXPECT_THAT(wire, ElementsAreArray(expected_wire)); } TEST(ChunkHeaderSerializerTest, VerifyMagicBytes) { - const auto s = SerializeChunkHeader(MakeSampleHeaderV1()); + const auto s = SerializeChunkHeader(MakeSampleHeaderV2()); ASSERT_GE(s.size(), 2); ASSERT_EQ(s[0], 'R'); ASSERT_EQ(s[1], 'D'); } -TEST(ChunkHeaderSerializerTest, DeserializeLittleEndian) { +TEST(ChunkHeaderSerializerTest, DeserializeV1BackwardsCompatible) { alignas(8) const uint8_t raw_wire[64] = { 0x52, 0x44, 0x01, 0x00, 0xAB, 0xCD, 0x34, 0x12, 0x78, 0x56, 0x18, 0x00, 0x78, 0x56, 0x34, 0x12, 0xF0, 0xDE, 0xBC, 0x9A, 0x44, 0x33, - 0x22, 0x11, 0xEF, 0xCD, 0xAB, 0x89, 0x67, 0x45, 0x23, 0x01, + 0x22, 0x11, 0xEF, 0xCD, 0xAB, 0x89, 0x67, 0x45, 0x23, 0x01, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, }; const auto wire = absl::MakeConstSpan(reinterpret_cast(raw_wire), @@ -96,29 +118,60 @@ TEST(ChunkHeaderSerializerTest, DeserializeLittleEndian) { } TEST(ChunkHeaderSerializerTest, DeserializeRejectsInvalidMagic) { - auto bytes = SerializeChunkHeader(MakeSampleHeaderV1()); + auto bytes = SerializeChunkHeader(MakeSampleHeaderV2()); bytes[0] ^= 0xFF; // Corrupt the magic field. EXPECT_THAT(DeserializeChunkHeader(bytes), StatusIs(absl::StatusCode::kInvalidArgument)); } -TEST(ChunkHeaderSerializerTest, DeserializeRejectsInvalidVersion) { - auto bytes = SerializeChunkHeader(MakeSampleHeaderV1()); - // The `ver` field is a little-endian uint16. - bytes[2] = 0x02; +TEST(ChunkHeaderSerializerTest, DeserializeRejectsVersionDiffByTwoOrMore) { + auto bytes = SerializeChunkHeader(MakeSampleHeaderV2()); + + // Version 0: diff by 2 from current v2 -> rejected + bytes[2] = 0x00; bytes[3] = 0x00; + EXPECT_THAT(DeserializeChunkHeader(bytes), + StatusIs(absl::StatusCode::kFailedPrecondition)); + // Version 3: diff by 2 from supported v1 -> rejected + bytes[2] = 0x03; + bytes[3] = 0x00; + EXPECT_THAT(DeserializeChunkHeader(bytes), + StatusIs(absl::StatusCode::kFailedPrecondition)); + + // Version 4: diff by 2 from current v2 -> rejected + bytes[2] = 0x04; + bytes[3] = 0x00; EXPECT_THAT(DeserializeChunkHeader(bytes), StatusIs(absl::StatusCode::kFailedPrecondition)); } +TEST(ChunkHeaderSerializerTest, VersionCheckHelpers) { + EXPECT_TRUE(IsSupportedChunkHeaderVersion(1)); + EXPECT_TRUE(IsSupportedChunkHeaderVersion(2)); + EXPECT_FALSE(IsSupportedChunkHeaderVersion(0)); + EXPECT_FALSE(IsSupportedChunkHeaderVersion(3)); + EXPECT_FALSE(IsSupportedChunkHeaderVersion(4)); + + EXPECT_OK(ValidateChunkHeaderVersion(1)); + EXPECT_OK(ValidateChunkHeaderVersion(2)); + EXPECT_THAT(ValidateChunkHeaderVersion(0), + StatusIs(absl::StatusCode::kFailedPrecondition)); + EXPECT_THAT(ValidateChunkHeaderVersion(3), + StatusIs(absl::StatusCode::kFailedPrecondition)); + EXPECT_THAT(ValidateChunkHeaderVersion(4), + StatusIs(absl::StatusCode::kFailedPrecondition)); +} + TEST(ChunkMetadataSerializerTest, SerializeAndDeserialize) { const ChunkMetadata original = MakeSampleMetadataV1(); const auto bytes = SerializeChunkMetadata(original); EXPECT_THAT(DeserializeChunkMetadata(bytes, /*ver=*/1), IsOkAndHolds(original)); + EXPECT_THAT(DeserializeChunkMetadata(bytes, /*ver=*/2), + IsOkAndHolds(original)); } TEST(ChunkMetadataSerializerTest, SerializeToLittleEndian) {