diff --git a/src/brpc/adapter_transport.cpp b/src/brpc/adapter_transport.cpp index f0e2534376..777e2282bf 100644 --- a/src/brpc/adapter_transport.cpp +++ b/src/brpc/adapter_transport.cpp @@ -113,6 +113,26 @@ const AdapterTransport* AdapterTransport::Get(const Socket* socket) { return static_cast(socket->_transport.get()); } +bool AdapterTransport::upgrade_capable(SocketMode mode) const { + if (_mode != mode || _high_speed_transport == NULL) { + return false; + } + switch (mode) { +#if BRPC_WITH_RDMA + case SOCKET_MODE_RDMA: + return static_cast( + _high_speed_transport.get())->UpgradeReady(); +#endif +#if BRPC_WITH_UBRING + case SOCKET_MODE_UBRING: + return static_cast( + _high_speed_transport.get())->UpgradeReady(); +#endif + default: + return false; + } +} + int AdapterTransport::StartClientUpgrade(const Socket* socket, void (*done)(int, void*), void* data) { @@ -337,7 +357,7 @@ void AdapterTransport::Init(Socket* socket, const SocketOptions& options) { options.user != static_cast( get_client_side_messenger())) { // UBSHM server handshake is parsed by InputMessenger. - _on_edge_trigger = InputMessenger::OnNewMessages; + _on_edge_trigger = OnNewMessagesAfterUpgrade; #endif } else { _on_edge_trigger = OnNewDataFromTcp; @@ -385,7 +405,7 @@ int AdapterTransport::Reset(int32_t expected_nref) { } std::shared_ptr AdapterTransport::Connect() { - if (_high_speed_transport) { + if (upgrade_capable(_mode)) { return std::make_shared(_default_connect); } return _tcp_transport->Connect(); @@ -475,14 +495,11 @@ void AdapterTransport::SetHighSpeedAvailable(bool available) { } void AdapterTransport::OnNewMessagesAfterUpgrade(Socket* socket) { -#if BRPC_WITH_RDMA AdapterTransport* adapter = Get(socket); - if (adapter->_mode == SOCKET_MODE_RDMA && - adapter->_handshake.phase() == handshake::ESTABLISHED) { + if (adapter->_handshake.phase() == handshake::ESTABLISHED) { adapter->CheckUnexpectedTcpData(); return; } -#endif InputMessenger::OnNewMessages(socket); diff --git a/src/brpc/adapter_transport.h b/src/brpc/adapter_transport.h index 1cbc9803bf..76041847d8 100644 --- a/src/brpc/adapter_transport.h +++ b/src/brpc/adapter_transport.h @@ -58,7 +58,7 @@ class AdapterTransport : public Transport { Transport* high_speed_transport() const { return _high_speed_transport.get(); } - bool upgrade_capable() const { return _high_speed_transport != NULL; } + bool upgrade_capable(SocketMode mode) const; static AdapterTransport* Get(Socket* socket); static const AdapterTransport* Get(const Socket* socket); diff --git a/src/brpc/handshake/handshake_io.cpp b/src/brpc/handshake/handshake_io.cpp index 53beb06083..ca5e7c17c5 100644 --- a/src/brpc/handshake/handshake_io.cpp +++ b/src/brpc/handshake/handshake_io.cpp @@ -70,6 +70,9 @@ static int ReadExactLoop(butil::atomic* read_butex, const timespec duetime = butil::milliseconds_from_now(WAIT_TIMEOUT_MS); const ssize_t nr = read_once(received, len - received); if (nr < 0) { + if (errno == EINTR) { + continue; + } if (errno != EAGAIN) { return -1; } @@ -112,6 +115,9 @@ static int WriteAllLoop(size_t len, WriteOnce write_once, errno = EPIPE; return -1; } + if (errno == EINTR) { + continue; + } if (errno != EAGAIN) { return -1; } diff --git a/src/brpc/handshake/rdma_handshake.cpp b/src/brpc/handshake/rdma_handshake.cpp index 3cf19d0901..dd53e83f0a 100644 --- a/src/brpc/handshake/rdma_handshake.cpp +++ b/src/brpc/handshake/rdma_handshake.cpp @@ -565,7 +565,8 @@ StepResult RdmaServerHandshakeAdapter::RunRdmaServerHandshake( StepResult RdmaServerHandshakeAdapter::RunServerStep( butil::IOBuf* source, Socket* socket) { #if BRPC_WITH_RDMA - if (AdapterTransport::Get(socket)->upgrade_capable()) { + if (AdapterTransport::Get(socket)->upgrade_capable( + SOCKET_MODE_RDMA)) { return RunRdmaServerHandshake(source, socket); } #endif diff --git a/src/brpc/handshake/ubshm_handshake.cpp b/src/brpc/handshake/ubshm_handshake.cpp index d7d4bb7531..e31431b682 100644 --- a/src/brpc/handshake/ubshm_handshake.cpp +++ b/src/brpc/handshake/ubshm_handshake.cpp @@ -166,7 +166,13 @@ handshake::StepResult UBShmHandshakeAdapter::BuildHello( message.hello_ver = handshake::ubshm_wire::HELLO_VERSION; message.impl_ver = handshake::ubshm_wire::IMPL_VERSION; message.len = len; - memcpy(message.shm_name, shm_name, SHM_MAX_NAME_BUFF_LEN); + if (shm_name == NULL) { + errno = EINVAL; + return handshake::STEP_ERROR; + } + const size_t shm_name_len = + strnlen(shm_name, SHM_MAX_NAME_LEN); + memcpy(message.shm_name, shm_name, shm_name_len); } payload->assign( handshake::ubshm_wire::HELLO_LEN - @@ -365,8 +371,7 @@ StepResult UBShmServerHandshakeAdapter::RunUBShmServerHandshake( }; callbacks.transport.negotiate_resources = []() { return STEP_OK; }; callbacks.validate_established = [&]() { - if (!source->empty() || - !transport->UpgradeActive()) { + if (!source->empty()) { return STEP_ERROR; } return STEP_OK; @@ -396,7 +401,8 @@ StepResult UBShmServerHandshakeAdapter::RunUBShmServerHandshake( StepResult UBShmServerHandshakeAdapter::RunServerStep( butil::IOBuf* source, Socket* socket) { #if BRPC_WITH_UBRING - if (AdapterTransport::Get(socket)->upgrade_capable()) { + if (AdapterTransport::Get(socket)->upgrade_capable( + SOCKET_MODE_UBRING)) { return RunUBShmServerHandshake(source, socket); } #endif diff --git a/src/brpc/rdma_transport.cpp b/src/brpc/rdma_transport.cpp index 894637380e..81ebd26b85 100644 --- a/src/brpc/rdma_transport.cpp +++ b/src/brpc/rdma_transport.cpp @@ -101,7 +101,12 @@ RdmaTransport::CreateServerHandshakeAdapters() { void RdmaTransport::ActivateUpgrade() { SetHighSpeedAvailable(true); } -void RdmaTransport::DeactivateUpgrade() { SetHighSpeedAvailable(false); } +void RdmaTransport::DeactivateUpgrade() { + SetHighSpeedAvailable(false); + if (_rdma_ep != nullptr) { + _rdma_ep->Reset(); + } +} int RdmaTransport::CutFromIOBuf(butil::IOBuf *buf) { butil::IOBuf *data[1] = {buf}; diff --git a/src/brpc/rdma_transport.h b/src/brpc/rdma_transport.h index 1590374189..2aaf5fadab 100644 --- a/src/brpc/rdma_transport.h +++ b/src/brpc/rdma_transport.h @@ -65,6 +65,7 @@ class RdmaTransport : public Transport { void ActivateUpgrade(); void DeactivateUpgrade(); bool UpgradeActive() const { return _rdma_state == RDMA_ON; } + bool UpgradeReady() const { return _rdma_ep != nullptr; } private: void SetHighSpeedAvailable(bool available); diff --git a/src/brpc/transport_handshake.h b/src/brpc/transport_handshake.h index 3899a3aa44..1794ac3bef 100644 --- a/src/brpc/transport_handshake.h +++ b/src/brpc/transport_handshake.h @@ -136,7 +136,7 @@ class HandshakeSession { } void SetPhase(int phase) { - _phase.store(phase, butil::memory_order_relaxed); + _phase.store(phase, butil::memory_order_release); } int protocol_version() const { return _protocol_version; } diff --git a/src/brpc/ubshm_transport.cpp b/src/brpc/ubshm_transport.cpp index db1e6c6352..bcc0fec7a4 100644 --- a/src/brpc/ubshm_transport.cpp +++ b/src/brpc/ubshm_transport.cpp @@ -98,7 +98,12 @@ int UBShmTransport::PrepareServerUpgradeResources(ubring::SHM *remote_trx_shm, void UBShmTransport::ActivateUpgrade() { SetHighSpeedAvailable(true); } -void UBShmTransport::DeactivateUpgrade() { SetHighSpeedAvailable(false); } +void UBShmTransport::DeactivateUpgrade() { + SetHighSpeedAvailable(false); + if (_ub_ep != nullptr) { + _ub_ep->Reset(); + } +} void UBShmTransport::FinishUpgrade() { if (_ub_ep != NULL && _ub_ep->_ub_ring != NULL) { diff --git a/src/brpc/ubshm_transport.h b/src/brpc/ubshm_transport.h index 6de2e8917c..b8d840595e 100644 --- a/src/brpc/ubshm_transport.h +++ b/src/brpc/ubshm_transport.h @@ -57,6 +57,7 @@ class UBShmTransport : public Transport { void DeactivateUpgrade(); void FinishUpgrade(); bool UpgradeActive() const { return _ub_state == UB_ON; } + bool UpgradeReady() const { return _ub_ep != nullptr; } private: void SetHighSpeedAvailable(bool available); diff --git a/test/brpc_transport_handshake_unittest.cpp b/test/brpc_transport_handshake_unittest.cpp index 2c084c3990..032befe218 100644 --- a/test/brpc_transport_handshake_unittest.cpp +++ b/test/brpc_transport_handshake_unittest.cpp @@ -198,6 +198,31 @@ TEST(HandshakeFrameTest, rejects_lengths_outside_bounds) { ASSERT_EQ(header.size(), source.size()); } +#if BRPC_WITH_RDMA && BRPC_WITH_UBRING +TEST(TransportHandshakeTest, upgrade_capability_is_mode_specific) { + SocketOptions options; + SocketId id; + SocketUniquePtr socket; + + options.socket_mode = SOCKET_MODE_RDMA; + ASSERT_EQ(0, Socket::Create(options, &id)); + ASSERT_EQ(0, Socket::Address(id, &socket)); + AdapterTransport* rdma_adapter = AdapterTransport::Get(socket.get()); + ASSERT_TRUE(rdma_adapter->upgrade_capable(SOCKET_MODE_RDMA)); + ASSERT_FALSE(rdma_adapter->upgrade_capable(SOCKET_MODE_UBRING)); + socket->SetFailed(); + socket.reset(); + + options.socket_mode = SOCKET_MODE_UBRING; + ASSERT_EQ(0, Socket::Create(options, &id)); + ASSERT_EQ(0, Socket::Address(id, &socket)); + AdapterTransport* ubshm_adapter = AdapterTransport::Get(socket.get()); + ASSERT_TRUE(ubshm_adapter->upgrade_capable(SOCKET_MODE_UBRING)); + ASSERT_FALSE(ubshm_adapter->upgrade_capable(SOCKET_MODE_RDMA)); + socket->SetFailed(); +} +#endif + TEST(TransportHandshakeTest, publish_fallback_after_tcp_state) { HandshakeSession session; int tcp_active = 0; diff --git a/test/brpc_ubring_unittest.cpp b/test/brpc_ubring_unittest.cpp index 8f1c545dfa..39bc4f3e9a 100644 --- a/test/brpc_ubring_unittest.cpp +++ b/test/brpc_ubring_unittest.cpp @@ -174,6 +174,26 @@ TEST(UBShmHandshakeAdapterTest, codec_preserves_v2_wire_format) { EXPECT_EQ("UB", frame.substr(0, 2)); } +TEST(UBShmHandshakeAdapterTest, short_name_is_zero_padded) { + brpc::ubring::UBShmHandshakeAdapter adapter; + char short_name[SHM_MAX_NAME_BUFF_LEN]; + memset(short_name, 0x5a, sizeof(short_name)); + short_name[0] = 'x'; + short_name[1] = '\0'; + + std::string payload; + ASSERT_EQ(brpc::handshake::STEP_OK, + adapter.BuildHello(true, 4096, short_name, &payload)); + + brpc::ubring::HelloMessage decoded{}; + ASSERT_EQ(brpc::handshake::STEP_OK, + adapter.ParseHello(payload, &decoded)); + EXPECT_EQ('x', decoded.shm_name[0]); + for (size_t i = 1; i < SHM_MAX_NAME_BUFF_LEN; ++i) { + EXPECT_EQ('\0', decoded.shm_name[i]) << "index=" << i; + } +} + TEST(UBShmHandshakeAdapterTest, disabled_hello_requests_tcp_fallback) { brpc::ubring::UBShmHandshakeAdapter adapter; std::string payload;