diff --git a/src/brpc/adapter_transport.cpp b/src/brpc/adapter_transport.cpp index ffd8e46521..1fb093a65d 100644 --- a/src/brpc/adapter_transport.cpp +++ b/src/brpc/adapter_transport.cpp @@ -53,6 +53,21 @@ class AdapterConnect : public AppConnect { explicit AdapterConnect(const std::shared_ptr& app_connect) : _app_connect(app_connect) {} + static std::shared_ptr Wrap( + const std::shared_ptr& app_connect) { + if (std::dynamic_pointer_cast(app_connect)) { + return app_connect; + } + return std::make_shared(app_connect); + } + + static std::shared_ptr Unwrap( + const std::shared_ptr& app_connect) { + const std::shared_ptr adapter = + std::dynamic_pointer_cast(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{ @@ -427,7 +442,14 @@ int AdapterTransport::Reset(int32_t expected_nref) { std::shared_ptr AdapterTransport::Connect() { if (upgrade_capable(_mode)) { - return std::make_shared(_default_connect); + return AdapterConnect::Wrap(_default_connect); + } + SocketUser* const client_messenger = + static_cast(get_client_side_messenger()); + if (client_messenger != NULL && + _socket->user() == client_messenger) { + FallbackToTcp(); + return AdapterConnect::Unwrap(_tcp_transport->Connect()); } return _tcp_transport->Connect(); } diff --git a/src/brpc/handshake/rdma_handshake.cpp b/src/brpc/handshake/rdma_handshake.cpp index dd53e83f0a..61bc5fc762 100644 --- a/src/brpc/handshake/rdma_handshake.cpp +++ b/src/brpc/handshake/rdma_handshake.cpp @@ -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 diff --git a/src/brpc/handshake/ubshm_handshake.cpp b/src/brpc/handshake/ubshm_handshake.cpp index e31431b682..cc33ac81c5 100644 --- a/src/brpc/handshake/ubshm_handshake.cpp +++ b/src/brpc/handshake/ubshm_handshake.cpp @@ -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( @@ -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(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(ubring::FLAGS_data_queue_size) * MB_TO_BYTE; @@ -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(); diff --git a/test/brpc_transport_handshake_unittest.cpp b/test/brpc_transport_handshake_unittest.cpp index 97dbb45c57..6562ba90ab 100644 --- a/test/brpc_transport_handshake_unittest.cpp +++ b/test/brpc_transport_handshake_unittest.cpp @@ -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" @@ -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, @@ -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 original = + std::make_shared(); + + SocketOptions first_options; + first_options.fd = first_fds[0]; + first_options.user = + static_cast(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 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 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 original = + std::make_shared(); + + SocketOptions options; + options.fd = fds[0]; + options.user = + static_cast(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(adapter->high_speed_transport()); + ASSERT_NE(nullptr, rdma_transport); + rdma_transport->Release(); + + const std::shared_ptr connect = adapter->Connect(); + EXPECT_EQ(original.get(), connect.get()); + EXPECT_EQ(FALLBACK_TCP, adapter->handshake_phase()); + + socket->SetFailed(); +} +#endif + } // namespace handshake } // namespace brpc diff --git a/test/brpc_ubring_unittest.cpp b/test/brpc_ubring_unittest.cpp index 39bc4f3e9a..f4b416303d 100644 --- a/test/brpc_ubring_unittest.cpp +++ b/test/brpc_ubring_unittest.cpp @@ -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;