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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 24 additions & 5 deletions tpu_sync/transport/lib/chunk.fbs
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -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)
Expand Down
18 changes: 11 additions & 7 deletions tpu_sync/transport/lib/chunk.h
Original file line number Diff line number Diff line change
Expand Up @@ -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 <cstdint>

Expand All @@ -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
Expand All @@ -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
Expand All @@ -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_
68 changes: 49 additions & 19 deletions tpu_sync/transport/lib/chunk_serializer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -33,21 +33,22 @@ namespace tpu_raiden::transport::lib {

namespace {

absl::InlinedVector<char, kChunkHeaderSize> SerializeHeaderV1(

absl::InlinedVector<char, kChunkHeaderSize> 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<char, kChunkHeaderSize> bytes(sizeof(h));
std::memcpy(bytes.data(), &h, sizeof(h));
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();
Expand All @@ -61,7 +62,21 @@ void DeserializeHeaderV1(const flatbuf::ChunkHeader& h, ChunkHeader& header) {
header.uuid = h.uuid();
}

absl::InlinedVector<char, kMaxMetadataSize> 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<char, kMaxMetadataSize> SerializeMetadata(
const ChunkMetadata& meta) {
const flatbuf::ChunkMetadata m(meta.layer_idx, meta.dst_shard_idx,
meta.dst_offset_bytes, meta.size_bytes);
Expand All @@ -70,8 +85,8 @@ absl::InlinedVector<char, kMaxMetadataSize> 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();
Expand All @@ -82,18 +97,17 @@ void DeserializeMetadataV1(const flatbuf::ChunkMetadata& m,

absl::InlinedVector<char, kChunkHeaderSize> 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<ChunkHeader> DeserializeChunkHeader(absl::Span<const char> 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) {
Expand All @@ -103,41 +117,57 @@ absl::StatusOr<ChunkHeader> DeserializeChunkHeader(absl::Span<const char> 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<char, kMaxMetadataSize> SerializeChunkMetadata(
const ChunkMetadata& meta) {
return SerializeMetadataV1(meta);
return SerializeMetadata(meta);
}

absl::StatusOr<ChunkMetadata> DeserializeChunkMetadata(absl::Span<const char> 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");
}

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;
}
Expand Down
29 changes: 29 additions & 0 deletions tpu_sync/transport/lib/chunk_serializer.h
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -37,12 +38,40 @@ inline constexpr uint16_t kRaidenMagic =
static_cast<uint16_t>(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;
Expand Down
Loading
Loading