Skip to content
Merged
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
26 changes: 23 additions & 3 deletions tpu_sync/transport/block_transport.cc
Original file line number Diff line number Diff line change
Expand Up @@ -187,6 +187,25 @@ absl::Status ForEachPayload(MajorOrder major_order,
return absl::InvalidArgumentError("Unknown block transport major order");
}

absl::StatusOr<lib::Request> BuildBufferRequest(
size_t buffer_id, size_t dst_shard_idx, size_t dst_offset_bytes,
const uint8_t* data_ptr, size_t size_bytes, uint64_t uuid) {
return lib::Request{
.socket_opcode = lib::kOpBufferPush,
.laddr = const_cast<uint8_t*>(data_ptr),
.raddr = nullptr,
.len = size_bytes,
.major_order = 0,
.layer_idx = static_cast<int>(buffer_id),
.parallelism = 1,
.remote_id = static_cast<uint32_t>(dst_offset_bytes),
.local_id = static_cast<uint32_t>(dst_shard_idx),
.count_or_size = static_cast<uint32_t>(size_bytes),
.uuid = uuid,
.request_id = 0,
};
}

} // namespace

BlockTransport::BlockTransport(BlockTransportDelegate* delegate, int local_port,
Expand Down Expand Up @@ -1346,9 +1365,10 @@ absl::Status BlockTransport::PushBuffer(absl::string_view peer,
size_t dst_offset_bytes,
const uint8_t* data_ptr,
size_t size_bytes, uint64_t uuid) {
absl::Status status =
raw_transport_.PushBuffer(peer, buffer_id, dst_shard_idx,
dst_offset_bytes, data_ptr, size_bytes, uuid);
ASSIGN_OR_RETURN(
auto req, BuildBufferRequest(buffer_id, dst_shard_idx, dst_offset_bytes,
data_ptr, size_bytes, uuid));
absl::Status status = raw_transport_.ProcessSocketBufferPush(peer, req);
if (!status.ok()) {
RecordTransferFailure(status, metric_labels::kDirectionPush);
}
Expand Down
66 changes: 66 additions & 0 deletions tpu_sync/transport/block_transport_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,12 @@

#include "tpu_sync/transport/block_transport.h"

#include <unistd.h>

#include <algorithm>
#include <atomic>
#include <chrono> // NOLINT
#include <csignal>
#include <cstddef>
#include <cstdint>
#include <cstring>
Expand Down Expand Up @@ -51,8 +54,11 @@ namespace transport {
namespace {

using ::absl_testing::StatusIs;
using ::testing::Each;
using ::testing::Eq;
using ::testing::HasSubstr;
using ::testing::Not;
using ::testing::Pointwise;

constexpr absl::Duration kMetricPollingTimeout = absl::Seconds(5);
constexpr absl::Duration kMetricPollingInterval = absl::Milliseconds(10);
Expand Down Expand Up @@ -1083,6 +1089,66 @@ TEST(BlockTransportTest, MultiShardPushLayerMajor) {
}
}

TEST(BlockTransportTest, PushBufferCorrectness) {
constexpr size_t size = 64 * 1024;
MockDelegate src(size);
MockDelegate dst(size);

BlockTransport src_transport(&src, 0);
BlockTransport dst_transport(&dst, 0);
std::this_thread::sleep_for(std::chrono::milliseconds(50));

constexpr size_t kLen = 62 * 1024;
constexpr size_t kDstOffset = 512;
std::vector<uint8_t> push_payload(kLen);
for (size_t i = 0; i < kLen; ++i) {
push_payload[i] = static_cast<uint8_t>((i % 255) + 1);
}
const std::string dst_addr =
absl::StrCat("localhost:", dst_transport.local_port());
const auto push_res = src_transport.PushBuffer(
dst_addr, /*buffer_id=*/0, /*dst_shard_idx=*/0,
/*dst_offset_bytes=*/kDstOffset, push_payload.data(), push_payload.size(),
/*uuid=*/0);
EXPECT_OK(push_res) << push_res.message();

const uint8_t* dst_buf = dst.GetHostPointer(0, 0);
EXPECT_THAT(absl::MakeConstSpan(dst_buf, kDstOffset), Each(Eq(0)));
EXPECT_THAT(absl::MakeConstSpan(dst_buf + kDstOffset, kLen),
Pointwise(Eq(), absl::MakeConstSpan(push_payload)));
EXPECT_THAT(absl::MakeConstSpan(dst_buf + kDstOffset + kLen,
size - kDstOffset - kLen),
Each(Eq(0)));
}

TEST(BlockTransportTest, PollEINTRIsBenign) {
// Set up src/dst buffers.
constexpr size_t size = 4096;
MockDelegate src(size);
MockDelegate dst(size);

// Create two transports.
BlockTransport src_transport(&src, 0);
BlockTransport dst_transport(&dst, 0);
std::this_thread::sleep_for(std::chrono::milliseconds(50));

// Register a dummy signal handler.
signal(SIGUSR1, [](int) {});
// Send a signal to the process to interrupt some poll() calls with EINTR.
kill(getpid(), SIGUSR1);

// Perform a push to verify the connection worker didn't die.
const std::string dst_addr =
absl::StrCat("localhost:", dst_transport.local_port());
const std::vector<uint8_t> push_payload(1024, 0xAB);
constexpr size_t kDstOffset = 512;
const auto push_res = src_transport.PushBuffer(
dst_addr, /*buffer_id=*/0, /*dst_shard_idx=*/0,
/*dst_offset_bytes=*/kDstOffset, push_payload.data(), push_payload.size(),
/*uuid=*/0);
EXPECT_OK(push_res) << push_res.message();
}

} // namespace
} // namespace transport
} // namespace tpu_raiden
1 change: 1 addition & 0 deletions tpu_sync/transport/lib/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,7 @@ cc_library(
":chunk",
":chunk_serializer",
":raw_buffer_transport_delegate",
":transport_adapter",
"//tpu_sync/core:status_macros",
"//tpu_sync/transport:buffer_push_task",
"//tpu_sync/transport/lib/conn:pool",
Expand Down
31 changes: 18 additions & 13 deletions tpu_sync/transport/lib/raw_buffer_transport.cc
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@
#include "absl/synchronization/mutex.h"
#include "absl/types/span.h"
#include "tpu_sync/transport/buffer_push_task.h"
#include "tpu_sync/transport/lib/transport_adapter.h"

#ifndef IOV_MAX
#define IOV_MAX 1024
Expand Down Expand Up @@ -557,12 +558,8 @@ absl::Status RawBufferTransport::RegisterExpectedLayerChunks(
return absl::OkStatus();
}

absl::Status RawBufferTransport::PushBuffer(absl::string_view peer,
size_t buffer_id,
size_t dst_shard_idx,
size_t dst_offset_bytes,
const uint8_t* data_ptr,
size_t size_bytes, uint64_t uuid) {
absl::Status RawBufferTransport::ProcessSocketBufferPush(
absl::string_view peer, const Request& request) {
if (peer.empty()) {
return absl::InvalidArgumentError(
"Destination peer address cannot be empty");
Expand All @@ -573,23 +570,31 @@ absl::Status RawBufferTransport::PushBuffer(absl::string_view peer,
auto fd_cleaner =
absl::MakeCleanup([&] { conn_pool_.Return(ok_to_pool, fd, peer); });

const uint8_t opcode = request.socket_opcode;
const uint64_t uuid = request.uuid;

if (opcode != kOpBufferPush) {
return absl::InvalidArgumentError(
absl::StrCat("Unsupported buffer push opcode: ", opcode));
}

ChunkHeader header = {};
header.version = 1;
header.op = kOpBufferPush;
header.buffer_id = static_cast<uint16_t>(buffer_id);
header.remote_id = static_cast<uint32_t>(dst_offset_bytes);
header.local_id = static_cast<uint32_t>(dst_shard_idx);
header.count_or_size = static_cast<uint32_t>(size_bytes);
header.buffer_id = static_cast<uint16_t>(request.layer_idx);
header.remote_id = request.remote_id;
header.local_id = request.local_id;
header.count_or_size = static_cast<uint32_t>(request.len);
header.uuid = uuid;

VLOG(1) << "Pushing chunk to peer=" << peer << " uuid=" << uuid
<< " dst_shard=" << dst_shard_idx
<< " dst_offset=" << dst_offset_bytes << " size=" << size_bytes;
<< " dst_shard=" << request.local_id
<< " dst_offset=" << request.remote_id << " size=" << request.len;

const auto s_header = SerializeChunkHeader(header);
const std::array<struct iovec, 2> iovs = {
iovec(const_cast<char*>(s_header.data()), s_header.size()),
iovec(const_cast<uint8_t*>(data_ptr), size_bytes),
iovec(request.laddr, request.len),
};
RETURN_IF_ERROR(WriteVExact(fd, iovs));

Expand Down
13 changes: 6 additions & 7 deletions tpu_sync/transport/lib/raw_buffer_transport.h
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@
#include "tpu_sync/transport/lib/chunk.h"
#include "tpu_sync/transport/lib/conn/pool.h"
#include "tpu_sync/transport/lib/raw_buffer_transport_delegate.h"
#include "tpu_sync/transport/lib/transport_adapter.h"

namespace tpu_raiden::transport::lib {

Expand Down Expand Up @@ -86,13 +87,11 @@ class RawBufferTransport final {
size_t dst_shard_idx, size_t dst_offset_bytes,
size_t size_bytes);

// Synchronously pushes a buffer identified by `buffer_id` to the remote
// `peer`, by sending out a `kOpBufferPush ChunkHeader` followed by the data.
// It waits for a one-byte ack from the `peer` before it returns.
absl::Status PushBuffer(absl::string_view peer, size_t buffer_id,
size_t dst_shard_idx, size_t dst_offset_bytes,
const uint8_t* data_ptr, size_t size_bytes,
uint64_t uuid);
// Synchronously pushes a buffer (Op 5) by sending out a
// `kOpBufferPush ChunkHeader` followed by the data. It waits for a one-byte
// ack from the `peer` before it returns.
absl::Status ProcessSocketBufferPush(absl::string_view peer,
const Request& request);

// Pushes a vector of buffers to multiple peers using `PushBatch()`.
absl::Status PushBuffers(const std::vector<BufferPushTask>& tasks,
Expand Down
62 changes: 0 additions & 62 deletions tpu_sync/transport/lib/raw_buffer_transport_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -126,41 +126,6 @@ TEST(RawBufferTransportTest, PullBufferCorrectness) {
Each(Eq(0)));
}

TEST(RawBufferTransportTest, PushBufferCorrectness) {
// Set up src/dst buffers.
constexpr size_t size = 64 * 1024;
RawMockDelegate src(size);
RawMockDelegate dst(size);
RandomNonZero(src.DataSpan());

// Pre-condition: all the dst bytes are not equal to the src.
ASSERT_THAT(dst.DataSpan(), Pointwise(Ne(), src.DataSpan()));

// Create two transports.
RawBufferTransport src_transport(&src, kLocalPort);
RawBufferTransport dst_transport(&dst, kLocalPort);
std::this_thread::sleep_for(std::chrono::milliseconds(50));

// Push a buffer segment from src to dst.
constexpr size_t kLen = 62 * 1024;
constexpr size_t kDstOffset = 512;
std::vector<uint8_t> push_payload(kLen);
RandomNonZero(absl::MakeSpan(push_payload));
const std::string dst_addr = GetIpPort(dst_transport);
const auto push_res =
src_transport.PushBuffer(dst_addr, kBufferId, kDstShardIdx, kDstOffset,
push_payload.data(), push_payload.size(),
/*uuid=*/0);
EXPECT_OK(push_res) << push_res.message();

// Post-condition: only the copied dst bytes are equal to the src.
EXPECT_THAT(dst.DataSpan(0, kDstOffset), Each(Eq(0)));
EXPECT_THAT(dst.DataSpan(kDstOffset, kLen),
Pointwise(Eq(), absl::MakeConstSpan(push_payload)));
EXPECT_THAT(dst.DataSpan(kDstOffset + kLen, size - kDstOffset - kLen),
Each(Eq(0)));
}

TEST(RawBufferTransportTest, PushBuffersCorrectness) {
// Set up src/dst buffers.
constexpr size_t size = 128 * 1024;
Expand Down Expand Up @@ -239,33 +204,6 @@ TEST(RawBufferTransportTest, PushBuffersCorrectness) {
EXPECT_TRUE(dst2.on_data_received());
}

TEST(RawBufferTransportTest, PollEINTRIsBenign) {
// Set up src/dst buffers.
constexpr size_t size = 4096;
RawMockDelegate src(size);
RawMockDelegate dst(size);

// Create two transports.
RawBufferTransport src_transport(&src, 0);
RawBufferTransport dst_transport(&dst, 0);
std::this_thread::sleep_for(std::chrono::milliseconds(50));

// Register a dummy signal handler.
signal(SIGUSR1, [](int) {});
// Send a signal to the process to interrupt some poll() calls with EINTR.
kill(getpid(), SIGUSR1);

// Perform a push to verify the connection worker didn't die.
const std::string dst_addr = GetIpPort(dst_transport);
const std::vector<uint8_t> push_payload(1024, 0xAB);
constexpr size_t kDstOffset = 512;
const auto push_res =
src_transport.PushBuffer(dst_addr, kBufferId, kDstShardIdx, kDstOffset,
push_payload.data(), push_payload.size(),
/*uuid=*/0);
EXPECT_OK(push_res) << push_res.message();
}

TEST(RawBufferTransportTest, RejectsOutOfBounds) {
// Set up src/dst buffers.
constexpr size_t size = 1024;
Expand Down
Loading