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
24 changes: 23 additions & 1 deletion src/brpc/adapter_transport.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,21 @@ class AdapterConnect : public AppConnect {
explicit AdapterConnect(const std::shared_ptr<AppConnect>& app_connect)
: _app_connect(app_connect) {}

static std::shared_ptr<AppConnect> Wrap(
const std::shared_ptr<AppConnect>& app_connect) {
if (std::dynamic_pointer_cast<AdapterConnect>(app_connect)) {
return app_connect;
}
return std::make_shared<AdapterConnect>(app_connect);
}

static std::shared_ptr<AppConnect> Unwrap(
const std::shared_ptr<AppConnect>& app_connect) {
const std::shared_ptr<AdapterConnect> adapter =
std::dynamic_pointer_cast<AdapterConnect>(app_connect);
return adapter ? adapter->_app_connect : app_connect;
}

void StartConnect(const Socket* socket,
void (*done)(int, void*), void* data) override {
ApplicationConnectTask* task = new ApplicationConnectTask{
Expand Down Expand Up @@ -427,7 +442,14 @@ int AdapterTransport::Reset(int32_t expected_nref) {

std::shared_ptr<AppConnect> AdapterTransport::Connect() {
if (upgrade_capable(_mode)) {
return std::make_shared<AdapterConnect>(_default_connect);
return AdapterConnect::Wrap(_default_connect);
}
SocketUser* const client_messenger =
static_cast<SocketUser*>(get_client_side_messenger());
if (client_messenger != NULL &&
_socket->user() == client_messenger) {
FallbackToTcp();
return AdapterConnect::Unwrap(_tcp_transport->Connect());
}
return _tcp_transport->Connect();
}
Expand Down
4 changes: 3 additions & 1 deletion src/brpc/handshake/rdma_handshake.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -557,7 +557,9 @@ StepResult RdmaServerHandshakeAdapter::RunRdmaServerHandshake(
callbacks.transport.set_tcp_active = [transport]() {
transport->DeactivateUpgrade();
};
callbacks.transport.on_failed = []() {};
callbacks.transport.on_failed = [transport]() {
transport->DeactivateUpgrade();
};
return GetSession(socket)->RunServer(callbacks);
}
#endif
Expand Down
22 changes: 17 additions & 5 deletions src/brpc/handshake/ubshm_handshake.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -194,8 +194,15 @@ handshake::StepResult UBShmHandshakeAdapter::ParseHello(
errno = EPROTO;
return handshake::STEP_ERROR;
}
return NegotiationValid(*message) ?
handshake::STEP_OK : handshake::STEP_FALLBACK;
if (!NegotiationValid(*message)) {
return handshake::STEP_FALLBACK;
}
if (strnlen(message->shm_name, SHM_MAX_NAME_BUFF_LEN) ==
SHM_MAX_NAME_BUFF_LEN) {
errno = EPROTO;
return handshake::STEP_ERROR;
}
return handshake::STEP_OK;
}

bool UBShmHandshakeAdapter::NegotiationValid(
Expand Down Expand Up @@ -334,11 +341,14 @@ StepResult UBShmServerHandshakeAdapter::RunUBShmServerHandshake(
transport->DeactivateUpgrade();
return STEP_FALLBACK;
}
const size_t remote_name_len =
strnlen(remote.shm_name, SHM_MAX_NAME_BUFF_LEN);

ubring::SHM remote_trx_shm = {
NULL, remote.len, 0, {0},
static_cast<uint32_t>(socket->fd())};
strncpy(remote_trx_shm.name, remote.shm_name,
SHM_MAX_NAME_BUFF_LEN);
memcpy(remote_trx_shm.name, remote.shm_name,
remote_name_len + 1);

const size_t local_shm_len =
static_cast<size_t>(ubring::FLAGS_data_queue_size) * MB_TO_BYTE;
Expand Down Expand Up @@ -382,7 +392,9 @@ StepResult UBShmServerHandshakeAdapter::RunUBShmServerHandshake(
callbacks.transport.set_tcp_active = [transport]() {
transport->DeactivateUpgrade();
};
callbacks.transport.on_failed = []() {};
callbacks.transport.on_failed = [transport]() {
transport->DeactivateUpgrade();
};
const StepResult result = GetSession(socket)->RunServer(callbacks);
if (result == STEP_OK) {
transport->FinishUpgrade();
Expand Down
99 changes: 99 additions & 0 deletions test/brpc_transport_handshake_unittest.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,11 @@
#include "butil/sys_byteorder.h"
#include "brpc/adapter_transport.h"
#include "brpc/handshake/handshake_io.h"
#include "brpc/input_messenger.h"
#include "brpc/policy/transport_handshake_protocol.h"
#if BRPC_WITH_RDMA
#include "brpc/rdma_transport.h"
#endif
#include "brpc/socket.h"
#include "brpc/transport_handshake.h"

Expand Down Expand Up @@ -85,6 +89,16 @@ static void* RunBlockingHandshakeRead(void* arg) {
return NULL;
}

class TestAppConnect : public AppConnect {
public:
void StartConnect(const Socket*, void (*done)(int, void*),
void* data) override {
done(0, data);
}

void StopConnect(Socket*) override {}
};

static FrameSpec FixedSpec(const char* magic, size_t magic_len,
size_t total_len) {
return FrameSpec(magic, magic_len, total_len, total_len,
Expand Down Expand Up @@ -605,5 +619,90 @@ TEST(TransportHandshakeTest,
socket->SetFailed();
}

#if BRPC_WITH_RDMA
TEST(TransportHandshakeTest,
inherited_adapter_connect_is_not_wrapped_twice) {
int first_fds[2];
ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, first_fds));
butil::fd_guard first_peer(first_fds[1]);

const std::shared_ptr<AppConnect> original =
std::make_shared<TestAppConnect>();

SocketOptions first_options;
first_options.fd = first_fds[0];
first_options.user =
static_cast<SocketUser*>(get_or_new_client_side_messenger());
first_options.need_on_edge_trigger = true;
first_options.socket_mode = SOCKET_MODE_RDMA;
first_options.app_connect = original;

SocketId first_id;
ASSERT_EQ(0, Socket::Create(first_options, &first_id));
SocketUniquePtr first_socket;
ASSERT_EQ(0, Socket::Address(first_id, &first_socket));

const std::shared_ptr<AppConnect> wrapped =
AdapterTransport::Get(first_socket.get())->Connect();
ASSERT_NE(nullptr, wrapped);
EXPECT_NE(original.get(), wrapped.get());

int second_fds[2];
ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, second_fds));
butil::fd_guard second_peer(second_fds[1]);

SocketOptions second_options = first_options;
second_options.fd = second_fds[0];
second_options.app_connect = wrapped;

SocketId second_id;
ASSERT_EQ(0, Socket::Create(second_options, &second_id));
SocketUniquePtr second_socket;
ASSERT_EQ(0, Socket::Address(second_id, &second_socket));

const std::shared_ptr<AppConnect> inherited =
AdapterTransport::Get(second_socket.get())->Connect();
EXPECT_EQ(wrapped.get(), inherited.get());

first_socket->SetFailed();
second_socket->SetFailed();
}

TEST(TransportHandshakeTest,
unavailable_client_upgrade_publishes_tcp_fallback) {
int fds[2];
ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds));
butil::fd_guard peer_fd(fds[1]);

const std::shared_ptr<AppConnect> original =
std::make_shared<TestAppConnect>();

SocketOptions options;
options.fd = fds[0];
options.user =
static_cast<SocketUser*>(get_or_new_client_side_messenger());
options.need_on_edge_trigger = true;
options.socket_mode = SOCKET_MODE_RDMA;
options.app_connect = original;

SocketId id;
ASSERT_EQ(0, Socket::Create(options, &id));
SocketUniquePtr socket;
ASSERT_EQ(0, Socket::Address(id, &socket));

AdapterTransport* adapter = AdapterTransport::Get(socket.get());
RdmaTransport* rdma_transport =
static_cast<RdmaTransport*>(adapter->high_speed_transport());
ASSERT_NE(nullptr, rdma_transport);
rdma_transport->Release();

const std::shared_ptr<AppConnect> connect = adapter->Connect();
EXPECT_EQ(original.get(), connect.get());
EXPECT_EQ(FALLBACK_TCP, adapter->handshake_phase());

socket->SetFailed();
}
#endif

} // namespace handshake
} // namespace brpc
20 changes: 20 additions & 0 deletions test/brpc_ubring_unittest.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -194,6 +194,26 @@ TEST(UBShmHandshakeAdapterTest, short_name_is_zero_padded) {
}
}

TEST(UBShmHandshakeAdapterTest, rejects_unterminated_remote_name) {
brpc::ubring::HelloMessage message{};
message.msg_len = 64;
message.hello_ver = 2;
message.impl_ver = 1;
message.len = 4096;
memset(message.shm_name, 'A', SHM_MAX_NAME_BUFF_LEN);

std::string payload(
sizeof(HelloMessageLayout), '\0');
message.Serialize(&payload[0]);

brpc::ubring::UBShmHandshakeAdapter adapter;
brpc::ubring::HelloMessage decoded{};
errno = 0;
EXPECT_EQ(brpc::handshake::STEP_ERROR,
adapter.ParseHello(payload, &decoded));
EXPECT_EQ(EPROTO, errno);
}

TEST(UBShmHandshakeAdapterTest, disabled_hello_requests_tcp_fallback) {
brpc::ubring::UBShmHandshakeAdapter adapter;
std::string payload;
Expand Down
Loading