diff --git a/BUILD.bazel b/BUILD.bazel index db73605e6e..f6d495dbe4 100644 --- a/BUILD.bazel +++ b/BUILD.bazel @@ -525,7 +525,6 @@ filegroup( srcs = glob([ "src/brpc/*.proto", "src/brpc/policy/*.proto", - "src/brpc/rdma/*.proto", "src/brpc/urma/*.proto", ]), visibility = ["//visibility:public"], diff --git a/CMakeLists.txt b/CMakeLists.txt index f9e1aa8fc3..7b023ec23e 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -720,7 +720,7 @@ set(PROTO_FILES idl_options.proto brpc/trackme.proto brpc/streaming_rpc_meta.proto brpc/proto_base.proto - brpc/rdma/rdma_handshake.proto + brpc/rdma_handshake.proto brpc/urma/urma_handshake.proto) file(MAKE_DIRECTORY ${PROJECT_BINARY_DIR}/output/include/brpc) set(PROTOC_FLAGS ${PROTOC_FLAGS} -I${PROTOBUF_INCLUDE_DIR}) diff --git a/Makefile b/Makefile index 271b518ae6..dd33e0dc6a 100644 --- a/Makefile +++ b/Makefile @@ -203,7 +203,7 @@ JSON2PB_DIRS = src/json2pb JSON2PB_SOURCES = $(foreach d,$(JSON2PB_DIRS),$(wildcard $(addprefix $(d)/*,$(SRCEXTS)))) JSON2PB_OBJS = $(addsuffix .o, $(basename $(JSON2PB_SOURCES))) -BRPC_DIRS = src/brpc src/brpc/details src/brpc/builtin src/brpc/policy src/brpc/policy/mysql src/brpc/rdma +BRPC_DIRS = src/brpc src/brpc/details src/brpc/builtin src/brpc/handshake src/brpc/policy src/brpc/policy/mysql src/brpc/rdma ifeq ($(WITH_URMA),1) BRPC_DIRS += src/brpc/urma endif diff --git a/docs/cn/handshake_common_design.md b/docs/cn/handshake_common_design.md new file mode 100644 index 0000000000..afc8e7a580 --- /dev/null +++ b/docs/cn/handshake_common_design.md @@ -0,0 +1,50 @@ +# RDMA 与 UBSHM 公共握手 + +## 范围 + +本次重构将 RDMA 和 UBSHM 的握手流程收敛到公共层。`Socket` 持有 +`AdapterTransport`;后者以 TCP 作为握手控制通道,在升级成功后选择高速 +数据面,协商不可用时继续使用 TCP。URMA 目前仍使用独立的 `UrmaTransport` +和握手实现,不在本次迁移范围内。 + +## 分工 + +| 组件 | 职责 | +| --- | --- | +| `AdapterTransport` | 持有 TCP 与候选高速 Transport,启动客户端握手、分发服务端输入,并选择数据面 | +| `HandshakeSession` | 执行公共客户端和服务端状态机,管理帧收发、增量解析和终态发布 | +| `HandshakeProtocol` | 定义 hello、ACK 和扩展帧,编解码协议字段并保存协议状态 | +| `HandshakeTransport` | 准备和协商资源,处理升级成功、TCP 回退及失败清理 | +| RDMA / UBSHM protocol | 实现各自的 wire format、版本兼容和字段校验 | +| `RdmaEndpoint` / `UBShmEndpoint` | 管理各自的资源和数据面,不驱动完整握手流程 | + +客户端在 TCP 连接建立后启动握手 bthread;服务端通过 `InputMessenger` +增量解析 hello、扩展字段(若协议需要)和 ACK。`AdapterTransport` 选择具体的 +protocol 和 transport participant,`HandshakeSession` 只负责调用参与者并编排流程。 +服务端通过协议版本恢复选中的 protocol,因为 ACK 本身不带 magic。 + +## 状态与回退 + +握手终态为 `ESTABLISHED`、`FALLBACK_TCP` 或 `FAILED`。升级成功后,TCP +仍是控制连接,业务数据走 RDMA/UBSHM;控制连接上出现额外业务数据视为 +协议错误。资源不可用或协商拒绝时,释放本次升级资源,设置 TCP 为数据面, +再以 release 顺序发布 `FALLBACK_TCP`。事件线程以 acquire 顺序读取终态, +继续解析已缓存及后续的 TCP 业务数据。已确认属于握手协议的畸形帧则失败, +不作为普通 RPC 数据回放。 + +增量解析中,不完整帧返回 `NOT_ENOUGH_DATA` 且不消费输入;不可能匹配 +握手 magic 的前缀返回 `TRY_OTHERS`。TCP 分片、握手帧与后续数据粘连都由 +帧边界处理,不能假设一次 read 恰好得到一帧。 + +## 协议兼容 + +- RDMA 保持既有 v2/v3 wire format、协议 ID 和注册名称。 +- UBSHM 使用 v3 hello。64 字节 hello 后,双方交换 4 字节网络序格式扩展; + 只有双方选择 `LEGACY_64` 才能升级,否则回退 TCP。ACK 仍为 4 字节。 +- 公共 framing 仅负责帧边界,协议字段由各自 `HandshakeProtocol` 解释。 + +## 验证重点 + +测试覆盖分片和粘连输入、magic 前缀分流、协议版本与格式校验、TCP 回退、 +资源释放及终态发布,并分别运行 RDMA 与 UBSHM 回归测试。URMA 的公共握手 +迁移需要单独设计和测试,不应仅凭本次重构推断其已完成。 diff --git a/src/brpc/adapter_transport.cpp b/src/brpc/adapter_transport.cpp new file mode 100644 index 0000000000..ed51326426 --- /dev/null +++ b/src/brpc/adapter_transport.cpp @@ -0,0 +1,673 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include "brpc/adapter_transport.h" + +#include +#include +#include +#include +#include + +#include "brpc/input_messenger.h" +#include "brpc/destroyable.h" +#include "brpc/handshake/rdma_handshake.h" +#include "brpc/handshake/ubshm_handshake.h" +#if BRPC_WITH_RDMA +#include "brpc/rdma/rdma_helper.h" +#endif +#if BRPC_WITH_UBRING +#include "brpc/ubshm/ub_endpoint.h" +#include "brpc/ubshm/ub_helper.h" +#include "brpc/ubshm/ubr_trx.h" +#endif +#include "brpc/rdma_transport.h" +#include "brpc/tcp_transport.h" +#include "brpc/ubshm_transport.h" + +namespace brpc { + +namespace { + +bool MatchesMagicPrefix(const char *prefix, size_t prefix_len, + const char *magic, size_t magic_len) { + const size_t compare_len = std::min(prefix_len, magic_len); + return memcmp(prefix, magic, compare_len) == 0; +} + +class AdapterConnect : public AppConnect { +public: + 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{ + socket, _app_connect, done, data}; + if (AdapterTransport::StartClientUpgrade( + socket, OnUpgradeComplete, task) != 0) { + AdapterTransport::Get(const_cast(socket))->CompleteConnection( + handshake::FAILED); + const int error = errno != 0 ? errno : EAGAIN; + delete task; + done(error, data); + } + } + + void StopConnect(Socket*) override {} + +private: + struct ApplicationConnectTask { + const Socket* socket; + std::shared_ptr app_connect; + void (*done)(int, void*); + void* data; + }; + + static void OnApplicationComplete(int error, void* arg) { + std::unique_ptr task( + static_cast(arg)); + task->done(error, task->data); + } + + static void OnUpgradeComplete(int error, void* arg) { + ApplicationConnectTask* task = + static_cast(arg); + if (error != 0 || !task->app_connect) { + std::unique_ptr owned(task); + task->done(error, task->data); + return; + } + task->app_connect->StartConnect( + task->socket, OnApplicationComplete, task); + } + + std::shared_ptr _app_connect; +}; + +struct ClientHandshakeTask { + AdapterTransport* adapter; + void (*done)(int, void*); + void* data; + SocketUniquePtr socket; +}; + +#if BRPC_WITH_RDMA +class RdmaClientHandshakeTransport : public handshake::HandshakeTransport { +public: + RdmaClientHandshakeTransport( + RdmaTransport* transport, rdma::RdmaHandshakeAdapter* protocol, + Socket* socket, int* connect_error) + : _transport(transport), _protocol(protocol), _socket(socket), + _connect_error(connect_error) {} + + handshake::StepResult PrepareResources() override { + if (_transport->PrepareUpgradeResources() == 0) { + return handshake::STEP_OK; + } + errno = 0; + return handshake::STEP_FALLBACK; + } + + handshake::StepResult NegotiateResources() override { + return _transport->NegotiateUpgradeResources( + _protocol->remote(), false) == 0 + ? handshake::STEP_OK : handshake::STEP_FALLBACK; + } + + void OnEstablished() override { _transport->ActivateUpgrade(); } + + void OnFallback() override { _transport->DeactivateUpgrade(); } + + void OnFailed() override { + _transport->DeactivateUpgrade(); + const int saved_errno = errno != 0 ? errno : EPROTO; + *_connect_error = saved_errno; + _socket->SetFailed(saved_errno, + "Fail to complete rdma handshake from %s: %s", + _socket->description().c_str(), + berror(saved_errno)); + } + +private: + RdmaTransport* _transport; + rdma::RdmaHandshakeAdapter* _protocol; + Socket* _socket; + int* _connect_error; +}; +#endif + +#if BRPC_WITH_UBRING +class UBShmClientHandshakeTransport : public handshake::HandshakeTransport { +public: + UBShmClientHandshakeTransport( + UBShmTransport* transport, ubring::SHM* local_shm, + const std::string& shm_name, Socket* socket, int* connect_error) + : _transport(transport), _local_shm(local_shm), + _shm_name(shm_name), _socket(socket), + _connect_error(connect_error) {} + + handshake::StepResult PrepareResources() override { + return _transport->PrepareUpgradeResources( + _local_shm, _shm_name.c_str()) == 0 + ? handshake::STEP_OK : handshake::STEP_FALLBACK; + } + + handshake::StepResult NegotiateResources() override { + return _transport->NegotiateUpgradeResources( + _local_shm, _shm_name.c_str()) == 0 + ? handshake::STEP_OK : handshake::STEP_FALLBACK; + } + + void OnEstablished() override { _transport->ActivateUpgrade(); } + void OnFallback() override { _transport->DeactivateUpgrade(); } + + void OnFailed() override { + _transport->DeactivateUpgrade(); + const int saved_errno = errno != 0 ? errno : EPROTO; + *_connect_error = saved_errno; + _socket->SetFailed(saved_errno, + "Fail to complete ubring handshake from %s: %s", + _socket->description().c_str(), + berror(saved_errno)); + } + +private: + UBShmTransport* _transport; + ubring::SHM* _local_shm; + std::string _shm_name; + Socket* _socket; + int* _connect_error; +}; +#endif + +} // namespace + +AdapterTransport::AdapterTransport(SocketMode mode) + : _mode(mode), _connection_completed(0) {} +AdapterTransport::~AdapterTransport() = default; + +AdapterTransport* AdapterTransport::Get(Socket* socket) { + CHECK(socket != NULL); + return static_cast(socket->_transport.get()); +} + +const AdapterTransport* AdapterTransport::Get(const Socket* socket) { + CHECK(socket != NULL); + 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) { + AdapterTransport* adapter = Get(const_cast(socket)); + ClientHandshakeTask* task = new ClientHandshakeTask{adapter, done, data, SocketUniquePtr()}; + if (Socket::Address(socket->id(), &task->socket) != 0) { + delete task; + return -1; + } + bthread_t tid; + bthread_attr_t attr = BTHREAD_ATTR_NORMAL; + bthread_attr_set_name(&attr, "StartClientUpgrade"); + if (bthread_start_background(&tid, &attr, + ProcessClientHandshake, task) < 0) { + delete task; + return -1; + } + return 0; +} + +ParseResult AdapterTransport::ProcessUpgradeReadable(butil::IOBuf* source) { + ParseResult result(PARSE_ERROR_NOT_ENOUGH_DATA); + if (_socket->parsing_context() != NULL) { + handshake::ServerHandshakeContext* context = + static_cast( + _socket->parsing_context()); + CHECK(context->adapter() != NULL); + result = context->adapter()->ExecuteServerHandshake(source, _socket); + } else if (!source->empty()) { + static const size_t MAX_MAGIC_LEN = 4; + char prefix[MAX_MAGIC_LEN] = {}; + const size_t prefix_len = std::min(source->size(), MAX_MAGIC_LEN); + source->copy_to(prefix, prefix_len); + + const bool matches_ub = + MatchesMagicPrefix(prefix, prefix_len, "UB", 2); + const bool matches_rdma = + MatchesMagicPrefix(prefix, prefix_len, "RDMA", 4) || + MatchesMagicPrefix(prefix, prefix_len, "RDM3", 4); + if (!matches_ub && !matches_rdma) { + result = ParseResult(PARSE_ERROR_TRY_OTHERS); + } else { + handshake::HandshakeAdapter* adapter = + matches_ub + ? handshake::GetUBShmServerHandshakeAdapter() + : handshake::GetRdmaServerHandshakeAdapter(); + result = adapter->ExecuteServerHandshake(source, _socket); + } + } + const int phase = _handshake.phase(); + if (!connection_completed() && + (phase == handshake::ESTABLISHED || + phase == handshake::FALLBACK_TCP || phase == handshake::FAILED)) { + CompleteConnection(static_cast(phase)); + } + return result; +} + +void AdapterTransport::CompleteConnection(handshake::Phase terminal_phase) { + CHECK(terminal_phase == handshake::ESTABLISHED || + terminal_phase == handshake::FALLBACK_TCP || + terminal_phase == handshake::FAILED); + if (terminal_phase == handshake::FAILED && + _handshake.phase() != handshake::FAILED) { + _handshake.MarkFailed(); + } + int expected = 0; + _connection_completed.compare_exchange_strong( + expected, 1, butil::memory_order_release, + butil::memory_order_relaxed); +} + +void* AdapterTransport::ProcessClientHandshake(void* arg) { + std::unique_ptr task( + static_cast(arg)); + AdapterTransport* adapter = task->adapter; + Socket* socket = task->socket.get(); + int connect_error = 0; + (void)connect_error; + +#if BRPC_WITH_RDMA + if (adapter->_mode == SOCKET_MODE_RDMA) { + RdmaTransport* transport = static_cast( + adapter->_high_speed_transport.get()); + if (!rdma::IsRdmaAvailable()) { + adapter->FallbackToTcp(); + adapter->CompleteConnection(handshake::FALLBACK_TCP); + task->done(0, task->data); + return NULL; + } + + std::unique_ptr protocol = + transport->CreateClientHandshakeAdapter(); + CHECK(protocol != NULL); + RdmaClientHandshakeTransport participant( + transport, protocol.get(), socket, &connect_error); + const handshake::StepResult result = adapter->_handshake.RunClient( + protocol.get(), &participant); + if (result == handshake::STEP_OK && + transport->StartUpgradeEvents() < 0) { + const int saved_errno = errno != 0 ? errno : ERDMA; + transport->DeactivateUpgrade(); + adapter->_handshake.MarkFailed(); + socket->SetFailed( + saved_errno, + "Fail to start RDMA CQ events from %s: %s", + socket->description().c_str(), berror(saved_errno)); + connect_error = saved_errno; + } + if (result == handshake::STEP_ERROR && connect_error == 0) { + connect_error = errno != 0 ? errno : EPROTO; + } + adapter->CompleteConnection(static_cast( + adapter->_handshake.phase())); + task->done(connect_error, task->data); + return NULL; + } +#endif + +#if BRPC_WITH_UBRING + if (adapter->_mode == SOCKET_MODE_UBRING) { + UBShmTransport* transport = static_cast( + adapter->_high_speed_transport.get()); + if (!ubring::IsUBAvailable()) { + adapter->FallbackToTcp(); + adapter->CompleteConnection(handshake::FALLBACK_TCP); + task->done(0, task->data); + return NULL; + } + + const size_t local_shm_len = + static_cast(ubring::FLAGS_data_queue_size) * MB_TO_BYTE; + ubring::SHM local_trx_shm = { + NULL, local_shm_len, 0, {0}, static_cast(socket->fd())}; + const auto shm_name_str = + butil::endpoint2str(socket->local_side()); + ubring::UBShmHandshakeAdapter wire; + wire.ConfigureClientHello(local_shm_len, shm_name_str.c_str()); + UBShmClientHandshakeTransport participant( + transport, &local_trx_shm, shm_name_str.c_str(), socket, + &connect_error); + const handshake::StepResult result = adapter->_handshake.RunClient( + &wire, &participant); + if (result == handshake::STEP_OK) { + transport->GetUBShmEp()->SetNegotiatedDataFormat( + ubring::UBR_DATA_FORMAT_LEGACY_64); + transport->FinishUpgrade(); + } + if (result == handshake::STEP_ERROR && connect_error == 0) { + connect_error = errno != 0 ? errno : EPROTO; + } + adapter->CompleteConnection(static_cast( + adapter->_handshake.phase())); + task->done(connect_error, task->data); + return NULL; + } +#endif + + socket->SetFailed(EPROTO, "Unsupported client transport handshake"); + adapter->CompleteConnection(handshake::FAILED); + task->done(EPROTO, task->data); + return NULL; +} + +void AdapterTransport::Init(Socket* socket, const SocketOptions& options) { + CHECK_EQ(_mode, options.socket_mode); + _socket = socket; + _default_connect = options.app_connect; + _on_edge_trigger = options.on_edge_triggered_events; + if (options.need_on_edge_trigger && _on_edge_trigger == NULL) { + if (_mode == SOCKET_MODE_TCP) { + _on_edge_trigger = InputMessenger::OnNewMessages; +#if BRPC_WITH_RDMA + } else if (_mode == SOCKET_MODE_RDMA && + options.user != static_cast( + get_client_side_messenger())) { + // RDMA server handshake is parsed by InputMessenger. + _on_edge_trigger = OnNewMessagesAfterUpgrade; +#endif +#if BRPC_WITH_UBRING + } else if (_mode == SOCKET_MODE_UBRING && + options.user != static_cast( + get_client_side_messenger())) { + // UBSHM server handshake is parsed by InputMessenger. + _on_edge_trigger = OnNewMessagesAfterUpgrade; +#endif + } else { + _on_edge_trigger = OnNewDataFromTcp; + } + } + _handshake.Reset(socket); + _connection_completed.store(0, butil::memory_order_relaxed); + _tcp_transport.reset(new TcpTransport); + _tcp_transport->Init(socket, options); + + switch (_mode) { +#if BRPC_WITH_RDMA + case SOCKET_MODE_RDMA: + _high_speed_transport.reset(new RdmaTransport); + break; +#endif +#if BRPC_WITH_UBRING + case SOCKET_MODE_UBRING: + _high_speed_transport.reset(new UBShmTransport); + break; +#endif + default: + break; + } + if (_high_speed_transport) { + _high_speed_transport->Init(socket, options); + } +} + +void AdapterTransport::Release() { + if (_high_speed_transport) { + _high_speed_transport->Release(); + } + _tcp_transport->Release(); +} + +int AdapterTransport::Reset(int32_t expected_nref) { + if (_high_speed_transport) { + _high_speed_transport->Reset(expected_nref); + } + _tcp_transport->Reset(expected_nref); + _handshake.Reset(_socket); + _connection_completed.store(0, butil::memory_order_relaxed); + return 0; +} + +std::shared_ptr AdapterTransport::Connect() { + if (upgrade_capable(_mode)) { + 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(); +} + +Transport* AdapterTransport::ActiveTransport() const { + if (_high_speed_transport && + _handshake.phase() == handshake::ESTABLISHED) { + return _high_speed_transport.get(); + } + return _tcp_transport.get(); +} + +int AdapterTransport::CutFromIOBuf(butil::IOBuf* buf) { + return ActiveTransport()->CutFromIOBuf(buf); +} + +ssize_t AdapterTransport::CutFromIOBufList( + butil::IOBuf** buf, size_t ndata) { + return ActiveTransport()->CutFromIOBufList(buf, ndata); +} + +int AdapterTransport::WaitEpollOut(butil::atomic* epollout_butex, + bool pollin, timespec duetime) { + return ActiveTransport()->WaitEpollOut( + epollout_butex, pollin, duetime); +} + +void AdapterTransport::ProcessEvent(bthread_attr_t attr) { + ActiveTransport()->ProcessEvent(attr); +} + +void AdapterTransport::QueueMessage(InputMessageClosure& input_msg, + int* num_bthread_created, + bool last_msg) { + ActiveTransport()->QueueMessage( + input_msg, num_bthread_created, last_msg); +} + +void AdapterTransport::Debug(std::ostream& os) { + if (_high_speed_transport) { + _high_speed_transport->Debug(os); + } + const char* state = "UNKNOWN"; + switch (_handshake.phase()) { + case handshake::UNINITIALIZED: state = "UNINITIALIZED"; break; + case handshake::PREPARING: state = "PREPARING"; break; + case handshake::HELLO_SEND: state = "HELLO_SEND"; break; + case handshake::HELLO_WAIT: state = "HELLO_WAIT"; break; + case handshake::NEGOTIATING: state = "NEGOTIATING"; break; + case handshake::ACK_SEND: state = "ACK_SEND"; break; + case handshake::ACK_WAIT: state = "ACK_WAIT"; break; + case handshake::EXTENSION_SEND: state = "EXTENSION_SEND"; break; + case handshake::EXTENSION_WAIT: state = "EXTENSION_WAIT"; break; + case handshake::ESTABLISHED: state = "ESTABLISHED"; break; + case handshake::FALLBACK_TCP: state = "FALLBACK_TCP"; break; + case handshake::FAILED: state = "FAILED"; break; + } + os << "\nhandshake_state=" << state + << "\nhandshake_version=" << _handshake.protocol_version(); +} + +void AdapterTransport::FallbackToTcp() { + _handshake.PublishFallback([this]() { + SetHighSpeedAvailable(false); + }); +} + +void AdapterTransport::SetHighSpeedAvailable(bool available) { + if (!_high_speed_transport) { + return; + } + switch (_mode) { +#if BRPC_WITH_RDMA + case SOCKET_MODE_RDMA: + static_cast(_high_speed_transport.get()) + ->SetHighSpeedAvailable(available); + break; +#endif +#if BRPC_WITH_UBRING + case SOCKET_MODE_UBRING: + static_cast(_high_speed_transport.get()) + ->SetHighSpeedAvailable(available); + break; +#endif + default: + break; + } +} + +void AdapterTransport::OnNewMessagesAfterUpgrade(Socket* socket) { + AdapterTransport* adapter = Get(socket); + if (adapter->_handshake.phase() == handshake::ESTABLISHED) { + adapter->CheckUnexpectedTcpData(); + return; + } + + InputMessenger::OnNewMessages(socket); + +#if BRPC_WITH_RDMA + if (adapter->_mode == SOCKET_MODE_RDMA && + adapter->_handshake.phase() == handshake::ESTABLISHED) { + RdmaTransport* transport = static_cast( + adapter->_high_speed_transport.get()); + if (transport->StartUpgradeEvents() < 0) { + const int saved_errno = errno != 0 ? errno : ERDMA; + transport->DeactivateUpgrade(); + adapter->_handshake.MarkFailed(); + adapter->CompleteConnection(handshake::FAILED); + socket->SetFailed( + saved_errno, + "Fail to start RDMA CQ events from %s: %s", + socket->description().c_str(), berror(saved_errno)); + } + } +#endif +} + +void AdapterTransport::OnNewDataFromTcp(Socket* socket) { + static_cast(socket->_transport.get())->ProcessTcpEvent(); +} + +void AdapterTransport::ProcessTcpEvent() { + int progress = Socket::PROGRESS_INIT; + while (true) { + const int phase = _handshake.phase(); + if (phase != handshake::UNINITIALIZED && + phase < handshake::ESTABLISHED) { + _handshake.NotifyReadable(); + } else if (phase == handshake::FALLBACK_TCP) { + InputMessenger::OnNewMessages(_socket); + return; + } else if (phase == handshake::ESTABLISHED) { + CheckUnexpectedTcpData(); + return; + } + if (!_socket->MoreReadEvents(&progress)) { + break; + } + } +} + +void AdapterTransport::CheckUnexpectedTcpData() { + int progress = Socket::PROGRESS_INIT; + while (true) { + uint8_t byte; + const ssize_t nr = read(_socket->fd(), &byte, 1); + if (nr == 0) { + _socket->SetEOF(); + return; + } + if (nr > 0) { + _socket->SetFailed(EPROTO, "Read unexpected data from %s", + _socket->description().c_str()); + return; + } + if (errno == EINTR) { + continue; + } + if (errno != EAGAIN) { + const int saved_errno = errno; + _socket->SetFailed(saved_errno, "Fail to read from %s: %s", + _socket->description().c_str(), + berror(saved_errno)); + return; + } + if (!_socket->MoreReadEvents(&progress)) { + return; + } + } +} + +void AdapterTransport::TryReadOnTcp() { + if (_socket->_nevent.fetch_add(1, butil::memory_order_acq_rel) != 0) { + return; + } + const int phase = _handshake.phase(); + if (phase == handshake::FALLBACK_TCP) { + InputMessenger::OnNewMessages(_socket); + } else if (phase == handshake::ESTABLISHED) { + CheckUnexpectedTcpData(); + } +} + +} // namespace brpc diff --git a/src/brpc/adapter_transport.h b/src/brpc/adapter_transport.h new file mode 100644 index 0000000000..76041847d8 --- /dev/null +++ b/src/brpc/adapter_transport.h @@ -0,0 +1,101 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#ifndef BRPC_ADAPTER_TRANSPORT_H +#define BRPC_ADAPTER_TRANSPORT_H + +#include + +#include "brpc/socket_mode.h" +#include "brpc/transport.h" +#include "brpc/transport_handshake.h" +#include "brpc/parse_result.h" + +namespace brpc { + +class TcpTransport; +class RdmaTransport; +class UBShmTransport; + +// The top-level Transport installed in Socket. It starts on TcpTransport and +// may switch to an independent RDMA/URMA/UBSHM Transport after a successful +// handshake. TCP remains usable before negotiation and after fallback. +class AdapterTransport : public Transport { + friend class TransportFactory; + friend class RdmaTransport; + friend class UBShmTransport; +public: + void Init(Socket* socket, const SocketOptions& options) override; + void Release() override; + int Reset(int32_t expected_nref) override; + std::shared_ptr Connect() override; + int CutFromIOBuf(butil::IOBuf* buf) override; + ssize_t CutFromIOBufList(butil::IOBuf** buf, size_t ndata) override; + int WaitEpollOut(butil::atomic* epollout_butex, + bool pollin, timespec duetime) override; + void ProcessEvent(bthread_attr_t attr) override; + void QueueMessage(InputMessageClosure& input_msg, + int* num_bthread_created, bool last_msg) override; + void Debug(std::ostream& os) override; + + int handshake_phase() const { return _handshake.phase(); } + int handshake_version() const { return _handshake.protocol_version(); } + handshake::HandshakeSession* handshake_session() { return &_handshake; } + Transport* high_speed_transport() const { + return _high_speed_transport.get(); + } + bool upgrade_capable(SocketMode mode) const; + + static AdapterTransport* Get(Socket* socket); + static const AdapterTransport* Get(const Socket* socket); + + // The only client-side upgrade entry point. Concrete transports provide + // resources; AdapterTransport owns the handshake orchestration. + static int StartClientUpgrade(const Socket* socket, + void (*done)(int, void*), void* data); + + ParseResult ProcessUpgradeReadable(butil::IOBuf* source); + void CompleteConnection(handshake::Phase terminal_phase); + bool connection_completed() const { + return _connection_completed.load(butil::memory_order_acquire) != 0; + } + + static void OnNewDataFromTcp(Socket* socket); + +private: + explicit AdapterTransport(SocketMode mode); + ~AdapterTransport() override; + + Transport* ActiveTransport() const; + void SetHighSpeedAvailable(bool available); + void FallbackToTcp(); + void TryReadOnTcp(); + void ProcessTcpEvent(); + void CheckUnexpectedTcpData(); + static void OnNewMessagesAfterUpgrade(Socket* socket); + static void* ProcessClientHandshake(void* arg); + + SocketMode _mode; + handshake::HandshakeSession _handshake; + std::unique_ptr _tcp_transport; + std::unique_ptr _high_speed_transport; + butil::atomic _connection_completed; +}; + +} // namespace brpc + +#endif // BRPC_ADAPTER_TRANSPORT_H diff --git a/src/brpc/global.cpp b/src/brpc/global.cpp index 0a0837e096..d64ae9f60a 100644 --- a/src/brpc/global.cpp +++ b/src/brpc/global.cpp @@ -69,7 +69,7 @@ // Protocols #include "brpc/protocol.h" -#include "brpc/policy/rdma_handshake_protocol.h" +#include "brpc/policy/transport_handshake_protocol.h" #include "brpc/policy/baidu_rpc_protocol.h" #include "brpc/policy/http_rpc_protocol.h" #include "brpc/policy/http2_rpc_protocol.h" @@ -438,12 +438,16 @@ static void GlobalInitializeOrDieImpl() { } // Protocols - Protocol rdma_handshake_protocol = { - ParseRdmaHandshake, nullptr, nullptr, - ProcessRdmaHandshake, nullptr, + Protocol transport_handshake_protocol = { + ParseTransportHandshake, nullptr, nullptr, + ProcessTransportHandshake, nullptr, nullptr, nullptr, nullptr, CONNECTION_TYPE_ALL, "rdma_handshake" }; - if (RegisterProtocol(PROTOCOL_RDMA_HANDSHAKE, rdma_handshake_protocol) != 0) { + // Retain the existing enum value and registered name to avoid changing + // public protocol identifiers while widening the implementation from RDMA + // to all transport upgrades. + if (RegisterProtocol(PROTOCOL_RDMA_HANDSHAKE, + transport_handshake_protocol) != 0) { exit(1); } diff --git a/src/brpc/handshake/handshake_adapter.cpp b/src/brpc/handshake/handshake_adapter.cpp new file mode 100644 index 0000000000..0c9c72765a --- /dev/null +++ b/src/brpc/handshake/handshake_adapter.cpp @@ -0,0 +1,55 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include "brpc/handshake/handshake_adapter.h" + +#include "brpc/socket.h" + +namespace brpc { +namespace handshake { + +// InputMessenger may call this entry repeatedly while bytes arrive. Keep all +// connection state in HandshakeSession and Socket rather than in the adapter, +// so a single stateless adapter can serve every connection. The parsing +// context retains that selected adapter between the server hello and the peer +// ACK, whose frame has no protocol magic of its own. +ParseResult StandardHandshakeAdapter::ExecuteServerHandshake( + butil::IOBuf* source, Socket* socket) { + const StepResult result = RunServerStep(source, socket); + if (result == STEP_NEED_MORE) { + if (GetSession(socket)->phase() == ACK_WAIT && + socket->parsing_context() == NULL) { + ServerHandshakeContext* context = + ServerHandshakeContext::Create(this); + if (context == NULL) { + GetSession(socket)->MarkFailed(); + return MakeParseError(PARSE_ERROR_ABSOLUTELY_WRONG); + } + socket->reset_parsing_context(context); + } + return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); + } + + socket->reset_parsing_context(NULL); + if (result == STEP_ERROR) { + return MakeParseError(PARSE_ERROR_ABSOLUTELY_WRONG); + } + return MakeParseError(PARSE_ERROR_TRY_OTHERS); +} + +} // namespace handshake +} // namespace brpc diff --git a/src/brpc/handshake/handshake_adapter.h b/src/brpc/handshake/handshake_adapter.h new file mode 100644 index 0000000000..2a344c9ed3 --- /dev/null +++ b/src/brpc/handshake/handshake_adapter.h @@ -0,0 +1,73 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#ifndef BRPC_HANDSHAKE_HANDSHAKE_ADAPTER_H +#define BRPC_HANDSHAKE_HANDSHAKE_ADAPTER_H + +#include "butil/macros.h" +#include "brpc/parse_result.h" +#include "brpc/transport_handshake.h" + +namespace butil { +class IOBuf; +} + +namespace brpc { + +class Socket; + +namespace handshake { + +// The minimal seam between an upgrade protocol and InputMessenger. Protocols +// that cannot use the standard parser lifecycle implement this interface +// directly. +class HandshakeAdapter { +public: + virtual ~HandshakeAdapter() = default; + + virtual ParseResult ExecuteServerHandshake( + butil::IOBuf* source, Socket* socket) = 0; + +protected: + HandshakeAdapter() = default; + +private: + DISALLOW_COPY_AND_ASSIGN(HandshakeAdapter); +}; + +// Reusable InputMessenger implementation. Protocol adapters provide only the +// protocol-specific server step; the common session owns all phases. +class StandardHandshakeAdapter : public HandshakeAdapter { +public: + ParseResult ExecuteServerHandshake( + butil::IOBuf* source, Socket* socket) override; + +protected: + StandardHandshakeAdapter() = default; + + virtual StepResult RunServerStep( + butil::IOBuf* source, Socket* socket) = 0; + virtual HandshakeSession* GetSession(Socket* socket) const = 0; + +private: + DISALLOW_COPY_AND_ASSIGN(StandardHandshakeAdapter); +}; + +} // namespace handshake +} // namespace brpc + +#endif // BRPC_HANDSHAKE_HANDSHAKE_ADAPTER_H diff --git a/src/brpc/handshake/handshake_frame.cpp b/src/brpc/handshake/handshake_frame.cpp new file mode 100644 index 0000000000..f25e50c15c --- /dev/null +++ b/src/brpc/handshake/handshake_frame.cpp @@ -0,0 +1,235 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include "brpc/handshake/handshake_frame.h" + +#include +#include +#include + +#include "butil/sys_byteorder.h" + +namespace brpc { +namespace handshake { + +size_t FrameCodec::LengthFieldSize(const FrameSpec& spec) { + switch (spec.length_encoding) { + case FrameSpec::FIXED: return 0; + case FrameSpec::U16_TOTAL_LENGTH: return sizeof(uint16_t); + case FrameSpec::U32_BODY_LENGTH: return sizeof(uint32_t); + } + return 0; +} + +FrameResult FrameCodec::DecodeLength(const FrameSpec& spec, + const void* header, + size_t* frame_len) { + const size_t length_size = LengthFieldSize(spec); + const size_t header_len = spec.magic_len + length_size; + if (spec.min_frame_len < header_len || + spec.max_frame_len < spec.min_frame_len) { + return FRAME_PROTOCOL_ERROR; + } + + if (spec.length_encoding == FrameSpec::FIXED) { + if (spec.min_frame_len != spec.max_frame_len) { + return FRAME_PROTOCOL_ERROR; + } + *frame_len = spec.min_frame_len; + } else if (spec.length_encoding == FrameSpec::U16_TOTAL_LENGTH) { + uint16_t total_be = 0; + memcpy(&total_be, static_cast(header) + spec.magic_len, + sizeof(total_be)); + *frame_len = butil::NetToHost16(total_be); + } else { + uint32_t body_be = 0; + memcpy(&body_be, static_cast(header) + spec.magic_len, + sizeof(body_be)); + const size_t body_len = butil::NetToHost32(body_be); + if (body_len > std::numeric_limits::max() - header_len) { + return FRAME_PROTOCOL_ERROR; + } + *frame_len = header_len + body_len; + } + + if (*frame_len < spec.min_frame_len || + *frame_len > spec.max_frame_len) { + return FRAME_PROTOCOL_ERROR; + } + return FRAME_OK; +} + +FrameResult FrameCodec::Encode(const FrameSpec& spec, + const std::string& payload, + std::string* frame) { + if (frame == NULL || (spec.magic_len != 0 && spec.magic == NULL)) { + return FRAME_PROTOCOL_ERROR; + } + const size_t length_size = LengthFieldSize(spec); + const size_t header_len = spec.magic_len + length_size; + if (payload.size() > std::numeric_limits::max() - header_len) { + return FRAME_PROTOCOL_ERROR; + } + const size_t total_len = header_len + payload.size(); + if (total_len < spec.min_frame_len || total_len > spec.max_frame_len) { + return FRAME_PROTOCOL_ERROR; + } + if (spec.length_encoding == FrameSpec::FIXED && + spec.min_frame_len != spec.max_frame_len) { + return FRAME_PROTOCOL_ERROR; + } + if (spec.length_encoding == FrameSpec::U16_TOTAL_LENGTH && + total_len > std::numeric_limits::max()) { + return FRAME_PROTOCOL_ERROR; + } + if (spec.length_encoding == FrameSpec::U32_BODY_LENGTH && + payload.size() > std::numeric_limits::max()) { + return FRAME_PROTOCOL_ERROR; + } + + frame->clear(); + frame->reserve(total_len); + if (spec.magic_len != 0) { + frame->append(spec.magic, spec.magic_len); + } + if (spec.length_encoding == FrameSpec::U16_TOTAL_LENGTH) { + const uint16_t total_be = + butil::HostToNet16(static_cast(total_len)); + frame->append(reinterpret_cast(&total_be), + sizeof(total_be)); + } else if (spec.length_encoding == FrameSpec::U32_BODY_LENGTH) { + const uint32_t body_be = + butil::HostToNet32(static_cast(payload.size())); + frame->append(reinterpret_cast(&body_be), + sizeof(body_be)); + } + frame->append(payload); + return FRAME_OK; +} + +FrameResult FrameCodec::ReadFrame(HandshakeIO* io, const FrameSpec& spec, + bool push_back_on_not_mine, + std::string* payload) { + if (io == NULL || payload == NULL || + (spec.magic_len != 0 && spec.magic == NULL)) { + return FRAME_PROTOCOL_ERROR; + } + + std::string header(spec.magic_len + LengthFieldSize(spec), '\0'); + if (spec.magic_len != 0 && + io->ReadExact(&header[0], spec.magic_len) < 0) { + return FRAME_IO_ERROR; + } + if (spec.magic_len != 0 && + memcmp(header.data(), spec.magic, spec.magic_len) != 0) { + if (push_back_on_not_mine && + io->PushBack(header.data(), spec.magic_len) < 0) { + return FRAME_IO_ERROR; + } + return FRAME_NOT_MINE; + } + + const size_t length_size = LengthFieldSize(spec); + if (length_size != 0 && + io->ReadExact(&header[spec.magic_len], length_size) < 0) { + return FRAME_IO_ERROR; + } + size_t frame_len = 0; + FrameResult result = DecodeLength(spec, header.data(), &frame_len); + if (result != FRAME_OK) { + return result; + } + const size_t body_len = frame_len - header.size(); + payload->assign(body_len, '\0'); + if (body_len != 0 && io->ReadExact(&(*payload)[0], body_len) < 0) { + return FRAME_IO_ERROR; + } + return FRAME_OK; +} + +FrameResult FrameCodec::ParseBufferedFrame(HandshakeInput* input, + const FrameSpec& spec, + std::string* payload, + bool* magic_matched) { + if (magic_matched != NULL) { + *magic_matched = false; + } + if (input == NULL || payload == NULL || + (spec.magic_len != 0 && spec.magic == NULL)) { + return FRAME_PROTOCOL_ERROR; + } + const size_t header_len = spec.magic_len + LengthFieldSize(spec); + if (input->Size() < spec.magic_len) { + return FRAME_NEED_MORE; + } + std::string header(header_len, '\0'); + if (spec.magic_len != 0 && + !input->CopyTo(&header[0], spec.magic_len)) { + return FRAME_NEED_MORE; + } + if (spec.magic_len != 0 && + memcmp(header.data(), spec.magic, spec.magic_len) != 0) { + return FRAME_NOT_MINE; + } + if (magic_matched != NULL) { + *magic_matched = true; + } + if (input->Size() < header_len || + (header_len != 0 && !input->CopyTo(&header[0], header_len))) { + return FRAME_NEED_MORE; + } + size_t frame_len = 0; + FrameResult result = DecodeLength(spec, header.data(), &frame_len); + if (result != FRAME_OK) { + return result; + } + if (input->Size() < frame_len) { + return FRAME_NEED_MORE; + } + + std::string frame(frame_len, '\0'); + if (!input->CopyTo(&frame[0], frame_len)) { + return FRAME_NEED_MORE; + } + if (!input->Consume(frame_len)) { + return FRAME_PROTOCOL_ERROR; + } + payload->assign(frame.data() + header_len, frame_len - header_len); + return FRAME_OK; +} + +FrameResult FrameCodec::WriteFrame(HandshakeIO* io, const FrameSpec& spec, + const std::string& payload) { + if (io == NULL) { + return FRAME_PROTOCOL_ERROR; + } + std::string frame; + const FrameResult result = Encode(spec, payload, &frame); + if (result != FRAME_OK) { + return result; + } + return io->WriteAll(frame.data(), frame.size()) == 0 + ? FRAME_OK : FRAME_IO_ERROR; +} + +FrameResult FrameCodec::DrainFrame(HandshakeIO* io, const FrameSpec& spec) { + std::string ignored; + return ReadFrame(io, spec, false, &ignored); +} + +} // namespace handshake +} // namespace brpc diff --git a/src/brpc/handshake/handshake_frame.h b/src/brpc/handshake/handshake_frame.h new file mode 100644 index 0000000000..dbe399f738 --- /dev/null +++ b/src/brpc/handshake/handshake_frame.h @@ -0,0 +1,92 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#ifndef BRPC_HANDSHAKE_HANDSHAKE_FRAME_H +#define BRPC_HANDSHAKE_HANDSHAKE_FRAME_H + +#include +#include + +#include "brpc/handshake/handshake_io.h" + +namespace brpc { +namespace handshake { + +enum FrameResult { + FRAME_OK = 0, + FRAME_NOT_MINE, + FRAME_NEED_MORE, + FRAME_IO_ERROR, + FRAME_PROTOCOL_ERROR, +}; + +struct FrameSpec { + enum LengthEncoding { + FIXED, + U16_TOTAL_LENGTH, + U32_BODY_LENGTH, + }; + + FrameSpec() + : magic(NULL), magic_len(0), min_frame_len(0), max_frame_len(0), + length_encoding(FIXED) {} + + FrameSpec(const char* magic_in, size_t magic_len_in, + size_t min_frame_len_in, size_t max_frame_len_in, + LengthEncoding length_encoding_in) + : magic(magic_in), magic_len(magic_len_in), + min_frame_len(min_frame_len_in), + max_frame_len(max_frame_len_in), + length_encoding(length_encoding_in) {} + + const char* magic; + size_t magic_len; + size_t min_frame_len; + size_t max_frame_len; + LengthEncoding length_encoding; +}; + +// Handles only framing. Protocol implementations receive and produce payloads +// after magic/length fields and remain responsible for their own business +// fields and version semantics. +class FrameCodec { +public: + static FrameResult Encode(const FrameSpec& spec, + const std::string& payload, + std::string* frame); + static FrameResult ReadFrame(HandshakeIO* io, const FrameSpec& spec, + bool push_back_on_not_mine, + std::string* payload); + static FrameResult ParseBufferedFrame(HandshakeInput* input, + const FrameSpec& spec, + std::string* payload, + bool* magic_matched = NULL); + static FrameResult WriteFrame(HandshakeIO* io, const FrameSpec& spec, + const std::string& payload); + static FrameResult DrainFrame(HandshakeIO* io, const FrameSpec& spec); + +private: + static size_t LengthFieldSize(const FrameSpec& spec); + static FrameResult DecodeLength(const FrameSpec& spec, + const void* header, + size_t* frame_len); +}; + +} // namespace handshake +} // namespace brpc + +#endif // BRPC_HANDSHAKE_HANDSHAKE_FRAME_H diff --git a/src/brpc/handshake/handshake_io.cpp b/src/brpc/handshake/handshake_io.cpp new file mode 100644 index 0000000000..67cc95e37c --- /dev/null +++ b/src/brpc/handshake/handshake_io.cpp @@ -0,0 +1,162 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include "brpc/handshake/handshake_io.h" + +#include +#include +#include + +#include "bthread/butex.h" +#include "butil/time.h" +#include "brpc/errno.pb.h" +#include "brpc/socket.h" + +namespace brpc { +namespace handshake { + +size_t IOBufHandshakeInput::Size() const { + return _source != NULL ? _source->size() : 0; +} + +bool IOBufHandshakeInput::CopyTo(void* data, size_t len) const { + return _source != NULL && _source->copy_to(data, len) == len; +} + +bool IOBufHandshakeInput::Consume(size_t len) { + return _source != NULL && _source->pop_front(len) == len; +} + +static const int WAIT_TIMEOUT_MS = 50; + +SocketHandshakeIO::SocketHandshakeIO(Socket* socket) + : _socket(socket) + , _read_butex(bthread::butex_create_checked >()) { +} + +SocketHandshakeIO::~SocketHandshakeIO() { + bthread::butex_destroy(_read_butex); +} + +void SocketHandshakeIO::Reset(Socket* socket) { + _socket = socket; +} + +void SocketHandshakeIO::NotifyReadable() { + _read_butex->fetch_add(1, butil::memory_order_release); + bthread::butex_wake(_read_butex); +} + +template +static int ReadExactLoop(butil::atomic* read_butex, + SocketId socket_id, + size_t len, ReadOnce read_once) { + size_t received = 0; + while (received < len) { + const int expected = read_butex->load(butil::memory_order_acquire); + 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; + } + SocketUniquePtr alive; + if (Socket::Address(socket_id, &alive) != 0) { + errno = EFAILEDSOCKET; + return -1; + } + if (bthread::butex_wait(read_butex, expected, &duetime) < 0 && + errno != EWOULDBLOCK && errno != ETIMEDOUT) { + return -1; + } + } else if (nr == 0) { + errno = EEOF; + return -1; + } else { + received += nr; + } + } + return 0; +} + +int SocketHandshakeIO::ReadExact(void* data, size_t len) { + CHECK(data != NULL); + CHECK(_socket != NULL); + const int fd = _socket->fd(); + return ReadExactLoop(_read_butex, _socket->id(), len, + [data, fd](size_t offset, size_t remaining) { + return read(fd, static_cast(data) + offset, remaining); + }); +} + +template +static int WriteAllLoop(size_t len, WriteOnce write_once, + WaitWritable wait_writable) { + size_t written = 0; + while (written < len) { + const timespec duetime = butil::milliseconds_from_now(WAIT_TIMEOUT_MS); + const ssize_t nw = write_once(written, len - written); + if (nw > 0) { + written += nw; + continue; + } + if (nw == 0) { + errno = EPIPE; + return -1; + } + if (errno == EINTR) { + continue; + } + if (errno != EAGAIN) { + return -1; + } + if (wait_writable(&duetime) < 0 && + errno != ETIMEDOUT) { + return -1; + } + } + return 0; +} + +int SocketHandshakeIO::WriteAll(const void* data, size_t len) { + CHECK(data != NULL); + CHECK(_socket != NULL); + const int fd = _socket->fd(); + return WriteAllLoop(len, + [data, fd](size_t offset, size_t remaining) { + return write(fd, static_cast(data) + offset, + remaining); + }, + [this](const timespec* duetime) { + return _socket->WaitEpollOut( + _socket->fd(), true, duetime); + }); +} + +int SocketHandshakeIO::PushBack(const void* data, size_t len) { + CHECK(_socket != NULL); + if (len != 0) { + return _socket->fd_input_processor().read_buf().append(data, len) == 0 ? 0 : -1; + } + return 0; +} + +} // namespace handshake +} // namespace brpc diff --git a/src/brpc/handshake/handshake_io.h b/src/brpc/handshake/handshake_io.h new file mode 100644 index 0000000000..82f83ab5cd --- /dev/null +++ b/src/brpc/handshake/handshake_io.h @@ -0,0 +1,89 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#ifndef BRPC_HANDSHAKE_HANDSHAKE_IO_H +#define BRPC_HANDSHAKE_HANDSHAKE_IO_H + +#include + +#include "butil/atomicops.h" +#include "butil/iobuf.h" +#include "butil/macros.h" + +namespace brpc { + +class Socket; + +namespace handshake { + +// Blocking byte-stream interface used by client handshakes and by protocols +// whose server handshake still runs in a dedicated bthread. +class HandshakeIO { +public: + virtual ~HandshakeIO() = default; + + virtual int ReadExact(void* data, size_t len) = 0; + virtual int WriteAll(const void* data, size_t len) = 0; + virtual int PushBack(const void* data, size_t len) = 0; +}; + +// Non-blocking input used by the standard InputMessenger parser path. +// Consume is called only after a complete frame has been validated. +class HandshakeInput { +public: + virtual ~HandshakeInput() = default; + + virtual size_t Size() const = 0; + virtual bool CopyTo(void* data, size_t len) const = 0; + virtual bool Consume(size_t len) = 0; +}; + +class IOBufHandshakeInput : public HandshakeInput { +public: + explicit IOBufHandshakeInput(butil::IOBuf* source) : _source(source) {} + + size_t Size() const override; + bool CopyTo(void* data, size_t len) const override; + bool Consume(size_t len) override; + +private: + butil::IOBuf* _source; +}; + +class SocketHandshakeIO : public HandshakeIO { +public: + explicit SocketHandshakeIO(Socket* socket = NULL); + ~SocketHandshakeIO() override; + + void Reset(Socket* socket); + void NotifyReadable(); + + int ReadExact(void* data, size_t len) override; + int WriteAll(const void* data, size_t len) override; + int PushBack(const void* data, size_t len) override; + +private: + Socket* _socket; + butil::atomic* _read_butex; + + DISALLOW_COPY_AND_ASSIGN(SocketHandshakeIO); +}; + +} // namespace handshake +} // namespace brpc + +#endif // BRPC_HANDSHAKE_HANDSHAKE_IO_H diff --git a/src/brpc/handshake/rdma_handshake.cpp b/src/brpc/handshake/rdma_handshake.cpp new file mode 100644 index 0000000000..26c622a1ac --- /dev/null +++ b/src/brpc/handshake/rdma_handshake.cpp @@ -0,0 +1,550 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include "brpc/handshake/rdma_handshake.h" + +#include +#include +#include + +#include "butil/logging.h" +#include "butil/raw_pack.h" +#include "butil/sys_byteorder.h" +#include "brpc/adapter_transport.h" +#include "brpc/handshake/rdma_handshake_constants.h" +#include "brpc/rdma_handshake.pb.h" +#include "brpc/socket.h" + +#if BRPC_WITH_RDMA + +#include + +#include + +#include "brpc/rdma_transport.h" + +namespace brpc { +namespace rdma { + +DEFINE_int32(rdma_client_handshake_version, 2, + "RDMA handshake protocol version used by client. " + "2 = legacy 'RDMA' magic (default, compatible with all servers); " + "3 = new 'RDM3' protobuf-based handshake " + "(MUST only be enabled after target servers support v3)."); +DECLARE_bool(rdma_trace_verbose); + +extern const uint16_t MIN_QP_SIZE; +extern const uint16_t MIN_BLOCK_SIZE; +extern bool g_skip_rdma_init; + +DEFINE_bool(rdma_ece, false, + "Enable end-to-end ECE negotiation in the RDMA v3 handshake"); + +void RdmaHandshakeAdapter::FillLocalHello(ParsedHello* local) const { + _ep->GetLocalConnectionInfo(local); +} + +void RdmaHandshakeAdapter::PrepareClientEce() { + if (!FLAGS_rdma_ece) { + return; + } + ibv_ece ece; + const int rc = _ep->QueryLocalEce(&ece); + if (rc == 0) { + _ep->SetOutgoingEce(ece); + } else if (rc < 0) { + LOG_IF(WARNING, FLAGS_rdma_trace_verbose) + << "Fail to IbvQueryEce on client, ECE not advertised"; + } +} + +const handshake::FrameSpec& RdmaHandshakeAdapter::AckFrameSpec() const { + return RdmaAckFrameSpec(); +} + +handshake::StepResult RdmaHandshakeAdapter::BuildHello( + bool enabled, std::string* payload) { + return BuildLocalHello(enabled, payload); +} + +handshake::StepResult RdmaHandshakeAdapter::ParseHello( + const std::string& payload) { + return ParseRemoteHello(payload, &_remote); +} + +namespace v2_wire { + +void HelloMessage::Serialize(void* data) const { + butil::RawPacker(data) + .pack16(msg_len) + .pack16(hello_ver) + .pack16(impl_ver) + .pack32(block_size) + .pack16(sq_size) + .pack16(rq_size) + .pack16(lid) + .pack_bytes(gid.raw, sizeof(gid.raw)) + .pack32(qp_num); +} + +void HelloMessage::Deserialize(const void* data) { + butil::RawUnpacker(data) + .unpack16(msg_len) + .unpack16(hello_ver) + .unpack16(impl_ver) + .unpack32(block_size) + .unpack16(sq_size) + .unpack16(rq_size) + .unpack16(lid) + .unpack_bytes(gid.raw, sizeof(gid.raw)) + .unpack32(qp_num); +} + +static bool ValidHelloMessage(const HelloMessage& msg) { + return msg.hello_ver == HELLO_V2_VERSION && + msg.impl_ver == IMPL_V2_VERSION && + msg.block_size >= MIN_BLOCK_SIZE && + msg.sq_size >= MIN_QP_SIZE && + msg.rq_size >= MIN_QP_SIZE; +} + +static void TranslateHello(const HelloMessage& msg, ParsedHello* out) { + out->block_size = msg.block_size; + out->sq_size = msg.sq_size; + out->rq_size = msg.rq_size; + out->lid = msg.lid; + out->gid = msg.gid; + out->qp_num = msg.qp_num; +} + +static void FillMessage(const ParsedHello& local, HelloMessage* msg) { + msg->msg_len = HELLO_V2_MSG_LEN_MIN; + msg->hello_ver = HELLO_V2_VERSION; + msg->impl_ver = IMPL_V2_VERSION; + msg->block_size = local.block_size; + msg->sq_size = local.sq_size; + msg->rq_size = local.rq_size; + msg->lid = local.lid; + msg->gid = local.gid; + msg->qp_num = local.qp_num; +} + +static handshake::StepResult SerializePayload( + const HelloMessage& msg, std::string* payload) { + uint8_t body[HELLO_V2_MSG_LEN_MIN - HELLO_MAGIC_LEN]; + msg.Serialize(body); + // FrameCodec owns msg_len, so the protocol payload starts after it. + payload->assign(reinterpret_cast(body + sizeof(uint16_t)), + sizeof(body) - sizeof(uint16_t)); + return handshake::STEP_OK; +} + +static handshake::StepResult ParsePayload( + const std::string& payload, ParsedHello* remote) { + const size_t base_payload_len = + HELLO_V2_MSG_LEN_MIN - HELLO_MAGIC_LEN - sizeof(uint16_t); + if (payload.size() < base_payload_len) { + errno = EPROTO; + return handshake::STEP_ERROR; + } + uint8_t body[HELLO_V2_MSG_LEN_MIN - HELLO_MAGIC_LEN]; + const uint16_t total_be = butil::HostToNet16( + static_cast(HELLO_MAGIC_LEN + sizeof(uint16_t) + + payload.size())); + memcpy(body, &total_be, sizeof(total_be)); + memcpy(body + sizeof(total_be), payload.data(), base_payload_len); + + HelloMessage msg{}; + msg.Deserialize(body); + if (!ValidHelloMessage(msg)) { + return handshake::STEP_FALLBACK; + } + TranslateHello(msg, remote); + return handshake::STEP_OK; +} + +} // namespace v2_wire + +const handshake::FrameSpec& +RdmaClientHandshakeAdapterV2::HelloFrameSpec() const { + return RdmaHelloFrameSpec(2); +} + +handshake::StepResult RdmaClientHandshakeAdapterV2::BuildLocalHello( + bool enabled, std::string* payload) { + CHECK(enabled); + ParsedHello local{}; + FillLocalHello(&local); + v2_wire::HelloMessage msg{}; + v2_wire::FillMessage(local, &msg); + return v2_wire::SerializePayload(msg, payload); +} + +handshake::StepResult RdmaClientHandshakeAdapterV2::ParseRemoteHello( + const std::string& payload, ParsedHello* remote) { + return v2_wire::ParsePayload(payload, remote); +} + +const handshake::FrameSpec& +RdmaServerHandshakeAdapterV2::HelloFrameSpec() const { + return RdmaHelloFrameSpec(2); +} + +handshake::StepResult RdmaServerHandshakeAdapterV2::BuildLocalHello( + bool enabled, std::string* payload) { + v2_wire::HelloMessage msg{}; + msg.msg_len = HELLO_V2_MSG_LEN_MIN; + if (enabled) { + ParsedHello local{}; + FillLocalHello(&local); + v2_wire::FillMessage(local, &msg); + } + return v2_wire::SerializePayload(msg, payload); +} + +handshake::StepResult RdmaServerHandshakeAdapterV2::ParseRemoteHello( + const std::string& payload, ParsedHello* remote) { + return v2_wire::ParsePayload(payload, remote); +} + +namespace v3_wire { + +static bool ValidRdmaHello(const RdmaHello& msg) { + if (msg.gid().size() != sizeof(ibv_gid)) { + return false; + } + const uint16_t max_uint16 = std::numeric_limits::max(); + if (msg.sq_size() > max_uint16 || msg.rq_size() > max_uint16 || + msg.lid() > max_uint16) { + return false; + } + if (msg.block_size() < MIN_BLOCK_SIZE || msg.sq_size() < MIN_QP_SIZE || + msg.rq_size() < MIN_QP_SIZE) { + return false; + } + return msg.qp_num() != 0 || g_skip_rdma_init; +} + +static void FillLocalRdmaHello(const ParsedHello& local, RdmaHello* msg) { + msg->set_block_size(local.block_size); + msg->set_sq_size(local.sq_size); + msg->set_rq_size(local.rq_size); + msg->set_lid(local.lid); + msg->set_gid(reinterpret_cast(local.gid.raw), + sizeof(local.gid.raw)); + msg->set_qp_num(local.qp_num); + if (FLAGS_rdma_ece && local.ece.has_value()) { + RdmaEce* ece = msg->mutable_ece(); + ece->set_vendor_id(local.ece->vendor_id); + ece->set_options(local.ece->options); + ece->set_comp_mask(local.ece->comp_mask); + } +} + +static void TranslateHello(const RdmaHello& msg, ParsedHello* out) { + out->block_size = msg.block_size(); + out->sq_size = static_cast(msg.sq_size()); + out->rq_size = static_cast(msg.rq_size()); + out->lid = static_cast(msg.lid()); + fast_memcpy(out->gid.raw, msg.gid().data(), sizeof(out->gid.raw)); + out->qp_num = msg.qp_num(); + if (FLAGS_rdma_ece && msg.has_ece()) { + ibv_ece ece; + ece.vendor_id = msg.ece().vendor_id(); + ece.options = msg.ece().options(); + ece.comp_mask = msg.ece().comp_mask(); + out->ece = ece; + } +} + +static handshake::StepResult SerializePayload( + const RdmaHello& msg, std::string* payload) { + if (!msg.SerializeToString(payload) || + payload->size() > HELLO_V3_MAX_PB_SIZE) { + errno = EPROTO; + return handshake::STEP_ERROR; + } + return handshake::STEP_OK; +} + +static handshake::StepResult ParsePayload( + const std::string& payload, ParsedHello* remote) { + RdmaHello msg; + if (!msg.ParseFromArray(payload.data(), static_cast(payload.size()))) { + errno = EPROTO; + return handshake::STEP_ERROR; + } + if (!ValidRdmaHello(msg)) { + return handshake::STEP_FALLBACK; + } + TranslateHello(msg, remote); + return handshake::STEP_OK; +} + +static void FillDisabledHello(RdmaHello* msg) { + msg->set_block_size(0); + msg->set_sq_size(0); + msg->set_rq_size(0); + msg->set_lid(0); + msg->set_gid(std::string(sizeof(ibv_gid), '\0')); + msg->set_qp_num(0); +} + +} // namespace v3_wire + +const handshake::FrameSpec& +RdmaClientHandshakeAdapterV3::HelloFrameSpec() const { + return RdmaHelloFrameSpec(3); +} + +handshake::StepResult RdmaClientHandshakeAdapterV3::BuildLocalHello( + bool enabled, std::string* payload) { + CHECK(enabled); + PrepareClientEce(); + ParsedHello local{}; + FillLocalHello(&local); + RdmaHello msg; + v3_wire::FillLocalRdmaHello(local, &msg); + return v3_wire::SerializePayload(msg, payload); +} + +handshake::StepResult RdmaClientHandshakeAdapterV3::ParseRemoteHello( + const std::string& payload, ParsedHello* remote) { + return v3_wire::ParsePayload(payload, remote); +} + +const handshake::FrameSpec& +RdmaServerHandshakeAdapterV3::HelloFrameSpec() const { + return RdmaHelloFrameSpec(3); +} + +handshake::StepResult RdmaServerHandshakeAdapterV3::BuildLocalHello( + bool enabled, std::string* payload) { + RdmaHello msg; + if (enabled) { + ParsedHello local{}; + FillLocalHello(&local); + v3_wire::FillLocalRdmaHello(local, &msg); + } else { + v3_wire::FillDisabledHello(&msg); + } + return v3_wire::SerializePayload(msg, payload); +} + +handshake::StepResult RdmaServerHandshakeAdapterV3::ParseRemoteHello( + const std::string& payload, ParsedHello* remote) { + return v3_wire::ParsePayload(payload, remote); +} + +std::unique_ptr CreateClientHandshakeAdapter( + RdmaEndpoint* ep) { + if (FLAGS_rdma_client_handshake_version == 3) { + return std::unique_ptr( + new RdmaClientHandshakeAdapterV3(ep)); + } + return std::unique_ptr( + new RdmaClientHandshakeAdapterV2(ep)); +} + +std::vector > +CreateServerHandshakeAdapters(RdmaEndpoint* ep) { + std::vector > adapters; + adapters.emplace_back(new RdmaServerHandshakeAdapterV2(ep)); + adapters.emplace_back(new RdmaServerHandshakeAdapterV3(ep)); + return adapters; +} + +} // namespace rdma +} // namespace brpc + +#endif // BRPC_WITH_RDMA + +namespace brpc { +namespace handshake { + +class RdmaServerHandshakeAdapter : public StandardHandshakeAdapter { +public: + RdmaServerHandshakeAdapter() = default; + +protected: + StepResult RunServerStep( + butil::IOBuf* source, Socket* socket) override; + HandshakeSession* GetSession(Socket* socket) const override; + +private: + StepResult RunFallbackServerHandshake( + butil::IOBuf* source, Socket* socket); +#if BRPC_WITH_RDMA + StepResult RunRdmaServerHandshake( + butil::IOBuf* source, Socket* socket); +#endif + + DISALLOW_COPY_AND_ASSIGN(RdmaServerHandshakeAdapter); +}; + +static constexpr uint16_t V2_HELLO_VERSION_INVALID = + std::numeric_limits::max(); +static constexpr size_t V3_GID_LEN = 16; + +class RdmaFallbackProtocol : public FallbackHandshakeProtocol { +public: + explicit RdmaFallbackProtocol(int version) : _version(version) {} + + int ProtocolVersion() const override { return _version; } + const FrameSpec& HelloFrameSpec() const override { + return rdma::RdmaHelloFrameSpec(_version); + } + const FrameSpec& AckFrameSpec() const override { + return rdma::RdmaAckFrameSpec(); + } + StepResult BuildHello(bool enabled, std::string* payload) override { + if (enabled) { + errno = EPROTO; + return STEP_ERROR; + } + if (_version == 2) { + payload->assign( + rdma::HELLO_V2_MSG_LEN_MIN - rdma::HELLO_MAGIC_LEN - + sizeof(uint16_t), + '\0'); + butil::RawPacker(&(*payload)[0]) + .pack16(V2_HELLO_VERSION_INVALID); + return STEP_OK; + } + + rdma::RdmaHello reply; + reply.set_block_size(0); + reply.set_sq_size(0); + reply.set_rq_size(0); + reply.set_lid(0); + reply.set_gid(std::string(V3_GID_LEN, '\0')); + reply.set_qp_num(0); + if (!reply.SerializeToString(payload)) { + errno = EPROTO; + return STEP_ERROR; + } + return STEP_OK; + } + +private: + int _version; +}; + +HandshakeAdapter* GetRdmaServerHandshakeAdapter() { + static RdmaServerHandshakeAdapter adapter; + return &adapter; +} + +HandshakeSession* RdmaServerHandshakeAdapter::GetSession( + Socket* socket) const { + return AdapterTransport::Get(socket)->handshake_session(); +} + +StepResult RdmaServerHandshakeAdapter::RunFallbackServerHandshake( + butil::IOBuf* source, Socket* socket) { + IOBufHandshakeInput input(source); + RdmaFallbackProtocol v2(2); + RdmaFallbackProtocol v3(3); + std::vector protocols; + protocols.push_back(&v2); + protocols.push_back(&v3); + FallbackHandshakeTransport transport; + return GetSession(socket)->RunServer( + protocols, &input, &transport, false); +} + +#if BRPC_WITH_RDMA +class RdmaServerHandshakeTransport : public HandshakeTransport { +public: + RdmaServerHandshakeTransport( + RdmaTransport* transport, Socket* socket, butil::IOBuf* source) + : _transport(transport), _socket(socket), _source(source), + _protocol(NULL) {} + + void OnProtocolSelected(HandshakeProtocol* protocol) override { + _protocol = static_cast(protocol); + } + + StepResult PrepareResources() override { + if (_transport->PrepareUpgradeResources() == 0) { + return STEP_OK; + } + PLOG(WARNING) << "Fail to allocate rdma resources, fallback to tcp:" + << _socket->description(); + _transport->DeactivateUpgrade(); + return STEP_FALLBACK; + } + + StepResult NegotiateResources() override { + CHECK(_protocol != NULL); + if (_transport->NegotiateUpgradeResources( + _protocol->remote(), true) == 0) { + return STEP_OK; + } + PLOG(WARNING) << "Fail to negotiate rdma resources, fallback to tcp:" + << _socket->description(); + _transport->DeactivateUpgrade(); + return STEP_FALLBACK; + } + + void OnEstablished() override { _transport->ActivateUpgrade(); } + void OnFallback() override { _transport->DeactivateUpgrade(); } + void OnFailed() override { _transport->DeactivateUpgrade(); } + + StepResult ValidateEstablished() override { + return _source->empty() ? STEP_OK : STEP_ERROR; + } + +private: + RdmaTransport* _transport; + Socket* _socket; + butil::IOBuf* _source; + rdma::RdmaHandshakeAdapter* _protocol; +}; + +StepResult RdmaServerHandshakeAdapter::RunRdmaServerHandshake( + butil::IOBuf* source, Socket* socket) { + RdmaTransport* transport = RdmaTransport::Get(socket); + CHECK(transport->GetRdmaEp() != NULL); + + std::vector > protocols = + transport->CreateServerHandshakeAdapters(); + std::vector protocol_ptrs; + protocol_ptrs.reserve(protocols.size()); + for (size_t i = 0; i < protocols.size(); ++i) { + protocol_ptrs.push_back(protocols[i].get()); + } + IOBufHandshakeInput input(source); + RdmaServerHandshakeTransport participant(transport, socket, source); + return GetSession(socket)->RunServer( + protocol_ptrs, &input, &participant, false); +} +#endif + +StepResult RdmaServerHandshakeAdapter::RunServerStep( + butil::IOBuf* source, Socket* socket) { +#if BRPC_WITH_RDMA + if (AdapterTransport::Get(socket)->upgrade_capable( + SOCKET_MODE_RDMA)) { + return RunRdmaServerHandshake(source, socket); + } +#endif + return RunFallbackServerHandshake(source, socket); +} + +} // namespace handshake +} // namespace brpc diff --git a/src/brpc/handshake/rdma_handshake.h b/src/brpc/handshake/rdma_handshake.h new file mode 100644 index 0000000000..80057e687e --- /dev/null +++ b/src/brpc/handshake/rdma_handshake.h @@ -0,0 +1,160 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#ifndef BRPC_HANDSHAKE_RDMA_HANDSHAKE_H +#define BRPC_HANDSHAKE_RDMA_HANDSHAKE_H + +#include "brpc/handshake/handshake_adapter.h" + +namespace brpc { +namespace handshake { + +// Returns the RDMA adapter used by policy::ParseTransportHandshake. The +// concrete type is private to the implementation; callers only learn the +// common HandshakeAdapter interface. +HandshakeAdapter* GetRdmaServerHandshakeAdapter(); + +} // namespace handshake +} // namespace brpc + +#if BRPC_WITH_RDMA + +#include +#include +#include + +#include + +#include "butil/containers/optional.h" +#include "butil/macros.h" +#include "brpc/rdma/rdma_endpoint.h" +#include "brpc/handshake/rdma_handshake_constants.h" +#include "brpc/transport_handshake.h" + +namespace brpc { +namespace rdma { + +using ParsedHello = RdmaConnectionInfo; + +namespace v2_wire { + +struct HelloMessage { + void Serialize(void* data) const; + void Deserialize(const void* data); + + uint16_t msg_len; + uint16_t hello_ver; + uint16_t impl_ver; + uint32_t block_size; + uint16_t sq_size; + uint16_t rq_size; + uint16_t lid; + ibv_gid gid; + uint32_t qp_num; +}; + +} // namespace v2_wire + +// RDMA adapters implement only protocol fields. HandshakeSession owns frame +// I/O, length validation, ACK exchange and resource callback ordering. +class RdmaHandshakeAdapter : public handshake::HandshakeProtocol { +public: + RdmaHandshakeAdapter(RdmaEndpoint* ep, int version) + : _ep(ep), _version(version), _remote() {} + ~RdmaHandshakeAdapter() override = default; + + int ProtocolVersion() const override { return _version; } + const handshake::FrameSpec& AckFrameSpec() const override; + handshake::StepResult BuildHello( + bool enabled, std::string* payload) override; + handshake::StepResult ParseHello( + const std::string& payload) override; + const ParsedHello& remote() const { return _remote; } + + const handshake::FrameSpec& HelloFrameSpec() const override = 0; + virtual handshake::StepResult BuildLocalHello( + bool enabled, std::string* payload) = 0; + virtual handshake::StepResult ParseRemoteHello( + const std::string& payload, ParsedHello* remote) = 0; + +protected: + void FillLocalHello(ParsedHello* local) const; + void PrepareClientEce(); + + RdmaEndpoint* _ep; + int _version; + ParsedHello _remote; + +private: + DISALLOW_COPY_AND_ASSIGN(RdmaHandshakeAdapter); +}; + +class RdmaClientHandshakeAdapterV2 : public RdmaHandshakeAdapter { +public: + explicit RdmaClientHandshakeAdapterV2(RdmaEndpoint* ep) + : RdmaHandshakeAdapter(ep, 2) {} + const handshake::FrameSpec& HelloFrameSpec() const override; + handshake::StepResult BuildLocalHello( + bool enabled, std::string* payload) override; + handshake::StepResult ParseRemoteHello( + const std::string& payload, ParsedHello* remote) override; +}; + +class RdmaServerHandshakeAdapterV2 : public RdmaHandshakeAdapter { +public: + explicit RdmaServerHandshakeAdapterV2(RdmaEndpoint* ep) + : RdmaHandshakeAdapter(ep, 2) {} + const handshake::FrameSpec& HelloFrameSpec() const override; + handshake::StepResult BuildLocalHello( + bool enabled, std::string* payload) override; + handshake::StepResult ParseRemoteHello( + const std::string& payload, ParsedHello* remote) override; +}; + +class RdmaClientHandshakeAdapterV3 : public RdmaHandshakeAdapter { +public: + explicit RdmaClientHandshakeAdapterV3(RdmaEndpoint* ep) + : RdmaHandshakeAdapter(ep, 3) {} + const handshake::FrameSpec& HelloFrameSpec() const override; + handshake::StepResult BuildLocalHello( + bool enabled, std::string* payload) override; + handshake::StepResult ParseRemoteHello( + const std::string& payload, ParsedHello* remote) override; +}; + +class RdmaServerHandshakeAdapterV3 : public RdmaHandshakeAdapter { +public: + explicit RdmaServerHandshakeAdapterV3(RdmaEndpoint* ep) + : RdmaHandshakeAdapter(ep, 3) {} + const handshake::FrameSpec& HelloFrameSpec() const override; + handshake::StepResult BuildLocalHello( + bool enabled, std::string* payload) override; + handshake::StepResult ParseRemoteHello( + const std::string& payload, ParsedHello* remote) override; +}; + +std::unique_ptr CreateClientHandshakeAdapter( + RdmaEndpoint* ep); + +std::vector > +CreateServerHandshakeAdapters(RdmaEndpoint* ep); + +} // namespace rdma +} // namespace brpc + +#endif // BRPC_WITH_RDMA +#endif // BRPC_HANDSHAKE_RDMA_HANDSHAKE_H diff --git a/src/brpc/rdma/rdma_handshake_constants.h b/src/brpc/handshake/rdma_handshake_constants.h similarity index 69% rename from src/brpc/rdma/rdma_handshake_constants.h rename to src/brpc/handshake/rdma_handshake_constants.h index aa9811b98e..af60d2e9d3 100644 --- a/src/brpc/rdma/rdma_handshake_constants.h +++ b/src/brpc/handshake/rdma_handshake_constants.h @@ -15,8 +15,13 @@ // specific language governing permissions and limitations // under the License. -#ifndef BRPC_RDMA_RDMA_HANDSHAKE_CONSTANTS_H -#define BRPC_RDMA_RDMA_HANDSHAKE_CONSTANTS_H +#ifndef BRPC_HANDSHAKE_RDMA_HANDSHAKE_CONSTANTS_H +#define BRPC_HANDSHAKE_RDMA_HANDSHAKE_CONSTANTS_H + +#include +#include + +#include "brpc/handshake/handshake_frame.h" namespace brpc { namespace rdma { @@ -50,7 +55,27 @@ constexpr size_t HELLO_V3_MAX_PB_SIZE = 8192; constexpr size_t HELLO_ACK_LEN = 4; constexpr uint32_t HELLO_ACK_RDMA_OK = 0x1; +inline const handshake::FrameSpec& RdmaHelloFrameSpec(int version) { + static const handshake::FrameSpec v2( + HELLO_MAGIC, HELLO_MAGIC_LEN, + HELLO_V2_MSG_LEN_MIN, HELLO_V2_MSG_LEN_MAX, + handshake::FrameSpec::U16_TOTAL_LENGTH); + static const handshake::FrameSpec v3( + HELLO_MAGIC_V3, HELLO_MAGIC_LEN, + HELLO_MAGIC_LEN + HELLO_V3_PB_SIZE_LEN + 1, + HELLO_MAGIC_LEN + HELLO_V3_PB_SIZE_LEN + HELLO_V3_MAX_PB_SIZE, + handshake::FrameSpec::U32_BODY_LENGTH); + return version == 2 ? v2 : v3; +} + +inline const handshake::FrameSpec& RdmaAckFrameSpec() { + static const handshake::FrameSpec spec( + NULL, 0, HELLO_ACK_LEN, HELLO_ACK_LEN, + handshake::FrameSpec::FIXED); + return spec; +} + } // namespace rdma } // namespace brpc -#endif // BRPC_RDMA_RDMA_HANDSHAKE_CONSTANTS_H +#endif // BRPC_HANDSHAKE_RDMA_HANDSHAKE_CONSTANTS_H diff --git a/src/brpc/handshake/ubshm_handshake.cpp b/src/brpc/handshake/ubshm_handshake.cpp new file mode 100644 index 0000000000..a409915d0f --- /dev/null +++ b/src/brpc/handshake/ubshm_handshake.cpp @@ -0,0 +1,454 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include "brpc/handshake/ubshm_handshake.h" + +#include +#include + +#include "butil/raw_pack.h" +#include "butil/sys_byteorder.h" +#include "brpc/adapter_transport.h" +#include "brpc/socket.h" + +#if BRPC_WITH_UBRING + +#include +#include + +#include "butil/logging.h" +#include "brpc/reloadable_flags.h" +#include "brpc/ubshm/common/common.h" +#include "brpc/ubshm/ub_endpoint.h" +#include "brpc/ubshm/ub_helper.h" +#include "brpc/ubshm/ubr_trx.h" +#include "brpc/ubshm_transport.h" + +#endif + +namespace brpc { +namespace handshake { +namespace ubshm_wire { + +static const char* const MAGIC = "UB"; +static const size_t MAGIC_LEN = 2; +static const size_t HELLO_LEN = 64; +static const size_t ACK_LEN = 4; +#if BRPC_WITH_UBRING +static const uint16_t HELLO_VERSION = 3; +static const uint16_t IMPL_VERSION = 1; +#endif // BRPC_WITH_UBRING +static const FrameSpec& HelloFrameSpec() { + static const FrameSpec spec( + MAGIC, MAGIC_LEN, HELLO_LEN, HELLO_LEN, FrameSpec::FIXED); + return spec; +} + +static const FrameSpec& AckFrameSpec() { + static const FrameSpec spec( + NULL, 0, ACK_LEN, ACK_LEN, FrameSpec::FIXED); + return spec; +} + +} // namespace ubshm_wire +} // namespace handshake +} // namespace brpc + +#if BRPC_WITH_UBRING + +namespace brpc { +namespace ubring { + +DEFINE_int32(data_queue_size, 4, "data queue size for UB"); +DEFINE_bool(ub_trace_verbose, false, "Print log message verbosely"); +BRPC_VALIDATE_GFLAG(ub_trace_verbose, brpc::PassValidate); + +void HelloMessage::Serialize(void* data) const { + char* current_pos = static_cast(data); + const uint16_t net_msg_len = butil::HostToNet16(msg_len); + memcpy(current_pos, &net_msg_len, sizeof(net_msg_len)); + current_pos += sizeof(net_msg_len); + const uint16_t net_hello_ver = butil::HostToNet16(hello_ver); + memcpy(current_pos, &net_hello_ver, sizeof(net_hello_ver)); + current_pos += sizeof(net_hello_ver); + const uint16_t net_impl_ver = butil::HostToNet16(impl_ver); + memcpy(current_pos, &net_impl_ver, sizeof(net_impl_ver)); + current_pos += sizeof(net_impl_ver); + const uint64_t net_len = butil::HostToNet64(len); + memcpy(current_pos, &net_len, sizeof(net_len)); + current_pos += sizeof(net_len); + memcpy(current_pos, shm_name, SHM_MAX_NAME_BUFF_LEN); +} + +void HelloFormatExtension::Serialize(void* data) const { + char* current = static_cast(data); + const uint16_t length = butil::HostToNet16(extension_len); + const uint16_t format = butil::HostToNet16(format_id); + memcpy(current, &length, sizeof(length)); + memcpy(current + sizeof(length), &format, sizeof(format)); +} + +void HelloFormatExtension::Deserialize(const void* data) { + const char* current = static_cast(data); + uint16_t length; + uint16_t format; + memcpy(&length, current, sizeof(length)); + memcpy(&format, current + sizeof(length), sizeof(format)); + extension_len = butil::NetToHost16(length); + format_id = butil::NetToHost16(format); +} + +void HelloMessage::Deserialize(const void* data) { + const char* current_pos = static_cast(data); + uint16_t net_msg_len; + memcpy(&net_msg_len, current_pos, sizeof(net_msg_len)); + msg_len = butil::NetToHost16(net_msg_len); + current_pos += sizeof(net_msg_len); + uint16_t net_hello_ver; + memcpy(&net_hello_ver, current_pos, sizeof(net_hello_ver)); + hello_ver = butil::NetToHost16(net_hello_ver); + current_pos += sizeof(net_hello_ver); + uint16_t net_impl_ver; + memcpy(&net_impl_ver, current_pos, sizeof(net_impl_ver)); + impl_ver = butil::NetToHost16(net_impl_ver); + current_pos += sizeof(net_impl_ver); + uint64_t net_len; + memcpy(&net_len, current_pos, sizeof(net_len)); + len = butil::NetToHost64(net_len); + current_pos += sizeof(net_len); + memcpy(shm_name, current_pos, SHM_MAX_NAME_BUFF_LEN); +} + +std::string HelloMessage::toString() const { + constexpr size_t MAX_LEN = + 16 + 6 + 16 + 6 + 16 + 6 + 20 + 6 + SHM_MAX_NAME_BUFF_LEN + 32; + std::array buf; + const int n = snprintf( + buf.data(), buf.size(), + "msg_len=%u, hello_ver=%u, impl_ver=%u, len=%lu, shm_name=%.*s", + msg_len, hello_ver, impl_ver, + static_cast(len), + static_cast(SHM_MAX_NAME_BUFF_LEN), shm_name); + return std::string(buf.data(), static_cast(n)); +} + +const handshake::FrameSpec& UBShmHandshakeAdapter::HelloFrameSpec() const { + return handshake::ubshm_wire::HelloFrameSpec(); +} + +const handshake::FrameSpec& UBShmHandshakeAdapter::AckFrameSpec() const { + return handshake::ubshm_wire::AckFrameSpec(); +} + +const handshake::FrameSpec& +UBShmHandshakeAdapter::ExtensionFrameSpec() const { + static const handshake::FrameSpec spec( + NULL, 0, HelloFormatExtension::WIRE_SIZE, + HelloFormatExtension::WIRE_SIZE, handshake::FrameSpec::FIXED); + return spec; +} + +void UBShmHandshakeAdapter::ConfigureClientHello( + uint64_t len, const char* shm_name) { + _local_len = len; + _local_name = shm_name != NULL ? shm_name : ""; + _server_reply = false; +} + +handshake::StepResult UBShmHandshakeAdapter::BuildHello( + bool enabled, std::string* payload) { + const char* name = NULL; + if (enabled) { + name = _server_reply ? _remote.shm_name : _local_name.c_str(); + } + return BuildHello(enabled, enabled ? _local_len : 0, name, payload); +} + +handshake::StepResult UBShmHandshakeAdapter::ParseHello( + const std::string& payload) { + return ParseHello(payload, &_remote); +} + +handshake::StepResult UBShmHandshakeAdapter::BuildExtension( + bool enabled, std::string* payload) { + const HelloFormatExtension extension = { + HelloFormatExtension::WIRE_SIZE, + static_cast(enabled ? UBR_DATA_FORMAT_LEGACY_64 + : UBR_DATA_FORMAT_NONE)}; + payload->resize(HelloFormatExtension::WIRE_SIZE); + extension.Serialize(&(*payload)[0]); + return handshake::STEP_OK; +} + +handshake::StepResult UBShmHandshakeAdapter::ParseExtension( + const std::string& payload) { + if (payload.size() != HelloFormatExtension::WIRE_SIZE) { + errno = EPROTO; + return handshake::STEP_ERROR; + } + HelloFormatExtension extension{}; + extension.Deserialize(payload.data()); + return extension.extension_len == HelloFormatExtension::WIRE_SIZE && + extension.format_id == UBR_DATA_FORMAT_LEGACY_64 + ? handshake::STEP_OK : handshake::STEP_FALLBACK; +} + +handshake::StepResult UBShmHandshakeAdapter::BuildHello( + bool enabled, uint64_t len, const char* shm_name, + std::string* payload) const { + HelloMessage message{}; + message.msg_len = static_cast( + handshake::ubshm_wire::HELLO_LEN); + if (enabled) { + message.hello_ver = handshake::ubshm_wire::HELLO_VERSION; + message.impl_ver = handshake::ubshm_wire::IMPL_VERSION; + message.len = 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 - + handshake::ubshm_wire::MAGIC_LEN, + '\0'); + message.Serialize(&(*payload)[0]); + return handshake::STEP_OK; +} + +handshake::StepResult UBShmHandshakeAdapter::ParseHello( + const std::string& payload, HelloMessage* message) const { + if (payload.size() != handshake::ubshm_wire::HELLO_LEN - + handshake::ubshm_wire::MAGIC_LEN) { + errno = EPROTO; + return handshake::STEP_ERROR; + } + message->Deserialize(payload.data()); + if (message->msg_len != handshake::ubshm_wire::HELLO_LEN) { + errno = EPROTO; + return handshake::STEP_ERROR; + } + 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( + const HelloMessage& message) const { + return message.hello_ver == handshake::ubshm_wire::HELLO_VERSION && + message.impl_ver == handshake::ubshm_wire::IMPL_VERSION; +} + +} // namespace ubring +} // namespace brpc + +#endif // BRPC_WITH_UBRING + +namespace brpc { +namespace handshake { + +class UBShmServerHandshakeAdapter : public StandardHandshakeAdapter { +public: + UBShmServerHandshakeAdapter() = default; + +protected: + StepResult RunServerStep( + butil::IOBuf* source, Socket* socket) override; + HandshakeSession* GetSession(Socket* socket) const override; + +private: + StepResult RunFallbackServerHandshake( + butil::IOBuf* source, Socket* socket); +#if BRPC_WITH_UBRING + StepResult RunUBShmServerHandshake( + butil::IOBuf* source, Socket* socket); +#endif + + DISALLOW_COPY_AND_ASSIGN(UBShmServerHandshakeAdapter); +}; + + +class UBShmFallbackProtocol : public FallbackHandshakeProtocol { +public: + int ProtocolVersion() const override { return 2; } + const FrameSpec& HelloFrameSpec() const override { + return ubshm_wire::HelloFrameSpec(); + } + const FrameSpec& AckFrameSpec() const override { + return ubshm_wire::AckFrameSpec(); + } + StepResult BuildHello(bool enabled, std::string* payload) override { + if (enabled) { + errno = EPROTO; + return STEP_ERROR; + } + payload->assign( + ubshm_wire::HELLO_LEN - ubshm_wire::MAGIC_LEN, '\0'); + butil::RawPacker(&(*payload)[0]) + .pack16(static_cast(ubshm_wire::HELLO_LEN)); + return STEP_OK; + } +}; + +HandshakeAdapter* GetUBShmServerHandshakeAdapter() { + static UBShmServerHandshakeAdapter adapter; + return &adapter; +} + +HandshakeSession* UBShmServerHandshakeAdapter::GetSession( + Socket* socket) const { + return AdapterTransport::Get(socket)->handshake_session(); +} + +StepResult UBShmServerHandshakeAdapter::RunFallbackServerHandshake( + butil::IOBuf* source, Socket* socket) { + IOBufHandshakeInput input(source); + UBShmFallbackProtocol protocol; + std::vector protocols(1, &protocol); + FallbackHandshakeTransport transport; + return GetSession(socket)->RunServer( + protocols, &input, &transport, false); +} + +#if BRPC_WITH_UBRING +class UBShmServerHandshakeTransport : public HandshakeTransport { +public: + UBShmServerHandshakeTransport( + UBShmTransport* transport, ubring::UBShmHandshakeAdapter* protocol, + Socket* socket, butil::IOBuf* source) + : _transport(transport), _protocol(protocol), _socket(socket), + _source(source) {} + + void OnProtocolSelected(HandshakeProtocol*) override { + LOG_IF(INFO, ubring::FLAGS_ub_trace_verbose) + << "server receive handshake message : " + << _protocol->remote().toString(); + } + + StepResult PrepareResources() override { + if (!ubring::IsUBAvailable()) { + _transport->DeactivateUpgrade(); + return STEP_FALLBACK; + } + const ubring::HelloMessage& remote = _protocol->remote(); + 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())}; + 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; + ubring::SHM local_trx_shm = { + NULL, local_shm_len, 0, {0}, + static_cast(_socket->fd())}; + char client_name[SHM_MAX_NAME_BUFF_LEN + 1]; + memcpy(client_name, remote.shm_name, SHM_MAX_NAME_BUFF_LEN); + client_name[SHM_MAX_NAME_BUFF_LEN] = '\0'; + char* client_ip_port = strrchr(client_name, '_'); + if (client_ip_port != NULL) { + *client_ip_port = '\0'; + } + const int result = snprintf( + local_trx_shm.name, SHM_MAX_NAME_BUFF_LEN, "%s_%s", + client_name, SERVER_SHM_NAME_SUFFIX); + if (UNLIKELY(result < 0)) { + _transport->DeactivateUpgrade(); + return STEP_FALLBACK; + } + if (_transport->PrepareServerUpgradeResources( + &remote_trx_shm, &local_trx_shm) < 0) { + LOG(WARNING) + << "Fail to allocate ub resources, fallback to tcp:" + << _socket->description(); + _transport->DeactivateUpgrade(); + return STEP_FALLBACK; + } + return STEP_OK; + } + + StepResult NegotiateResources() override { return STEP_OK; } + void OnEstablished() override { _transport->ActivateUpgrade(); } + void OnFallback() override { _transport->DeactivateUpgrade(); } + void OnFailed() override { _transport->DeactivateUpgrade(); } + + StepResult ValidateEstablished() override { + return _source->empty() ? STEP_OK : STEP_ERROR; + } + +private: + UBShmTransport* _transport; + ubring::UBShmHandshakeAdapter* _protocol; + Socket* _socket; + butil::IOBuf* _source; +}; + +StepResult UBShmServerHandshakeAdapter::RunUBShmServerHandshake( + butil::IOBuf* source, Socket* socket) { + UBShmTransport* transport = UBShmTransport::Get(socket); + CHECK(transport->GetUBShmEp() != NULL); + + ubring::UBShmHandshakeAdapter wire; + const uint64_t local_shm_len = + static_cast(ubring::FLAGS_data_queue_size) * MB_TO_BYTE; + wire.ConfigureServerReply(local_shm_len); + IOBufHandshakeInput input(source); + std::vector protocols(1, &wire); + UBShmServerHandshakeTransport participant( + transport, &wire, socket, source); + const StepResult result = GetSession(socket)->RunServer( + protocols, &input, &participant, false); + if (result == STEP_OK) { + transport->GetUBShmEp()->SetNegotiatedDataFormat( + ubring::UBR_DATA_FORMAT_LEGACY_64); + transport->FinishUpgrade(); + LOG_IF(INFO, ubring::FLAGS_ub_trace_verbose) + << "Server handshake ends (use ubring) on " + << socket->description(); + } else if (result == STEP_FALLBACK) { + LOG_IF(INFO, ubring::FLAGS_ub_trace_verbose) + << "Server handshake ends (use tcp) on " + << socket->description(); + } + return result; +} +#endif + +StepResult UBShmServerHandshakeAdapter::RunServerStep( + butil::IOBuf* source, Socket* socket) { +#if BRPC_WITH_UBRING + if (AdapterTransport::Get(socket)->upgrade_capable( + SOCKET_MODE_UBRING)) { + return RunUBShmServerHandshake(source, socket); + } +#endif + return RunFallbackServerHandshake(source, socket); +} + +} // namespace handshake +} // namespace brpc diff --git a/src/brpc/handshake/ubshm_handshake.h b/src/brpc/handshake/ubshm_handshake.h new file mode 100644 index 0000000000..a41af31243 --- /dev/null +++ b/src/brpc/handshake/ubshm_handshake.h @@ -0,0 +1,123 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#ifndef BRPC_HANDSHAKE_UBSHM_HANDSHAKE_H +#define BRPC_HANDSHAKE_UBSHM_HANDSHAKE_H + +#include "brpc/handshake/handshake_adapter.h" + +namespace brpc { +namespace handshake { + +// Returns the adapter used by the common transport-handshake policy parser. +// The concrete server executor is private to the implementation. +HandshakeAdapter* GetUBShmServerHandshakeAdapter(); + +} // namespace handshake +} // namespace brpc + +#if BRPC_WITH_UBRING + +#include +#include + +#include + +#include "butil/macros.h" +#include "brpc/transport_handshake.h" +#include "brpc/ubshm/shm/shm_def.h" + +namespace brpc { +namespace ubring { + +DECLARE_int32(data_queue_size); +DECLARE_bool(ub_trace_verbose); + +// UBSHM v3 wire payload. HandshakeSession owns framing and ACK exchange; +// this type and UBShmHandshakeAdapter only handle protocol fields. +struct HelloMessage { + void Serialize(void* data) const; + void Deserialize(const void* data); + std::string toString() const; + + uint16_t msg_len; + uint16_t hello_ver; + uint16_t impl_ver; + uint64_t len; + char shm_name[SHM_MAX_NAME_BUFF_LEN]; +}; + +enum UbrDataFormat { + UBR_DATA_FORMAT_NONE = 0, + UBR_DATA_FORMAT_LEGACY_64 = 1, +}; + +struct HelloFormatExtension { + static const uint16_t WIRE_SIZE = 4; + uint16_t extension_len; + uint16_t format_id; + + void Serialize(void* data) const; + void Deserialize(const void* data); +}; + +class UBShmHandshakeAdapter : public handshake::HandshakeProtocol { +public: + UBShmHandshakeAdapter() + : _local_len(0), _server_reply(false), _remote() {} + + int ProtocolVersion() const override { return 3; } + const handshake::FrameSpec& HelloFrameSpec() const override; + const handshake::FrameSpec& AckFrameSpec() const override; + handshake::StepResult BuildHello( + bool enabled, std::string* payload) override; + handshake::StepResult ParseHello( + const std::string& payload) override; + bool HasExtension() const override { return true; } + const handshake::FrameSpec& ExtensionFrameSpec() const override; + handshake::StepResult BuildExtension( + bool enabled, std::string* payload) override; + handshake::StepResult ParseExtension( + const std::string& payload) override; + + void ConfigureClientHello(uint64_t len, const char* shm_name); + void ConfigureServerReply(uint64_t len) { + _local_len = len; + _server_reply = true; + } + const HelloMessage& remote() const { return _remote; } + + handshake::StepResult BuildHello( + bool enabled, uint64_t len, const char* shm_name, + std::string* payload) const; + handshake::StepResult ParseHello( + const std::string& payload, HelloMessage* message) const; + +private: + bool NegotiationValid(const HelloMessage& message) const; + uint64_t _local_len; + std::string _local_name; + bool _server_reply; + HelloMessage _remote; + DISALLOW_COPY_AND_ASSIGN(UBShmHandshakeAdapter); +}; + +} // namespace ubring +} // namespace brpc + +#endif // BRPC_WITH_UBRING +#endif // BRPC_HANDSHAKE_UBSHM_HANDSHAKE_H diff --git a/src/brpc/input_messenger.h b/src/brpc/input_messenger.h index c0041d6f8a..2464034885 100644 --- a/src/brpc/input_messenger.h +++ b/src/brpc/input_messenger.h @@ -40,6 +40,7 @@ class UBShmEndpoint; } class TcpTransport; class RdmaTransport; +class AdapterTransport; struct InputMessageHandler { // The callback to cut a message from `source'. // Returned message will be passed to process_request or process_response @@ -102,6 +103,7 @@ class InputMessageClosure { class InputMessenger : public SocketUser { friend class TcpTransport; friend class RdmaTransport; +friend class AdapterTransport; friend class rdma::RdmaEndpoint; friend class urma::UrmaEndpoint; friend class ubring::UBShmEndpoint; diff --git a/src/brpc/policy/rdma_handshake_protocol.cpp b/src/brpc/policy/rdma_handshake_protocol.cpp index 580abdda5d..fb1e183f1e 100644 --- a/src/brpc/policy/rdma_handshake_protocol.cpp +++ b/src/brpc/policy/rdma_handshake_protocol.cpp @@ -17,24 +17,16 @@ #include "brpc/policy/rdma_handshake_protocol.h" -#include "butil/logging.h" -#include "brpc/destroyable.h" -#include "brpc/rdma/rdma_handshake_server.h" - namespace brpc { namespace policy { ParseResult ParseRdmaHandshake(butil::IOBuf* source, Socket* socket, - bool /*read_eof*/, const void* /*arg*/) { - return rdma::ExecuteServerHandshake(source, socket); + bool read_eof, const void* arg) { + return ParseTransportHandshake(source, socket, read_eof, arg); } void ProcessRdmaHandshake(InputMessageBase* msg) { - // ParseRdmaHandshake replies inline and only ever returns - // NOT_ENOUGH_DATA / TRY_OTHERS / hard errors, never a real message, so this - // must never run. Keep a placeholder (required for server registration). - DestroyingPtr destroying_msg(msg); - CHECK(false) << "ProcessRdmaHandshake should never be called"; + ProcessTransportHandshake(msg); } } // namespace policy diff --git a/src/brpc/policy/rdma_handshake_protocol.h b/src/brpc/policy/rdma_handshake_protocol.h index e569a4c333..51416e2eee 100644 --- a/src/brpc/policy/rdma_handshake_protocol.h +++ b/src/brpc/policy/rdma_handshake_protocol.h @@ -18,36 +18,15 @@ #ifndef BRPC_POLICY_RDMA_HANDSHAKE_PROTOCOL_H #define BRPC_POLICY_RDMA_HANDSHAKE_PROTOCOL_H -// NOTE: This file is intentionally INDEPENDENT of BRPC_WITH_RDMA. A server may -// run in TCP mode either because it was built without RDMA, or because RDMA was -// not enabled at runtime. In both cases an RDMA client that connects to it will -// send an RDMA handshake magic ("RDMA" for v2, "RDM3" for v3) first. Without -// special handling the server treats those bytes as an unknown protocol and -// closes the connection, so the client (blocked reading the server hello) only -// sees EOF and cannot fall back to TCP. -// -// To let the client fall back on the SAME connection, the server recognizes the -// RDMA handshake as a "first-class" protocol (magic in the first 4 bytes, -// PROTOCOL_RDMA_HANDSHAKE is ordered before PROTOCOL_HTTP), replies a hello with -// an incompatible version so the client rejects it and downgrades to TCP, then -// drains the client's subsequent ACK and lets normal RPC parsing continue. - -#include "butil/iobuf.h" -#include "brpc/input_message_base.h" -#include "brpc/parse_result.h" -#include "brpc/socket.h" +// Compatibility facade. New code should include +// transport_handshake_protocol.h and use ParseTransportHandshake. +#include "brpc/policy/transport_handshake_protocol.h" namespace brpc { namespace policy { -// Parse binary format of rdma handshake. ParseResult ParseRdmaHandshake(butil::IOBuf* source, Socket* socket, bool read_eof, const void* arg); - -// Actions to a rdma handshake request, which is left unimplemented. -// All requests are processed in the parsing process. This function -// must be declared since server only enables rdma handshake as a -// server-side protocol when this function is declared. void ProcessRdmaHandshake(InputMessageBase* msg); } // namespace policy diff --git a/src/brpc/policy/transport_handshake_protocol.cpp b/src/brpc/policy/transport_handshake_protocol.cpp new file mode 100644 index 0000000000..5d343e0da7 --- /dev/null +++ b/src/brpc/policy/transport_handshake_protocol.cpp @@ -0,0 +1,37 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include "brpc/policy/transport_handshake_protocol.h" + +#include "butil/logging.h" +#include "brpc/adapter_transport.h" + +namespace brpc { +namespace policy { + +ParseResult ParseTransportHandshake(butil::IOBuf* source, Socket* socket, + bool /*read_eof*/, const void* /*arg*/) { + return AdapterTransport::Get(socket)->ProcessUpgradeReadable(source); +} + +void ProcessTransportHandshake(InputMessageBase* msg) { + DestroyingPtr destroying_msg(msg); + CHECK(false) << "ProcessTransportHandshake should never be called"; +} + +} // namespace policy +} // namespace brpc diff --git a/src/brpc/policy/transport_handshake_protocol.h b/src/brpc/policy/transport_handshake_protocol.h new file mode 100644 index 0000000000..f4a1fe9acc --- /dev/null +++ b/src/brpc/policy/transport_handshake_protocol.h @@ -0,0 +1,44 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#ifndef BRPC_POLICY_TRANSPORT_HANDSHAKE_PROTOCOL_H +#define BRPC_POLICY_TRANSPORT_HANDSHAKE_PROTOCOL_H + +// This policy is intentionally independent of BRPC_WITH_RDMA and +// BRPC_WITH_UBRING. A plain TCP server must recognize an upgrade hello and +// return a disabled hello so the client can continue with TCP on the same +// connection. + +#include "butil/iobuf.h" +#include "brpc/input_message_base.h" +#include "brpc/parse_result.h" +#include "brpc/socket.h" + +namespace brpc { +namespace policy { + +ParseResult ParseTransportHandshake(butil::IOBuf* source, Socket* socket, + bool read_eof, const void* arg); + +// Upgrade handshakes are completed inline by the parser. This placeholder is +// required for server-side protocol registration and must never be invoked. +void ProcessTransportHandshake(InputMessageBase* msg); + +} // namespace policy +} // namespace brpc + +#endif // BRPC_POLICY_TRANSPORT_HANDSHAKE_PROTOCOL_H diff --git a/src/brpc/rdma/rdma_endpoint.cpp b/src/brpc/rdma/rdma_endpoint.cpp index a602122e17..2db9179897 100644 --- a/src/brpc/rdma/rdma_endpoint.cpp +++ b/src/brpc/rdma/rdma_endpoint.cpp @@ -17,22 +17,20 @@ #if BRPC_WITH_RDMA -#include -#include "butil/fd_utility.h" -#include "butil/logging.h" // CHECK, LOG -#include "butil/sys_byteorder.h" // HostToNet,NetToHost -#include "bthread/bthread.h" +#include "brpc/rdma/rdma_endpoint.h" #include "brpc/errno.pb.h" #include "brpc/event_dispatcher.h" #include "brpc/input_messenger.h" -#include "brpc/socket.h" -#include "brpc/reloadable_flags.h" #include "brpc/rdma/block_pool.h" #include "brpc/rdma/rdma_helper.h" -#include "brpc/rdma/rdma_endpoint.h" #include "brpc/rdma_transport.h" -#include "brpc/rdma/rdma_handshake.h" -#include "brpc/rdma/rdma_handshake_constants.h" +#include "brpc/reloadable_flags.h" +#include "brpc/socket.h" +#include "bthread/bthread.h" +#include "butil/fd_utility.h" +#include "butil/logging.h" // CHECK, LOG +#include "butil/sys_byteorder.h" // HostToNet,NetToHost +#include DECLARE_int32(task_group_ntags); @@ -111,8 +109,6 @@ RdmaResource::~RdmaResource() { RdmaEndpoint::RdmaEndpoint(Socket* s) : _socket(s) - , _state(UNINIT) - , _handshake_version(0) , _resource(nullptr) , _send_cq_events(0) , _recv_cq_events(0) @@ -146,20 +142,16 @@ RdmaEndpoint::RdmaEndpoint(Socket* s) if (_rq_size > MAX_QP_SIZE) { _rq_size = MAX_QP_SIZE; } - _read_butex = bthread::butex_create_checked >(); _input_processor.Init(s, InputMessengerProcessor::STREAM_RDMA_QP); } RdmaEndpoint::~RdmaEndpoint() { Reset(); - bthread::butex_destroy(_read_butex); } void RdmaEndpoint::Reset() { DeallocateResources(); - _state.store(UNINIT, butil::memory_order_relaxed); - _handshake_version = 0; _outgoing_ece.reset(); _resource = nullptr; _send_cq_events = 0; @@ -167,8 +159,8 @@ void RdmaEndpoint::Reset() { _cq_sid = INVALID_SOCKET_ID; _sbuf.clear(); _rbuf.clear(); - _rbuf_data.clear(); _input_processor.Reset(); + _rbuf_data.clear(); _remote_recv_block_size = 0; _accumulated_ack = 0; _unsolicited = 0; @@ -185,583 +177,41 @@ void RdmaEndpoint::Reset() { _new_rq_wrs.store(0, butil::memory_order_relaxed); } -void RdmaConnect::StartConnect(const Socket* socket, - void (*done)(int err, void* data), - void* data) { - auto* rdma_transport = static_cast(socket->_transport.get()); - CHECK(rdma_transport->_rdma_ep != nullptr); - SocketUniquePtr s; - if (Socket::Address(socket->id(), &s) != 0) { - return; - } - if (!IsRdmaAvailable()) { - rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF; - rdma_transport->_rdma_ep->_state.store( - RdmaEndpoint::FALLBACK_TCP, butil::memory_order_release); - done(0, data); - return; - } - _done = done; - _data = data; - bthread_t tid; - bthread_attr_t attr = BTHREAD_ATTR_NORMAL; - bthread_attr_set_name(&attr, "RdmaProcessHandshakeAtClient"); - if (bthread_start_background(&tid, &attr, - RdmaEndpoint::ProcessHandshakeAtClient, - rdma_transport->_rdma_ep) < 0) { - LOG(FATAL) << "Fail to start handshake bthread"; - Run(); - } else { - s.release(); - } -} - -void RdmaConnect::StopConnect(Socket* socket) { } - -void RdmaConnect::Run() { - _done(errno, _data); -} - -void RdmaEndpoint::OnNewDataFromTcp(Socket* s) { - if (s->CreatedByConnect()) { - OnNewDataFromTcpAtClient(s); - } else { - OnNewDataFromTcpAtServer(s); - } -} - -void RdmaEndpoint::OnNewDataFromTcpAtClient(Socket* s) { - auto* rdma_transport = static_cast(s->_transport.get()); - RdmaEndpoint* ep = rdma_transport->GetRdmaEp(); - CHECK(ep != nullptr); - - int progress = Socket::PROGRESS_INIT; - while (true) { - // Pair with release stores of FALLBACK_TCP so RDMA_OFF is visible - // before normal TCP message processing starts. - const State state = ep->_state.load(butil::memory_order_acquire); - if (state == UNINIT) { - // The connection may be closed or reset before the client starts - // handshake. This will be handled by client handshake. Ignore here. - } else if (state < ESTABLISHED) { // during handshake - ep->_read_butex->fetch_add(1, butil::memory_order_release); - bthread::butex_wake(ep->_read_butex); - } else if (state == FALLBACK_TCP){ // handshake finishes - InputMessenger::OnNewMessages(s); - return; - } else if (state == ESTABLISHED) { - if (!ep->HandleTcpEventAfterEstablished()) { - return; - } - } - if (!s->MoreReadEvents(&progress)) { - break; - } - } -} - -void RdmaEndpoint::OnNewDataFromTcpAtServer(Socket* s) { - auto* rdma_transport = static_cast(s->_transport.get()); - RdmaEndpoint* ep = rdma_transport->GetRdmaEp(); - CHECK(ep != nullptr); - - int progress = Socket::PROGRESS_INIT; - while (true) { - if (s->Failed()) { - return; - } - - // Pair with the release stores of ESTABLISHED / FALLBACK_TCP. - if (ep->_state.load(butil::memory_order_acquire) != ESTABLISHED) { - InputMessenger::OnNewMessages(s); - // That call may have just finished the handshake and turned RDMA - // on. Start consuming CQ events here rather than inside the parse - // callback: by now OnNewMessages is done with the Socket's - // `parsing_context` / `preferred_index`, so the QP stream can take - // them over without ever overlapping with the fd stream. This is - // the ordering StartCqEvents() asks for. - if (!s->Failed() && - ep->_state.load(butil::memory_order_acquire) == ESTABLISHED && - ep->StartCqEvents() < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to start cq events on " << *s; - ep->_state.store(FAILED, butil::memory_order_relaxed); - s->SetFailed(saved_errno, "Fail to start cq events on %s: %s", - s->description().c_str(), berror(saved_errno)); - } - return; - } - // RDMA carries the RPCs now, so the fd is watched for EOF only and must - // not be parsed: `preferred_index' / `parsing_context' live on the Socket - // and the QP stream is driving them (https://github.com/apache/brpc/issues/3479). - if (!ep->HandleTcpEventAfterEstablished()) { - return; - } - if (!s->MoreReadEvents(&progress)) { - break; - } - } -} - -bool RdmaEndpoint::HandleTcpEventAfterEstablished() { - uint8_t tmp; - ssize_t nr = read(_socket->fd(), &tmp, 1); - if (nr == 0) { - _socket->SetEOF(); - return false; - } - if (nr > 0) { - LOG(WARNING) << "Read unexpected data from " << *_socket; - _socket->SetFailed(EPROTO, "Read unexpected data from %s", - _socket->description().c_str()); - return false; - } - - if (errno != EAGAIN) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to read from " << *_socket; - _socket->SetFailed(saved_errno, "Fail to read from %s: %s", - _socket->description().c_str(), - berror(saved_errno)); - // The socket is dead now, so do not come back for another read of it. - return false; - } - return true; -} - -static const int WAIT_TIMEOUT_MS = 50; - -// Drive an EAGAIN-aware read loop to completion (exactly `len` bytes). -// `read_once(offset, remaining)` performs ONE underlying read attempt: -// returns > 0 : number of bytes consumed (added to running total); -// returns = 0 : end-of-stream (the loop fails with EEOF); -// returns < 0 : errno set; EAGAIN is handled here via butex_wait, -// any other errno bubbles up. -// `offset` is bytes already received in THIS call (initially 0); the -// callable uses it to choose the next write target (e.g. `(char*)buf -// + offset`). Callables that don't need offset (e.g. IOPortal append) -// can ignore it. -// -// Centralizes the EAGAIN/butex/EOF loop so the two ReadFromFd -// overloads below stay one-liners; any future read source (memory- -// mapped, scatter-vector, etc.) can plug in by passing its own -// `read_once`. -template -static int ReadFromFdLoop(butil::atomic* read_butex, - size_t len, ReadOnce&& read_once) { - size_t received = 0; - while (received < len) { - const int expected_val = read_butex->load(butil::memory_order_acquire); - const timespec duetime = butil::milliseconds_from_now(WAIT_TIMEOUT_MS); - ssize_t nr = read_once(received, len - received); - if (nr < 0) { - if (errno == EAGAIN) { - if (bthread::butex_wait(read_butex, expected_val, &duetime) < 0) { - if (errno != EWOULDBLOCK && errno != ETIMEDOUT) { - return -1; - } - } - } else { - return -1; - } - } else if (nr == 0) { // Got EOF - errno = EEOF; - return -1; - } else { - received += nr; - } - } - return 0; -} - -int RdmaEndpoint::ReadFromFd(void* data, size_t len) { - CHECK(data != nullptr); - const int fd = _socket->fd(); - return ReadFromFdLoop(_read_butex, len, - [data, fd](size_t offset, size_t remaining) { - return read(fd, (uint8_t*)data + offset, remaining); - }); -} - -int RdmaEndpoint::ReadFromFd(butil::IOPortal* data, size_t len) { - CHECK(data != nullptr); - const int fd = _socket->fd(); - return ReadFromFdLoop(_read_butex, len, - [data, fd](size_t /*offset*/, size_t remaining) { - return data->append_from_file_descriptor(fd, remaining); - }); -} - -// Drive an EAGAIN-aware write loop to completion (exactly `len` bytes). -// -// `write_once(offset, remaining)` performs ONE underlying write attempt: -// - returns >= 0 : number of bytes consumed (added to running total); -// - returns < 0 : errno set; EAGAIN triggers `wait_writable(duetime)`, -// any other errno bubbles up. -// `offset` is bytes already written in THIS call (initially 0); the -// callable uses it to choose the next read source (e.g. `(char*)buf -// + offset`). Callables that drain a self-tracking sink (e.g. -// IOBuf::cut_into_file_descriptor) can ignore both args. -// -// `wait_writable(duetime)` is invoked on EAGAIN to park until the fd -// becomes writable again. It returns 0 on wake-up (or ETIMEDOUT), -// non-zero on hard failure. -template -static int WriteToFdLoop(size_t len, WriteOnce&& write_once, WaitWritable&& wait_writable) { - size_t written = 0; - while (written < len) { - const timespec duetime = butil::milliseconds_from_now(WAIT_TIMEOUT_MS); - ssize_t nw = write_once(written, len - written); - if (nw >= 0) { - written += nw; - continue; - } - - if (errno != EAGAIN) { - return -1; - } - if (!wait_writable(&duetime)) { - return -1; - } - } - return 0; -} - -int RdmaEndpoint::WriteToFd(void* data, size_t len) { - CHECK(data != nullptr); - Socket* s = _socket; - const int fd = s->fd(); - return WriteToFdLoop(len, - [data, fd](size_t offset, size_t remaining) { - return write(fd, (uint8_t*)data + offset, remaining); - }, - [s, fd](const timespec* duetime) { - return s->WaitEpollOut(fd, true, duetime) == 0 || errno == ETIMEDOUT; - }); -} - -int RdmaEndpoint::WriteToFd(butil::IOBuf* data) { - CHECK(data != nullptr); - Socket* s = _socket; - const int fd = s->fd(); - return WriteToFdLoop(data->size(), - [data, fd](size_t /*offset*/, size_t /*remaining*/) { - return data->cut_into_file_descriptor(fd); - }, - [s, fd](const timespec* duetime) { - return s->WaitEpollOut(fd, true, duetime) == 0 || errno == ETIMEDOUT; - }); -} - -void RdmaEndpoint::ApplyRemoteHello(const ParsedHello& remote) { +void RdmaEndpoint::ApplyRemoteInfo(const RdmaConnectionInfo &remote) { _remote_recv_block_size = remote.block_size; _local_window_capacity = std::min(_sq_size, remote.rq_size) - RESERVED_WR_NUM; - _remote_window_capacity = std::min(_rq_size, remote.sq_size) - RESERVED_WR_NUM; + _remote_window_capacity = + std::min(_rq_size, remote.sq_size) - RESERVED_WR_NUM; _sq_imm_window_size = RESERVED_WR_NUM; - _remote_rq_window_size.store(_local_window_capacity, butil::memory_order_relaxed); + _remote_rq_window_size.store(_local_window_capacity, + butil::memory_order_relaxed); _sq_window_size.store(_local_window_capacity, butil::memory_order_relaxed); } -// Client-side handshake entry: the state machine. -// -// C_ALLOC_QPCQ -// | -// v -// C_HELLO_SEND (hs->SendLocalHello) -// | -// v -// C_HELLO_WAIT (hs->ReceiveAndParseRemoteHello) -// | -// v -// [negotiation: ApplyRemoteHello + C_BRINGUP_QP] -// | -// v -// C_ACK_SEND -// | -// v -// ESTABLISHED / FALLBACK_TCP -void* RdmaEndpoint::ProcessHandshakeAtClient(void* arg) { - auto ep = static_cast(arg); - SocketUniquePtr s(ep->_socket); - RdmaConnect::RunGuard rg((RdmaConnect*)s->_app_connect.get()); - auto rdma_transport = static_cast(s->_transport.get()); - - LOG_IF(INFO, FLAGS_rdma_trace_verbose) - << "Start handshake on " << s->description(); - - std::unique_ptr handshake = CreateClientHandshake(ep); - CHECK(handshake != nullptr); - ep->_handshake_version = handshake->ProtocolVersion(); - - // First initialize CQ and QP resources. - ep->_state.store(C_ALLOC_QPCQ, butil::memory_order_relaxed); - if (ep->AllocateResources() < 0) { - PLOG(WARNING) << "Fail to allocate rdma resources, fallback to tcp:" - << s->description(); - errno = 0; - rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF; - ep->_state.store(FALLBACK_TCP, butil::memory_order_release); - return nullptr; - } - - // Send hello message to server - ep->_state.store(C_HELLO_SEND, butil::memory_order_relaxed); - if (handshake->SendLocalHello() < 0) { - int saved_errno = errno; - PLOG(WARNING) << "Fail to send hello message to server:" - << s->description(); - s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state.store(FAILED, butil::memory_order_relaxed); - return nullptr; - } - - // Receive and parse remote hello. - ep->_state.store(C_HELLO_WAIT, butil::memory_order_relaxed); - ParsedHello remote{}; - const RemoteHelloResult r = handshake->ReceiveAndParseRemoteHello(&remote); - if (r == RemoteHelloResult::ERROR) { - int saved_errno = errno; - PLOG(WARNING) << "Fail to receive hello from server:" - << s->description(); - s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state.store(FAILED, butil::memory_order_relaxed); - return nullptr; +void RdmaEndpoint::GetLocalConnectionInfo(RdmaConnectionInfo *local) const { + CHECK(local != NULL); + local->block_size = g_rdma_recv_block_size; + local->sq_size = _sq_size; + local->rq_size = _rq_size; + local->lid = GetRdmaLid(); + local->gid = GetRdmaGid(); + local->qp_num = BAIDU_LIKELY(_resource) ? _resource->qp->qp_num : 0; + local->ece.reset(); + if (_outgoing_ece.has_value()) { + local->ece = _outgoing_ece; } - - if (r != RemoteHelloResult::NEGOTIATED) { - LOG(WARNING) << "Fail to negotiate with server, fallback to tcp:" - << s->description(); - rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF; - } else { - ep->ApplyRemoteHello(remote); - ep->_state.store(C_BRINGUP_QP, butil::memory_order_relaxed); - if (ep->BringUpQp(remote, /*is_server=*/false) < 0) { - LOG(WARNING) << "Fail to bringup QP, fallback to tcp:" - << s->description(); - rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF; - } else { - rdma_transport->_rdma_state = RdmaTransport::RDMA_ON; - } - } - - // Send ACK message to server - ep->_state.store(C_ACK_SEND, butil::memory_order_relaxed); - bool rdma_on = rdma_transport->_rdma_state == RdmaTransport::RDMA_ON; - uint32_t flags = rdma_on ? HELLO_ACK_RDMA_OK : 0; - uint32_t flags_be = butil::HostToNet32(flags); - if (ep->WriteToFd(&flags_be, HELLO_ACK_LEN) < 0) { - int saved_errno = errno; - PLOG(WARNING) << "Fail to send Ack Message to server:" - << s->description(); - s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state.store(FAILED, butil::memory_order_relaxed); - return nullptr; - } - - if (rdma_transport->_rdma_state == RdmaTransport::RDMA_ON) { - ep->_state.store(ESTABLISHED, butil::memory_order_release); - // The handshake is over, so the QP stream may start parsing now. - if (ep->StartCqEvents() < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to start cq events on " << s->description(); - s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state.store(FAILED, butil::memory_order_relaxed); - return nullptr; - } - LOG_IF(INFO, FLAGS_rdma_trace_verbose) - << "Client handshake ends (use rdma v" << ep->_handshake_version - << ") on " << s->description(); - } else { - ep->_state.store(FALLBACK_TCP, butil::memory_order_release); - LOG_IF(INFO, FLAGS_rdma_trace_verbose) - << "Client handshake ends (use tcp) on " << s->description(); - } - - errno = 0; - - return nullptr; } -// Server-side handshake entry: the state machine. -// -// S_HELLO_WAIT (read magic + dispatch + hs->ReceiveAndParseRemoteHello) -// | -// v -// [negotiation: ApplyRemoteHello + S_ALLOC_QPCQ + S_BRINGUP_QP] -// | -// v -// S_HELLO_SEND (hs->SendLocalHello) -// | -// v -// S_ACK_WAIT -// | -// v -// ESTABLISHED / FALLBACK_TCP -ParseResult RdmaEndpoint::ExecuteServerHandshake(butil::IOBuf* source, Socket* s) { - RdmaTransport* rdma_transport = static_cast(s->_transport.get()); - RdmaEndpoint* ep = rdma_transport->_rdma_ep; - CHECK(ep != nullptr); - - const State state = ep->_state.load(butil::memory_order_acquire); - if (state >= ESTABLISHED) { - // The handshake is over (ESTABLISHED / FALLBACK_TCP / FAILED). Data - // arriving now belongs to a real protocol, yet CutInputMessage() still - // reaches us. - if (state == ESTABLISHED && - s->parsing_stream_type() == InputMessengerProcessor::STREAM_TCP_FD) { - // RDMA is on, so the fd is not an RPC channel any more and whatever - // shows up on it is a protocol error. Reached even though - // OnNewDataFromTcpAtServer() stops handing the fd to OnNewMessages() - // once RDMA is on, because the handshake completes inside OnNewMessages(): - // that round keeps reading the fd until it goes quiet. - if (source->empty()) { - // Nothing to reject yet. Asking for more data keeps the pin, so - // the rest of this round comes back here rather than reaching a - // real protocol, and lets OnNewMessages() report EOF as usual. - return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); - } - LOG(WARNING) << "Unexpected " << source->size() << " bytes on the tcp " - "fd of an RDMA connection, drop connection: " - << s->description(); - ep->_state.store(FAILED, butil::memory_order_relaxed); - return MakeParseError(PARSE_ERROR_ABSOLUTELY_WRONG); - } - // Anything else, for the real protocol to parse: the stream carried by - // the QP, or an fd that stayed a normal RPC stream because the handshake - // fell back or failed. - return MakeParseError(PARSE_ERROR_TRY_OTHERS); +int RdmaEndpoint::QueryLocalEce(ibv_ece *ece) const { + if (ece == NULL || IbvQueryEce == NULL || _resource == NULL || + _resource->qp == NULL) { + return 1; } + return IbvQueryEce(_resource->qp, ece) == 0 ? 0 : -1; +} - if (s->parsing_context() == nullptr) { - // Phase 1: read the client hello, negotiate, reply server hello. - if (source->size() < HELLO_MAGIC_LEN) { - return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); - } - uint8_t magic[HELLO_MAGIC_LEN]; - CHECK_EQ(source->copy_to(magic, HELLO_MAGIC_LEN), HELLO_MAGIC_LEN); - - // Pick the version-specific server handshake from the peeked magic (the - // magic is NOT consumed; ReceiveAndParseRemoteHello() reads it again - // from `source`). - std::unique_ptr hs = CreateServerHandshakeByMagic(ep, source, magic); - if (hs == nullptr) { - return MakeParseError(PARSE_ERROR_TRY_OTHERS); - } - ep->_handshake_version = hs->ProtocolVersion(); - ep->_state.store(S_HELLO_WAIT, butil::memory_order_relaxed); - - ParsedHello remote{}; - const RemoteHelloResult r = hs->ReceiveAndParseRemoteHello(&remote); - if (r == RemoteHelloResult::NEED_MORE) { - return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); - } - if (r == RemoteHelloResult::ERROR) { - ep->_state.store(FAILED, butil::memory_order_relaxed); - return MakeParseError(PARSE_ERROR_ABSOLUTELY_WRONG); - } - - // Negotiate + allocate resources. - bool negotiated = r == RemoteHelloResult::NEGOTIATED; - if (negotiated) { - ep->ApplyRemoteHello(remote); - ep->_state.store(S_ALLOC_QPCQ, butil::memory_order_relaxed); - if (ep->AllocateResources() < 0) { - PLOG(WARNING) << "Fail to allocate rdma resources, fallback to tcp:" - << s->description(); - negotiated = false; - } else { - ep->_state.store(S_BRINGUP_QP, butil::memory_order_relaxed); - if (ep->BringUpQp(remote, /*is_server=*/true) < 0) { - LOG(WARNING) << "Fail to bringup QP, fallback to tcp:" - << s->description(); - negotiated = false; - } - } - } - if (!negotiated) { - rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF; - } - - // Reply the server hello. - // Emits a real hello when _rdma_state != RDMA_OFF; - // an un-negotiable one otherwise. - ep->_state.store(S_HELLO_SEND, butil::memory_order_relaxed); - if (hs->SendLocalHello() < 0) { - PLOG(WARNING) << "Fail to send server hello to " << s->description(); - ep->_state.store(FAILED, butil::memory_order_relaxed); - return MakeParseError(PARSE_ERROR_ABSOLUTELY_WRONG); - } - - // Enter the wait-ACK phase. Whether negotiation succeeded is already - // recorded in rdma_transport->_rdma_state (RDMA_OFF iff negotiation - // failed), so the context itself needs no extra flag. - s->reset_parsing_context(ServerHandshakeContext::Create()); - ep->_state.store(S_ACK_WAIT, butil::memory_order_relaxed); - return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); - } - - // Phase 2: drain the 4B ACK and finalize. - if (source->size() < HELLO_ACK_LEN) { - return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); - } - - uint32_t flags_be = 0; - CHECK_EQ(source->cutn(&flags_be, HELLO_ACK_LEN), HELLO_ACK_LEN); - uint32_t flags = butil::NetToHost32(flags_be); - bool client_ack_ok = (flags & HELLO_ACK_RDMA_OK) != 0; - if (!client_ack_ok) { - LOG_IF(INFO, FLAGS_rdma_trace_verbose) - << "Server handshake ends (use tcp) on " << s->description(); - rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF; - ep->_state.store(FALLBACK_TCP, butil::memory_order_release); - s->reset_parsing_context(nullptr); - return MakeParseError(PARSE_ERROR_TRY_OTHERS); - } - - if (rdma_transport->_rdma_state == RdmaTransport::RDMA_OFF) { - LOG(WARNING) << "Client wants RDMA in ACK but server fell back: " - << s->description(); - ep->_state.store(FAILED, butil::memory_order_relaxed); - s->reset_parsing_context(nullptr); - return MakeParseError(PARSE_ERROR_ABSOLUTELY_WRONG); - } - - if (!source->empty()) { - // RDMA is on, so the TCP fd is no longer an RPC channel. Anything - // trailing the ACK on it can only be a protocol error. This catches what - // arrived in the same read as the ACK. - LOG(WARNING) << "Unexpected " << source->size() << " bytes after the " - "handshake ACK of an RDMA connection, drop connection: " - << s->description(); - ep->_state.store(FAILED, butil::memory_order_relaxed); - s->reset_parsing_context(nullptr); - return MakeParseError(PARSE_ERROR_ABSOLUTELY_WRONG); - } - - LOG_IF(INFO, FLAGS_rdma_trace_verbose) - << "Server handshake ends (use rdma v" << ep->_handshake_version - << ") on " << s->description(); - rdma_transport->_rdma_state = RdmaTransport::RDMA_ON; - ep->_state.store(ESTABLISHED, butil::memory_order_release); - s->reset_parsing_context(nullptr); - - // Two things are deliberately not done here. - // - // The CQ events are not started: this runs inside CutInputMessage, which - // keeps touching `preferred_index` / `parsing_context` after we return, - // and PollCq would race it for those. OnNewDataFromTcpAtServer() starts - // them once that is over. - // - // TRY_OTHERS is not returned: it would hand `preferred_index` to the real - // protocol, and the remaining reads of this OnNewMessages() round would - // parse the fd as an RPC stream although RDMA has just taken over. Asking - // for more data keeps this handler pinned, so those reads come back to the - // guard at the top of this function and are rejected there. - return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); +void RdmaEndpoint::SetOutgoingEce(const ibv_ece &ece) { + _outgoing_ece = ece; } bool RdmaEndpoint::IsWritable() const { @@ -778,6 +228,7 @@ bool RdmaEndpoint::IsWritable() const { // The reason is that we need to use some protected member function of IOBuf. class RdmaIOBuf : public butil::IOBuf { friend class RdmaEndpoint; + private: // Cut the current IOBuf to ibv_sge list and `to' for at most first max_sge // blocks or first max_len bytes. @@ -800,7 +251,8 @@ friend class RdmaEndpoint; lkey = (uint32_t)meta; } } - if (BAIDU_UNLIKELY(lkey == 0)) { // only happens when meta is not specified + if (BAIDU_UNLIKELY(lkey == + 0)) { // only happens when meta is not specified lkey = GetLKey((char*)start - r.offset); } if (lkey == 0) { @@ -844,8 +296,7 @@ ssize_t RdmaEndpoint::CutFromIOBufList(butil::IOBuf** from, size_t ndata) { size_t current = 0; uint32_t remote_rq_window_size = _remote_rq_window_size.load(butil::memory_order_relaxed); - uint32_t sq_window_size = - _sq_window_size.load(butil::memory_order_relaxed); + uint32_t sq_window_size = _sq_window_size.load(butil::memory_order_relaxed); ibv_send_wr wr; int max_sge = GetRdmaMaxSge(); ibv_sge sglist[max_sge]; @@ -899,7 +350,8 @@ ssize_t RdmaEndpoint::CutFromIOBufList(butil::IOBuf** from, size_t ndata) { wr.imm_data = butil::HostToNet32(imm); // Avoid too much recv completion event to reduce the cpu overhead bool solicited = false; - if (remote_rq_window_size == 1 || sq_window_size == 1 || current + 1 >= ndata) { + if (remote_rq_window_size == 1 || sq_window_size == 1 || + current + 1 >= ndata) { // Only last message in the write queue or last message in the // current window will be flagged as solicited. solicited = true; @@ -943,7 +395,8 @@ ssize_t RdmaEndpoint::CutFromIOBufList(butil::IOBuf** from, size_t ndata) { // So we just consider this error as an unrecoverable error. std::ostringstream oss; DebugInfo(oss, ", "); - LOG(WARNING) << "Fail to ibv_post_send: " << berror(err) << " " << oss.str(); + LOG(WARNING) << "Fail to ibv_post_send: " << berror(err) << " " + << oss.str(); errno = err; return -1; } @@ -960,14 +413,16 @@ ssize_t RdmaEndpoint::CutFromIOBufList(butil::IOBuf** from, size_t ndata) { // counters. remote_rq_window_size = _remote_rq_window_size.fetch_sub(1, butil::memory_order_relaxed) - 1; - sq_window_size = _sq_window_size.fetch_sub(1, butil::memory_order_relaxed) - 1; + sq_window_size = + _sq_window_size.fetch_sub(1, butil::memory_order_relaxed) - 1; } return total_len; } int RdmaEndpoint::SendAck(int num) { - if (_new_rq_wrs.fetch_add(num, butil::memory_order_relaxed) > _remote_window_capacity / 2 && + if (_new_rq_wrs.fetch_add(num, butil::memory_order_relaxed) > + _remote_window_capacity / 2 && _sq_imm_window_size > 0) { return SendImm(_new_rq_wrs.exchange(0, butil::memory_order_relaxed)); } @@ -993,7 +448,8 @@ int RdmaEndpoint::SendImm(uint32_t imm) { DebugInfo(oss, ", "); // We use other way to guarantee the Send Queue is not full. // So we just consider this error as an unrecoverable error. - LOG(WARNING) << "Fail to ibv_post_send: " << berror(err) << " " << oss.str(); + LOG(WARNING) << "Fail to ibv_post_send: " << berror(err) << " " + << oss.str(); return -1; } @@ -1039,13 +495,11 @@ ssize_t RdmaEndpoint::HandleCompletion(ibv_wc& wc) { if (wc.byte_len < (uint32_t)FLAGS_rdma_zerocopy_min_size) { zerocopy = false; } - CHECK_NE(_state.load(butil::memory_order_relaxed), FALLBACK_TCP); - butil::IOPortal& read_buf = _input_processor.read_buf(); if (zerocopy) { - _rbuf[_rq_received].cutn(&read_buf, wc.byte_len); + _rbuf[_rq_received].cutn(&_input_processor.read_buf(), wc.byte_len); } else { // Copy data when the receive data is really small - read_buf.append(_rbuf_data[_rq_received], wc.byte_len); + _input_processor.read_buf().append(_rbuf_data[_rq_received], wc.byte_len); } } if (0 != (wc.wc_flags & IBV_WC_WITH_IMM) && wc.imm_data > 0) { @@ -1105,7 +559,8 @@ int RdmaEndpoint::PostRecv(uint32_t num, bool zerocopy) { if (zerocopy) { _rbuf[_rq_received].clear(); butil::IOBufAsZeroCopyOutputStream os(&_rbuf[_rq_received], - g_rdma_recv_block_size + IOBUF_BLOCK_HEADER_LEN); + g_rdma_recv_block_size + + IOBUF_BLOCK_HEADER_LEN); int size = 0; if (!os.Next(&_rbuf_data[_rq_received], &size)) { // Memory is not enough for preparing a block @@ -1128,7 +583,8 @@ int RdmaEndpoint::PostRecv(uint32_t num, bool zerocopy) { return 0; } -static ibv_qp* AllocateQp(ibv_cq* send_cq, ibv_cq* recv_cq, uint32_t sq_size, uint32_t rq_size) { +static ibv_qp *AllocateQp(ibv_cq *send_cq, ibv_cq *recv_cq, uint32_t sq_size, + uint32_t rq_size) { ibv_qp_init_attr attr; memset(&attr, 0, sizeof(attr)); attr.send_cq = send_cq; @@ -1159,34 +615,36 @@ static RdmaResource* AllocateQpCq(uint16_t sq_size, uint16_t rq_size) { return nullptr; } - resource->send_cq = IbvCreateCq(GetRdmaContext(), FLAGS_rdma_prepared_qp_size, - nullptr, resource->comp_channel, GetRdmaCompVector()); + resource->send_cq = + IbvCreateCq(GetRdmaContext(), FLAGS_rdma_prepared_qp_size, nullptr, + resource->comp_channel, GetRdmaCompVector()); if (nullptr == resource->send_cq) { PLOG(WARNING) << "Fail to create send CQ"; return nullptr; } - resource->recv_cq = IbvCreateCq(GetRdmaContext(), FLAGS_rdma_prepared_qp_size, - nullptr, resource->comp_channel, GetRdmaCompVector()); + resource->recv_cq = + IbvCreateCq(GetRdmaContext(), FLAGS_rdma_prepared_qp_size, nullptr, + resource->comp_channel, GetRdmaCompVector()); if (nullptr == resource->recv_cq) { PLOG(WARNING) << "Fail to create recv CQ"; return nullptr; } - resource->qp = AllocateQp(resource->send_cq, resource->recv_cq, sq_size, rq_size); + resource->qp = + AllocateQp(resource->send_cq, resource->recv_cq, sq_size, rq_size); if (nullptr == resource->qp) { PLOG(WARNING) << "Fail to create QP"; return nullptr; } } else { - resource->polling_cq = - IbvCreateCq(GetRdmaContext(), 2 * FLAGS_rdma_prepared_qp_size, nullptr, nullptr, 0); + resource->polling_cq = IbvCreateCq( + GetRdmaContext(), 2 * FLAGS_rdma_prepared_qp_size, nullptr, nullptr, 0); if (nullptr == resource->polling_cq) { PLOG(WARNING) << "Fail to create polling CQ"; return nullptr; } - resource->qp = AllocateQp(resource->polling_cq, - resource->polling_cq, + resource->qp = AllocateQp(resource->polling_cq, resource->polling_cq, sq_size, rq_size); if (nullptr == resource->qp) { PLOG(WARNING) << "Fail to create QP"; @@ -1231,12 +689,12 @@ int RdmaEndpoint::DoAllocateResources() { g_rdma_resource_list = g_rdma_resource_list->next; } } - if (_resource == nullptr) { + if (!_resource) { _resource = AllocateQpCq(_sq_size, _rq_size); } else { _resource->next = nullptr; } - if (_resource == nullptr) { + if (!_resource) { return -1; } @@ -1266,23 +724,23 @@ int RdmaEndpoint::DoAllocateResources() { } int RdmaEndpoint::StartCqEvents() { - if (InputMessengerProcessor::STREAM_NONE != _socket->parsing_stream_type()) { - LOG(WARNING) << "StartCqEvents() called while " << *_socket << " is parsing"; + if (InputMessengerProcessor::STREAM_NONE != + _socket->parsing_stream_type()) { + LOG(WARNING) << "StartCqEvents() called while " << *_socket + << " is parsing"; errno = ERDMA; return -1; } if (_cq_sid != INVALID_SOCKET_ID) { - // Already started. return 0; } if (_resource == nullptr) { if (BAIDU_UNLIKELY(g_skip_rdma_init)) { - // For UT: AllocateResources() succeeds without allocating anything. return 0; } - - LOG(WARNING) << "No RDMA resource to start CQ events on, " << *_socket; + LOG(WARNING) << "No RDMA resource to start CQ events on, " + << *_socket; errno = ERDMA; return -1; } @@ -1298,15 +756,13 @@ int RdmaEndpoint::StartCqEvents() { PLOG(WARNING) << "Fail to create socket for cq"; return -1; } - if (FLAGS_rdma_use_polling) { PollerAddCqSid(); } - return 0; } -int RdmaEndpoint::BringUpQp(const ParsedHello& remote, bool is_server) { +int RdmaEndpoint::BringUpQp(const RdmaConnectionInfo &remote, bool is_server) { if (BAIDU_UNLIKELY(g_skip_rdma_init)) { // For UT return 0; @@ -1318,11 +774,9 @@ int RdmaEndpoint::BringUpQp(const ParsedHello& remote, bool is_server) { attr.pkey_index = 0; // TODO: support more pkey use in future attr.port_num = GetRdmaPortNum(); attr.qp_access_flags = IBV_ACCESS_REMOTE_WRITE; - int err = IbvModifyQp(_resource->qp, &attr, (ibv_qp_attr_mask)( - IBV_QP_STATE | - IBV_QP_PKEY_INDEX | - IBV_QP_PORT | - IBV_QP_ACCESS_FLAGS)); + int err = IbvModifyQp(_resource->qp, &attr, + (ibv_qp_attr_mask)(IBV_QP_STATE | IBV_QP_PKEY_INDEX | + IBV_QP_PORT | IBV_QP_ACCESS_FLAGS)); if (err != 0) { LOG(WARNING) << "Fail to modify QP from RESET to INIT: " << berror(err); return -1; @@ -1334,7 +788,7 @@ int RdmaEndpoint::BringUpQp(const ParsedHello& remote, bool is_server) { // End-to-end model: // Server: `remote->ece' is the client's queried ECE; set it here, // then after RTS we query the reduced/negotiated ECE and - // send it back in the server hello. + // return it in the server negotiation response. // Client: `remote->ece' is the server's reduced ECE; // just set it here. bool use_ece = true; @@ -1370,14 +824,11 @@ int RdmaEndpoint::BringUpQp(const ParsedHello& remote, bool is_server) { attr.rq_psn = 0; attr.max_dest_rd_atomic = 0; attr.min_rnr_timer = 0; // We do not allow rnr error - err = IbvModifyQp(_resource->qp, &attr, (ibv_qp_attr_mask)( - IBV_QP_STATE | - IBV_QP_PATH_MTU | - IBV_QP_MIN_RNR_TIMER | - IBV_QP_AV | - IBV_QP_MAX_DEST_RD_ATOMIC | - IBV_QP_DEST_QPN | - IBV_QP_RQ_PSN)); + err = IbvModifyQp(_resource->qp, &attr, + (ibv_qp_attr_mask)(IBV_QP_STATE | IBV_QP_PATH_MTU | + IBV_QP_MIN_RNR_TIMER | IBV_QP_AV | + IBV_QP_MAX_DEST_RD_ATOMIC | + IBV_QP_DEST_QPN | IBV_QP_RQ_PSN)); if (err != 0) { LOG(WARNING) << "Fail to modify QP from INIT to RTR: " << berror(err); return -1; @@ -1389,29 +840,29 @@ int RdmaEndpoint::BringUpQp(const ParsedHello& remote, bool is_server) { attr.rnr_retry = 0; // We do not allow rnr error attr.sq_psn = 0; attr.max_rd_atomic = 0; - err = IbvModifyQp(_resource->qp, &attr, (ibv_qp_attr_mask)( - IBV_QP_STATE | - IBV_QP_RNR_RETRY | - IBV_QP_RETRY_CNT | - IBV_QP_TIMEOUT | - IBV_QP_SQ_PSN | - IBV_QP_MAX_QP_RD_ATOMIC)); + err = + IbvModifyQp(_resource->qp, &attr, + (ibv_qp_attr_mask)(IBV_QP_STATE | IBV_QP_RNR_RETRY | + IBV_QP_RETRY_CNT | IBV_QP_TIMEOUT | + IBV_QP_SQ_PSN | IBV_QP_MAX_QP_RD_ATOMIC)); if (err != 0) { LOG(WARNING) << "Fail to modify QP from RTR to RTS: " << berror(err); return -1; } - // On the server side, now that the QP reached RTS, query the reduced/negotiated - // ECE (the subset of enhancements supported by both peers) so it can be returned - // to the client in the server hello. - if (is_server && use_ece && IbvQueryEce != nullptr && remote.ece.has_value()) { + // On the server side, now that the QP reached RTS, query the + // reduced/negotiated ECE (the subset of enhancements supported by both peers) + // so it can be returned to the client in the server hello. + if (is_server && use_ece && IbvQueryEce != nullptr && + remote.ece.has_value()) { ibv_ece ece; int qerr = IbvQueryEce(_resource->qp, &ece); if (qerr == 0) { _outgoing_ece = ece; } else { LOG(WARNING) << "Fail to IbvQueryEce(negotiated), " - "continue without ECE: " << berror(qerr); + "continue without ECE: " + << berror(qerr); } } @@ -1443,7 +894,7 @@ static int DrainCq(ibv_cq* cq) { } void RdmaEndpoint::DeallocateResources() { - if (_resource == nullptr) { + if (!_resource) { return; } if (FLAGS_rdma_use_polling) { @@ -1470,7 +921,7 @@ void RdmaEndpoint::DeallocateResources() { bool remove_consumer = true; _reclaim: if (!move_to_rdma_resource_list) { - if (_resource->qp != nullptr) { + if (nullptr != _resource->qp) { int err = IbvDestroyQp(_resource->qp); LOG_IF(WARNING, 0 != err) << "Fail to destroy QP: " << berror(err); _resource->qp = nullptr; @@ -1485,11 +936,14 @@ void RdmaEndpoint::DeallocateResources() { // Destroy send_comp_channel will destroy this fd, // so that we should remove it from epoll fd first int fd = _resource->comp_channel->fd; - GetGlobalEventDispatcher(fd, _socket->_io_event.bthread_tag()).RemoveConsumer(fd); + GetGlobalEventDispatcher( + fd, _socket->_io_event.bthread_tag()) + .RemoveConsumer(fd); remove_consumer = false; } int err = IbvDestroyCompChannel(_resource->comp_channel); - LOG_IF(WARNING, 0 != err) << "Fail to destroy CQ channel: " << berror(err); + LOG_IF(WARNING, 0 != err) + << "Fail to destroy CQ channel: " << berror(err); } _resource->polling_cq = nullptr; @@ -1570,8 +1024,8 @@ int RdmaEndpoint::GetAndAckEvents(SocketUniquePtr& s) { } else { // Unexpected CQ event that does not belong to // this endpoint's send/recv CQs. - LOG(WARNING) << "Unexpected CQ event from cq=" << cq - << " of " << s->description(); + LOG(WARNING) << "Unexpected CQ event from cq=" << cq << " of " + << s->description(); // Acknowledge this single event immediately // to avoid leaking unacknowledged events. IbvAckCqEvents(cq, 1); @@ -1590,16 +1044,15 @@ int RdmaEndpoint::GetAndAckEvents(SocketUniquePtr& s) { int RdmaEndpoint::ReqNotifyCq(bool send_cq, bool fatal_on_error) { const int err = ibv_req_notify_cq( - send_cq ? _resource->send_cq : _resource->recv_cq, - send_cq ? 0 : 1); + send_cq ? _resource->send_cq : _resource->recv_cq, send_cq ? 0 : 1); if (0 != err) { errno = err; PLOG(WARNING) << "Fail to arm " << (send_cq ? "send" : "recv") << " CQ comp channel from " << _socket->description(); if (fatal_on_error) { _socket->SetFailed(err, "Fail to arm %s CQ channel from %s: %s", - send_cq ? "send" : "recv", _socket->description().c_str(), - berror(err)); + send_cq ? "send" : "recv", + _socket->description().c_str(), berror(err)); } // The logging and SetFailed() above may clobber errno. errno = err; @@ -1624,9 +1077,8 @@ void RdmaEndpoint::PollCq(Socket* m) { if (m->id() != ep->_cq_sid) { return; } - auto* rdma_transport = static_cast(s->_transport.get()); + RdmaTransport *rdma_transport = RdmaTransport::Get(s.get()); CHECK(ep == rdma_transport->_rdma_ep); - CHECK_GE(ep->_state.load(butil::memory_order_acquire), ESTABLISHED); bool send = false; ibv_cq* cq = ep->_resource->recv_cq; @@ -1744,50 +1196,30 @@ void RdmaEndpoint::PollCq(Socket* m) { // Otherwise it may call too many bthread_flush to affect performance. const int64_t received_us = butil::cpuwide_time_us(); const int64_t base_realtime = butil::gettimeofday_us() - received_us; - if (ep->_input_processor.ProcessNewMessage(bytes, false, received_us, - base_realtime, last_msg) < 0) { + if (ep->_input_processor.ProcessNewMessage( + bytes, false, received_us, base_realtime, last_msg) < 0) { return; } } } -std::string RdmaEndpoint::GetStateStr() const { - switch (_state.load(butil::memory_order_relaxed)) { - case UNINIT: return "UNINIT"; - case C_ALLOC_QPCQ: return "C_ALLOC_QPCQ"; - case C_HELLO_SEND: return "C_HELLO_SEND"; - case C_HELLO_WAIT: return "C_HELLO_WAIT"; - case C_BRINGUP_QP: return "C_BRINGUP_QP"; - case C_ACK_SEND: return "C_ACK_SEND"; - case S_HELLO_WAIT: return "S_HELLO_WAIT"; - case S_ALLOC_QPCQ: return "S_ALLOC_QPCQ"; - case S_BRINGUP_QP: return "S_BRINGUP_QP"; - case S_HELLO_SEND: return "S_HELLO_SEND"; - case S_ACK_WAIT: return "S_ACK_WAIT"; - case ESTABLISHED: return "ESTABLISHED"; - case FALLBACK_TCP: return "FALLBACK_TCP"; - case FAILED: return "FAILED"; - default: return "UNKNOWN"; - } -} - -void RdmaEndpoint::DebugInfo(std::ostream& os, butil::StringPiece connector) const { - os << "rdma_state=ON" - << connector << "handshake_state=" << GetStateStr() - << connector << "handshake_version=" << static_cast(_handshake_version) - << connector << "rdma_sq_imm_window_size=" << _sq_imm_window_size - << connector << "rdma_remote_rq_window_size=" << _remote_rq_window_size.load(butil::memory_order_relaxed) - << connector << "rdma_sq_window_size=" << _sq_window_size.load(butil::memory_order_relaxed) - << connector << "rdma_local_window_capacity=" << _local_window_capacity - << connector << "rdma_remote_window_capacity=" << _remote_window_capacity - << connector << "rdma_sbuf_head=" << _sq_current - << connector << "rdma_sbuf_tail=" << _sq_sent - << connector << "rdma_rbuf_head=" << _rq_received - << connector << "rdma_unacked_rq_wr=" << _new_rq_wrs.load(butil::memory_order_relaxed) - << connector << "rdma_received_ack=" << _accumulated_ack - << connector << "rdma_unsolicited_sent=" << _unsolicited - << connector << "rdma_unsignaled_sq_wr=" << _sq_unsignaled - << connector << "rdma_read_buf=" << _input_processor.read_buf().size(); +void RdmaEndpoint::DebugInfo(std::ostream &os, + butil::StringPiece connector) const { + os << "rdma_state=ON" << connector + << "rdma_sq_imm_window_size=" << _sq_imm_window_size << connector + << "rdma_remote_rq_window_size=" + << _remote_rq_window_size.load(butil::memory_order_relaxed) << connector + << "rdma_sq_window_size=" + << _sq_window_size.load(butil::memory_order_relaxed) << connector + << "rdma_local_window_capacity=" << _local_window_capacity << connector + << "rdma_remote_window_capacity=" << _remote_window_capacity << connector + << "rdma_sbuf_head=" << _sq_current << connector + << "rdma_sbuf_tail=" << _sq_sent << connector + << "rdma_rbuf_head=" << _rq_received << connector + << "rdma_unacked_rq_wr=" << _new_rq_wrs.load(butil::memory_order_relaxed) + << connector << "rdma_received_ack=" << _accumulated_ack << connector + << "rdma_unsolicited_sent=" << _unsolicited << connector + << "rdma_unsignaled_sq_wr=" << _sq_unsignaled; } int RdmaEndpoint::GlobalInitialize() { @@ -1801,8 +1233,8 @@ int RdmaEndpoint::GlobalInitialize() { g_rdma_resource_mutex = new butil::Mutex; for (int i = 0; i < FLAGS_rdma_prepared_qp_cnt; ++i) { - RdmaResource* res = AllocateQpCq(FLAGS_rdma_prepared_qp_size, - FLAGS_rdma_prepared_qp_size); + RdmaResource *res = + AllocateQpCq(FLAGS_rdma_prepared_qp_size, FLAGS_rdma_prepared_qp_size); if (!res) { return -1; } @@ -1897,8 +1329,8 @@ int RdmaEndpoint::PollingModeInitialize(bthread_tag_t tag, }; for (int i = 0; i < FLAGS_rdma_poller_num; ++i) { auto args = new FnArgs{&pollers[i], &running}; - auto attr = FLAGS_rdma_disable_bthread ? BTHREAD_ATTR_PTHREAD - : BTHREAD_ATTR_NORMAL; + auto attr = + FLAGS_rdma_disable_bthread ? BTHREAD_ATTR_PTHREAD : BTHREAD_ATTR_NORMAL; attr.tag = tag; bthread_attr_set_name(&attr, "RdmaPolling"); pollers[i].callback = callback; @@ -1927,27 +1359,23 @@ void RdmaEndpoint::PollingModeRelease(bthread_tag_t tag) { } void RdmaEndpoint::PollerAddCqSid() { - if (_cq_sid == INVALID_SOCKET_ID) { - return; - } - auto index = butil::fmix32(_cq_sid) % FLAGS_rdma_poller_num; auto& group = _poller_groups[bthread_self_tag()]; auto& pollers = group.pollers; auto& poller = pollers[index]; - poller.op_queue.Enqueue(CqSidOp{_cq_sid, CqSidOp::ADD}); + if (INVALID_SOCKET_ID != _cq_sid) { + poller.op_queue.Enqueue(CqSidOp{_cq_sid, CqSidOp::ADD}); + } } void RdmaEndpoint::PollerRemoveCqSid() { - if (INVALID_SOCKET_ID == _cq_sid) { - return; - } - auto index = butil::fmix32(_cq_sid) % FLAGS_rdma_poller_num; auto& group = _poller_groups[bthread_self_tag()]; auto& pollers = group.pollers; auto& poller = pollers[index]; - poller.op_queue.Enqueue(CqSidOp{_cq_sid, CqSidOp::REMOVE}); + if (INVALID_SOCKET_ID != _cq_sid) { + poller.op_queue.Enqueue(CqSidOp{_cq_sid, CqSidOp::REMOVE}); + } } } // namespace rdma diff --git a/src/brpc/rdma/rdma_endpoint.h b/src/brpc/rdma/rdma_endpoint.h index 6d6ac391cd..f9223234fb 100644 --- a/src/brpc/rdma/rdma_endpoint.h +++ b/src/brpc/rdma/rdma_endpoint.h @@ -31,51 +31,27 @@ #include "butil/containers/mpsc_queue.h" #include "butil/containers/optional.h" #include "brpc/socket.h" -#include "brpc/rdma/rdma_handshake_server.h" namespace brpc { class Socket; +class RdmaTransport; namespace rdma { DECLARE_bool(rdma_use_polling); DECLARE_int32(rdma_poller_num); DECLARE_bool(rdma_disable_bthread); -class RdmaHandshakeClientV2; -class RdmaHandshakeServerV2; -class RdmaHandshakeClientV3; -class RdmaHandshakeServerV3; -struct ParsedHello; -enum class RemoteHelloResult; -class RdmaHello; -class RdmaEndpoint; -namespace v2_wire { - RemoteHelloResult ReadBodyAndNegotiate(RdmaEndpoint* ep, ParsedHello* remote); - int DrainBytes(RdmaEndpoint* ep, size_t n); -} // namespace v2_wire - -namespace v3_wire { - void FillLocalRdmaHello(const RdmaEndpoint* ep, RdmaHello* msg); - int ReadAndParseV3Hello(RdmaEndpoint* ep, RdmaHello* out); - int WriteV3Hello(RdmaEndpoint* ep, const RdmaHello& msg); -} // namespace v3_wire - -class RdmaConnect : public AppConnect { -public: - void StartConnect(const Socket* socket, - void (*done)(int err, void* data), void* data) override; - void StopConnect(Socket*) override; - struct RunGuard { - RunGuard(RdmaConnect* rc) { this_rc = rc; } - ~RunGuard() { if (this_rc) this_rc->Run(); } - RdmaConnect* this_rc; - }; - -private: - void Run(); - void (*_done)(int, void*){nullptr}; - void* _data{nullptr}; +// Wire-independent RDMA connection parameters consumed by resource setup. +// Transport adapters translate their protocol-specific payload into this DTO. +struct RdmaConnectionInfo { + uint32_t block_size; + uint16_t sq_size; + uint16_t rq_size; + uint16_t lid; + ibv_gid gid; + uint32_t qp_num; + butil::optional ece; }; struct RdmaResource { @@ -93,17 +69,8 @@ struct RdmaResource { }; class BAIDU_CACHELINE_ALIGNMENT RdmaEndpoint : public SocketUser { -friend class RdmaConnect; friend class Socket; -friend class RdmaHandshakeClientV2; -friend class RdmaHandshakeServerV2; -friend class RdmaHandshakeClientV3; -friend class RdmaHandshakeServerV3; -friend RemoteHelloResult v2_wire::ReadBodyAndNegotiate(RdmaEndpoint*, ParsedHello*); -friend int v2_wire::DrainBytes(RdmaEndpoint*, size_t); -friend void v3_wire::FillLocalRdmaHello(const RdmaEndpoint*, RdmaHello*); -friend int v3_wire::ReadAndParseV3Hello(RdmaEndpoint*, RdmaHello*); -friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&); +friend class ::brpc::RdmaTransport; public: explicit RdmaEndpoint(Socket* s); ~RdmaEndpoint() override; @@ -124,16 +91,16 @@ friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&); // Whether the endpoint can send more data bool IsWritable() const; + // Resource information consumed by the transport-level RDMA adapter. + void GetLocalConnectionInfo(RdmaConnectionInfo* local) const; + // Returns 0 on success, 1 when ECE query is unavailable, -1 on error. + int QueryLocalEce(ibv_ece* ece) const; + void SetOutgoingEce(const ibv_ece& ece); + // For debug void DebugInfo(std::ostream& os, butil::StringPiece connector = "\n") const; - // Callback when there is new epollin event on TCP fd. - static void OnNewDataFromTcp(Socket* m); - - // Real handshake for RDMA-mode sockets. - static ParseResult ExecuteServerHandshake(butil::IOBuf* source, Socket* socket); - // Initialize polling mode static int PollingModeInitialize(bthread_tag_t tag, std::function callback, @@ -143,33 +110,8 @@ friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&); static void PollingModeRelease(bthread_tag_t tag); private: - enum State { - UNINIT = 0x0, - C_ALLOC_QPCQ = 0x1, - C_HELLO_SEND = 0x2, - C_HELLO_WAIT = 0x3, - C_BRINGUP_QP = 0x4, - C_ACK_SEND = 0x5, - S_HELLO_WAIT = 0x11, - S_ALLOC_QPCQ = 0x12, - S_BRINGUP_QP = 0x13, - S_HELLO_SEND = 0x14, - S_ACK_WAIT = 0x15, - ESTABLISHED = 0x100, - FALLBACK_TCP = 0x200, - FAILED = 0x300 - }; - - // Process handshake at the client - static void* ProcessHandshakeAtClient(void* arg); - - static void OnNewDataFromTcpAtClient(Socket* m); - static void OnNewDataFromTcpAtServer(Socket* m); - - bool HandleTcpEventAfterEstablished(); - // Allocate resources. On failure the endpoint is left with no RDMA - // resource attached, so that the handshake can safely fall back to TCP. + // resource attached, so the caller can safely continue without RDMA. // Return 0 if success, -1 if failed and errno set int AllocateResources(); @@ -177,31 +119,12 @@ friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&); // in the middle with resources partially allocated. // Return 0 if success, -1 if failed and errno set int DoAllocateResources(); + // Start consuming CQ events only after handshake parsing has finished. + int StartCqEvents(); // Release resources void DeallocateResources(); - // Create the Socket wrapping the CQ (and register it with the poller in - // polling mode), which is what makes CQ events reachable and thus starts - // PollCq. - // - // Must not be called before the handshake has reached ESTABLISHED, nor - // from within the fd stream's parsing path: PollCq() parses the input - // stream carried by the QP, and the Socket's `parsing_context` and - // `preferred_index` belong to the fd stream until the handshake is over - // and CutInputMessage has returned. It keeps writing both after the - // handshake handler hands the stream back. Those two are per-Socket, - // so letting PollCq in early makes two streams parse through one context. - // The server therefore calls this from OnNewDataFromTcpAtServer(), after - // OnNewMessages() returns, not from ExecuteServerHandshake(). - // - // No CQE is lost by deferring: BringUpQp() fills the RQ before the QP - // reaches RTS, both CQs are armed by DoAllocateResources(), and adding an - // already readable fd to an edge-triggered epoll reports it immediately. - // - // Return 0 if success, -1 if failed and errno set - int StartCqEvents(); - // Send Imm data to the remote side // Arguments: // imm: imm data in the WR @@ -239,37 +162,18 @@ friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&); // -1: failed, errno set int DoPostRecv(void* block, size_t block_size); - // Read at most len bytes from fd in _socket to data - // wait for _read_butex if encounter EAGAIN - // return -1 if encounter other errno (including EOF) - int ReadFromFd(void* data, size_t len); - int ReadFromFd(butil::IOPortal* data, size_t len); - - - // Write at most len bytes from data to fd in _socket - // wait for _epollout_butex if encounter EAGAIN - // return -1 if encounter other errno - int WriteToFd(void* data, size_t len); - - // Write data to fd in _socket. - // wait for _epollout_butex if encounter EAGAIN. - // return -1 if encounter other errno. - int WriteToFd(butil::IOBuf* data); - - // Copy negotiated remote parameters into the endpoint and compute - // the SQ/RQ window capacities. Called by both - // ProcessHandshakeAtClient and ProcessHandshakeAtServer after the - // peer's hello has been validated. - void ApplyRemoteHello(const ParsedHello& remote); + // Copy negotiated remote parameters into the endpoint and compute the + // SQ/RQ window capacities. + void ApplyRemoteInfo(const RdmaConnectionInfo& remote); // Bringup the QP from RESET state to RTS state. // Arguments: - // remote: parsed remote hello. Provides the remote LID/GID/QP + // remote: negotiated peer parameters. Provides the remote LID/GID/QP // number for the RTR transition, and (on v3) the peer's // ECE to set during the INIT->RTR transition. // is_server: true on the server side, false on the client side. // Returns 0 on success, -1 on failed and errno set. - int BringUpQp(const ParsedHello& remote, bool is_server); + int BringUpQp(const RdmaConnectionInfo& remote, bool is_server); // Get event from comp channel and ack the events int GetAndAckEvents(SocketUniquePtr& s); @@ -280,9 +184,6 @@ friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&); // Poll CQ and get the work completion static void PollCq(Socket* m); - // Get the description of current handshake state - std::string GetStateStr() const; - // Add cq socket id to poller void PollerAddCqSid(); @@ -291,25 +192,10 @@ friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&); // Not owner Socket* _socket; + // Input state dedicated to the stream carried by the RDMA QP. + InputMessengerProcessor _input_processor; - // State of Handshake. FALLBACK_TCP publishes RdmaTransport::_rdma_state - // with release ordering and is consumed by OnNewDataFromTcpAtClient with acquire - // ordering. Other state accesses do not publish data and use relaxed - // ordering. - butil::atomic _state; - - // Wire-level handshake protocol version (set by dispatch in - // ProcessHandshakeAtClient/Server). Aligned with the protocol code: - // 0 = unnegotiated - // 2 = v2 "RDMA" - // 3 = v3 "RDM3" - int _handshake_version; - - // ECE payload to advertise in the next local hello: - // Client: the locally queried ECE capabilities (filled - // before C_HELLO_SEND); - // Server: the reduced/negotiated ECE queried after the - // QP reached RTS (filled in BringUpQp). + // ECE payload prepared by resource setup and consumed by the RDMA adapter. butil::optional _outgoing_ece; // rdma resource @@ -326,9 +212,6 @@ friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&); uint16_t _sq_size; uint16_t _rq_size; - // The input stream carried by the QP. - InputMessengerProcessor _input_processor; - // Act as sendbuf and recvbuf, but requires no memcpy std::vector _sbuf; std::vector _rbuf; @@ -364,9 +247,6 @@ friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&); // The number of new WRs posted in the local Recv Queue butil::atomic _new_rq_wrs; - // butex for inform read events on TCP fd during handshake - butil::atomic *_read_butex; - DISALLOW_COPY_AND_ASSIGN(RdmaEndpoint); // Cq socket id operation type diff --git a/src/brpc/rdma/rdma_handshake.cpp b/src/brpc/rdma/rdma_handshake.cpp deleted file mode 100644 index 17c6715562..0000000000 --- a/src/brpc/rdma/rdma_handshake.cpp +++ /dev/null @@ -1,514 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, -// software distributed under the License is distributed on an -// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -// KIND, either express or implied. See the License for the -// specific language governing permissions and limitations -// under the License. - -#if BRPC_WITH_RDMA - -#include "brpc/rdma/rdma_handshake.h" -#include "brpc/rdma/rdma_handshake_constants.h" - -#include -#include // std::min -#include -#include -#include -#include "butil/iobuf.h" // IOBuf, IOPortal, IOBufAsZeroCopy*Stream -#include "butil/sys_byteorder.h" -#include "butil/raw_pack.h" // RawPacker, RawUnpacker -#include "brpc/socket.h" -#include "brpc/rdma/rdma_endpoint.h" -#include "brpc/rdma/rdma_helper.h" -#include "brpc/rdma_transport.h" -#include "brpc/rdma/rdma_handshake.pb.h" -#include "butil/object_pool.h" - -namespace brpc { -namespace rdma { - -DEFINE_int32(rdma_client_handshake_version, 2, - "RDMA handshake protocol version used by client. " - "2 = legacy 'RDMA' magic (default, compatible with all servers); " - "3 = new 'RDM3' protobuf-based handshake " - "(MUST only be enabled after target servers support v3)."); -DECLARE_bool(rdma_trace_verbose); - -extern const uint16_t MIN_QP_SIZE; -extern const uint16_t MIN_BLOCK_SIZE; -extern uint32_t g_rdma_recv_block_size; -extern bool g_skip_rdma_init; - -extern int (*IbvQueryEce)(ibv_qp*, ibv_ece*); - -DEFINE_bool(rdma_ece, false, "Enable end-to-end ECE (Enhanced Connection Establishment) " - "negotiation in the RDMA v3 handshake. Automatically degrades " - "to no-ECE when the peer, the local libibverbs, or set_ece " - "does not support it. Acts as a kill switch (default off)."); - -DECLARE_bool(rdma_trace_verbose); - -namespace v2_wire { - -void HelloMessage::Serialize(void* data) const { - butil::RawPacker(data) - .pack16(msg_len) - .pack16(hello_ver) - .pack16(impl_ver) - .pack32(block_size) - .pack16(sq_size) - .pack16(rq_size) - .pack16(lid) - // gid is a raw 16-byte identifier and must NOT be byte-swapped. - .pack_bytes(gid.raw, sizeof(gid.raw)) - .pack32(qp_num); -} - -void HelloMessage::Deserialize(void* data) { - butil::RawUnpacker(data) - .unpack16(msg_len) - .unpack16(hello_ver) - .unpack16(impl_ver) - .unpack32(block_size) - .unpack16(sq_size) - .unpack16(rq_size) - .unpack16(lid) - // gid is a raw 16-byte identifier and must NOT be byte-swapped. - .unpack_bytes(gid.raw, sizeof(gid.raw)) - .unpack32(qp_num); -} - -static bool ValidHelloMessage(const HelloMessage& msg) { - return msg.hello_ver == HELLO_V2_VERSION && - msg.impl_ver == IMPL_V2_VERSION && - msg.block_size >= MIN_BLOCK_SIZE && - msg.sq_size >= MIN_QP_SIZE && - msg.rq_size >= MIN_QP_SIZE; -} - -static void TranslateV2Hello(const HelloMessage& msg, ParsedHello* out) { - out->block_size = msg.block_size; - out->sq_size = msg.sq_size; - out->rq_size = msg.rq_size; - out->lid = msg.lid; - out->gid = msg.gid; - out->qp_num = msg.qp_num; -} - -RemoteHelloResult ReadBodyAndNegotiate(RdmaEndpoint* ep, ParsedHello* remote) { - uint8_t data[HELLO_V2_MSG_LEN_MIN]; - if (ep->ReadFromFd(data, HELLO_V2_MSG_LEN_MIN - HELLO_MAGIC_LEN) < 0) { - return RemoteHelloResult::ERROR; - } - HelloMessage remote_msg{}; - remote_msg.Deserialize(data); - if (remote_msg.msg_len < HELLO_V2_MSG_LEN_MIN || - remote_msg.msg_len > HELLO_V2_MSG_LEN_MAX) { - errno = EPROTO; - return RemoteHelloResult::ERROR; - } - if (remote_msg.msg_len > HELLO_V2_MSG_LEN_MIN) { - // Drain unknown trailing bytes so they don't pollute subsequent - // reads (e.g. the upcoming ACK message). v2 base fields already - // carry enough information for negotiation; unknown trailing - // bytes are treated as optional hints that v2 safely ignores. - size_t ext_len = remote_msg.msg_len - HELLO_V2_MSG_LEN_MIN; - if (DrainBytes(ep, ext_len) < 0) { - return RemoteHelloResult::ERROR; - } - } - if (!ValidHelloMessage(remote_msg)) { - return RemoteHelloResult::FALLBACK; - } - TranslateV2Hello(remote_msg, remote); - return RemoteHelloResult::NEGOTIATED; -} - -int DrainBytes(RdmaEndpoint* ep, size_t n) { - uint8_t scratch[64]; - while (n > 0) { - size_t chunk = std::min(n, sizeof(scratch)); - if (ep->ReadFromFd(scratch, chunk) < 0) { - return -1; - } - n -= chunk; - } - return 0; -} - -} // namespace v2_wire - -int RdmaHandshakeClientV2::SendLocalHello() { - RdmaEndpoint* ep = _ep; - uint8_t data[HELLO_V2_MSG_LEN_MIN]; - - v2_wire::HelloMessage local_msg{}; - local_msg.msg_len = HELLO_V2_MSG_LEN_MIN; - local_msg.hello_ver = HELLO_V2_VERSION; - local_msg.impl_ver = IMPL_V2_VERSION; - local_msg.block_size = g_rdma_recv_block_size; - local_msg.sq_size = ep->_sq_size; - local_msg.rq_size = ep->_rq_size; - local_msg.lid = GetRdmaLid(); - local_msg.gid = GetRdmaGid(); - if (BAIDU_LIKELY(ep->_resource)) { - local_msg.qp_num = ep->_resource->qp->qp_num; - } else { - // Only happens in UT - local_msg.qp_num = 0; - } - fast_memcpy(data, HELLO_MAGIC, 4); - local_msg.Serialize((char*)data + 4); - return ep->WriteToFd(data, HELLO_V2_MSG_LEN_MIN); -} - -RemoteHelloResult RdmaHandshakeClientV2::ReceiveAndParseRemoteHello(ParsedHello* remote) { - uint8_t magic[HELLO_MAGIC_LEN]; - if (_ep->ReadFromFd(magic, HELLO_MAGIC_LEN) < 0) { - return RemoteHelloResult::ERROR; - } - if (memcmp(magic, HELLO_MAGIC, HELLO_MAGIC_LEN) != 0) { - errno = EPROTO; - return RemoteHelloResult::ERROR; - } - - return v2_wire::ReadBodyAndNegotiate(_ep, remote); -} - -// Parse one complete v2 client hello out of `_source` (non-blocking). -// v2 hello: [ "RDMA" 4B ][ msg_len 2B ][ 34B ... ], base total = 40B. -RemoteHelloResult RdmaHandshakeServerV2::ReceiveAndParseRemoteHello(ParsedHello* remote) { - butil::IOBuf* source = _source; - constexpr size_t HDR_LEN = HELLO_MAGIC_LEN + 2; - if (source->size() < HDR_LEN) { - // msg_len has not fully arrived yet. - return RemoteHelloResult::NEED_MORE; - } - - uint8_t hdr[HDR_LEN]; - CHECK_EQ(source->copy_to(hdr, sizeof(hdr)), sizeof(hdr)); - - uint16_t msg_len = 0; - butil::RawUnpacker(hdr + HELLO_MAGIC_LEN).unpack16(msg_len); - if (msg_len < HELLO_V2_MSG_LEN_MIN || msg_len > HELLO_V2_MSG_LEN_MAX) { - errno = EPROTO; - return RemoteHelloResult::ERROR; - } - if (source->size() < msg_len) { - // Full message has not fully arrived yet. - return RemoteHelloResult::NEED_MORE; - } - - // Consume the whole hello: magic + 36B base body + optional extension. - CHECK_EQ(source->pop_front(HELLO_MAGIC_LEN), HELLO_MAGIC_LEN); - uint8_t body[HELLO_V2_MSG_LEN_MIN - HELLO_MAGIC_LEN]; // 36B - CHECK_EQ(source->cutn(body, sizeof(body)), sizeof(body)); - if (!source->empty()) { - // Drain unknown trailing bytes. - source->clear(); - } - - v2_wire::HelloMessage remote_msg{}; - remote_msg.Deserialize(body); - if (!v2_wire::ValidHelloMessage(remote_msg)) { - return RemoteHelloResult::FALLBACK; - } - v2_wire::TranslateV2Hello(remote_msg, remote); - return RemoteHelloResult::NEGOTIATED; -} - -int RdmaHandshakeServerV2::SendLocalHello() { - uint8_t data[HELLO_V2_MSG_LEN_MIN]; - v2_wire::HelloMessage local_msg{}; - local_msg.msg_len = HELLO_V2_MSG_LEN_MIN; - auto rdma_transport = static_cast(_ep->_socket->_transport.get()); - if (rdma_transport->_rdma_state == RdmaTransport::RDMA_OFF) { - local_msg.hello_ver = 0; - local_msg.impl_ver = 0; - local_msg.block_size = 0; - local_msg.sq_size = 0; - local_msg.rq_size = 0; - local_msg.lid = 0; - memset(local_msg.gid.raw, 0, sizeof(local_msg.gid.raw)); - local_msg.qp_num = 0; - } else { - local_msg.hello_ver = HELLO_V2_VERSION; - local_msg.impl_ver = IMPL_V2_VERSION; - local_msg.block_size = g_rdma_recv_block_size; - local_msg.sq_size = _ep->_sq_size; - local_msg.rq_size = _ep->_rq_size; - local_msg.lid = GetRdmaLid(); - local_msg.gid = GetRdmaGid(); - if (BAIDU_LIKELY(_ep->_resource)) { - local_msg.qp_num = _ep->_resource->qp->qp_num; - } else { - // Only happens in UT - local_msg.qp_num = 0; - } - } - fast_memcpy(data, HELLO_MAGIC, 4); - local_msg.Serialize((char*)data + 4); - return _ep->WriteToFd(data, HELLO_V2_MSG_LEN_MIN); -} - -namespace v3_wire { - -bool ValidRdmaHello(const RdmaHello& msg) { - if (msg.gid().size() != sizeof(ibv_gid)) { - return false; - } - // ParsedHello stores these as uint16_t; reject values that would truncate. - constexpr uint16_t MAX_UINT16 = std::numeric_limits::max(); - if (msg.sq_size() > MAX_UINT16 || msg.rq_size() > MAX_UINT16 || msg.lid() > MAX_UINT16) { - return false; - } - if (msg.block_size() < MIN_BLOCK_SIZE) { - return false; - } - if (msg.sq_size() < MIN_QP_SIZE) { - return false; - } - if (msg.rq_size() < MIN_QP_SIZE) { - return false; - } - // qp_num == 0 only happens in UT (no real QP allocated). - if (msg.qp_num() == 0 && !g_skip_rdma_init) { - return false; - } - return true; -} - -void FillLocalRdmaHello(const RdmaEndpoint* ep, RdmaHello* msg) { - msg->set_block_size(g_rdma_recv_block_size); - msg->set_sq_size(ep->_sq_size); - msg->set_rq_size(ep->_rq_size); - msg->set_lid(GetRdmaLid()); - ibv_gid gid = GetRdmaGid(); - msg->set_gid(reinterpret_cast(gid.raw), sizeof(gid.raw)); - if (BAIDU_LIKELY(ep->_resource)) { - msg->set_qp_num(ep->_resource->qp->qp_num); - } else { - // Only happens in UT - msg->set_qp_num(0); - } - - // Advertise ECE only when enabled. Role-dependent payload: - // Client hello: the locally queried ECE capabilities; - // Server hello: the reduced/negotiated ECE queried after RTS. - // When the relevant ECE is not valid (disabled, unsupported, or query - // failed) the field is simply omitted and the peer degrades to no-ECE. - // Advertise ECE if there is anything to advertise. The endpoint pre-fills - // _outgoing_ece in a role-specific way: client side stores its locally - // queried capabilities; server side stores the reduced/negotiated ECE - // after RTS. nullopt -> omit the field (peer degrades to no-ECE). - if (FLAGS_rdma_ece && ep->_outgoing_ece.has_value()) { - RdmaEce* ece = msg->mutable_ece(); - ece->set_vendor_id(ep->_outgoing_ece->vendor_id); - ece->set_options(ep->_outgoing_ece->options); - ece->set_comp_mask(ep->_outgoing_ece->comp_mask); - } -} - -int ReadAndParseV3Hello(RdmaEndpoint* ep, RdmaHello* out) { - uint8_t size_buf[HELLO_V3_PB_SIZE_LEN]; - if (ep->ReadFromFd(size_buf, HELLO_V3_PB_SIZE_LEN) < 0) { - return -1; - } - uint32_t pb_size = butil::NetToHost32( - *reinterpret_cast(size_buf)); - if (pb_size == 0 || pb_size > HELLO_V3_MAX_PB_SIZE) { - errno = EPROTO; - return -1; - } - butil::IOPortal body; - if (ep->ReadFromFd(&body, pb_size) < 0) { - return -1; - } - - butil::IOBufAsZeroCopyInputStream input(body); - if (!out->ParseFromZeroCopyStream(&input)) { - LOG(ERROR) << "Failed to parse RdmaHello"; - errno = EPROTO; - return -1; - } - return 0; -} - -int WriteV3Hello(RdmaEndpoint* ep, const RdmaHello& msg) { - uint32_t pb_size = static_cast(msg.ByteSizeLong()); - if (pb_size > HELLO_V3_MAX_PB_SIZE) { - errno = EPROTO; - return -1; - } - - // [ "RDM3" 4B ][ pb_size 4B (big-endian) ][ RdmaHello protobuf bytes ] - butil::IOBuf packet; - packet.append(HELLO_MAGIC_V3, HELLO_MAGIC_LEN); - uint32_t pb_size_be = butil::HostToNet32(pb_size); - packet.append(&pb_size_be, HELLO_V3_PB_SIZE_LEN); - butil::IOBufAsZeroCopyOutputStream output(&packet); - if (!msg.SerializeToZeroCopyStream(&output)) { - LOG(ERROR) << "Failed to serialize RdmaHello"; - errno = EPROTO; - return -1; - } - return ep->WriteToFd(&packet); -} - -void TranslateHello(const RdmaHello& msg, ParsedHello* out) { - out->block_size = msg.block_size(); - out->sq_size = static_cast(msg.sq_size()); - out->rq_size = static_cast(msg.rq_size()); - out->lid = static_cast(msg.lid()); - fast_memcpy(out->gid.raw, msg.gid().data(), sizeof(out->gid.raw)); - out->qp_num = msg.qp_num(); - if (FLAGS_rdma_ece && msg.has_ece()) { - ibv_ece ece; - ece.vendor_id = msg.ece().vendor_id(); - ece.options = msg.ece().options(); - ece.comp_mask = msg.ece().comp_mask(); - out->ece = ece; - } -} - -} // namespace v3_wire - -int RdmaHandshakeClientV3::SendLocalHello() { - // Query local ECE capabilities so they can be advertised in the client - // hello. v3-only. Best-effort: any failure or missing API just means we - // won't advertise ECE (the peer then degrades to no-ECE establishment). - if (FLAGS_rdma_ece && IbvQueryEce != nullptr && - _ep->_resource && _ep->_resource->qp) { - ibv_ece ece; - if (IbvQueryEce(_ep->_resource->qp, &ece) == 0) { - _ep->_outgoing_ece = ece; - } else { - LOG_IF(WARNING, FLAGS_rdma_trace_verbose) - << "Fail to IbvQueryEce on client, ECE not advertised: " - << _ep->_socket->description(); - } - } - - RdmaHello local_msg{}; - v3_wire::FillLocalRdmaHello(_ep, &local_msg); - return v3_wire::WriteV3Hello(_ep, local_msg); -} - -RemoteHelloResult RdmaHandshakeClientV3::ReceiveAndParseRemoteHello(ParsedHello* remote) { - uint8_t magic[HELLO_MAGIC_LEN]; - if (_ep->ReadFromFd(magic, HELLO_MAGIC_LEN) < 0) { - return RemoteHelloResult::ERROR; - } - if (memcmp(magic, HELLO_MAGIC_V3, HELLO_MAGIC_LEN) != 0) { - errno = EPROTO; - return RemoteHelloResult::ERROR; - } - - RdmaHello remote_msg{}; - if (v3_wire::ReadAndParseV3Hello(_ep, &remote_msg) < 0) { - return RemoteHelloResult::ERROR; - } - if (!v3_wire::ValidRdmaHello(remote_msg)) { - return RemoteHelloResult::FALLBACK; - } - v3_wire::TranslateHello(remote_msg, remote); - return RemoteHelloResult::NEGOTIATED; -} - -// Parse one complete v3 client hello out of `_source` (non-blocking). -// v3 hello: [ "RDM3" 4B ][ pb_size 4B (big-endian) ][ RdmaHello ] -RemoteHelloResult RdmaHandshakeServerV3::ReceiveAndParseRemoteHello(ParsedHello* remote) { - constexpr size_t HDR_LEN = HELLO_MAGIC_LEN + HELLO_V3_PB_SIZE_LEN; - if (_source->size() < HDR_LEN) { - // pb_size has not fully arrived yet. - return RemoteHelloResult::NEED_MORE; - } - - uint8_t hdr[HDR_LEN]; - CHECK_EQ(_source->copy_to(hdr, sizeof(hdr)), sizeof(hdr)); - - uint32_t pb_size = butil::NetToHost32( - *reinterpret_cast(hdr + HELLO_MAGIC_LEN)); - if (pb_size == 0 || pb_size > HELLO_V3_MAX_PB_SIZE) { - errno = EPROTO; - return RemoteHelloResult::ERROR; - } - size_t total = HDR_LEN + pb_size; - if (_source->size() < total) { - // Full message has not fully arrived yet. - return RemoteHelloResult::NEED_MORE; - } - - CHECK_EQ(_source->cutn(hdr, HDR_LEN), HDR_LEN); - butil::IOBuf pb; - CHECK_EQ(_source->cutn(&pb, pb_size), pb_size); - RdmaHello remote_msg; - butil::IOBufAsZeroCopyInputStream input(pb); - if (!remote_msg.ParseFromZeroCopyStream(&input)) { - LOG(ERROR) << "Failed to parse RdmaHello"; - errno = EPROTO; - return RemoteHelloResult::ERROR; - } - if (!v3_wire::ValidRdmaHello(remote_msg)) { - return RemoteHelloResult::FALLBACK; - } - v3_wire::TranslateHello(remote_msg, remote); - return RemoteHelloResult::NEGOTIATED; -} - -int RdmaHandshakeServerV3::SendLocalHello() { - RdmaHello local_msg{}; - auto rdma_transport = static_cast(_ep->_socket->_transport.get()); - if (rdma_transport->_rdma_state == RdmaTransport::RDMA_OFF) { - // Un-negotiable hello: all body fields are zero so the client's - // rejects it and downgrades to TCP on the same connection. - local_msg.set_block_size(0); - local_msg.set_sq_size(0); - local_msg.set_rq_size(0); - local_msg.set_lid(0); - local_msg.set_gid(std::string(sizeof(ibv_gid), '\0')); - local_msg.set_qp_num(0); - } else { - v3_wire::FillLocalRdmaHello(_ep, &local_msg); - } - return v3_wire::WriteV3Hello(_ep, local_msg); -} - -std::unique_ptr CreateClientHandshake(RdmaEndpoint* ep) { - switch (FLAGS_rdma_client_handshake_version) { - case 3: - return std::unique_ptr(new RdmaHandshakeClientV3(ep)); - case 2: - default: - return std::unique_ptr(new RdmaHandshakeClientV2(ep)); - } -} - -std::unique_ptr CreateServerHandshakeByMagic( - RdmaEndpoint* ep, butil::IOBuf* source, const uint8_t magic[HELLO_MAGIC_LEN]) { - if (memcmp(magic, HELLO_MAGIC, HELLO_MAGIC_LEN) == 0) { - return std::unique_ptr( - new RdmaHandshakeServerV2(ep, source)); - } - if (memcmp(magic, HELLO_MAGIC_V3, HELLO_MAGIC_LEN) == 0) { - return std::unique_ptr( - new RdmaHandshakeServerV3(ep, source)); - } - return nullptr; -} - -} // namespace rdma -} // namespace brpc - -#endif // BRPC_WITH_RDMA diff --git a/src/brpc/rdma/rdma_handshake.h b/src/brpc/rdma/rdma_handshake.h deleted file mode 100644 index 2d10220ab9..0000000000 --- a/src/brpc/rdma/rdma_handshake.h +++ /dev/null @@ -1,185 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, -// software distributed under the License is distributed on an -// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -// KIND, either express or implied. See the License for the -// specific language governing permissions and limitations -// under the License. - -#ifndef BRPC_RDMA_HANDSHAKE_H -#define BRPC_RDMA_HANDSHAKE_H - -#if BRPC_WITH_RDMA - -#include -#include -#include "butil/macros.h" -#include "butil/containers/optional.h" -#include "brpc/rdma/rdma_handshake_constants.h" - -namespace butil { -class IOBuf; -} - -namespace brpc { -namespace rdma { - -class RdmaEndpoint; - -// Wire-format-agnostic representation of a peer's hello message. -// Each protocol version (v2 binary, v3 protobuf) translates its own -// wire format into this struct so the state-machine driver in -// RdmaEndpoint::ProcessHandshakeAt{Client,Server} stays free of any -// wire-format details. -struct ParsedHello { - uint32_t block_size; - uint16_t sq_size; - uint16_t rq_size; - uint16_t lid; - ibv_gid gid; - uint32_t qp_num; - // ECE (Enhanced Connection Establishment), v3 handshake only. - // nullopt means the peer did not advertise any ECE (either it disabled - // ECE, its lib does not support ECE, or it is a v2 peer). When engaged: - // - on the server side: the client's queried ECE capabilities; - // - on the client side: the server's reduced/negotiated ECE. - butil::optional ece; -}; - -// Result of reading/parsing a peer's hello (see ReceiveAndParseRemoteHello). -enum class RemoteHelloResult { - // A full hello was read and negotiation succeeded. - NEGOTIATED, - // A full hello was read but negotiation failed. - FALLBACK, - // (server only) Not enough data yet. - NEED_MORE, - // IO/protocol error (errno set). - ERROR, -}; - -namespace v2_wire { - -// v2 binary HelloMessage. -struct HelloMessage { - void Serialize(void* data) const; - void Deserialize(void* data); - - uint16_t msg_len; - uint16_t hello_ver; - uint16_t impl_ver; - uint32_t block_size; - uint16_t sq_size; - uint16_t rq_size; - uint16_t lid; - ibv_gid gid; - uint32_t qp_num; -}; - -} // namespace v2_wire - -// Base class of an RDMA handshake, shared by both roles. -class RdmaHandshake { -public: - RdmaHandshake(RdmaEndpoint* ep, int version) : _ep(ep), _version(version) {} - virtual ~RdmaHandshake() = default; - - DISALLOW_COPY_AND_ASSIGN(RdmaHandshake); - - // Wire-level protocol version (2 for "RDMA", 3 for "RDM3"). - int ProtocolVersion() const { return _version; } - - // Build and send the local hello (including the protocol magic). - // Returns 0 on success, -1 on IO error (errno set). - // - // For a server in fallback state, implementations MUST still - // produce a sendable message; each version uses its own wire - // convention to signal "I am falling back" to the peer: - // - v2: zero hello_ver/impl_ver so the peer's HelloNegotiationValid - // rejects it; - // - v3: qp_num==0 so the peer's ValidRdmaHello rejects it. - virtual int SendLocalHello() = 0; - - // Read and parse the peer's hello into *remote. - virtual RemoteHelloResult ReceiveAndParseRemoteHello(ParsedHello* remote) = 0; - -protected: - RdmaEndpoint* _ep; - int _version; -}; - -// Server-side handshake base: parses the remote hello non-blockingly out of -// `_source` (an IOBuf filled by InputMessenger), never touching the fd. -class ServerRdmaHandshake : public RdmaHandshake { -public: - ServerRdmaHandshake(RdmaEndpoint* ep, butil::IOBuf* source, int version) - : RdmaHandshake(ep, version), _source(source) {} - -protected: - butil::IOBuf* _source; -}; - -// v2 handshake (legacy "RDMA" magic, 36B binary HelloMessage). -class RdmaHandshakeClientV2 : public RdmaHandshake { -public: - explicit RdmaHandshakeClientV2(RdmaEndpoint* ep) : RdmaHandshake(ep, 2) {} - int SendLocalHello() override; - RemoteHelloResult ReceiveAndParseRemoteHello(ParsedHello* remote) override; -}; - -class RdmaHandshakeServerV2 : public ServerRdmaHandshake { -public: - RdmaHandshakeServerV2(RdmaEndpoint* ep, butil::IOBuf* source) - : ServerRdmaHandshake(ep, source, 2) {} - int SendLocalHello() override; - RemoteHelloResult ReceiveAndParseRemoteHello(ParsedHello* remote) override; -}; - -// v3 handshake (new "RDM3" magic, protobuf RdmaHello). -// [ "RDM3" 4B ][ pb_size 4B (big-endian) ][ RdmaHello protobuf bytes ] -class RdmaHandshakeClientV3 : public RdmaHandshake { -public: - explicit RdmaHandshakeClientV3(RdmaEndpoint* ep) : RdmaHandshake(ep, 3) {} - int SendLocalHello() override; - RemoteHelloResult ReceiveAndParseRemoteHello(ParsedHello* remote) override; -}; - -class RdmaHandshakeServerV3 : public ServerRdmaHandshake { -public: - RdmaHandshakeServerV3(RdmaEndpoint* ep, butil::IOBuf* source) - : ServerRdmaHandshake(ep, source, 3) {} - int SendLocalHello() override; - RemoteHelloResult ReceiveAndParseRemoteHello(ParsedHello* remote) override; -}; - -// Factory methods -// -// Pick the client-side handshake based on -// FLAGS_rdma_client_handshake_version: -// 2 (default) -> RdmaHandshakeClientV2 -// 3 -> RdmaHandshakeClientV3 -// Other values fall back to V2. -std::unique_ptr CreateClientHandshake(RdmaEndpoint* ep); - -// Pick the server-side handshake based on the 4B magic already read. -// Returns nullptr if `magic` is not a recognized RDMA magic -// (the caller should then fallback to TCP). -// "RDMA" -> RdmaHandshakeServerV2 -// "RDM3" -> RdmaHandshakeServerV3 -std::unique_ptr CreateServerHandshakeByMagic( - RdmaEndpoint* ep, butil::IOBuf* source, const uint8_t magic[HELLO_MAGIC_LEN]); - -} // namespace rdma -} // namespace brpc - -#endif // BRPC_WITH_RDMA -#endif // BRPC_RDMA_HANDSHAKE_H diff --git a/src/brpc/rdma/rdma_handshake_server.cpp b/src/brpc/rdma/rdma_handshake_server.cpp deleted file mode 100644 index 1b0ca4226a..0000000000 --- a/src/brpc/rdma/rdma_handshake_server.cpp +++ /dev/null @@ -1,219 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, -// software distributed under the License is distributed on an -// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -// KIND, either express or implied. See the License for the -// specific language governing permissions and limitations -// under the License. - -#include "brpc/rdma/rdma_handshake_server.h" - -#include -#include -#include -#include "butil/iobuf.h" -#include "butil/logging.h" -#include "butil/object_pool.h" -#include "butil/raw_pack.h" -#include "butil/sys_byteorder.h" -#include "brpc/socket.h" -#include "brpc/rdma/rdma_handshake.pb.h" -#include "brpc/rdma/rdma_handshake_constants.h" -#if BRPC_WITH_RDMA -#include "brpc/rdma/rdma_endpoint.h" -#endif - -namespace brpc { -namespace rdma { - -ServerHandshakeContext* ServerHandshakeContext::Create() { - return butil::get_object(); -} - -void ServerHandshakeContext::Destroy() { - butil::return_object(this); -} - -// Fallback-only server handshake. Used for any connection that is NOT in RDMA -// mode: builds without RDMA (no RdmaEndpoint exists at all), and RDMA-enabled -// builds where this particular connection is plain TCP. Every RDMA client that -// reaches here is answered with an un-negotiable hello and asked to downgrade -// to TCP on the same connection, then its ACK is drained. - -// An intentionally-invalid v2 hello_ver: the client's ValidHelloMessage() -// requires hello_ver==2, so it rejects this and falls back. -static constexpr uint16_t V2_HELLO_VERSION_INVALID = std::numeric_limits::max(); -// Length of the gid field (== sizeof(ibv_gid)), spelled as a literal so this -// compile-switch-independent file stays free of . -static constexpr size_t V3_GID_LEN = 16; - -// Consume one complete v2 client hello ("RDMA" + 2B msg_len + body) from -// `source` without interpreting its content (fallback does not negotiate). -// Returns 1 (consumed), 0 (not enough data yet, nothing consumed) or -1 (error). -static int DrainClientHelloV2(butil::IOBuf* source) { - constexpr size_t HDR_LEN = HELLO_MAGIC_LEN + 2; - if (source->size() < HDR_LEN) { - return 0; - } - - uint8_t hdr[HDR_LEN]; - CHECK_EQ(source->copy_to(hdr, sizeof(hdr)), sizeof(hdr)); - - uint16_t msg_len = 0; - butil::RawUnpacker(hdr + HELLO_MAGIC_LEN).unpack16(msg_len); - if (msg_len < HELLO_V2_MSG_LEN_MIN || msg_len > HELLO_V2_MSG_LEN_MAX) { - return -1; - } - - if (source->size() < msg_len) { - return 0; - } - - CHECK_EQ(source->pop_front(msg_len), msg_len); - return 1; -} - -// Consume one complete v3 client hello ("RDM3" + 4B pb_size + protobuf body) -// from `source` without interpreting its content. -// Returns 1 (consumed), 0 (not enough data yet, nothing consumed) or -1 (error). -static int DrainClientHelloV3(butil::IOBuf* source) { - constexpr size_t HDR_LEN = HELLO_MAGIC_LEN + HELLO_V3_PB_SIZE_LEN; - if (source->size()< HDR_LEN) { - return 0; - } - - uint8_t hdr[HDR_LEN]; - CHECK_EQ(source->copy_to(hdr, sizeof(hdr)), sizeof(hdr)); - - uint32_t pb_size = butil::NetToHost32( - *reinterpret_cast(hdr + HELLO_MAGIC_LEN)); - if (pb_size == 0 || pb_size > HELLO_V3_MAX_PB_SIZE) { - return -1; - } - - const size_t total = HDR_LEN + pb_size; - if (source->size() < total) { - return 0; - } - - CHECK_EQ(source->pop_front(total), total); - return 1; -} - -// Reply an un-negotiable hello so the client downgrades to TCP. All body fields -// are zero/invalid so the client's validity check rejects it. -// Returns 0 on success, -1 otherwise. -static int SendUnnegotiableHello(Socket* socket, int version) { - butil::IOBuf packet; - if (version == 2) { - // magic "RDMA" + msg_len(=40) + invalid hello_ver; rest stays zero. - // NOTE: msg_len is the length of the WHOLE hello INCLUDING the 4B magic - // (== HELLO_V2_MSG_LEN_MIN), not just the body; the client rejects any - // msg_len < HELLO_V2_MSG_LEN_MIN as a protocol error. - packet.append(HELLO_MAGIC, HELLO_MAGIC_LEN); - char reply[HELLO_V2_MSG_LEN_MIN - HELLO_MAGIC_LEN]{}; - butil::RawPacker(reply).pack16(HELLO_V2_MSG_LEN_MIN) - .pack16(V2_HELLO_VERSION_INVALID); - packet.append(reply, sizeof(reply)); - } else { - // "RDM3" + pb_size + RdmaHello with block_size==0 & qp_num==0 so the - // client's ValidRdmaHello() returns false. - RdmaHello reply_msg; - reply_msg.set_block_size(0); - reply_msg.set_sq_size(0); - reply_msg.set_rq_size(0); - reply_msg.set_lid(0); - reply_msg.set_gid(std::string(V3_GID_LEN, '\0')); - reply_msg.set_qp_num(0); - packet.append(HELLO_MAGIC_V3, HELLO_MAGIC_LEN); - uint32_t pb_size_be = - butil::HostToNet32(static_cast(reply_msg.ByteSizeLong())); - packet.append(&pb_size_be, sizeof(pb_size_be)); - butil::IOBufAsZeroCopyOutputStream output(&packet); - if (!reply_msg.SerializeToZeroCopyStream(&output)) { - LOG(WARNING) << "Fail to serialize RDMA v3 fallback hello"; - return -1; - } - } - - if (socket->Write(&packet) != 0) { - PLOG(WARNING) << "Fail to send RDMA fallback hello to " << socket->description(); - return -1; - } - return 0; -} - -// Fallback handshake for connections that are NOT in RDMA mode. -// -// Unlike the RDMA-mode path, which turns handshake bytes away once its endpoint -// has left the handshake (the state >= ESTABLISHED guard in -// RdmaEndpoint::ExecuteServerHandshake), this one keeps no record of having -// run. See the tail of phase 2. -static ParseResult FallbackServerHandshake(butil::IOBuf* source, Socket* socket) { - if (socket->parsing_context() == nullptr) { - if (source->size() < HELLO_MAGIC_LEN) { - return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); - } - // Phase 1: consume the client hello and reply an un-negotiable hello. - uint8_t magic[HELLO_MAGIC_LEN]; - CHECK_EQ(source->copy_to(magic, HELLO_MAGIC_LEN), HELLO_MAGIC_LEN); - - int version; - if (memcmp(magic, HELLO_MAGIC, HELLO_MAGIC_LEN) == 0) { - version = 2; - } else if (memcmp(magic, HELLO_MAGIC_V3, HELLO_MAGIC_LEN) == 0) { - version = 3; - } else { - return MakeParseError(PARSE_ERROR_TRY_OTHERS); - } - - const int r = version == 2 ? DrainClientHelloV2(source) : DrainClientHelloV3(source); - if (r == 0) { - // Hello not complete yet; keep the buffer intact and retry later. - return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); - } - if (r < 0) { - return MakeParseError(PARSE_ERROR_ABSOLUTELY_WRONG); - } - if (SendUnnegotiableHello(socket, version) < 0) { - return MakeParseError(PARSE_ERROR_ABSOLUTELY_WRONG); - } - // Wait for the client ACK across subsequent reads. - socket->reset_parsing_context(ServerHandshakeContext::Create()); - return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); - } - - // Phase 2: drain the 4B ACK. - if (source->size() < HELLO_ACK_LEN) { - return MakeParseError(PARSE_ERROR_NOT_ENOUGH_DATA); - } - CHECK_EQ(source->pop_front(HELLO_ACK_LEN), HELLO_ACK_LEN); - // Handshake done. - // Drop the context and let InputMessenger parse the following real RPC. - socket->reset_parsing_context(nullptr); - return MakeParseError(PARSE_ERROR_TRY_OTHERS); -} - -ParseResult ExecuteServerHandshake(butil::IOBuf* source, Socket* socket) { - // Only RDMA-mode connections carry a live RdmaEndpoint and run the real - // handshake. A connection that is not in RDMA mode (RDMA compiled in but - // this connection is plain TCP, or RDMA not compiled at all) falls back. -#if BRPC_WITH_RDMA - if (socket->socket_mode() == SOCKET_MODE_RDMA) { - return RdmaEndpoint::ExecuteServerHandshake(source, socket); - } -#endif - return FallbackServerHandshake(source, socket); -} - -} // namespace rdma -} // namespace brpc diff --git a/src/brpc/rdma/rdma_handshake_server.h b/src/brpc/rdma/rdma_handshake_server.h deleted file mode 100644 index 705aa893ef..0000000000 --- a/src/brpc/rdma/rdma_handshake_server.h +++ /dev/null @@ -1,48 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, -// software distributed under the License is distributed on an -// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -// KIND, either express or implied. See the License for the -// specific language governing permissions and limitations -// under the License. - -#ifndef BRPC_RDMA_RDMA_HANDSHAKE_SERVER_H -#define BRPC_RDMA_RDMA_HANDSHAKE_SERVER_H - -#include "brpc/destroyable.h" -#include "brpc/parse_result.h" - -namespace butil { -class IOBuf; -} - -namespace brpc { -class Socket; -namespace rdma { - -// State kept across multiple parse calls of the server handshake. -struct ServerHandshakeContext : public Destroyable { - static ServerHandshakeContext* Create(); - void Destroy() override; -}; - -// The single server-side RDMA handshake entry for policy::ParseRdmaHandshake. -// Returns a ParseResult ready to be handed back from the protocol parser: -// - not an RDMA magic / handshake finished -> PARSE_ERROR_TRY_OTHERS; -// - an RDMA magic but not enough bytes yet -> PARSE_ERROR_NOT_ENOUGH_DATA; -// - IO/protocol error -> PARSE_ERROR_ABSOLUTELY_WRONG. -ParseResult ExecuteServerHandshake(butil::IOBuf* source, Socket* socket); - -} // namespace rdma -} // namespace brpc - -#endif // BRPC_RDMA_RDMA_HANDSHAKE_SERVER_H diff --git a/src/brpc/rdma/rdma_handshake.proto b/src/brpc/rdma_handshake.proto similarity index 100% rename from src/brpc/rdma/rdma_handshake.proto rename to src/brpc/rdma_handshake.proto diff --git a/src/brpc/rdma_transport.cpp b/src/brpc/rdma_transport.cpp index a2037aa725..c7933eeb3c 100644 --- a/src/brpc/rdma_transport.cpp +++ b/src/brpc/rdma_transport.cpp @@ -18,9 +18,8 @@ #if BRPC_WITH_RDMA #include "brpc/rdma_transport.h" +#include "brpc/adapter_transport.h" #include "brpc/event_dispatcher.h" -#include "brpc/tcp_transport.h" -#include "brpc/input_messenger.h" #include "brpc/rdma/rdma_endpoint.h" #include "brpc/rdma/rdma_helper.h" @@ -30,27 +29,26 @@ DECLARE_bool(usercode_in_pthread); extern SocketVarsCollector *g_vars; +RdmaTransport *RdmaTransport::Get(const Socket *socket) { + const AdapterTransport *adapter = AdapterTransport::Get(socket); + Transport *transport = adapter->high_speed_transport(); + CHECK(transport != NULL); + return static_cast(transport); +} + void RdmaTransport::Init(Socket *socket, const SocketOptions &options) { CHECK(_rdma_ep == nullptr); - if (options.socket_mode == SOCKET_MODE_RDMA) { - _rdma_ep = new rdma::RdmaEndpoint(socket); - _rdma_state = RDMA_UNKNOWN; - } else { - _rdma_state = RDMA_OFF; - socket->_socket_mode = SOCKET_MODE_TCP; - } _socket = socket; _default_connect = options.app_connect; - _on_edge_trigger = options.on_edge_triggered_events; - if (options.need_on_edge_trigger && _on_edge_trigger == nullptr) { - if (_rdma_ep != nullptr) { - _on_edge_trigger = rdma::RdmaEndpoint::OnNewDataFromTcp; - } else { - _on_edge_trigger = InputMessenger::OnNewMessages; - } + _on_edge_trigger = nullptr; + _rdma_state = RDMA_UNKNOWN; + _rdma_ep = new (std::nothrow) rdma::RdmaEndpoint(socket); + if (!_rdma_ep) { + const int saved_errno = errno != 0 ? errno : ENOMEM; + errno = saved_errno; + PLOG(WARNING) << "Fail to create RdmaEndpoint, disable RDMA upgrade"; + _rdma_state = RDMA_OFF; } - _tcp_transport = std::make_shared(); - _tcp_transport->Init(socket, options); } void RdmaTransport::Release() { @@ -70,31 +68,56 @@ int RdmaTransport::Reset(int32_t expected_nref) { } std::shared_ptr RdmaTransport::Connect() { - if (_default_connect == nullptr) { - return std::make_shared(); - } return _default_connect; } -int RdmaTransport::CutFromIOBuf(butil::IOBuf *buf) { - // Only send over the RDMA channel once the handshake has NEGOTIATED it - // (RDMA_ON). While the state is still RDMA_UNKNOWN (handshake in progress, - // or a server connection that turned out to be plain TCP and never - // handshook) or RDMA_OFF (fell back), the QP is not usable and everything - // must go over the TCP fd. Mirrors the RDMA_ON check in WaitEpollOut(). - if (_rdma_ep && _rdma_state == RDMA_ON) { - butil::IOBuf *data_arr[1] = {buf}; - return _rdma_ep->CutFromIOBufList(data_arr, 1); - } else { - return _tcp_transport->CutFromIOBuf(buf); +void RdmaTransport::SetHighSpeedAvailable(bool available) { + _rdma_state = available ? RDMA_ON : RDMA_OFF; +} + +int RdmaTransport::PrepareUpgradeResources() { + return _rdma_ep->AllocateResources(); +} + +int RdmaTransport::NegotiateUpgradeResources( + const rdma::RdmaConnectionInfo &remote, bool server) { + _rdma_ep->ApplyRemoteInfo(remote); + return _rdma_ep->BringUpQp(remote, server); +} + +int RdmaTransport::StartUpgradeEvents() { + return _rdma_ep->StartCqEvents(); +} + +std::unique_ptr +RdmaTransport::CreateClientHandshakeAdapter() { + return rdma::CreateClientHandshakeAdapter(_rdma_ep); +} + +std::vector> +RdmaTransport::CreateServerHandshakeAdapters() { + return rdma::CreateServerHandshakeAdapters(_rdma_ep); +} + +void RdmaTransport::ActivateUpgrade() { + SetHighSpeedAvailable(true); +} + +void RdmaTransport::DeactivateUpgrade() { + SetHighSpeedAvailable(false); + if (_rdma_ep != nullptr) { + _rdma_ep->Reset(); } } +int RdmaTransport::CutFromIOBuf(butil::IOBuf *buf) { + butil::IOBuf *data[1] = {buf}; + return static_cast(CutFromIOBufList(data, 1)); +} + ssize_t RdmaTransport::CutFromIOBufList(butil::IOBuf **buf, size_t ndata) { - if (_rdma_ep && _rdma_state == RDMA_ON) { - return _rdma_ep->CutFromIOBufList(buf, ndata); - } - return _tcp_transport->CutFromIOBufList(buf, ndata); + CHECK(_rdma_ep != nullptr); + return _rdma_ep->CutFromIOBufList(buf, ndata); } int RdmaTransport::WaitEpollOut(butil::atomic *_epollout_butex, @@ -108,8 +131,7 @@ int RdmaTransport::WaitEpollOut(butil::atomic *_epollout_butex, if (errno != EAGAIN && errno != ETIMEDOUT) { const int saved_errno = errno; PLOG(WARNING) << "Fail to wait rdma window of " << _socket; - _socket->SetFailed(saved_errno, - "Fail to wait rdma window of %s: %s", + _socket->SetFailed(saved_errno, "Fail to wait rdma window of %s: %s", _socket->description().c_str(), berror(saved_errno)); } @@ -122,8 +144,6 @@ int RdmaTransport::WaitEpollOut(butil::atomic *_epollout_butex, } } } - } else { - return _tcp_transport->WaitEpollOut(_epollout_butex, pollin, duetime); } return 0; } @@ -163,15 +183,16 @@ void RdmaTransport::QueueMessage(InputMessageClosure& input_msg, // TODO(gejun): Join threads. bthread_t th; - bthread_attr_t tmp = (FLAGS_usercode_in_pthread ? - BTHREAD_ATTR_PTHREAD : - BTHREAD_ATTR_NORMAL) | BTHREAD_NOSIGNAL; + bthread_attr_t tmp = + (FLAGS_usercode_in_pthread ? BTHREAD_ATTR_PTHREAD : BTHREAD_ATTR_NORMAL) | + BTHREAD_NOSIGNAL; tmp.keytable_pool = _socket->keytable_pool(); tmp.tag = bthread_self_tag(); bthread_attr_set_name(&tmp, "ProcessInputMessage"); - if (!FLAGS_usercode_in_coroutine && bthread_start_background( - &th, &tmp, ProcessInputMessage, to_run_msg) == 0) { + if (!FLAGS_usercode_in_coroutine && + bthread_start_background(&th, &tmp, ProcessInputMessage, to_run_msg) == + 0) { ++*num_bthread_created; } else { ProcessInputMessage(to_run_msg); @@ -186,15 +207,18 @@ void RdmaTransport::Debug(std::ostream &os) { int RdmaTransport::ContextInitOrDie(bool serverOrNot, const void* _options) { if (serverOrNot) { - if (!OptionsAvailableOverRdma(static_cast(_options))) { + if (!OptionsAvailableOverRdma( + static_cast(_options))) { return -1; } rdma::GlobalRdmaInitializeOrDie(); - if (!rdma::InitPollingModeWithTag(static_cast(_options)->bthread_tag)) { + if (!rdma::InitPollingModeWithTag( + static_cast(_options)->bthread_tag)) { return -1; } } else { - if (!OptionsAvailableForRdma(static_cast(_options))) { + if (!OptionsAvailableForRdma( + static_cast(_options))) { return -1; } rdma::GlobalRdmaInitializeOrDie(); @@ -213,8 +237,7 @@ bool RdmaTransport::OptionsAvailableForRdma(const ChannelOptions* opt) { return false; } if (!rdma::SupportedByRdma(opt->protocol.name())) { - LOG(WARNING) << "Cannot use " << opt->protocol.name() - << " over RDMA"; + LOG(WARNING) << "Cannot use " << opt->protocol.name() << " over RDMA"; return false; } return true; diff --git a/src/brpc/rdma_transport.h b/src/brpc/rdma_transport.h index 1d78fbb430..2aaf5fadab 100644 --- a/src/brpc/rdma_transport.h +++ b/src/brpc/rdma_transport.h @@ -22,14 +22,15 @@ #include "brpc/socket.h" #include "brpc/channel.h" #include "brpc/transport.h" +#include "brpc/rdma/rdma_endpoint.h" +#include "brpc/handshake/rdma_handshake.h" namespace brpc { +class AdapterTransport; class RdmaTransport : public Transport { -friend class TransportFactory; -friend class rdma::RdmaEndpoint; -friend class rdma::RdmaConnect; -friend class rdma::RdmaHandshakeServerV2; -friend class rdma::RdmaHandshakeServerV3; + friend class TransportFactory; + friend class AdapterTransport; + friend class rdma::RdmaEndpoint; public: void Init(Socket* socket, const SocketOptions& options) override; void Release() override; @@ -37,7 +38,8 @@ friend class rdma::RdmaHandshakeServerV3; std::shared_ptr Connect() override; int CutFromIOBuf(butil::IOBuf* buf) override; ssize_t CutFromIOBufList(butil::IOBuf** buf, size_t ndata) override; - int WaitEpollOut(butil::atomic* _epollout_butex, bool pollin, const timespec duetime) override; + int WaitEpollOut(butil::atomic* epollout_butex, + bool pollin, timespec duetime) override; void ProcessEvent(bthread_attr_t attr) override; void QueueMessage(InputMessageClosure& inputMsg, int* num_bthread_created, bool last_msg) override; void Debug(std::ostream &os) override; @@ -45,8 +47,28 @@ friend class rdma::RdmaHandshakeServerV3; CHECK(_rdma_ep != nullptr); return _rdma_ep; } + static RdmaTransport* Get(const Socket* socket); + static RdmaTransport* Get(const SocketUniquePtr& socket) { + return Get(socket.get()); + } static int ContextInitOrDie(bool serverOrNot, const void* _options); + + // Resource operations consumed by the upper-level handshake coordinator. + int PrepareUpgradeResources(); + int StartUpgradeEvents(); + int NegotiateUpgradeResources(const rdma::RdmaConnectionInfo& remote, + bool server); + std::unique_ptr + CreateClientHandshakeAdapter(); + std::vector> + CreateServerHandshakeAdapters(); + void ActivateUpgrade(); + void DeactivateUpgrade(); + bool UpgradeActive() const { return _rdma_state == RDMA_ON; } + bool UpgradeReady() const { return _rdma_ep != nullptr; } private: + void SetHighSpeedAvailable(bool available); + static bool OptionsAvailableForRdma(const ChannelOptions* opt); static bool OptionsAvailableOverRdma(const ServerOptions* opt); @@ -60,7 +82,6 @@ friend class rdma::RdmaHandshakeServerV3; rdma::RdmaEndpoint* _rdma_ep = nullptr; // Should use RDMA or not RdmaState _rdma_state; - std::shared_ptr _tcp_transport; }; } // namespace brpc #endif // BRPC_WITH_RDMA diff --git a/src/brpc/socket.h b/src/brpc/socket.h index da701e9ad7..194b8fd2ab 100644 --- a/src/brpc/socket.h +++ b/src/brpc/socket.h @@ -56,11 +56,6 @@ class ChannelBalancer; } namespace rdma { class RdmaEndpoint; -class RdmaConnect; -class RdmaHandshakeClientV2; -class RdmaHandshakeServerV2; -class RdmaHandshakeClientV3; -class RdmaHandshakeServerV3; } namespace urma { @@ -74,14 +69,16 @@ class UrmaHandshakeServerV3; namespace ubring { class UBShmEndpoint; - class UBConnect; } - +namespace handshake { +class SocketHandshakeIO; +} class Socket; class AuthContext; class EventDispatcher; class Stream; class Transport; +class AdapterTransport; // Set SO_SNDBUF/SO_RCVBUF according to socket_*_buffer_size flags. void SetSocketBufferOptions(int fd); @@ -342,14 +339,8 @@ friend class policy::ConsistentHashingLoadBalancer; friend class policy::RtmpContext; friend class schan::ChannelBalancer; friend class rdma::RdmaEndpoint; -friend class rdma::RdmaConnect; friend class ubring::UBShmEndpoint; -friend class ubring::UBConnect; friend class UBShmTransport; -friend class rdma::RdmaHandshakeClientV2; -friend class rdma::RdmaHandshakeServerV2; -friend class rdma::RdmaHandshakeClientV3; -friend class rdma::RdmaHandshakeServerV3; friend class urma::UrmaEndpoint; friend class urma::UrmaConnect; friend class urma::UrmaHandshakeClientV2; @@ -364,6 +355,8 @@ friend class VersionedRefWithId; friend class IOEvent; friend void DereferenceSocket(Socket*); friend class Transport; +friend class AdapterTransport; +friend class handshake::SocketHandshakeIO; friend class TcpTransport; friend class RdmaTransport; friend class UrmaTransport; @@ -997,8 +990,10 @@ friend class TransportFactory; SSL* _ssl_session; // owner std::shared_ptr _ssl_ctx; - // Should use SOCKET_MODE_RDMA or SOCKET_MODE_TCP or Other, default is SOCKET_MODE_TCP Transport + // Requested provider: SOCKET_MODE_TCP, SOCKET_MODE_RDMA or another mode. SocketMode _socket_mode; + // The top-level AdapterTransport, which selects TCP or the requested + // accelerated Transport. std::unique_ptr _transport; // Pass from controller, for progressive reading. diff --git a/src/brpc/transport_factory.cpp b/src/brpc/transport_factory.cpp index 31bb168801..4c450bca85 100644 --- a/src/brpc/transport_factory.cpp +++ b/src/brpc/transport_factory.cpp @@ -16,6 +16,7 @@ // under the License. #include "brpc/transport_factory.h" +#include "brpc/adapter_transport.h" #include "brpc/rdma_transport.h" #include "brpc/tcp_transport.h" #include "brpc/ubshm_transport.h" @@ -49,11 +50,11 @@ int TransportFactory::ContextInitOrDie( std::unique_ptr TransportFactory::CreateTransport(SocketMode mode) { if (mode == SOCKET_MODE_TCP) { - return std::unique_ptr(new TcpTransport()); + return std::unique_ptr(new AdapterTransport(mode)); } #if BRPC_WITH_RDMA if (mode == SOCKET_MODE_RDMA) { - return std::unique_ptr(new RdmaTransport()); + return std::unique_ptr(new AdapterTransport(mode)); } #endif #if BRPC_WITH_URMA @@ -63,7 +64,7 @@ std::unique_ptr TransportFactory::CreateTransport(SocketMode mode) { #endif #if BRPC_WITH_UBRING if (mode == SOCKET_MODE_UBRING) { - return std::unique_ptr(new UBShmTransport()); + return std::unique_ptr(new AdapterTransport(mode)); } #endif LOG(ERROR) << "Unknown transport type " << mode; diff --git a/src/brpc/transport_factory.h b/src/brpc/transport_factory.h index 84b047daac..0cb6f355b2 100644 --- a/src/brpc/transport_factory.h +++ b/src/brpc/transport_factory.h @@ -22,8 +22,8 @@ #include "brpc/transport.h" namespace brpc { - -// Creates transport instances for a SocketMode. +// Creates AdapterTransport for TCP, RDMA, and UBSHM sockets. URMA currently +// uses its concrete transport directly. class TransportFactory { public: static int ContextInitOrDie(SocketMode mode, bool server_or_not, diff --git a/src/brpc/transport_handshake.cpp b/src/brpc/transport_handshake.cpp new file mode 100644 index 0000000000..c99a2148a6 --- /dev/null +++ b/src/brpc/transport_handshake.cpp @@ -0,0 +1,424 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include "brpc/transport_handshake.h" + +#include +#include + +#include "butil/logging.h" +#include "butil/object_pool.h" +#include "butil/sys_byteorder.h" + +namespace brpc { +namespace handshake { + +namespace { + +const size_t COMMON_ACK_SIZE = sizeof(uint32_t); +const uint32_t COMMON_ACK_OK = 0x1; + +} // namespace + +StepResult HandshakeProtocol::BuildAck(bool enabled, std::string* payload) { + CHECK(payload != NULL); + const uint32_t flags = butil::HostToNet32(enabled ? COMMON_ACK_OK : 0); + payload->assign(reinterpret_cast(&flags), sizeof(flags)); + return STEP_OK; +} + +StepResult HandshakeProtocol::ParseAck(const std::string& payload, + bool* enabled) { + CHECK(enabled != NULL); + if (payload.size() != COMMON_ACK_SIZE) { + errno = EPROTO; + return STEP_ERROR; + } + uint32_t flags = 0; + memcpy(&flags, payload.data(), sizeof(flags)); + *enabled = (butil::NetToHost32(flags) & COMMON_ACK_OK) != 0; + return STEP_OK; +} + +const FrameSpec& HandshakeProtocol::ExtensionFrameSpec() const { + static const FrameSpec empty_spec( + NULL, 0, 0, 0, FrameSpec::FIXED); + return empty_spec; +} + +ServerHandshakeContext* ServerHandshakeContext::Create( + HandshakeAdapter* adapter) { + ServerHandshakeContext* context = + butil::get_object(); + if (context != NULL) { + context->_adapter = adapter; + } + return context; +} + +void ServerHandshakeContext::Destroy() { + _adapter = NULL; + butil::return_object(this); +} + +static StepResult FinishWithFailure(HandshakeSession* session, + HandshakeTransport* transport) { + transport->OnFailed(); + session->MarkFailed(); + return STEP_ERROR; +} + +static StepResult FinishWithFallback(HandshakeSession* session, + HandshakeTransport* transport) { + session->PublishFallback([transport]() { transport->OnFallback(); }); + return STEP_FALLBACK; +} + +static StepResult ConvertFrameResult(FrameResult result) { + switch (result) { + case FRAME_OK: return STEP_OK; + case FRAME_NOT_MINE: return STEP_NOT_MINE; + case FRAME_NEED_MORE: return STEP_NEED_MORE; + case FRAME_IO_ERROR: return STEP_ERROR; + case FRAME_PROTOCOL_ERROR: + errno = EPROTO; + return STEP_ERROR; + } + errno = EPROTO; + return STEP_ERROR; +} + +StepResult HandshakeSession::SendHello(HandshakeProtocol* protocol, + bool enabled) { + std::string payload; + const StepResult result = protocol->BuildHello(enabled, &payload); + if (result != STEP_OK) { + return result; + } + return ConvertFrameResult( + FrameCodec::WriteFrame(_io, protocol->HelloFrameSpec(), payload)); +} + +StepResult HandshakeSession::ReceiveHello(HandshakeProtocol* protocol, + HandshakeInput* input, + bool push_back_on_not_mine, + bool* magic_matched) { + std::string payload; + const FrameResult frame_result = input != NULL + ? FrameCodec::ParseBufferedFrame( + input, protocol->HelloFrameSpec(), &payload, magic_matched) + : FrameCodec::ReadFrame( + _io, protocol->HelloFrameSpec(), push_back_on_not_mine, + &payload); + const StepResult result = ConvertFrameResult(frame_result); + if (result != STEP_OK) { + return result; + } + set_protocol_version(protocol->ProtocolVersion()); + return protocol->ParseHello(payload); +} + +StepResult HandshakeSession::SendAck(HandshakeProtocol* protocol, + bool enabled) { + std::string payload; + const StepResult result = protocol->BuildAck(enabled, &payload); + if (result != STEP_OK) { + return result; + } + return ConvertFrameResult( + FrameCodec::WriteFrame(_io, protocol->AckFrameSpec(), payload)); +} + +StepResult HandshakeSession::ReceiveAck(HandshakeProtocol* protocol, + HandshakeInput* input, + bool* enabled) { + std::string payload; + const FrameResult frame_result = input != NULL + ? FrameCodec::ParseBufferedFrame( + input, protocol->AckFrameSpec(), &payload) + : FrameCodec::ReadFrame( + _io, protocol->AckFrameSpec(), false, &payload); + const StepResult result = ConvertFrameResult(frame_result); + if (result != STEP_OK) { + return result; + } + return protocol->ParseAck(payload, enabled); +} + +StepResult HandshakeSession::SendExtension(HandshakeProtocol* protocol, + bool enabled) { + std::string payload; + const StepResult result = protocol->BuildExtension(enabled, &payload); + if (result != STEP_OK) { + return result; + } + return ConvertFrameResult( + FrameCodec::WriteFrame( + _io, protocol->ExtensionFrameSpec(), payload)); +} + +StepResult HandshakeSession::ReceiveExtension(HandshakeProtocol* protocol, + HandshakeInput* input) { + std::string payload; + const FrameResult frame_result = input != NULL + ? FrameCodec::ParseBufferedFrame( + input, protocol->ExtensionFrameSpec(), &payload) + : FrameCodec::ReadFrame( + _io, protocol->ExtensionFrameSpec(), false, &payload); + const StepResult result = ConvertFrameResult(frame_result); + return result == STEP_OK ? protocol->ParseExtension(payload) : result; +} + +StepResult HandshakeSession::SelectAndReceiveHello( + const std::vector& protocols, HandshakeInput* input, + bool push_back_on_not_mine, HandshakeProtocol** selected) { + CHECK(!protocols.empty()); + CHECK(selected != NULL); + if (input == NULL) { + // A blocking byte stream cannot try a second codec after consuming + // bytes from the fd. Such protocols must select a single codec before + // entering the common session. + CHECK_EQ(1UL, protocols.size()); + *selected = protocols.front(); + return ReceiveHello(*selected, NULL, push_back_on_not_mine); + } + + bool need_more = false; + for (size_t i = 0; i < protocols.size(); ++i) { + bool magic_matched = false; + const StepResult result = ReceiveHello( + protocols[i], input, false, &magic_matched); + if (result == STEP_NOT_MINE) { + continue; + } + if (result == STEP_NEED_MORE) { + if (magic_matched) { + *selected = protocols[i]; + set_protocol_version(protocols[i]->ProtocolVersion()); + return STEP_NEED_MORE; + } + need_more = true; + continue; + } + *selected = protocols[i]; + return result; + } + return need_more ? STEP_NEED_MORE : STEP_NOT_MINE; +} + +StepResult HandshakeSession::RunClient(HandshakeProtocol* protocol, + HandshakeTransport* transport) { + CHECK(protocol != NULL); + CHECK(transport != NULL); + transport->OnProtocolSelected(protocol); + // A client handshake runs once on a potentially reused bthread. Do not + // let an errno left by earlier work override this handshake's result. + errno = 0; + + SetPhase(PREPARING); + StepResult result = transport->PrepareResources(); + if (result == STEP_FALLBACK) { + return FinishWithFallback(this, transport); + } + if (result != STEP_OK) { + return FinishWithFailure(this, transport); + } + + SetPhase(HELLO_SEND); + if (SendHello(protocol, true) != STEP_OK) { + return FinishWithFailure(this, transport); + } + + SetPhase(HELLO_WAIT); + result = ReceiveHello(protocol, NULL, false); + if (result == STEP_NOT_MINE || result == STEP_NEED_MORE) { + errno = EPROTO; + } + if (result == STEP_ERROR || result == STEP_NOT_MINE || + result == STEP_NEED_MORE) { + return FinishWithFailure(this, transport); + } + bool enabled = result == STEP_OK; + + if (enabled && protocol->HasExtension()) { + SetPhase(EXTENSION_SEND); + if (SendExtension(protocol, true) != STEP_OK) { + return FinishWithFailure(this, transport); + } + SetPhase(EXTENSION_WAIT); + result = ReceiveExtension(protocol, NULL); + if (result != STEP_OK && result != STEP_FALLBACK) { + return FinishWithFailure(this, transport); + } + enabled = result == STEP_OK; + } + + if (enabled) { + SetPhase(NEGOTIATING); + result = transport->NegotiateResources(); + if (result != STEP_OK && result != STEP_FALLBACK) { + return FinishWithFailure(this, transport); + } + enabled = result == STEP_OK; + } + + SetPhase(ACK_SEND); + if (SendAck(protocol, enabled) != STEP_OK) { + return FinishWithFailure(this, transport); + } + + if (enabled) { + transport->OnEstablished(); + MarkEstablished(); + return STEP_OK; + } + return FinishWithFallback(this, transport); +} + +StepResult HandshakeSession::RunServer( + const std::vector& protocols, HandshakeInput* input, + HandshakeTransport* transport, bool fallback_on_not_mine) { + CHECK(!protocols.empty()); + CHECK(transport != NULL); + + // Once TCP fallback has been published, subsequent bytes are application + // protocol data and must bypass every upgrade codec without changing the + // terminal state. + if (phase() == FALLBACK_TCP) { + return STEP_NOT_MINE; + } + + HandshakeProtocol* selected = NULL; + if (phase() != ACK_WAIT && phase() != EXTENSION_WAIT) { + const int previous_phase = phase(); + _local_enabled = false; + SetPhase(HELLO_WAIT); + StepResult result = SelectAndReceiveHello( + protocols, input, fallback_on_not_mine, &selected); + if (result == STEP_NOT_MINE) { + if (fallback_on_not_mine) { + return FinishWithFallback(this, transport); + } + SetPhase(UNINITIALIZED); + return STEP_NOT_MINE; + } + if (result == STEP_NEED_MORE) { + if (selected == NULL) { + SetPhase(previous_phase); + } + return STEP_NEED_MORE; + } + if (result == STEP_ERROR) { + return FinishWithFailure(this, transport); + } + CHECK(selected != NULL); + transport->OnProtocolSelected(selected); + if (result == STEP_FALLBACK) { + // Publish the disabled transport state immediately. The server + // still sends a disabled hello and consumes the peer ACK before + // publishing the terminal FALLBACK_TCP phase. + transport->OnFallback(); + } + bool enabled = result == STEP_OK; + + if (enabled) { + SetPhase(PREPARING); + result = transport->PrepareResources(); + if (result != STEP_OK && result != STEP_FALLBACK) { + return FinishWithFailure(this, transport); + } + enabled = result == STEP_OK; + } + + if (enabled) { + SetPhase(NEGOTIATING); + result = transport->NegotiateResources(); + if (result != STEP_OK && result != STEP_FALLBACK) { + return FinishWithFailure(this, transport); + } + enabled = result == STEP_OK; + } + + SetPhase(HELLO_SEND); + _local_enabled = enabled; + if (SendHello(selected, enabled) != STEP_OK) { + return FinishWithFailure(this, transport); + } + SetPhase(enabled && selected->HasExtension() + ? EXTENSION_WAIT + : ACK_WAIT); + } else { + for (size_t i = 0; i < protocols.size(); ++i) { + if (protocols[i]->ProtocolVersion() == protocol_version()) { + selected = protocols[i]; + break; + } + } + CHECK(selected != NULL); + } + + if (phase() == EXTENSION_WAIT) { + StepResult result = ReceiveExtension(selected, input); + if (result == STEP_NEED_MORE) { + return STEP_NEED_MORE; + } + if (result != STEP_OK && result != STEP_FALLBACK) { + return FinishWithFailure(this, transport); + } + if (result == STEP_FALLBACK) { + _local_enabled = false; + } + SetPhase(EXTENSION_SEND); + if (SendExtension(selected, _local_enabled) != STEP_OK) { + return FinishWithFailure(this, transport); + } + SetPhase(ACK_WAIT); + } + + // Always try the ACK callback once. For a non-blocking server it returns + // STEP_NEED_MORE when the ACK has not arrived; when Hello and ACK are + // coalesced in the input buffer this consumes the ACK without waiting for + // another socket edge. + bool peer_enabled = false; + StepResult result = ReceiveAck(selected, input, &peer_enabled); + if (result == STEP_NEED_MORE) { + return STEP_NEED_MORE; + } + if (result == STEP_ERROR || result == STEP_NOT_MINE) { + return FinishWithFailure(this, transport); + } + if (result == STEP_FALLBACK) { + return FinishWithFallback(this, transport); + } + if (!peer_enabled) { + return FinishWithFallback(this, transport); + } + if (!_local_enabled) { + errno = EPROTO; + return FinishWithFailure(this, transport); + } + if (transport->ValidateEstablished() != STEP_OK) { + return FinishWithFailure(this, transport); + } + + transport->OnEstablished(); + MarkEstablished(); + return STEP_OK; +} + +} // namespace handshake +} // namespace brpc diff --git a/src/brpc/transport_handshake.h b/src/brpc/transport_handshake.h new file mode 100644 index 0000000000..b558ee0ac6 --- /dev/null +++ b/src/brpc/transport_handshake.h @@ -0,0 +1,233 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#ifndef BRPC_TRANSPORT_HANDSHAKE_H +#define BRPC_TRANSPORT_HANDSHAKE_H + +#include +#include +#include + +#include "butil/atomicops.h" +#include "butil/macros.h" +#include "brpc/destroyable.h" +#include "brpc/handshake/handshake_frame.h" + +namespace brpc { + +class Socket; + +namespace handshake { + +class HandshakeAdapter; + +// Context retained by InputMessenger between the hello and ACK parse calls. +// Remembering the selected stateless adapter is necessary because ACK frames +// have no magic and cannot be dispatched from their bytes alone. +struct ServerHandshakeContext : public Destroyable { + ServerHandshakeContext() : _adapter(NULL) {} + static ServerHandshakeContext* Create(HandshakeAdapter* adapter); + HandshakeAdapter* adapter() const { return _adapter; } + void Destroy() override; + +private: + HandshakeAdapter* _adapter; +}; + +// Protocol adapters may use transport-specific intermediate values, but the +// terminal values are shared so that AdapterTransport can make the same +// acquire-side decision for RDMA, URMA and UBSHM. +enum Phase { + UNINITIALIZED = 0, + PREPARING = 1, + HELLO_SEND = 2, + HELLO_WAIT = 3, + NEGOTIATING = 4, + ACK_SEND = 5, + ACK_WAIT = 6, + EXTENSION_SEND = 7, + EXTENSION_WAIT = 8, + ESTABLISHED = 0x100, + FALLBACK_TCP = 0x200, + FAILED = 0x300, +}; + +enum StepResult { + STEP_OK = 0, + STEP_FALLBACK, + STEP_NEED_MORE, + STEP_NOT_MINE, + STEP_ERROR, +}; + +// Wire-level participant in a transport upgrade. Implementations own parsed +// protocol state; HandshakeSession owns framing and phase orchestration. +class HandshakeProtocol { +public: + virtual ~HandshakeProtocol() = default; + + virtual int ProtocolVersion() const = 0; + virtual const FrameSpec& HelloFrameSpec() const = 0; + virtual const FrameSpec& AckFrameSpec() const = 0; + virtual StepResult BuildHello(bool enabled, std::string* payload) = 0; + virtual StepResult ParseHello(const std::string& payload) = 0; + + // RDMA and UBSHM use the same four-byte, network-order ACK. Protocols + // with a different ACK format may override these methods. + virtual StepResult BuildAck(bool enabled, std::string* payload); + virtual StepResult ParseAck(const std::string& payload, bool* enabled); + + virtual bool HasExtension() const { return false; } + virtual const FrameSpec& ExtensionFrameSpec() const; + virtual StepResult BuildExtension(bool, std::string*) { + return STEP_ERROR; + } + virtual StepResult ParseExtension(const std::string&) { + return STEP_ERROR; + } +}; + +// Resource-level participant in a transport upgrade. Cleanup remains part of +// this contract: resources may already exist when negotiation falls back or +// a framing/I/O error terminates the handshake. +class HandshakeTransport { +public: + virtual ~HandshakeTransport() = default; + + virtual void OnProtocolSelected(HandshakeProtocol*) {} + virtual StepResult PrepareResources() = 0; + virtual StepResult NegotiateResources() = 0; + virtual void OnEstablished() = 0; + virtual void OnFallback() = 0; + virtual void OnFailed() = 0; + virtual StepResult ValidateEstablished() { return STEP_OK; } +}; + +// Participant for a server that recognizes an upgrade protocol only to +// negotiate TCP fallback. It owns no high-speed resources. +class FallbackHandshakeTransport : public HandshakeTransport { +public: + StepResult PrepareResources() override { return STEP_OK; } + StepResult NegotiateResources() override { return STEP_OK; } + void OnEstablished() override {} + void OnFallback() override {} + void OnFailed() override {} +}; + +class FallbackHandshakeProtocol : public HandshakeProtocol { +public: + StepResult ParseHello(const std::string&) override { + return STEP_FALLBACK; + } + StepResult ParseAck(const std::string& payload, bool* enabled) override { + bool ignored = false; + const StepResult result = HandshakeProtocol::ParseAck( + payload, &ignored); + *enabled = false; + return result; + } +}; + +// Owns one connection-upgrade attempt, invokes the protocol field codec and +// resource callbacks, and provides common framing, TCP control-plane I/O, +// lifecycle and publication ordering. +class HandshakeSession { +public: + explicit HandshakeSession(Socket* socket = NULL) + : _socket_io(socket), _io(&_socket_io), _phase(UNINITIALIZED), + _protocol_version(0), _local_enabled(false) {} + + void Reset(Socket* socket) { + _socket_io.Reset(socket); + _io = &_socket_io; + _protocol_version = 0; + _local_enabled = false; + _phase.store(UNINITIALIZED, butil::memory_order_relaxed); + } + + int phase(butil::memory_order order = butil::memory_order_acquire) const { + return _phase.load(order); + } + + void SetPhase(int phase) { + _phase.store(phase, butil::memory_order_release); + } + + int protocol_version() const { return _protocol_version; } + void set_protocol_version(int version) { _protocol_version = version; } + + void MarkEstablished() { + _phase.store(ESTABLISHED, butil::memory_order_release); + } + + void MarkFailed() { + _phase.store(FAILED, butil::memory_order_release); + } + + // The callback MUST publish the transport's TCP-active state. The release + // store then makes that state and any pushed-back bytes visible to the + // event thread that observes FALLBACK_TCP with an acquire load. This is + // the common form of the ordering fixes from #3347 and #3406. + template + void PublishFallback(PublishTcpActive publish_tcp_active) { + publish_tcp_active(); + _phase.store(FALLBACK_TCP, butil::memory_order_release); + } + + void NotifyReadable() { _socket_io.NotifyReadable(); } + + // Injects an in-memory stream in common-component unit tests. Reset() + // restores the Socket-backed implementation. + void SetIOForTest(HandshakeIO* io) { _io = io; } + + StepResult RunClient(HandshakeProtocol* protocol, + HandshakeTransport* transport); + StepResult RunServer(const std::vector& protocols, + HandshakeInput* input, + HandshakeTransport* transport, + bool fallback_on_not_mine); + +private: + StepResult SendHello(HandshakeProtocol* protocol, bool enabled); + StepResult ReceiveHello(HandshakeProtocol* protocol, + HandshakeInput* input, + bool push_back_on_not_mine, + bool* magic_matched = NULL); + StepResult SendAck(HandshakeProtocol* protocol, bool enabled); + StepResult ReceiveAck(HandshakeProtocol* protocol, + HandshakeInput* input, bool* enabled); + StepResult SendExtension(HandshakeProtocol* protocol, bool enabled); + StepResult ReceiveExtension(HandshakeProtocol* protocol, + HandshakeInput* input); + StepResult SelectAndReceiveHello( + const std::vector& protocols, + HandshakeInput* input, bool push_back_on_not_mine, + HandshakeProtocol** selected); + + SocketHandshakeIO _socket_io; + HandshakeIO* _io; + butil::atomic _phase; + int _protocol_version; + bool _local_enabled; + + DISALLOW_COPY_AND_ASSIGN(HandshakeSession); +}; + +} // namespace handshake +} // namespace brpc + +#endif // BRPC_TRANSPORT_HANDSHAKE_H diff --git a/src/brpc/ubshm/ub_endpoint.cpp b/src/brpc/ubshm/ub_endpoint.cpp index 44685eb302..3835376646 100644 --- a/src/brpc/ubshm/ub_endpoint.cpp +++ b/src/brpc/ubshm/ub_endpoint.cpp @@ -19,23 +19,19 @@ #include -#include -#include -#include "butil/fd_utility.h" -#include "butil/logging.h" // CHECK, LOG -#include "butil/sys_byteorder.h" // HostToNet,NetToHost -#include "bthread/bthread.h" #include "brpc/errno.pb.h" #include "brpc/event_dispatcher.h" #include "brpc/input_messenger.h" #include "brpc/socket.h" -#include "brpc/reloadable_flags.h" -#include "brpc/ubshm/ub_helper.h" -#include "brpc/ubshm/ub_endpoint.h" -#include "brpc/ubshm/shm/shm_def.h" #include "brpc/ubshm/common/common.h" -#include "brpc/ubshm_transport.h" +#include "brpc/ubshm/shm/shm_def.h" +#include "brpc/ubshm/ub_endpoint.h" +#include "brpc/ubshm/ub_helper.h" #include "brpc/ubshm/ubr_trx.h" +#include "brpc/ubshm_transport.h" +#include "bthread/bthread.h" +#include "butil/logging.h" // CHECK, LOG +#include DECLARE_int32(task_group_ntags); @@ -44,114 +40,24 @@ DECLARE_bool(log_connection_close); namespace ubring { extern bool g_skip_ub_init; -DEFINE_int32(data_queue_size, 4, "data queue size for UB"); -DEFINE_bool(ub_trace_verbose, false, "Print log message verbosely"); -BRPC_VALIDATE_GFLAG(ub_trace_verbose, brpc::PassValidate); DEFINE_int32(ub_poller_num, 1, "Poller number in ub polling mode."); DEFINE_bool(ub_poller_yield, false, "Yield thread in UBRing polling mode."); DEFINE_bool(ub_edisp_unsched, false, "Disable event dispatcher schedule"); DEFINE_bool(ub_disable_bthread, false, "Disable bthread in UBRing polling mode."); +static const size_t MIN_ONCE_READ = 4096; +static const size_t MAX_ONCE_READ = 524288; static const size_t IOBUF_IOV_MAX = 256; -static const char* MAGIC_STR = "UB"; -static const size_t MAGIC_STR_LEN = 2; -static const size_t HELLO_MSG_LEN_MIN = 64; -static const size_t ACK_MSG_LEN = 4; -static uint16_t g_ub_hello_msg_len = 64; -static uint16_t g_ub_hello_version = 3; -static uint16_t g_ub_impl_version = 1; - -static const uint32_t ACK_MSG_UB_OK = 0x1; - -static butil::Mutex* g_ubring_resource_mutex = nullptr; - -void HelloFormatExtension::Serialize(void* data) const { - char* current_pos = static_cast(data); - const uint16_t net_extension_len = butil::HostToNet16(extension_len); - memcpy(current_pos, &net_extension_len, sizeof(net_extension_len)); - current_pos += sizeof(net_extension_len); - const uint16_t net_format_id = butil::HostToNet16(format_id); - memcpy(current_pos, &net_format_id, sizeof(net_format_id)); -} - -void HelloFormatExtension::Deserialize(const void* data) { - const char* current_pos = static_cast(data); - uint16_t net_extension_len; - memcpy(&net_extension_len, current_pos, sizeof(net_extension_len)); - extension_len = butil::NetToHost16(net_extension_len); - current_pos += sizeof(net_extension_len); - uint16_t net_format_id; - memcpy(&net_format_id, current_pos, sizeof(net_format_id)); - format_id = butil::NetToHost16(net_format_id); -} - -void HelloMessage::Serialize(void* data) const { - char* current_pos = static_cast(data); - const uint16_t net_msg_len = butil::HostToNet16(msg_len); - memcpy(current_pos, &net_msg_len, sizeof(net_msg_len)); - current_pos += sizeof(net_msg_len); - const uint16_t net_hello_ver = butil::HostToNet16(hello_ver); - memcpy(current_pos, &net_hello_ver, sizeof(net_hello_ver)); - current_pos += sizeof(net_hello_ver); - const uint16_t net_impl_ver = butil::HostToNet16(impl_ver); - memcpy(current_pos, &net_impl_ver, sizeof(net_impl_ver)); - current_pos += sizeof(net_impl_ver); - const uint64_t net_len = butil::HostToNet64(len); - memcpy(current_pos, &net_len, sizeof(net_len)); - current_pos += sizeof(net_len); - memcpy(current_pos, shm_name, SHM_MAX_NAME_BUFF_LEN); -} +static butil::Mutex *g_ubring_resource_mutex = NULL; -void HelloMessage::Deserialize(void* data) { - char* current_pos = static_cast(data); - uint16_t net_msg_len; - memcpy(&net_msg_len, current_pos, sizeof(net_msg_len)); - msg_len = butil::NetToHost16(net_msg_len); - current_pos += sizeof(net_msg_len); - uint16_t net_hello_ver; - memcpy(&net_hello_ver, current_pos, sizeof(net_hello_ver)); - hello_ver = butil::NetToHost16(net_hello_ver); - current_pos += sizeof(net_hello_ver); - uint16_t net_impl_ver; - memcpy(&net_impl_ver, current_pos, sizeof(net_impl_ver)); - impl_ver = butil::NetToHost16(net_impl_ver); - current_pos += sizeof(net_impl_ver); - uint64_t net_len; - memcpy(&net_len, current_pos, sizeof(net_len)); - len = butil::NetToHost64(net_len); - current_pos += sizeof(net_len); - memcpy(shm_name, current_pos, SHM_MAX_NAME_BUFF_LEN); -} - -std::string HelloMessage::toString() const { - constexpr size_t MAX_LEN = 16 + 6 + 16 + 6 + 16 + 6 + 20 + 6 + SHM_MAX_NAME_BUFF_LEN + 32; - std::array buf; - int n = snprintf(buf.data(), buf.size(), - "msg_len=%u, hello_ver=%u, impl_ver=%u, len=%lu, shm_name=%.*s", - msg_len, - hello_ver, - impl_ver, - static_cast(len), // compatible with 32/64-bit - static_cast(SHM_MAX_NAME_BUFF_LEN), // limit max output length - shm_name - ); - return std::string(buf.data(), static_cast(n)); -} - -UBShmEndpoint::UBShmEndpoint(Socket* s) - : _socket(s) - , _socket_id(s ? s->id() : INVALID_SOCKET_ID) - , _state(UNINIT) - , _ub_ring(nullptr) - , _poller_sid(INVALID_SOCKET_ID) -{ - _read_butex = bthread::butex_create_checked>(); +UBShmEndpoint::UBShmEndpoint(Socket *s) + : _socket(s), _socket_id(s ? s->id() : INVALID_SOCKET_ID), + _ub_ring(nullptr), _poller_sid(INVALID_SOCKET_ID) { } UBShmEndpoint::~UBShmEndpoint() { Reset(); - bthread::butex_destroy(_read_butex); } void UBShmEndpoint::Reset() { @@ -161,541 +67,6 @@ void UBShmEndpoint::Reset() { _ub_ring = nullptr; _poller_sid = INVALID_SOCKET_ID; _negotiated_data_format = UBR_DATA_FORMAT_NONE; - _state = UNINIT; -} - -void UBConnect::StartConnect(const Socket* socket, - void (*done)(int err, void* data), - void* data) { - auto* ub_transport = static_cast(socket->_transport.get()); - CHECK(ub_transport->_ub_ep != nullptr); - SocketUniquePtr s; - if (Socket::Address(socket->id(), &s) != 0) { - return; - } - if (!IsUBAvailable()) { - ub_transport->_ub_ep->_state = UBShmEndpoint::FALLBACK_TCP; - ub_transport->_ub_state = UBShmTransport::UB_OFF; - done(0, data); - return; - } - _done = done; - _data = data; - bthread_t tid; - bthread_attr_t attr = BTHREAD_ATTR_NORMAL; - bthread_attr_set_name(&attr, "UBProcessHandshakeAtClient"); - if (bthread_start_background(&tid, &attr, - UBShmEndpoint::ProcessHandshakeAtClient, ub_transport->_ub_ep) < 0) { - LOG(FATAL) << "Fail to start handshake bthread"; - Run(); - } else { - s.release(); - } -} - -void UBConnect::StopConnect(Socket* socket) { } - -void UBConnect::Run() { - _done(errno, _data); -} - -static void TryReadOnTcpDuringRdmaEst(Socket* s) { - int progress = Socket::PROGRESS_INIT; - while (true) { - uint8_t tmp; - ssize_t nr = read(s->fd(), &tmp, 1); - if (nr < 0) { - if (errno != EAGAIN) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to read from " << s; - s->SetFailed(saved_errno, "Fail to read from %s: %s", - s->description().c_str(), berror(saved_errno)); - return; - } - if (!s->MoreReadEvents(&progress)) { - break; - } - } else if (nr == 0) { - s->SetEOF(); - return; - } else { - LOG(WARNING) << "Read unexpected data from " << s; - s->SetFailed(EPROTO, "Read unexpected data from %s", - s->description().c_str()); - return; - } - } -} - -void UBShmEndpoint::OnNewDataFromTcp(Socket* m) { - auto* ub_transport = static_cast(m->_transport.get()); - UBShmEndpoint* ep = ub_transport->GetUBShmEp(); - CHECK(ep != nullptr); - - int progress = Socket::PROGRESS_INIT; - while (true) { - if (ep->_state == UNINIT) { - if (!m->CreatedByConnect()) { - if (!IsUBAvailable()) { - ep->_state = FALLBACK_TCP; - ub_transport->_ub_state = UBShmTransport::UB_OFF; - continue; - } - bthread_t tid; - ep->_state = S_HELLO_WAIT; - SocketUniquePtr s; - m->ReAddress(&s); - bthread_attr_t attr = BTHREAD_ATTR_NORMAL; - bthread_attr_set_name(&attr, "UBProcessHandshakeAtServer"); - if (bthread_start_background(&tid, &attr, - ProcessHandshakeAtServer, ep) < 0) { - ep->_state = UNINIT; - LOG(FATAL) << "Fail to start handshake bthread"; - } else { - s.release(); - } - } else { - // The connection may be closed or reset before the client - // starts handshake. This will be handled by client handshake. - // Ignore the exception here. - } - } else if (ep->_state < ESTABLISHED) { // during handshake - ep->_read_butex->fetch_add(1, butil::memory_order_release); - bthread::butex_wake(ep->_read_butex); - } else if (ep->_state == FALLBACK_TCP){ // handshake finishes - InputMessenger::OnNewMessages(m); - return; - } else if (ep->_state == ESTABLISHED) { - TryReadOnTcpDuringRdmaEst(ep->_socket); - return; - } - if (!m->MoreReadEvents(&progress)) { - break; - } - } -} -bool HelloNegotiationValid(HelloMessage& msg) { - if (msg.hello_ver == g_ub_hello_version && - msg.impl_ver == g_ub_impl_version) { - // This can be modified for future compatibility - return true; - } - return false; -} - -static const int WAIT_TIMEOUT_MS = 50; - -int UBShmEndpoint::ReadFromFd(void* data, size_t len) { - CHECK(data != nullptr); - int nr = 0; - size_t received = 0; - do { - const timespec duetime = butil::milliseconds_from_now(WAIT_TIMEOUT_MS); - nr = read(_socket->fd(), (uint8_t*)data + received, len - received); - if (nr < 0) { - if (errno == EAGAIN) { - const int expected_val = _read_butex->load(butil::memory_order_acquire); - if (bthread::butex_wait(_read_butex, expected_val, &duetime) < 0) { - if (errno != EWOULDBLOCK && errno != ETIMEDOUT) { - return -1; - } - } - } else { - return -1; - } - } else if (nr == 0) { - errno = EEOF; - return -1; - } else { - received += nr; - } - } while (received < len); - return 0; -} - -int UBShmEndpoint::WriteToFd(void* data, size_t len) { - CHECK(data != nullptr); - int nw = 0; - size_t written = 0; - do { - const timespec duetime = butil::milliseconds_from_now(WAIT_TIMEOUT_MS); - nw = write(_socket->fd(), (uint8_t*)data + written, len - written); - if (nw < 0) { - if (errno == EAGAIN) { - if (_socket->WaitEpollOut(_socket->fd(), true, &duetime) < 0) { - if (errno != ETIMEDOUT) { - return -1; - } - } - } else { - return -1; - } - } else { - written += nw; - } - } while (written < len); - return 0; -} - -inline void UBShmEndpoint::TryReadOnTcp() { - if (_socket->_nevent.fetch_add(1, butil::memory_order_acq_rel) == 0) { - if (_state == FALLBACK_TCP) { - InputMessenger::OnNewMessages(_socket); - } else if (_state == ESTABLISHED) { - TryReadOnTcpDuringRdmaEst(_socket); - } - } -} - -void* UBShmEndpoint::ProcessHandshakeAtClient(void* arg) { - UBShmEndpoint* ep = static_cast(arg); - ep->_negotiated_data_format = UBR_DATA_FORMAT_NONE; - SocketUniquePtr s(ep->_socket); - UBConnect::RunGuard rg((UBConnect*)s->_app_connect.get()); - - LOG_IF(INFO, FLAGS_ub_trace_verbose) - << "Start handshake on " << s->_local_side; - - uint8_t data[g_ub_hello_msg_len]; - - ep->_state = C_ALLOC_SHM; - auto* ub_transport = static_cast(s->_transport.get()); - size_t local_shm_len = (size_t)(FLAGS_data_queue_size) * MB_TO_BYTE; - SHM local_trx_shm = {nullptr, local_shm_len, 0, {0}, (uint32_t)s->fd()}; - auto shm_name_str = butil::endpoint2str(s->local_side()); - const char* shm_name = shm_name_str.c_str(); - if (ep->AllocateClientResources(&local_trx_shm, shm_name) < 0) { - LOG(WARNING) << "Fallback to tcp:" << s->description(); - ub_transport->_ub_state = UBShmTransport::UB_OFF; - ep->_state = FALLBACK_TCP; - return nullptr; - } - - ep->_state = C_HELLO_SEND; - HelloMessage local_msg{}; - local_msg.msg_len = g_ub_hello_msg_len; - local_msg.hello_ver = g_ub_hello_version; - local_msg.impl_ver = g_ub_impl_version; - local_msg.len = local_shm_len; - memcpy(local_msg.shm_name, local_trx_shm.name, SHM_MAX_NAME_BUFF_LEN); - memcpy(data, MAGIC_STR, MAGIC_STR_LEN); - local_msg.Serialize((char*)data + MAGIC_STR_LEN); - if (ep->WriteToFd(data, g_ub_hello_msg_len) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to send hello message to server:" << s->description(); - s->SetFailed(saved_errno, "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - LOG_IF(INFO, FLAGS_ub_trace_verbose) << "client handshake message : " << local_msg.toString(); - - ep->_state = C_HELLO_WAIT; - if (ep->ReadFromFd(data, MAGIC_STR_LEN) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to get hello message from server:" << s->description(); - s->SetFailed(saved_errno, "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - if (memcmp(data, MAGIC_STR, MAGIC_STR_LEN) != 0) { - LOG(WARNING) << "Read unexpected data during handshake:" << s->description(); - s->SetFailed(EPROTO, "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(EPROTO)); - ep->_state = FAILED; - return nullptr; - } - - if (ep->ReadFromFd(data, HELLO_MSG_LEN_MIN - MAGIC_STR_LEN) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to get Hello Message from server:" << s->description(); - s->SetFailed(saved_errno, "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - HelloMessage remote_msg; - remote_msg.Deserialize(data); - if (remote_msg.msg_len != HELLO_MSG_LEN_MIN) { - LOG(WARNING) << "Fail to parse Hello Message length from server:" - << s->description(); - s->SetFailed(EPROTO, "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(EPROTO)); - ep->_state = FAILED; - return nullptr; - } - - UbrDataFormat selected_format = UBR_DATA_FORMAT_NONE; - if (!HelloNegotiationValid(remote_msg)) { - LOG(WARNING) << "Fail to negotiate with server, fallback to tcp:" - << s->description(); - ub_transport->_ub_state = UBShmTransport::UB_OFF; - } else { - HelloFormatExtension local_extension = { - HelloFormatExtension::WIRE_SIZE, UBR_DATA_FORMAT_LEGACY_64}; - local_extension.Serialize(data); - ep->_state = C_FORMAT_SEND; - if (ep->WriteToFd(data, HelloFormatExtension::WIRE_SIZE) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to send format extension to server:" - << s->description(); - s->SetFailed(saved_errno, - "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - - ep->_state = C_FORMAT_WAIT; - if (ep->ReadFromFd(data, HelloFormatExtension::WIRE_SIZE) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to read format extension from server:" - << s->description(); - s->SetFailed(saved_errno, - "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - HelloFormatExtension remote_extension; - remote_extension.Deserialize(data); - if (remote_extension.extension_len != HelloFormatExtension::WIRE_SIZE || - remote_extension.format_id == UBR_DATA_FORMAT_NONE || - remote_extension.format_id != local_extension.format_id) { - LOG(WARNING) << "Fail to negotiate data format with server, " - << "fallback to tcp:" << s->description(); - ub_transport->_ub_state = UBShmTransport::UB_OFF; - } else { - selected_format = UBR_DATA_FORMAT_LEGACY_64; - ep->_state = C_MAP_REMOTE_SHM; - if (ep->_ub_ring->UbrMapRemoteShm(&local_trx_shm, shm_name) < 0) { - LOG(WARNING) << "Fail to map the remote shm, fallback to tcp:" - << s->description(); - ub_transport->_ub_state = UBShmTransport::UB_OFF; - } else { - ub_transport->_ub_state = UBShmTransport::UB_ON; - } - } - } - - ep->_state = C_ACK_SEND; - uint32_t flags = 0; - if (ub_transport->_ub_state != UBShmTransport::UB_OFF) { - flags |= ACK_MSG_UB_OK; - } - const uint32_t net_flags = butil::HostToNet32(flags); - memcpy(data, &net_flags, sizeof(net_flags)); - if (ep->WriteToFd(data, ACK_MSG_LEN) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to send Ack Message to server:" << s->description(); - s->SetFailed(saved_errno, "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - - if (ub_transport->_ub_state == UBShmTransport::UB_ON) { - ep->_negotiated_data_format = selected_format; - ep->_state = ESTABLISHED; - ep->_ub_ring->UbrUnlinkLocalShm(); - LOG_IF(INFO, FLAGS_ub_trace_verbose) - << "Client handshake ends (use ubring) on " << s->description(); - } else { - ep->_state = FALLBACK_TCP; - LOG_IF(INFO, FLAGS_ub_trace_verbose) - << "Client handshake ends (use tcp) on " << s->description(); - } - - errno = 0; - - return nullptr; -} - -void* UBShmEndpoint::ProcessHandshakeAtServer(void* arg) { - UBShmEndpoint* ep = static_cast(arg); - ep->_negotiated_data_format = UBR_DATA_FORMAT_NONE; - SocketUniquePtr s(ep->_socket); - - LOG_IF(INFO, FLAGS_ub_trace_verbose) - << "Start handshake on " << s->description(); - - uint8_t data[g_ub_hello_msg_len]; - - ep->_state = S_HELLO_WAIT; - if (ep->ReadFromFd(data, MAGIC_STR_LEN) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to read Hello Message from client:" << s->description() << " " << s->_remote_side; - s->SetFailed(saved_errno, "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - auto* ub_transport = static_cast(s->_transport.get()); - if (memcmp(data, MAGIC_STR, MAGIC_STR_LEN) != 0) { - LOG_IF(INFO, FLAGS_ub_trace_verbose) << "It seems that the " - << "client does not use RDMA, fallback to TCP:" - << s->description(); - s->fd_input_processor().read_buf().append(data, MAGIC_STR_LEN); - ep->_state = FALLBACK_TCP; - ub_transport->_ub_state = UBShmTransport::UB_OFF; - ep->TryReadOnTcp(); - return nullptr; - } - - if (ep->ReadFromFd(data, g_ub_hello_msg_len - MAGIC_STR_LEN) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to read Hello Message from client:" << s->description(); - s->SetFailed(saved_errno, "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - - HelloMessage remote_msg; - remote_msg.Deserialize(data); - LOG_IF(INFO, FLAGS_ub_trace_verbose) << "server receive handshake message : " << remote_msg.toString(); - if (remote_msg.msg_len != HELLO_MSG_LEN_MIN) { - LOG(WARNING) << "Fail to parse Hello Message length from client:" - << s->description(); - s->SetFailed(EPROTO, "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(EPROTO)); - ep->_state = FAILED; - return nullptr; - } - if (!HelloNegotiationValid(remote_msg)) { - LOG(WARNING) << "Fail to negotiate with client, fallback to tcp:" - << s->description(); - ub_transport->_ub_state = UBShmTransport::UB_OFF; - } else { - ep->_state = S_ALLOC_SHM; - ubring::SHM remote_trx_shm = {nullptr, remote_msg.len, 0, {0}, (uint32_t)ep->_socket->fd()}; - strncpy(remote_trx_shm.name, remote_msg.shm_name, SHM_MAX_NAME_BUFF_LEN); - - size_t local_shm_len = (size_t)(FLAGS_data_queue_size) * MB_TO_BYTE; - // server-side shared memory name - ubring::SHM local_trx_shm = {nullptr, local_shm_len, 0, {0}, (uint32_t)ep->_socket->fd()}; - char client_name[SHM_MAX_NAME_BUFF_LEN]; - strncpy(client_name, remote_msg.shm_name, SHM_MAX_NAME_BUFF_LEN); - - char *client_ip_port = strrchr(client_name, '_'); - if (client_ip_port != nullptr) { - *client_ip_port = '\0'; - } - int result = snprintf(local_trx_shm.name, SHM_MAX_NAME_BUFF_LEN, "%s_%s", - client_name, SERVER_SHM_NAME_SUFFIX); - if (UNLIKELY(result < 0)) { - LOG(WARNING) << "Copy client shared memory name failed, ret=" << result; - ub_transport->_ub_state = UBShmTransport::UB_OFF; - } - if (result >= 0 && ep->AllocateServerResources(&remote_trx_shm, &local_trx_shm) < 0) { - LOG(WARNING) << "Fail to allocate ub resources, fallback to tcp:" - << s->description(); - ub_transport->_ub_state = UBShmTransport::UB_OFF; - } - } - - ep->_state = S_HELLO_SEND; - HelloMessage local_msg{}; - local_msg.msg_len = g_ub_hello_msg_len; - if (ub_transport->_ub_state == UBShmTransport::UB_OFF) { - local_msg.impl_ver = 0; - local_msg.hello_ver = 0; - } else { - local_msg.hello_ver = g_ub_hello_version; - local_msg.impl_ver = g_ub_impl_version; - local_msg.len = (FLAGS_data_queue_size) * MB_TO_BYTE; - memcpy(local_msg.shm_name, remote_msg.shm_name, SHM_MAX_NAME_BUFF_LEN); - } - memcpy(data, MAGIC_STR, MAGIC_STR_LEN); - local_msg.Serialize((char*)data + MAGIC_STR_LEN); - if (ep->WriteToFd(data, g_ub_hello_msg_len) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to send Hello Message to client:" << s->description(); - s->SetFailed(saved_errno, "Fail to complete ub handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - - UbrDataFormat selected_format = UBR_DATA_FORMAT_NONE; - if (HelloNegotiationValid(remote_msg) && - HelloNegotiationValid(local_msg)) { - ep->_state = S_FORMAT_WAIT; - if (ep->ReadFromFd(data, HelloFormatExtension::WIRE_SIZE) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to read format extension from client:" - << s->description(); - s->SetFailed(saved_errno, - "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - HelloFormatExtension remote_extension; - remote_extension.Deserialize(data); - HelloFormatExtension local_extension = { - HelloFormatExtension::WIRE_SIZE, UBR_DATA_FORMAT_NONE}; - if (remote_extension.extension_len == HelloFormatExtension::WIRE_SIZE && - remote_extension.format_id == UBR_DATA_FORMAT_LEGACY_64) { - local_extension.format_id = UBR_DATA_FORMAT_LEGACY_64; - selected_format = UBR_DATA_FORMAT_LEGACY_64; - } else { - ub_transport->_ub_state = UBShmTransport::UB_OFF; - } - local_extension.Serialize(data); - ep->_state = S_FORMAT_SEND; - if (ep->WriteToFd(data, HelloFormatExtension::WIRE_SIZE) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to send format extension to client:" - << s->description(); - s->SetFailed(saved_errno, - "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - } - - ep->_state = S_ACK_WAIT; - if (ep->ReadFromFd(data, ACK_MSG_LEN) < 0) { - const int saved_errno = errno; - PLOG(WARNING) << "Fail to read ack message from client:" << s->description(); - s->SetFailed(saved_errno, "Fail to complete ubring handshake from %s: %s", - s->description().c_str(), berror(saved_errno)); - ep->_state = FAILED; - return nullptr; - } - - uint32_t net_flags; - memcpy(&net_flags, data, sizeof(net_flags)); - const uint32_t flags = butil::NetToHost32(net_flags); - if (flags & ACK_MSG_UB_OK) { - if (ub_transport->_ub_state == UBShmTransport::UB_OFF || - selected_format == UBR_DATA_FORMAT_NONE) { - LOG(WARNING) << "Invalid successful ACK from client:" - << s->description(); - s->SetFailed(EPROTO, "Fail to complete ub handshake from %s: %s", - s->description().c_str(), berror(EPROTO)); - ep->_state = FAILED; - return nullptr; - } else { - ub_transport->_ub_state = UBShmTransport::UB_ON; - ep->_negotiated_data_format = selected_format; - ep->_state = ESTABLISHED; - ep->_ub_ring->UbrUnlinkLocalShm(); - LOG_IF(INFO, FLAGS_ub_trace_verbose) - << "Server handshake ends (use ubring) on " << s->description(); - } - } else { - ub_transport->_ub_state = UBShmTransport::UB_OFF; - ep->_state = FALLBACK_TCP; - LOG_IF(INFO, FLAGS_ub_trace_verbose) - << "Server handshake ends (use tcp) on " << s->description(); - } - ep->TryReadOnTcp(); - - return nullptr; } bool UBShmEndpoint::IsWritable() const { @@ -736,9 +107,11 @@ ssize_t UBShmEndpoint::CutFromIOBufList(butil::IOBuf** from, size_t ndata) { nw = _ub_ring->UbrTrxWritev(vec, nvec); if (UNLIKELY(nw == -1)) { if (errno == EMSGSIZE) { - LOG(ERROR) << "Non-blocking send msg failed, message is larger than ubring capacity."; + LOG(ERROR) << "Non-blocking send msg failed, message is larger than " + "ubring capacity."; } else { - LOG(ERROR) << "Non-blocking send msg in failed, connection has been closed."; + LOG(ERROR) + << "Non-blocking send msg in failed, connection has been closed."; errno = EPIPE; } } else if (UNLIKELY(nw == UBRING_RETRY)) { @@ -758,7 +131,8 @@ ssize_t UBShmEndpoint::CutFromIOBufList(butil::IOBuf** from, size_t ndata) { return nw; } -int UBShmEndpoint::AllocateClientResources(ubring::SHM* local_trx_shm, const char* shm_name) { +int UBShmEndpoint::AllocateClientResources(ubring::SHM *local_trx_shm, + const char *shm_name) { if (BAIDU_UNLIKELY(g_skip_ub_init)) { // For UT return 0; @@ -772,18 +146,30 @@ int UBShmEndpoint::AllocateClientResources(ubring::SHM* local_trx_shm, const cha options.user = this; options.keytable_pool = _socket->_keytable_pool; if (Socket::Create(options, &_poller_sid) < 0) { + const int saved_errno = errno; PLOG(WARNING) << "Fail to create socket for UBRing poller"; + delete _ub_ring; + _ub_ring = NULL; + _poller_sid = INVALID_SOCKET_ID; + errno = saved_errno; return -1; } int ret = _ub_ring->UbrAllocateLocalShm(local_trx_shm, shm_name); if (ret != 0) { + const int saved_errno = errno; + DeallocateResources(); + delete _ub_ring; + _ub_ring = NULL; + _poller_sid = INVALID_SOCKET_ID; + errno = saved_errno; return ret; } PollerRegisterEvent(PollerSidOp::ADD, EPOLLIN); return 0; } -int UBShmEndpoint::AllocateServerResources(ubring::SHM* remote_trx_shm, ubring::SHM* local_trx_shm) { +int UBShmEndpoint::AllocateServerResources(ubring::SHM *remote_trx_shm, + ubring::SHM *local_trx_shm) { if (BAIDU_UNLIKELY(g_skip_ub_init)) { // For UT return 0; @@ -797,11 +183,22 @@ int UBShmEndpoint::AllocateServerResources(ubring::SHM* remote_trx_shm, ubring:: options.user = this; options.keytable_pool = _socket->_keytable_pool; if (Socket::Create(options, &_poller_sid) < 0) { + const int saved_errno = errno; PLOG(WARNING) << "Fail to create socket for UBRing poller"; + delete _ub_ring; + _ub_ring = NULL; + _poller_sid = INVALID_SOCKET_ID; + errno = saved_errno; return -1; } int ret = _ub_ring->UbrAllocateServerShm(remote_trx_shm, local_trx_shm); if (ret != 0) { + const int saved_errno = errno; + DeallocateResources(); + delete _ub_ring; + _ub_ring = NULL; + _poller_sid = INVALID_SOCKET_ID; + errno = saved_errno; return ret; } PollerRegisterEvent(PollerSidOp::ADD, EPOLLIN); @@ -829,7 +226,7 @@ void UBShmEndpoint::PollIn(UBShmEndpoint* ep, uint32_t ep_event) { if (Socket::Address(ep->_socket_id, &s) < 0) { return; } - auto* ub_transport = static_cast(s->_transport.get()); + UBShmTransport *ub_transport = UBShmTransport::Get(s.get()); CHECK(ep == ub_transport->_ub_ep); InputMessageClosure last_msg; @@ -850,9 +247,10 @@ void UBShmEndpoint::PollIn(UBShmEndpoint* ep, uint32_t ep_event) { if (nr <= 0) { if (0 == nr) { // Set `read_eof' flag and proceed to feed EOF into `Protocol' - // (implied by an empty processor.read_buf()), which may produce - // a new `InputMessageBase' under some protocols such as HTTP - LOG_IF(WARNING, FLAGS_log_connection_close) << *s << " was closed by remote side"; + // (implied by an empty processor.read_buf()), which may produce a new + // `InputMessageBase' under some protocols such as HTTP + LOG_IF(WARNING, FLAGS_log_connection_close) + << *s << " was closed by remote side"; read_eof = true; } else if (errno != EAGAIN) { if (errno == EINTR) { @@ -885,12 +283,11 @@ void UBShmEndpoint::PollOut(UBShmEndpoint* ep, uint32_t ep_event) { if (Socket::Address(ep->_socket_id, &s) < 0) { return; } - auto* ub_transport = static_cast(s->_transport.get()); + UBShmTransport *ub_transport = UBShmTransport::Get(s.get()); CHECK(ep == ub_transport->_ub_ep); if (ep->IsWritable()) { s->WakeAsEpollOut(); } - } int UBShmEndpoint::GlobalInitialize() { @@ -926,7 +323,7 @@ int UBShmEndpoint::PollingModeInitialize(bthread_tag_t tag, std::unique_ptr args(static_cast(p)); auto poller = args->poller; auto running = args->running; - std::unordered_set poller_sids; + std::unordered_set cq_sids; PollerSidOp op; if (poller->init_fn) { @@ -935,17 +332,18 @@ int UBShmEndpoint::PollingModeInitialize(bthread_tag_t tag, while (running->load(std::memory_order_relaxed)) { while (poller->op_queue.Dequeue(op)) { if (op.type == PollerSidOp::ADD) { - poller_sids.emplace(op); + cq_sids.emplace(op); } else if (op.type == PollerSidOp::REMOVE) { - poller_sids.erase(op); + cq_sids.erase(op); + } else if (op.type == PollerSidOp::MOD) { - poller_sids.erase(op); - poller_sids.emplace(op); + cq_sids.erase(op); + cq_sids.emplace(op); } } - for (const auto& poller_sid : poller_sids) { + for (auto cq : cq_sids) { SocketUniquePtr s; - if (Socket::Address(poller_sid.sid, &s) < 0) { + if (Socket::Address(cq.sid, &s) < 0) { continue; } UBShmEndpoint* ep = static_cast(s->user()); @@ -953,12 +351,12 @@ int UBShmEndpoint::PollingModeInitialize(bthread_tag_t tag, continue; } - if (poller_sid.events & EPOLLIN) { - PollIn(ep, poller_sid.events); + if (cq.event & EPOLLIN) { + PollIn(ep, cq.event); } - if (poller_sid.events & EPOLLOUT) { - PollOut(ep, poller_sid.events); + if (cq.event & EPOLLOUT) { + PollOut(ep, cq.event); } } if (poller->callback) { @@ -977,8 +375,8 @@ int UBShmEndpoint::PollingModeInitialize(bthread_tag_t tag, }; for (int i = 0; i < FLAGS_ub_poller_num; ++i) { auto args = new FnArgs{&pollers[i], &running}; - auto attr = FLAGS_ub_disable_bthread ? BTHREAD_ATTR_PTHREAD - : BTHREAD_ATTR_NORMAL; + auto attr = + FLAGS_ub_disable_bthread ? BTHREAD_ATTR_PTHREAD : BTHREAD_ATTR_NORMAL; attr.tag = tag; bthread_attr_set_name(&attr, "UBPolling"); pollers[i].callback = callback; @@ -1003,8 +401,7 @@ void UBShmEndpoint::PollingModeRelease(bthread_tag_t tag) { } } -void UBShmEndpoint::PollerRegisterEvent(PollerSidOp::OpType op, - uint32_t events) { +void UBShmEndpoint::PollerRegisterEvent(PollerSidOp::OpType op, uint32_t events) { auto index = butil::fmix32(_poller_sid) % FLAGS_ub_poller_num; auto& group = _poller_groups[bthread_self_tag()]; auto& pollers = group.pollers; diff --git a/src/brpc/ubshm/ub_endpoint.h b/src/brpc/ubshm/ub_endpoint.h index 2d7b1f5512..b9d6d196c8 100644 --- a/src/brpc/ubshm/ub_endpoint.h +++ b/src/brpc/ubshm/ub_endpoint.h @@ -20,78 +20,35 @@ #if BRPC_WITH_UBRING -#include -#include -#include -#include -#include -#include "butil/atomicops.h" -#include "butil/iobuf.h" -#include "butil/macros.h" -#include "butil/containers/mpsc_queue.h" +#include "brpc/handshake/ubshm_handshake.h" #include "brpc/socket.h" +#include "brpc/ubshm/shm/shm_def.h" #include "brpc/ubshm/ub_helper.h" #include "brpc/ubshm/ub_ring.h" -#include "brpc/ubshm/shm/shm_def.h" - +#include "butil/atomicops.h" +#include "butil/containers/mpsc_queue.h" +#include "butil/iobuf.h" +#include "butil/macros.h" +#include +#include namespace brpc { class Socket; +class UBShmTransport; +namespace handshake { +class UBShmServerHandshakeAdapter; +} namespace ubring { DECLARE_int32(ub_poller_num); DECLARE_bool(ub_edisp_unsched); DECLARE_bool(ub_disable_bthread); -enum UbrDataFormat { - UBR_DATA_FORMAT_NONE = 0, - UBR_DATA_FORMAT_LEGACY_64 = 1, -}; - -struct HelloFormatExtension { - // The V3 format extension is a fixed-size frame. A different wire size - // requires negotiation through a new hello version. - static const uint16_t WIRE_SIZE = 4; - - uint16_t extension_len; - uint16_t format_id; - - void Serialize(void* data) const; - void Deserialize(const void* data); -}; - -struct HelloMessage { - void Serialize(void* data) const; - void Deserialize(void* data); - std::string toString() const; - - uint16_t msg_len; - uint16_t hello_ver; - uint16_t impl_ver; - uint64_t len; - char shm_name[SHM_MAX_NAME_BUFF_LEN]; -}; - -class UBConnect : public AppConnect { -public: - void StartConnect(const Socket* socket, - void (*done)(int err, void* data), void* data) override; - void StopConnect(Socket*) override; - struct RunGuard { - RunGuard(UBConnect* rc) { this_rc = rc; } - ~RunGuard() { if (this_rc) this_rc->Run(); } - UBConnect* this_rc; - }; - -private: - void Run(); - void (*_done)(int, void*){nullptr}; - void* _data{nullptr}; -}; - class BAIDU_CACHELINE_ALIGNMENT UBShmEndpoint : public SocketUser { -friend class UBConnect; friend class Socket; +friend class ::brpc::UBShmTransport; +friend class ::brpc::handshake::UBShmServerHandshakeAdapter; + public: explicit UBShmEndpoint(Socket* s); ~UBShmEndpoint() override; @@ -104,6 +61,12 @@ friend class Socket; // Reset the endpoint (for next use) void Reset(); + void SetNegotiatedDataFormat(UbrDataFormat format) { + _negotiated_data_format = format; + } + UbrDataFormat negotiated_data_format() const { + return _negotiated_data_format; + } // Cut data from the given IOBuf list and use UBRING to send // Return bytes cut if success, -1 if failed and errno set @@ -130,9 +93,6 @@ friend class Socket; PollerRegisterEvent(PollerSidOp::REMOVE); } - // Callback when there is new epollin event on TCP fd - static void OnNewDataFromTcp(Socket* m); - // Initialize polling mode static int PollingModeInitialize(bthread_tag_t tag, std::function callback, @@ -141,37 +101,7 @@ friend class Socket; static void PollingModeRelease(bthread_tag_t tag); -#ifdef UNIT_TEST -public: -#else private: -#endif - enum State { - UNINIT = 0x0, - C_ALLOC_SHM = 0x1, - C_HELLO_SEND = 0x2, - C_HELLO_WAIT = 0x3, - C_FORMAT_SEND = 0x4, - C_FORMAT_WAIT = 0x5, - C_MAP_REMOTE_SHM = 0x6, - C_ACK_SEND = 0x7, - S_HELLO_WAIT = 0x11, - S_ALLOC_SHM = 0x12, - S_HELLO_SEND = 0x13, - S_FORMAT_WAIT = 0x14, - S_FORMAT_SEND = 0x15, - S_ACK_WAIT = 0x16, - ESTABLISHED = 0x100, - FALLBACK_TCP = 0x200, - FAILED = 0x300 - }; - - // Process handshake at the client - static void* ProcessHandshakeAtClient(void* arg); - - // Process handshake at the server - static void* ProcessHandshakeAtServer(void* arg); - // Allocate resources // Return 0 if success, -1 if failed and errno set int AllocateClientResources(SHM* local_trx_shm, const char* shm_name); @@ -181,58 +111,32 @@ friend class Socket; // Release resources void DeallocateResources(); - // Read at most len bytes from fd in _socket to data - // wait for _read_butex if encounter EAGAIN - // return -1 if encounter other errno (including EOF) - int ReadFromFd(void* data, size_t len); - - - // Write at most len bytes from data to fd in _socket - // wait for _epollout_butex if encounter EAGAIN - // return -1 if encounter other errno - int WriteToFd(void* data, size_t len); - - // Poll inbound and outbound UBRing events. + // Poll CQ and get the work completion static void PollIn(UBShmEndpoint* ep, uint32_t ep_event); static void PollOut(UBShmEndpoint* ep, uint32_t ep_event); - // Try to read data on TCP fd in _socket - inline void TryReadOnTcp(); - // Not owner Socket* _socket; SocketId _socket_id; - - State _state; UbrDataFormat _negotiated_data_format{UBR_DATA_FORMAT_NONE}; // ub resource ubring::UBRing* _ub_ring{nullptr}; - // Synthetic SocketId registered with the UBRing poller. SocketId _poller_sid; - // butex for inform read events on TCP fd during handshake - butil::atomic *_read_butex; - DISALLOW_COPY_AND_ASSIGN(UBShmEndpoint); struct PollerSidOp { - enum OpType { - ADD, - REMOVE, - MOD - }; + enum OpType { ADD, REMOVE, MOD }; SocketId sid; - uint32_t events; + uint32_t event; OpType type; }; struct PollerSidOpHash { - std::size_t operator()(const PollerSidOp& op) const { - return op.sid; - } + std::size_t operator()(const PollerSidOp& op) const { return op.sid; } }; struct PollerSidOpEqual { @@ -244,8 +148,7 @@ friend class Socket; // Poller instance struct BAIDU_CACHELINE_ALIGNMENT Poller { bthread_t tid{INVALID_BTHREAD}; - butil::MPSCQueue< - PollerSidOp, butil::ObjectPoolAllocator> op_queue; + butil::MPSCQueue> op_queue; // Callback used for io_uring/spdk etc std::function callback; // Init and Destroy function @@ -260,8 +163,7 @@ friend class Socket; }; static std::vector _poller_groups; - void PollerRegisterEvent(PollerSidOp::OpType op, - uint32_t events = EPOLLET); + void PollerRegisterEvent(PollerSidOp::OpType op, uint32_t events = EPOLLET); }; } // namespace ubring diff --git a/src/brpc/ubshm_transport.cpp b/src/brpc/ubshm_transport.cpp index 45f6a61fa6..9b0a0203e0 100644 --- a/src/brpc/ubshm_transport.cpp +++ b/src/brpc/ubshm_transport.cpp @@ -17,10 +17,15 @@ #if BRPC_WITH_UBRING -#include "brpc/ubshm_transport.h" -#include "brpc/tcp_transport.h" +#include + +#include "brpc/adapter_transport.h" +#include "brpc/errno.pb.h" +#include "brpc/ubshm/common/common.h" #include "brpc/ubshm/ub_endpoint.h" #include "brpc/ubshm/ub_helper.h" +#include "brpc/ubshm/ubr_trx.h" +#include "brpc/ubshm_transport.h" namespace brpc { DECLARE_bool(usercode_in_coroutine); @@ -28,23 +33,26 @@ DECLARE_bool(usercode_in_pthread); extern SocketVarsCollector *g_vars; +UBShmTransport *UBShmTransport::Get(const Socket *socket) { + const AdapterTransport *adapter = AdapterTransport::Get(socket); + Transport *transport = adapter->high_speed_transport(); + CHECK(transport != NULL); + return static_cast(transport); +} + void UBShmTransport::Init(Socket *socket, const SocketOptions &options) { CHECK(_ub_ep == nullptr); - if (options.socket_mode == SOCKET_MODE_UBRING) { - _ub_ep = new ubring::UBShmEndpoint(socket); - _ub_state = UB_UNKNOWN; - } else { - _ub_state = UB_OFF; - socket->_socket_mode = SOCKET_MODE_TCP; - } _socket = socket; _default_connect = options.app_connect; - _on_edge_trigger = options.on_edge_triggered_events; - if (options.need_on_edge_trigger && _on_edge_trigger == nullptr) { - _on_edge_trigger = ubring::UBShmEndpoint::OnNewDataFromTcp; + _on_edge_trigger = nullptr; + _ub_state = UB_UNKNOWN; + _ub_ep = new (std::nothrow) ubring::UBShmEndpoint(socket); + if (!_ub_ep) { + const int saved_errno = errno != 0 ? errno : ENOMEM; + errno = saved_errno; + PLOG(WARNING) << "Fail to create UBShmEndpoint, disable UBSHM upgrade"; + _ub_state = UB_OFF; } - _tcp_transport = std::make_shared(); - _tcp_transport->Init(socket, options); } void UBShmTransport::Release() { @@ -64,26 +72,53 @@ int UBShmTransport::Reset(int32_t expected_nref) { } std::shared_ptr UBShmTransport::Connect() { - if (_default_connect == nullptr) { - return std::make_shared(); - } return _default_connect; } -int UBShmTransport::CutFromIOBuf(butil::IOBuf *buf) { - if (_ub_ep && _ub_state != UB_OFF) { - butil::IOBuf *data_arr[1] = {buf}; - return _ub_ep->CutFromIOBufList(data_arr, 1); - } else { - return _tcp_transport->CutFromIOBuf(buf); +void UBShmTransport::SetHighSpeedAvailable(bool available) { + _ub_state = available ? UB_ON : UB_OFF; +} + +int UBShmTransport::PrepareUpgradeResources(ubring::SHM *local_trx_shm, + const char *shm_name) { + return _ub_ep->AllocateClientResources(local_trx_shm, shm_name); +} + +int UBShmTransport::NegotiateUpgradeResources(ubring::SHM *local_trx_shm, + const char *shm_name) { + return _ub_ep->_ub_ring->UbrMapRemoteShm(local_trx_shm, shm_name); +} + +int UBShmTransport::PrepareServerUpgradeResources(ubring::SHM *remote_trx_shm, + ubring::SHM *local_trx_shm) { + return _ub_ep->AllocateServerResources(remote_trx_shm, local_trx_shm); +} + +void UBShmTransport::ActivateUpgrade() { + SetHighSpeedAvailable(true); +} + +void UBShmTransport::DeactivateUpgrade() { + SetHighSpeedAvailable(false); + if (_ub_ep != nullptr) { + _ub_ep->Reset(); } } -ssize_t UBShmTransport::CutFromIOBufList(butil::IOBuf **buf, size_t ndata) { - if (_ub_ep && _ub_state != UB_OFF) { - return _ub_ep->CutFromIOBufList(buf, ndata); +void UBShmTransport::FinishUpgrade() { + if (_ub_ep != NULL && _ub_ep->_ub_ring != NULL) { + _ub_ep->_ub_ring->UbrUnlinkLocalShm(); } - return _tcp_transport->CutFromIOBufList(buf, ndata); +} + +int UBShmTransport::CutFromIOBuf(butil::IOBuf *buf) { + butil::IOBuf *data[1] = {buf}; + return static_cast(CutFromIOBufList(data, 1)); +} + +ssize_t UBShmTransport::CutFromIOBufList(butil::IOBuf **buf, size_t ndata) { + CHECK(_ub_ep != NULL); + return _ub_ep->CutFromIOBufList(buf, ndata); } int UBShmTransport::WaitEpollOut(butil::atomic *_epollout_butex, @@ -94,14 +129,13 @@ int UBShmTransport::WaitEpollOut(butil::atomic *_epollout_butex, if (!_ub_ep->IsWritable()) { g_vars->nwaitepollout << 1; _ub_ep->PollerRegisterEpollOut(pollin); - const int wait_rc = bthread::butex_wait( - _epollout_butex, expected_val, &duetime); + const int wait_rc = + bthread::butex_wait(_epollout_butex, expected_val, &duetime); if (wait_rc < 0) { if (errno != EAGAIN && errno != ETIMEDOUT) { const int saved_errno = errno; PLOG(WARNING) << "Fail to wait ub window of " << _socket; - _socket->SetFailed(saved_errno, - "Fail to wait ub window of %s: %s", + _socket->SetFailed(saved_errno, "Fail to wait ub window of %s: %s", _socket->description().c_str(), berror(saved_errno)); } @@ -115,8 +149,6 @@ int UBShmTransport::WaitEpollOut(butil::atomic *_epollout_butex, } _ub_ep->PollerUnRegisterEpollOut(pollin); } - } else { - return _tcp_transport->WaitEpollOut(_epollout_butex, pollin, duetime); } return 0; } @@ -156,15 +188,16 @@ void UBShmTransport::QueueMessage(InputMessageClosure& input_msg, // TODO(gejun): Join threads. bthread_t th; - bthread_attr_t tmp = (FLAGS_usercode_in_pthread ? - BTHREAD_ATTR_PTHREAD : - BTHREAD_ATTR_NORMAL) | BTHREAD_NOSIGNAL; + bthread_attr_t tmp = + (FLAGS_usercode_in_pthread ? BTHREAD_ATTR_PTHREAD : BTHREAD_ATTR_NORMAL) | + BTHREAD_NOSIGNAL; tmp.keytable_pool = _socket->keytable_pool(); tmp.tag = bthread_self_tag(); bthread_attr_set_name(&tmp, "ProcessInputMessage"); - if (!FLAGS_usercode_in_coroutine && bthread_start_background( - &th, &tmp, ProcessInputMessage, to_run_msg) == 0) { + if (!FLAGS_usercode_in_coroutine && + bthread_start_background(&th, &tmp, ProcessInputMessage, to_run_msg) == + 0) { ++*num_bthread_created; } else { ProcessInputMessage(to_run_msg); @@ -179,7 +212,8 @@ int UBShmTransport::ContextInitOrDie(bool serverOrNot, const void* _options) { return -1; } ubring::GlobalUBInitializeOrDie(); - if (!ubring::InitPollingModeWithTag(static_cast(_options)->bthread_tag)) { + if (!ubring::InitPollingModeWithTag( + static_cast(_options)->bthread_tag)) { return -1; } } else { @@ -202,8 +236,7 @@ bool UBShmTransport::OptionsAvailableForUB(const ChannelOptions* opt) { return false; } if (!ubring::SupportedByUB(opt->protocol.name())) { - LOG(WARNING) << "Cannot use " << opt->protocol.name() - << " over UB"; + LOG(WARNING) << "Cannot use " << opt->protocol.name() << " over UB"; return false; } return true; diff --git a/src/brpc/ubshm_transport.h b/src/brpc/ubshm_transport.h index b3d1e7c518..b8d840595e 100644 --- a/src/brpc/ubshm_transport.h +++ b/src/brpc/ubshm_transport.h @@ -21,12 +21,14 @@ #include "brpc/socket.h" #include "brpc/channel.h" #include "brpc/transport.h" +#include "brpc/ubshm/shm/shm_def.h" namespace brpc { +class AdapterTransport; class UBShmTransport : public Transport { friend class TransportFactory; + friend class AdapterTransport; friend class ubring::UBShmEndpoint; -friend class ubring::UBConnect; public: void Init(Socket* socket, const SocketOptions& options) override; void Release() override; @@ -34,7 +36,8 @@ friend class ubring::UBConnect; std::shared_ptr Connect() override; int CutFromIOBuf(butil::IOBuf* buf) override; ssize_t CutFromIOBufList(butil::IOBuf** buf, size_t ndata) override; - int WaitEpollOut(butil::atomic* _epollout_butex, bool pollin, const timespec duetime) override; + int WaitEpollOut(butil::atomic* epollout_butex, + bool pollin, timespec duetime) override; void ProcessEvent(bthread_attr_t attr) override; void QueueMessage(InputMessageClosure& inputMsg, int* num_bthread_created, bool last_msg) override; void Debug(std::ostream &os) override; @@ -42,8 +45,22 @@ friend class ubring::UBConnect; CHECK(_ub_ep != nullptr); return _ub_ep; } + static UBShmTransport* Get(const Socket* socket); static int ContextInitOrDie(bool serverOrNot, const void* _options); + int PrepareUpgradeResources(ubring::SHM* local_trx_shm, + const char* shm_name); + int NegotiateUpgradeResources(ubring::SHM* local_trx_shm, + const char* shm_name); + int PrepareServerUpgradeResources(ubring::SHM* remote_trx_shm, + ubring::SHM* local_trx_shm); + void ActivateUpgrade(); + void DeactivateUpgrade(); + void FinishUpgrade(); + bool UpgradeActive() const { return _ub_state == UB_ON; } + bool UpgradeReady() const { return _ub_ep != nullptr; } private: + void SetHighSpeedAvailable(bool available); + static bool OptionsAvailableForUB(const ChannelOptions* opt); static bool OptionsAvailableOverUB(const ServerOptions* opt); private: @@ -57,7 +74,6 @@ friend class ubring::UBConnect; ubring::UBShmEndpoint* _ub_ep = nullptr; // Should use UB or not UBState _ub_state; - std::shared_ptr _tcp_transport; }; } // namespace brpc #endif // BRPC_WITH_UBRING diff --git a/test/brpc_rdma_unittest.cpp b/test/brpc_rdma_unittest.cpp index 8886714268..97265ae0af 100644 --- a/test/brpc_rdma_unittest.cpp +++ b/test/brpc_rdma_unittest.cpp @@ -15,39 +15,36 @@ // specific language governing permissions and limitations // under the License. - +#include +#include +#include #include #include -#include -#include + #if BRPC_WITH_RDMA -#include -#include -#include -#include -#include -#include "butil/endpoint.h" -#include "butil/fd_guard.h" -#include "butil/iobuf.h" -#include "butil/sys_byteorder.h" -#include "butil/time.h" -#include "butil/files/temp_file.h" #include "brpc/acceptor.h" +#include "brpc/adapter_transport.h" #include "brpc/channel.h" #include "brpc/controller.h" -#include "brpc/server.h" -#include "brpc/socket.h" #include "brpc/errno.pb.h" +#include "brpc/handshake/rdma_handshake.h" +#include "brpc/handshake/rdma_handshake_constants.h" #include "brpc/parallel_channel.h" -#include "brpc/selective_channel.h" -#include "brpc/rdma_transport.h" #include "brpc/rdma/block_pool.h" #include "brpc/rdma/rdma_endpoint.h" -#include "brpc/rdma/rdma_handshake.h" -#include "brpc/rdma/rdma_handshake_constants.h" -#include "brpc/rdma/rdma_handshake.pb.h" #include "brpc/rdma/rdma_helper.h" +#include "brpc/rdma_handshake.pb.h" +#include "brpc/rdma_transport.h" +#include "brpc/selective_channel.h" +#include "brpc/server.h" +#include "brpc/socket.h" +#include "butil/endpoint.h" +#include "butil/fd_guard.h" +#include "butil/files/temp_file.h" +#include "butil/iobuf.h" +#include "butil/sys_byteorder.h" #include "echo.pb.h" +#include static const int PORT = 8713; @@ -62,24 +59,26 @@ DEFINE_bool(rdma_test_enable, false, "Enable tests requring rdma runtime."); namespace rdma { // HELLO_V2_VERSION / IMPL_V2_VERSION come from -// brpc/rdma/rdma_handshake_constants.h (shared wire constants). +// brpc/handshake/rdma_handshake_constants.h (shared wire constants). DECLARE_bool(rdma_trace_verbose); DECLARE_int32(rdma_memory_pool_max_regions); DECLARE_int32(rdma_client_handshake_version); DECLARE_bool(rdma_ece); -extern ibv_cq* (*IbvCreateCq)(ibv_context*, int, void*, ibv_comp_channel*, int); +extern ibv_cq* (*IbvCreateCq)(ibv_context*, int, void*, ibv_comp_channel*, + int); extern int (*IbvDestroyCq)(ibv_cq*); extern ibv_qp* (*IbvCreateQp)(ibv_pd*, ibv_qp_init_attr*); extern int (*IbvModifyQp)(ibv_qp*, ibv_qp_attr*, ibv_qp_attr_mask); -extern int (*IbvQueryQp)(ibv_qp*, ibv_qp_attr*, ibv_qp_attr_mask, ibv_qp_init_attr*); +extern int (*IbvQueryQp)(ibv_qp*, ibv_qp_attr*, ibv_qp_attr_mask, + ibv_qp_init_attr*); extern int (*IbvDestroyQp)(ibv_qp*); extern butil::atomic g_rdma_available; extern bool g_skip_rdma_init; extern bool g_fail_resource_alloc_for_test; -} // namespace rdma -} // namespace brpc +} // namespace rdma +} // namespace brpc static std::string g_ip = "127.0.0.1"; static butil::EndPoint g_ep; @@ -116,7 +115,7 @@ static bool WaitUntil(const std::function& pred, // whole buffer reaching the peer has to loop. static bool WriteAll(int fd, const void* buf, size_t len) { const uint8_t* p = (const uint8_t*)buf; - for (size_t done = 0; done < len; ) { + for (size_t done = 0; done < len;) { ssize_t n = write(fd, p + done, len - done); if (n < 0) { if (errno == EINTR) { @@ -133,7 +132,7 @@ static bool WriteAll(int fd, const void* buf, size_t len) { // and on EOF before the whole buffer arrived. static bool ReadAll(int fd, void* buf, size_t len) { uint8_t* p = (uint8_t*)buf; - for (size_t done = 0; done < len; ) { + for (size_t done = 0; done < len;) { ssize_t n = read(fd, p + done, len - done); if (n < 0) { if (errno == EINTR) { @@ -170,8 +169,7 @@ static void ConnectToServer(butil::fd_guard* sockfd) { class MyEchoService : public ::test::EchoService { void Echo(google::protobuf::RpcController* cntl_base, - const ::test::EchoRequest* req, - ::test::EchoResponse* res, + const ::test::EchoRequest* req, ::test::EchoResponse* res, google::protobuf::Closure* done) { Controller* cntl = static_cast(cntl_base); ClosureGuard done_guard(done); @@ -210,13 +208,11 @@ class RdmaTest : public ::testing::Test { _naming_url = std::string("File://") + _server_list.fname(); _server.AddService(&_svc, SERVER_DOESNT_OWN_SERVICE); } - ~RdmaTest() { } + ~RdmaTest() {} - virtual void SetUp() { } + virtual void SetUp() {} - virtual void TearDown() { - rdma::DumpMemoryPoolInfo(std::cout); - } + virtual void TearDown() { rdma::DumpMemoryPoolInfo(std::cout); } protected: void StartServer(bool use_rdma = true) { @@ -247,15 +243,13 @@ class RdmaTest : public ::testing::Test { return nullptr; } - // Accepting the connection happens in the server threads, so poll for it - // rather than sleeping. Returns nullptr if it never showed up. + // Server-side connection creation and teardown are asynchronous. Socket* WaitForServerSocket() { Socket* s = nullptr; WaitUntil([this, &s] { return (s = GetSocketFromServer(0)) != nullptr; }); return s; } - // Ditto for the connection going away. bool WaitForServerSocketGone() { return WaitUntil([this] { return GetSocketFromServer(0) == nullptr; }); } @@ -267,41 +261,29 @@ class RdmaTest : public ::testing::Test { MyEchoService _svc; }; -// Shorthand for the RDMA transport behind a Socket, which every endpoint state -// check below has to go through. +// Shorthand for the RDMA transport behind a Socket. static RdmaTransport* RdmaTransportOf(Socket* s) { - return static_cast(s->_transport.get()); -} -static RdmaTransport* RdmaTransportOf(const SocketUniquePtr& s) { - return RdmaTransportOf(s.get()); + return RdmaTransport::Get(s); } -// Polls until the endpoint reaches `expected` and returns the last state seen, -// so that ASSERT_RDMA_STATE() reports what the endpoint actually settled on. -static rdma::RdmaEndpoint::State WaitForRdmaState( - RdmaTransport* transport, rdma::RdmaEndpoint::State expected) { - rdma::RdmaEndpoint::State state = transport->_rdma_ep->_state; - WaitUntil([transport, expected, &state] { - state = transport->_rdma_ep->_state; - return state == expected; +static int WaitForHandshakePhase(Socket* s, handshake::Phase expected) { + int phase = AdapterTransport::Get(s)->handshake_phase(); + WaitUntil([s, expected, &phase] { + phase = AdapterTransport::Get(s)->handshake_phase(); + return phase == expected; }); - return state; + return phase; } -// Waits for `transport` to reach `expected`, failing the test if it does not. -#define ASSERT_RDMA_STATE(expected, transport) \ - ASSERT_EQ(expected, WaitForRdmaState(transport, expected)) +#define ASSERT_HANDSHAKE_PHASE(expected, socket) \ + ASSERT_EQ(expected, WaitForHandshakePhase(socket, expected)) -// Polls until the fd stream of `s` holds exactly `size` bytes. Tests asserting -// that a state did NOT change need this: waiting for the state itself would -// return before the peer had read anything at all. static bool WaitForFdReadBuf(Socket* s, size_t size) { return WaitUntil([s, size] { return s->fd_input_processor().read_buf().size() == size; }); } -// Build a well-formed v2 client hello: "RDMA" followed by the 36B body. static void MakeV2ClientHello(uint8_t (&data)[rdma::HELLO_V2_MSG_LEN_MIN]) { rdma::v2_wire::HelloMessage msg{}; msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; @@ -322,8 +304,7 @@ static void MakeV2ClientHello(uint8_t (&data)[rdma::HELLO_V2_MSG_LEN_MIN]) { // so every TEST_P below is automatically executed once per supported // version. Add a new version to INSTANTIATE_TEST_SUITE_P at the bottom // of this file and these RPC tests will gain coverage for free. -class RdmaRpcTest : public RdmaTest, - public ::testing::WithParamInterface { +class RdmaRpcTest : public RdmaTest, public ::testing::WithParamInterface { protected: void SetUp() override { RdmaTest::SetUp(); @@ -347,9 +328,10 @@ TEST_F(RdmaTest, stale_cq_callback_does_not_poll_new_generation) { SocketUniquePtr main_socket; ASSERT_EQ(0, Socket::Address(main_sid, &main_socket)); - RdmaTransport* transport = - static_cast(main_socket->_transport.get()); + RdmaTransport* transport = RdmaTransportOf(main_socket.get()); + ASSERT_NE(nullptr, transport); rdma::RdmaEndpoint* ep = transport->_rdma_ep; + ASSERT_NE(nullptr, ep); SocketOptions cq_options; cq_options.user = ep; @@ -376,13 +358,21 @@ TEST_F(RdmaTest, stale_cq_callback_does_not_poll_new_generation) { TEST_F(RdmaTest, client_close_before_hello_send) { StartServer(); - butil::fd_guard sockfd; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd)); - Socket* s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + sockaddr_in addr; + bzero((char*)&addr, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_port = htons(PORT); + + butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd >= 0); + ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + Socket* s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + close(sockfd); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -390,23 +380,29 @@ TEST_F(RdmaTest, client_close_before_hello_send) { TEST_F(RdmaTest, client_hello_msg_invalid_magic_str) { StartServer(); - butil::fd_guard sockfd; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd)); - Socket* s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + sockaddr_in addr; + bzero((char*)&addr, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_port = htons(PORT); + + butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd >= 0); + ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + Socket* s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; memcpy(data, "PRPC", 4); // send as normal baidu_std protocol - ASSERT_TRUE(WriteAll(sockfd, data, 4)); - // Wait for the bytes to show up in the fd stream (baidu_std wants 12B of - // header, so they stay buffered). Waiting on the state instead would prove - // nothing: it is already UNINIT before the server has read anything. - ASSERT_TRUE(WaitForFdReadBuf(s, 4)); - // A non-RDMA magic makes ParseRdmaHandshake return TRY_OTHERS and hand the - // bytes to other protocols; it does not touch the endpoint state, so it - // stays UNINIT (the old blocking handshake used to set FALLBACK_TCP here). - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + ASSERT_EQ(4, write(sockfd, data, 4)); + usleep(100000); // wait for server to handle the msg + // A non-RDMA magic makes the transport-handshake parser return TRY_OTHERS + // and hand the bytes to other protocols; it does not touch the endpoint + // state, so it stays UNINIT (the old blocking handshake used to set + // FALLBACK_TCP here). + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); StopServer(); } @@ -414,51 +410,66 @@ TEST_F(RdmaTest, client_hello_msg_invalid_magic_str) { TEST_F(RdmaTest, client_close_during_hello_send) { StartServer(); + sockaddr_in addr; + bzero((char*)&addr, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_port = htons(PORT); Socket* s = nullptr; uint8_t data[8]; - butil::fd_guard sockfd1; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd1)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + butil::fd_guard sockfd1(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd1 >= 0); + ASSERT_EQ(0, connect(sockfd1, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); memcpy(data, "RD", 2); - ASSERT_TRUE(WriteAll(sockfd1, data, 2)); // break in magic str - // Fewer than 4 magic bytes: ParseRdmaHandshake can't tell yet, returns - // NOT_ENOUGH_DATA and leaves the endpoint UNINIT (the old blocking - // handshake used to set S_HELLO_WAIT before reading the magic). Wait for - // the bytes to be buffered, the state alone would prove nothing. - ASSERT_TRUE(WaitForFdReadBuf(s, 2)); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - sockfd1.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - butil::fd_guard sockfd2; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd2)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + ASSERT_EQ(2, write(sockfd1, data, 2)); // break in magic str + usleep(100000); // wait for server to handle the msg + // Fewer than 4 magic bytes: the transport-handshake parser can't tell yet, + // returns NOT_ENOUGH_DATA and leaves the endpoint UNINIT (the old blocking + // the common handshake state remains uninitialized before reading magic). + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + close(sockfd1); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + butil::fd_guard sockfd2(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd2 >= 0); + ASSERT_EQ(0, connect(sockfd2, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); memcpy(data, "RDMA", 4); - ASSERT_TRUE(WriteAll(sockfd2, data, 4)); // break after magic str - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - sockfd2.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - butil::fd_guard sockfd3; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd3)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + ASSERT_EQ(4, write(sockfd2, data, 4)); // break after magic str + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); + close(sockfd2); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + butil::fd_guard sockfd3(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd3 >= 0); + ASSERT_EQ(0, connect(sockfd3, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); // Send the 4B magic plus a valid msg_len (=40) but no body, so the server // recognizes an RDMA v2 hello and waits for the remaining bytes. (A zero // msg_len would now be rejected up-front as a protocol error.) memcpy(data, "RDMA", 4); uint16_t v2_len = butil::HostToNet16(rdma::HELLO_V2_MSG_LEN_MIN); memcpy(data + 4, &v2_len, sizeof(v2_len)); - ASSERT_TRUE(WriteAll(sockfd3, data, 6)); // magic + msg_len, body missing - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - sockfd3.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + ASSERT_EQ(6, write(sockfd3, data, 6)); // magic + msg_len, body missing + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); + close(sockfd3); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -466,34 +477,46 @@ TEST_F(RdmaTest, client_close_during_hello_send) { TEST_F(RdmaTest, client_hello_msg_invalid_len) { StartServer(); + sockaddr_in addr; + bzero((char*)&addr, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_port = htons(PORT); Socket* s = nullptr; uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - butil::fd_guard sockfd1; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd1)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + butil::fd_guard sockfd1(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd1 >= 0); + ASSERT_EQ(0, connect(sockfd1, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); memcpy(data, "RDMA", 4); - ASSERT_TRUE(WriteAll(sockfd1, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(4, write(sockfd1, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); memset(data + 4, 0, 36); - ASSERT_TRUE(WriteAll(sockfd1, data + 4, 36)); // Write invalid length. - ASSERT_TRUE(WaitForServerSocketGone()); - - butil::fd_guard sockfd2; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd2)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + ASSERT_EQ(36, write(sockfd1, data + 4, 36)); // Write invalid length. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + butil::fd_guard sockfd2(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd2 >= 0); + ASSERT_EQ(0, connect(sockfd2, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); memcpy(data, "RDMA", 4); - ASSERT_TRUE(WriteAll(sockfd2, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(4, write(sockfd2, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); uint16_t len = butil::HostToNet16(35); memcpy(data + 4, &len, sizeof(len)); memset(data + 6, 0, 34); - ASSERT_TRUE(WriteAll(sockfd2, data + 4, 36)); // write invalid length - ASSERT_TRUE(WaitForServerSocketGone()); + ASSERT_EQ(36, write(sockfd2, data + 4, 36)); // write invalid length + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -501,19 +524,26 @@ TEST_F(RdmaTest, client_hello_msg_invalid_len) { TEST_F(RdmaTest, client_hello_msg_invalid_version) { StartServer(); + sockaddr_in addr; + bzero((char*)&addr, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_port = htons(PORT); Socket* s = nullptr; uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; uint16_t len = butil::HostToNet16(rdma::HELLO_V2_MSG_LEN_MIN); uint16_t ver = butil::HostToNet16(1); - butil::fd_guard sockfd1; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd1)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + butil::fd_guard sockfd1(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd1 >= 0); + ASSERT_EQ(0, connect(sockfd1, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); memcpy(data, "RDMA", 4); - ASSERT_TRUE(WriteAll(sockfd1, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(4, write(sockfd1, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); memcpy(data + 4, &len, 2); memset(data + 6, 0, 34); memcpy(data + 6, &ver, 2); // hello_ver == 1, impl_ver == 0 @@ -524,35 +554,46 @@ TEST_F(RdmaTest, client_hello_msg_invalid_version) { // hello_ver). Now that Step 1 enforces a HELLO_V2_MSG_LEN_MAX upper bound, // such an oversized msg_len would be rejected before reaching the // version check, breaking the intent of this UT. - ASSERT_TRUE(WriteAll(sockfd1, data + 4, 36)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); - ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransportOf(s)->_rdma_state); + ASSERT_EQ(36, write(sockfd1, data + 4, 36)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransport::Get(s)->_rdma_state); uint32_t flags = 0; - ASSERT_TRUE(WriteAll(sockfd1, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + ASSERT_EQ(sizeof(flags), write(sockfd1, &flags, sizeof(flags))); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); sockfd1.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - butil::fd_guard sockfd2; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd2)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + butil::fd_guard sockfd2(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd2 >= 0); + ASSERT_EQ(0, connect(sockfd2, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); memcpy(data, "RDMA", 4); - ASSERT_TRUE(WriteAll(sockfd2, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(4, write(sockfd2, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); memcpy(data + 4, &len, 2); memset(data + 6, 0, 32); memcpy(data + 8, &ver, 2); // hello_ver == 0, impl_ver == 1 - // See comment above on `WriteAll(sockfd1, data + 4, 36)` for why we + // See comment above on `write(sockfd1, data + 4, 36)` for why we // write from data + 4 instead of data. - ASSERT_TRUE(WriteAll(sockfd2, data + 4, 36)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); - ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransportOf(s)->_rdma_state); - ASSERT_TRUE(WriteAll(sockfd2, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + ASSERT_EQ(36, write(sockfd2, data + 4, 36)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransport::Get(s)->_rdma_state); + ASSERT_EQ(sizeof(flags), write(sockfd2, &flags, sizeof(flags))); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); sockfd2.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -560,6 +601,10 @@ TEST_F(RdmaTest, client_hello_msg_invalid_version) { TEST_F(RdmaTest, client_hello_msg_invalid_sq_rq_block_size) { StartServer(); + sockaddr_in addr; + bzero((char*)&addr, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_port = htons(PORT); Socket* s = nullptr; uint32_t flags = butil::HostToNet32(0); rdma::v2_wire::HelloMessage msg{}; @@ -573,60 +618,81 @@ TEST_F(RdmaTest, client_hello_msg_invalid_sq_rq_block_size) { msg.block_size = 8192; memcpy(data, "RDMA", 4); msg.Serialize(data + 4); - butil::fd_guard sockfd1; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd1)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - ASSERT_TRUE(WriteAll(sockfd1, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - ASSERT_TRUE(WriteAll(sockfd1, data + 4, 36)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); - ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransportOf(s)->_rdma_state); - ASSERT_TRUE(WriteAll(sockfd1, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + butil::fd_guard sockfd1(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd1 >= 0); + ASSERT_EQ(0, connect(sockfd1, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(4, write(sockfd1, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(36, write(sockfd1, data + 4, 36)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransport::Get(s)->_rdma_state); + ASSERT_EQ(sizeof(flags), write(sockfd1, &flags, sizeof(flags))); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); sockfd1.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); msg.sq_size = 16; msg.rq_size = 10; msg.block_size = 8192; memcpy(data, "RDMA", 4); msg.Serialize(data + 4); - butil::fd_guard sockfd2; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd2)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - ASSERT_TRUE(WriteAll(sockfd2, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - ASSERT_TRUE(WriteAll(sockfd2, data + 4, 36)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); - ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransportOf(s)->_rdma_state); - ASSERT_TRUE(WriteAll(sockfd2, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + butil::fd_guard sockfd2(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd2 >= 0); + ASSERT_EQ(0, connect(sockfd2, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(4, write(sockfd2, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(36, write(sockfd2, data + 4, 36)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransport::Get(s)->_rdma_state); + ASSERT_EQ(sizeof(flags), write(sockfd2, &flags, sizeof(flags))); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); sockfd2.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); msg.sq_size = 16; msg.rq_size = 16; msg.block_size = 1000; memcpy(data, "RDMA", 4); msg.Serialize(data + 4); - butil::fd_guard sockfd3; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd3)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - ASSERT_TRUE(WriteAll(sockfd3, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - ASSERT_TRUE(WriteAll(sockfd3, data + 4, 36)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); - ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransportOf(s)->_rdma_state); - ASSERT_TRUE(WriteAll(sockfd3, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + butil::fd_guard sockfd3(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd3 >= 0); + ASSERT_EQ(0, connect(sockfd3, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(4, write(sockfd3, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(36, write(sockfd3, data + 4, 36)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransport::Get(s)->_rdma_state); + ASSERT_EQ(sizeof(flags), write(sockfd3, &flags, sizeof(flags))); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); sockfd3.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -634,19 +700,37 @@ TEST_F(RdmaTest, client_hello_msg_invalid_sq_rq_block_size) { TEST_F(RdmaTest, client_close_after_qp_build) { StartServer(); + sockaddr_in addr; + bzero((char*)&addr, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_port = htons(PORT); Socket* s = nullptr; + rdma::v2_wire::HelloMessage msg{}; uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - MakeV2ClientHello(data); + msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; + msg.hello_ver = rdma::HELLO_V2_VERSION; + msg.impl_ver = rdma::IMPL_V2_VERSION; + msg.sq_size = 16; + msg.rq_size = 16; + msg.block_size = 8192; + msg.qp_num = 0; + msg.gid = rdma::GetRdmaGid(); + memcpy(data, "RDMA", 4); + msg.Serialize(data + 4); - butil::fd_guard sockfd1; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd1)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - ASSERT_TRUE(WriteAll(sockfd1, data, sizeof(data))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); - sockfd1.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + butil::fd_guard sockfd1(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd1 >= 0); + ASSERT_EQ(0, connect(sockfd1, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(40, write(sockfd1, data, 40)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); + close(sockfd1); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -654,24 +738,45 @@ TEST_F(RdmaTest, client_close_after_qp_build) { TEST_F(RdmaTest, client_close_during_ack_send) { StartServer(); + sockaddr_in addr; + bzero((char*)&addr, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_port = htons(PORT); Socket* s = nullptr; + rdma::v2_wire::HelloMessage msg{}; uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - MakeV2ClientHello(data); + msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; + msg.hello_ver = rdma::HELLO_V2_VERSION; + msg.impl_ver = rdma::IMPL_V2_VERSION; + msg.sq_size = 16; + msg.rq_size = 16; + msg.block_size = 8192; + msg.qp_num = 0; + msg.gid = rdma::GetRdmaGid(); + memcpy(data, "RDMA", 4); + msg.Serialize(data + 4); - butil::fd_guard sockfd1; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd1)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - ASSERT_TRUE(WriteAll(sockfd1, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - ASSERT_TRUE(WriteAll(sockfd1, data + 4, 36)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); + butil::fd_guard sockfd1(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd1 >= 0); + ASSERT_EQ(0, connect(sockfd1, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(4, write(sockfd1, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(36, write(sockfd1, data + 4, 36)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); uint32_t flags = butil::HostToNet32(1); - ASSERT_TRUE(WriteAll(sockfd1, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, RdmaTransportOf(s)); - sockfd1.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + ASSERT_EQ(sizeof(flags), write(sockfd1, &flags, sizeof(flags))); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ESTABLISHED, + AdapterTransport::Get(s)->handshake_phase()); + close(sockfd1); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -679,40 +784,68 @@ TEST_F(RdmaTest, client_close_during_ack_send) { TEST_F(RdmaTest, client_close_after_ack_send) { StartServer(); + sockaddr_in addr; + bzero((char*)&addr, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_port = htons(PORT); Socket* s = nullptr; + rdma::v2_wire::HelloMessage msg{}; uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - MakeV2ClientHello(data); + msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; + msg.hello_ver = rdma::HELLO_V2_VERSION; + msg.impl_ver = rdma::IMPL_V2_VERSION; + msg.sq_size = 16; + msg.rq_size = 16; + msg.block_size = 8192; + msg.qp_num = 0; + msg.gid = rdma::GetRdmaGid(); + memcpy(data, "RDMA", 4); + msg.Serialize(data + 4); - butil::fd_guard sockfd1; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd1)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - ASSERT_TRUE(WriteAll(sockfd1, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - ASSERT_TRUE(WriteAll(sockfd1, data + 4, 36)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); + butil::fd_guard sockfd1(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd1 >= 0); + ASSERT_EQ(0, connect(sockfd1, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(4, write(sockfd1, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(36, write(sockfd1, data + 4, 36)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); uint32_t flags = butil::HostToNet32(0); - ASSERT_TRUE(WriteAll(sockfd1, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); - ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransportOf(s)->_rdma_state); - sockfd1.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - butil::fd_guard sockfd2; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd2)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - ASSERT_TRUE(WriteAll(sockfd2, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - ASSERT_TRUE(WriteAll(sockfd2, data + 4, 36)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(sizeof(flags), write(sockfd1, &flags, sizeof(flags))); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransport::Get(s)->_rdma_state); + close(sockfd1); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + butil::fd_guard sockfd2(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd2 >= 0); + ASSERT_EQ(0, connect(sockfd2, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(4, write(sockfd2, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(36, write(sockfd2, data + 4, 36)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); flags = butil::HostToNet32(1); - ASSERT_TRUE(WriteAll(sockfd2, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, RdmaTransportOf(s)); - sockfd2.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + ASSERT_EQ(sizeof(flags), write(sockfd2, &flags, sizeof(flags))); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ESTABLISHED, + AdapterTransport::Get(s)->handshake_phase()); + close(sockfd2); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -720,48 +853,71 @@ TEST_F(RdmaTest, client_close_after_ack_send) { TEST_F(RdmaTest, client_send_data_on_tcp_after_ack_send) { StartServer(); + sockaddr_in addr; + bzero((char*)&addr, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_port = htons(PORT); Socket* s = nullptr; + rdma::v2_wire::HelloMessage msg{}; uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - MakeV2ClientHello(data); + msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; + msg.hello_ver = rdma::HELLO_V2_VERSION; + msg.impl_ver = rdma::IMPL_V2_VERSION; + msg.sq_size = 16; + msg.rq_size = 16; + msg.block_size = 8192; + msg.qp_num = 0; + msg.gid = rdma::GetRdmaGid(); + memcpy(data, "RDMA", 4); + msg.Serialize(data + 4); - butil::fd_guard sockfd1; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd1)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - ASSERT_TRUE(WriteAll(sockfd1, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - ASSERT_TRUE(WriteAll(sockfd1, data + 4, 36)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); + butil::fd_guard sockfd1(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd1 >= 0); + ASSERT_EQ(0, connect(sockfd1, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(4, write(sockfd1, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(36, write(sockfd1, data + 4, 36)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); uint32_t flags = butil::HostToNet32(0); - ASSERT_TRUE(WriteAll(sockfd1, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); - // 4 more bytes on a fd that fell back to TCP are not a protocol baidu_std - // knows, so the connection is dropped. - ASSERT_TRUE(WriteAll(sockfd1, &flags, sizeof(flags))); - ASSERT_TRUE(WaitForServerSocketGone()); - - butil::fd_guard sockfd2; - ASSERT_NO_FATAL_FAILURE(ConnectToServer(&sockfd2)); - s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - ASSERT_TRUE(WriteAll(sockfd2, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - ASSERT_TRUE(WriteAll(sockfd2, data + 4, 36)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(sizeof(flags), write(sockfd1, &flags, sizeof(flags))); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(sizeof(flags), write(sockfd1, &flags, sizeof(flags))); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + butil::fd_guard sockfd2(socket(AF_INET, SOCK_STREAM, 0)); + ASSERT_TRUE(sockfd2 >= 0); + ASSERT_EQ(0, connect(sockfd2, (sockaddr*)&addr, sizeof(sockaddr))); + usleep(100000); // wait for server to handle the msg + s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(4, write(sockfd2, data, 4)); // Write magic string. + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(36, write(sockfd2, data + 4, 36)); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); flags = butil::HostToNet32(1); - ASSERT_TRUE(WriteAll(sockfd2, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, RdmaTransportOf(s)); - // Once RDMA is on the fd carries no RPC data at all, so this is an error. - ASSERT_TRUE(WriteAll(sockfd2, &flags, sizeof(flags))); - ASSERT_TRUE(WaitForServerSocketGone()); + ASSERT_EQ(sizeof(flags), write(sockfd2, &flags, sizeof(flags))); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(handshake::ESTABLISHED, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(sizeof(flags), write(sockfd2, &flags, sizeof(flags))); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } -// Connect, push a well-formed v2 hello and read back the server's reply, which -// leaves the server in S_ACK_WAIT waiting for the 4B ACK. static void HandshakeUntilAckWait(butil::fd_guard* sockfd) { ASSERT_NO_FATAL_FAILURE(ConnectToServer(sockfd)); @@ -774,9 +930,6 @@ static void HandshakeUntilAckWait(butil::fd_guard* sockfd) { ASSERT_TRUE(ReadAll(*sockfd, reply, sizeof(reply))); } -// A client is free to pipeline its first request right behind the handshake -// ACK. Only the 4B ACK belongs to the handshake. Whatever follows it must be -// handed over to the real protocol instead of dropping the connection. TEST_F(RdmaTest, server_accepts_data_pipelined_behind_fallback_ack) { StartServer(); @@ -785,7 +938,7 @@ TEST_F(RdmaTest, server_accepts_data_pipelined_behind_fallback_ack) { Socket* s = WaitForServerSocket(); ASSERT_TRUE(s != nullptr); auto* transport = RdmaTransportOf(s); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, transport); + ASSERT_HANDSHAKE_PHASE(handshake::ACK_WAIT, s); // An ACK asking for TCP, plus the first 4 bytes of a baidu_std request. One // write, so that both end up in the same read on the server. @@ -799,7 +952,7 @@ TEST_F(RdmaTest, server_accepts_data_pipelined_behind_fallback_ack) { // now waiting for the rest of its 12B header. So the connection lives on // with those 4 bytes still buffered. Note that baidu_std gets them a moment // after the handshake gave up the stream, hence the wait. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, transport); + ASSERT_HANDSHAKE_PHASE(handshake::FALLBACK_TCP, s); ASSERT_EQ(RdmaTransport::RDMA_OFF, transport->_rdma_state); ASSERT_TRUE(GetSocketFromServer(0) != nullptr); ASSERT_TRUE(WaitForFdReadBuf(s, 4)); @@ -810,8 +963,6 @@ TEST_F(RdmaTest, server_accepts_data_pipelined_behind_fallback_ack) { StopServer(); } -// Once RDMA is on, the TCP fd is no longer an RPC channel, so bytes trailing -// the ACK can only be a protocol error. TEST_F(RdmaTest, server_rejects_data_pipelined_behind_rdma_ack) { StartServer(); @@ -819,8 +970,7 @@ TEST_F(RdmaTest, server_rejects_data_pipelined_behind_rdma_ack) { ASSERT_NO_FATAL_FAILURE(HandshakeUntilAckWait(&sockfd)); Socket* s = WaitForServerSocket(); ASSERT_TRUE(s != nullptr); - auto* transport = RdmaTransportOf(s); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, transport); + ASSERT_HANDSHAKE_PHASE(handshake::ACK_WAIT, s); uint8_t ack_and_data[rdma::HELLO_ACK_LEN + 4]; const uint32_t flags = butil::HostToNet32(rdma::HELLO_ACK_RDMA_OK); @@ -835,7 +985,6 @@ TEST_F(RdmaTest, server_rejects_data_pipelined_behind_rdma_ack) { StopServer(); } -// Once RDMA is on, the server must stop parsing its TCP fd altogether. TEST_F(RdmaTest, server_stops_parsing_tcp_fd_once_rdma_is_on) { StartServer(); @@ -844,13 +993,13 @@ TEST_F(RdmaTest, server_stops_parsing_tcp_fd_once_rdma_is_on) { Socket* s = WaitForServerSocket(); ASSERT_TRUE(s != nullptr); auto* transport = RdmaTransportOf(s); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, transport); + ASSERT_HANDSHAKE_PHASE(handshake::ACK_WAIT, s); // A bare ACK asking for RDMA. Nothing trails it, so the handshake ends in // ESTABLISHED instead of being rejected (see the test above). const uint32_t flags = butil::HostToNet32(rdma::HELLO_ACK_RDMA_OK); ASSERT_TRUE(WriteAll(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, transport); + ASSERT_HANDSHAKE_PHASE(handshake::ESTABLISHED, s); ASSERT_EQ(RdmaTransport::RDMA_ON, transport->_rdma_state); ASSERT_TRUE(GetSocketFromServer(0) != nullptr); @@ -860,8 +1009,6 @@ TEST_F(RdmaTest, server_stops_parsing_tcp_fd_once_rdma_is_on) { StopServer(); } -// The same bytes on the stream carried by the QP are a real RPC, and the handler -// must decline so that CutInputMessage() moves on to the protocol handlers. TEST_F(RdmaTest, server_parses_qp_stream_after_rdma_is_on) { StartServer(); @@ -870,11 +1017,11 @@ TEST_F(RdmaTest, server_parses_qp_stream_after_rdma_is_on) { Socket* s = WaitForServerSocket(); ASSERT_TRUE(s != nullptr); auto* transport = RdmaTransportOf(s); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, transport); + ASSERT_HANDSHAKE_PHASE(handshake::ACK_WAIT, s); const uint32_t flags = butil::HostToNet32(rdma::HELLO_ACK_RDMA_OK); ASSERT_TRUE(WriteAll(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, transport); + ASSERT_HANDSHAKE_PHASE(handshake::ESTABLISHED, s); InputMessengerProcessor& qp_stream = transport->_rdma_ep->_input_processor; ASSERT_TRUE(qp_stream.read_buf().empty()); @@ -891,15 +1038,12 @@ TEST_F(RdmaTest, server_parses_qp_stream_after_rdma_is_on) { ASSERT_EQ((int)PROTOCOL_BAIDU_STD, s->preferred_index()); ASSERT_EQ(4u, qp_stream.read_buf().size()); ASSERT_TRUE(s->fd_input_processor().read_buf().empty()); - ASSERT_EQ(rdma::RdmaEndpoint::ESTABLISHED, transport->_rdma_ep->_state); + ASSERT_EQ(handshake::ESTABLISHED, AdapterTransport::Get(s)->handshake_phase()); ASSERT_FALSE(s->Failed()); StopServer(); } -// After the handshake is over, CutInputMessage() still offers the data to every -// registered handler, this one included. It must decline instead of reading the -// data as a fresh client hello. TEST_F(RdmaTest, server_declines_handshake_bytes_after_fallback) { StartServer(); @@ -907,11 +1051,10 @@ TEST_F(RdmaTest, server_declines_handshake_bytes_after_fallback) { ASSERT_NO_FATAL_FAILURE(HandshakeUntilAckWait(&sockfd)); Socket* s = WaitForServerSocket(); ASSERT_TRUE(s != nullptr); - auto* transport = RdmaTransportOf(s); const uint32_t flags = butil::HostToNet32(0); ASSERT_TRUE(WriteAll(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, transport); + ASSERT_HANDSHAKE_PHASE(handshake::FALLBACK_TCP, s); // Replay a valid hello. baidu_std rejects it and no other protocol claims // it, so the connection is dropped. What must NOT happen is a second @@ -948,7 +1091,7 @@ TEST_F(RdmaTest, fd_and_qp_input_streams_are_separate) { // fd stream, and only there. ASSERT_TRUE(WriteAll(sockfd, "RD", 2)); ASSERT_TRUE(WaitForFdReadBuf(s, 2)); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, transport->_rdma_ep->_state); + ASSERT_EQ(handshake::UNINITIALIZED, AdapterTransport::Get(s)->handshake_phase()); ASSERT_EQ(2u, fd_stream.read_buf().size()); ASSERT_TRUE(qp_stream.read_buf().empty()); @@ -977,9 +1120,10 @@ TEST_F(RdmaTest, server_miss_before_hello_send) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); @@ -1007,16 +1151,19 @@ TEST_F(RdmaTest, server_close_before_hello_send) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); close(acc_fd); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FAILED, RdmaTransportOf(s)); + usleep(100000); + ASSERT_EQ(handshake::FAILED, AdapterTransport::Get(s.get())->handshake_phase()); bthread_id_join(cntl.call_id()); ASSERT_EQ(EEOF, cntl.ErrorCode()); @@ -1041,18 +1188,18 @@ TEST_F(RdmaTest, server_miss_during_magic_str) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); - ASSERT_TRUE(WriteAll(acc_fd, "RD", 2)); - // Half a magic is not enough to decide anything, so the client stays stuck - // in the handshake read and the RPC runs into its timeout. Joining below - // waits for exactly that, no sleeping needed. + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(2, write(acc_fd, "RD", 2)); + usleep(100000); bthread_id_join(cntl.call_id()); ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); @@ -1077,19 +1224,21 @@ TEST_F(RdmaTest, server_close_during_magic_str) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); - // Half a magic and then EOF. TCP keeps the order, so the client always sees - // the two bytes first and then the close, which is what this test is about. - ASSERT_TRUE(WriteAll(acc_fd, "RD", 2)); - acc_fd.reset(-1); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FAILED, RdmaTransportOf(s)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(2, write(acc_fd, "RD", 2)); + usleep(100000); + close(acc_fd); + usleep(100000); + ASSERT_EQ(handshake::FAILED, AdapterTransport::Get(s.get())->handshake_phase()); bthread_id_join(cntl.call_id()); ASSERT_EQ(EEOF, cntl.ErrorCode()); @@ -1114,16 +1263,19 @@ TEST_F(RdmaTest, server_hello_invalid_magic_str) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); ASSERT_EQ(4, write(acc_fd, "ABCD", 4)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FAILED, RdmaTransportOf(s)); + usleep(100000); + ASSERT_EQ(handshake::FAILED, AdapterTransport::Get(s.get())->handshake_phase()); bthread_id_join(cntl.call_id()); ASSERT_EQ(EPROTO, cntl.ErrorCode()); @@ -1148,16 +1300,21 @@ TEST_F(RdmaTest, server_miss_during_hello_msg) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); ASSERT_EQ(4, write(acc_fd, "RDMA", 4)); - ASSERT_EQ(2, write(acc_fd, "00", 2)); + const uint16_t msg_len = butil::HostToNet16( + static_cast(rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(static_cast(sizeof(msg_len)), + write(acc_fd, &msg_len, sizeof(msg_len))); bthread_id_join(cntl.call_id()); ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); @@ -1182,18 +1339,24 @@ TEST_F(RdmaTest, server_close_during_hello_msg) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); ASSERT_EQ(4, write(acc_fd, "RDMA", 4)); - ASSERT_EQ(2, write(acc_fd, "00", 2)); + const uint16_t msg_len = butil::HostToNet16( + static_cast(rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(static_cast(sizeof(msg_len)), + write(acc_fd, &msg_len, sizeof(msg_len))); close(acc_fd); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FAILED, RdmaTransportOf(s)); + usleep(100000); + ASSERT_EQ(handshake::FAILED, AdapterTransport::Get(s.get())->handshake_phase()); bthread_id_join(cntl.call_id()); ASSERT_EQ(EEOF, cntl.ErrorCode()); @@ -1218,20 +1381,24 @@ TEST_F(RdmaTest, server_hello_invalid_msg_len) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); memcpy(data, "RDMA", 4); uint16_t len = butil::HostToNet16(35); memcpy(data + 4, &len, 2); memset(data + 6, 0, 32); - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FAILED, RdmaTransportOf(s)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + usleep(100000); + ASSERT_EQ(handshake::FAILED, AdapterTransport::Get(s.get())->handshake_phase()); bthread_id_join(cntl.call_id()); ASSERT_EQ(EPROTO, cntl.ErrorCode()); @@ -1256,20 +1423,25 @@ TEST_F(RdmaTest, server_hello_invalid_version) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); memcpy(data, "RDMA", 4); uint16_t len = butil::HostToNet16(rdma::HELLO_V2_MSG_LEN_MIN); memcpy(data + 4, &len, 2); memset(data + 6, 0, 32); - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s.get())->handshake_phase()); ASSERT_EQ(4, read(acc_fd, data, 4)); uint32_t* tmp = (uint32_t*)data; ASSERT_EQ(0, butil::NetToHost32(*tmp)); @@ -1297,14 +1469,16 @@ TEST_F(RdmaTest, server_hello_invalid_sq_rq_size) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); rdma::v2_wire::HelloMessage msg{}; msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; @@ -1317,9 +1491,12 @@ TEST_F(RdmaTest, server_hello_invalid_sq_rq_size) { msg.gid = rdma::GetRdmaGid(); memcpy(data, "RDMA", 4); msg.Serialize(data + 4); - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s.get())->handshake_phase()); ASSERT_EQ(4, read(acc_fd, data, 4)); uint32_t* tmp = (uint32_t*)data; ASSERT_EQ(0, butil::NetToHost32(*tmp)); @@ -1347,14 +1524,16 @@ TEST_F(RdmaTest, server_miss_after_ack) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); rdma::v2_wire::HelloMessage msg{}; msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; @@ -1367,9 +1546,12 @@ TEST_F(RdmaTest, server_miss_after_ack) { msg.gid = rdma::GetRdmaGid(); memcpy(data, "RDMA", 4); msg.Serialize(data + 4); - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, RdmaTransportOf(s)); + usleep(100000); + ASSERT_EQ(handshake::ESTABLISHED, + AdapterTransport::Get(s.get())->handshake_phase()); ASSERT_EQ(4, read(acc_fd, data, 4)); uint32_t* tmp = (uint32_t*)data; ASSERT_EQ(1, butil::NetToHost32(*tmp)); @@ -1397,14 +1579,16 @@ TEST_F(RdmaTest, server_close_after_ack) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); rdma::v2_wire::HelloMessage msg{}; msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; @@ -1417,9 +1601,12 @@ TEST_F(RdmaTest, server_close_after_ack) { msg.gid = rdma::GetRdmaGid(); memcpy(data, "RDMA", 4); msg.Serialize(data + 4); - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, RdmaTransportOf(s)); + usleep(100000); + ASSERT_EQ(handshake::ESTABLISHED, + AdapterTransport::Get(s.get())->handshake_phase()); ASSERT_EQ(4, read(acc_fd, data, 4)); uint32_t* tmp = (uint32_t*)data; ASSERT_EQ(1, butil::NetToHost32(*tmp)); @@ -1448,14 +1635,16 @@ TEST_F(RdmaTest, server_send_data_on_tcp_after_ack) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(handshake::HELLO_WAIT, AdapterTransport::Get(s.get())->handshake_phase()); butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); rdma::v2_wire::HelloMessage msg{}; msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; @@ -1468,16 +1657,19 @@ TEST_F(RdmaTest, server_send_data_on_tcp_after_ack) { msg.gid = rdma::GetRdmaGid(); memcpy(data, "RDMA", 4); msg.Serialize(data + 4); - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, RdmaTransportOf(s)); - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + usleep(100000); + ASSERT_EQ(handshake::ESTABLISHED, + AdapterTransport::Get(s.get())->handshake_phase()); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + write(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); bthread_id_join(cntl.call_id()); ASSERT_EQ(EPROTO, cntl.ErrorCode()); } - TEST_F(RdmaTest, v2_client_hello_bytes_baseline) { butil::fd_guard sockfd(butil::tcp_listen(g_ep)); EXPECT_TRUE(sockfd >= 0); @@ -1497,6 +1689,7 @@ TEST_F(RdmaTest, v2_client_hello_bytes_baseline) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); @@ -1504,7 +1697,8 @@ TEST_F(RdmaTest, v2_client_hello_bytes_baseline) { ASSERT_TRUE(acc_fd >= 0); uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(acc_fd, data, rdma::HELLO_V2_MSG_LEN_MIN)); // [0..4) magic ASSERT_EQ(0, memcmp(data, "RDMA", 4)); @@ -1522,7 +1716,7 @@ TEST_F(RdmaTest, v2_client_hello_bytes_baseline) { msg.Deserialize(data + 4); ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, msg.msg_len); ASSERT_EQ(rdma::HELLO_V2_VERSION, msg.hello_ver); - ASSERT_EQ(rdma::IMPL_V2_VERSION, msg.impl_ver); + ASSERT_EQ(rdma::IMPL_V2_VERSION, msg.impl_ver); bthread_id_join(cntl.call_id()); } @@ -1538,10 +1732,12 @@ TEST_F(RdmaTest, v2_server_hello_bytes_baseline) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + usleep(100000); + Socket* s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); - // Send a well-formed v2 hello so the server enters S_ACK_WAIT. + // Send a well-formed v2 hello so the server enters the common ACK_WAIT. rdma::v2_wire::HelloMessage msg{}; msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; msg.hello_ver = rdma::HELLO_V2_VERSION; @@ -1555,12 +1751,15 @@ TEST_F(RdmaTest, v2_server_hello_bytes_baseline) { uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; memcpy(data, "RDMA", 4); msg.Serialize(data + 4); - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, write(sockfd, data, rdma::HELLO_V2_MSG_LEN_MIN)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + write(sockfd, data, rdma::HELLO_V2_MSG_LEN_MIN)); + usleep(100000); + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); // Read server's reply hello and assert its byte-level layout. uint8_t reply[rdma::HELLO_V2_MSG_LEN_MIN]; - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, read(sockfd, reply, rdma::HELLO_V2_MSG_LEN_MIN)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + read(sockfd, reply, rdma::HELLO_V2_MSG_LEN_MIN)); ASSERT_EQ(0, memcmp(reply, "RDMA", 4)); ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, @@ -1574,21 +1773,24 @@ TEST_F(RdmaTest, v2_server_hello_bytes_baseline) { reply_msg.Deserialize(reply + 4); ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, reply_msg.msg_len); ASSERT_EQ(rdma::HELLO_V2_VERSION, reply_msg.hello_ver); - ASSERT_EQ(rdma::IMPL_V2_VERSION, reply_msg.impl_ver); + ASSERT_EQ(rdma::IMPL_V2_VERSION, reply_msg.impl_ver); // Drive the server into FALLBACK_TCP via ACK flags=0 so the test ends // cleanly without requiring real RDMA hardware. uint32_t flags = butil::HostToNet32(0); ASSERT_EQ(sizeof(flags), write(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } -TEST_F(RdmaTest, v2_server_drains_tail_then_reads_ack) { +TEST_F(RdmaTest, v2_server_preserves_coalesced_ack_after_extension) { StartServer(); sockaddr_in addr; @@ -1598,9 +1800,11 @@ TEST_F(RdmaTest, v2_server_drains_tail_then_reads_ack) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + usleep(100000); + Socket* s = GetSocketFromServer(0); + ASSERT_TRUE(s != NULL); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); // Build a v2 hello with msg_len = 48 (40 base + 8B zero tail). rdma::v2_wire::HelloMessage msg{}; @@ -1613,23 +1817,23 @@ TEST_F(RdmaTest, v2_server_drains_tail_then_reads_ack) { msg.qp_num = 0; msg.gid = rdma::GetRdmaGid(); - uint8_t buf[48]; + uint8_t buf[52]; memcpy(buf, "RDMA", 4); msg.Serialize(buf + 4); memset(buf + 40, 0x00, 8); // 8B zero tail - ASSERT_TRUE(WriteAll(sockfd, buf, 48)); - // The tail is drained as part of the hello, so the server ends up waiting - // for the ACK. Wait for that before sending it, otherwise the ACK could - // ride along in the same read and this would no longer test the drain. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); - - // Send the real ACK (flags=1 = ACK_MSG_RDMA_OK). + // Coalesce the real ACK with the hello. The v2 parser must consume only + // msg_len bytes and preserve the ACK for the next handshake step. uint32_t flags = butil::HostToNet32(1); - ASSERT_EQ(sizeof(flags), write(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, RdmaTransportOf(s)); + memcpy(buf + 48, &flags, sizeof(flags)); + ASSERT_EQ(sizeof(buf), write(sockfd, buf, sizeof(buf))); + usleep(100000); + + ASSERT_EQ(handshake::ESTABLISHED, + AdapterTransport::Get(s)->handshake_phase()); sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -1644,9 +1848,11 @@ TEST_F(RdmaTest, v2_server_rejects_oversized_msg_len) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); + usleep(100000); + Socket* s = GetSocketFromServer(0); ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); // Build a v2 hello with msg_len = 4097 (HELLO_V2_MSG_LEN_MAX + 1). // We only send the 40B base; the server must reject before reading @@ -1664,10 +1870,14 @@ TEST_F(RdmaTest, v2_server_rejects_oversized_msg_len) { uint8_t buf[rdma::HELLO_V2_MSG_LEN_MIN]; memcpy(buf, "RDMA", 4); msg.Serialize(buf + 4); - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, write(sockfd, buf, rdma::HELLO_V2_MSG_LEN_MIN)); - ASSERT_TRUE(WaitForServerSocketGone()); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + write(sockfd, buf, rdma::HELLO_V2_MSG_LEN_MIN)); + usleep(100000); + + ASSERT_EQ(nullptr, GetSocketFromServer(0)); sockfd.reset(-1); + usleep(100000); StopServer(); } @@ -1684,6 +1894,7 @@ class HandshakeVersionFlag { ~HandshakeVersionFlag() { rdma::FLAGS_rdma_client_handshake_version = _saved; } + private: int _saved; }; @@ -1695,8 +1906,7 @@ std::string MakeV3Packet(const rdma::RdmaHello& msg) { std::string packet; packet.reserve(4 + 4 + body.size()); packet.append("RDM3", 4); - uint32_t pb_size_be = - butil::HostToNet32(static_cast(body.size())); + uint32_t pb_size_be = butil::HostToNet32(static_cast(body.size())); packet.append(reinterpret_cast(&pb_size_be), 4); packet.append(body); return packet; @@ -1715,13 +1925,12 @@ rdma::RdmaHello MakeValidV3Hello() { msg.set_rq_size(16); msg.set_lid(0); ibv_gid gid = rdma::GetRdmaGid(); - msg.set_gid(std::string(reinterpret_cast(gid.raw), - sizeof(gid.raw))); + msg.set_gid( + std::string(reinterpret_cast(gid.raw), sizeof(gid.raw))); msg.set_qp_num(0); return msg; } - TEST_F(RdmaTest, v3_client_hello_bytes_baseline) { HandshakeVersionFlag _hsv(3); @@ -1791,14 +2000,17 @@ TEST_F(RdmaTest, v3_server_hello_bytes_baseline) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + usleep(100000); + Socket* s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); // Send a valid v3 hello. std::string packet = MakeV3Packet(MakeValidV3Hello()); ASSERT_EQ((ssize_t)packet.size(), write(sockfd, packet.data(), packet.size())); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); + usleep(100000); + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); // Read server's reply hello: 4B magic + 4B pb_size + body. uint8_t reply_magic[4]; @@ -1825,12 +2037,14 @@ TEST_F(RdmaTest, v3_server_hello_bytes_baseline) { // Drive the server into FALLBACK_TCP via ACK flags=0 so the test ends // cleanly without requiring real RDMA hardware. uint32_t flags = butil::HostToNet32(0); - ASSERT_EQ((ssize_t)sizeof(flags), - write(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + ASSERT_EQ((ssize_t)sizeof(flags), write(sockfd, &flags, sizeof(flags))); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -1845,13 +2059,16 @@ TEST_F(RdmaTest, v3_server_rejects_zero_pb_size) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); + usleep(100000); + Socket* s = GetSocketFromServer(0); ASSERT_TRUE(s != nullptr); // "RDM3" + pb_size = 0 (4B big-endian zero). uint8_t buf[8] = {'R', 'D', 'M', '3', 0, 0, 0, 0}; ASSERT_EQ(8, write(sockfd, buf, 8)); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); + + ASSERT_EQ(nullptr, GetSocketFromServer(0)); sockfd.reset(-1); StopServer(); @@ -1867,7 +2084,8 @@ TEST_F(RdmaTest, v3_server_rejects_oversized_pb_size) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); + usleep(100000); + Socket* s = GetSocketFromServer(0); ASSERT_TRUE(s != nullptr); uint8_t buf[8]; @@ -1877,7 +2095,9 @@ TEST_F(RdmaTest, v3_server_rejects_oversized_pb_size) { butil::HostToNet32(static_cast(rdma::HELLO_V3_MAX_PB_SIZE + 1)); memcpy(buf + 4, &pb_size_be, 4); ASSERT_EQ(8, write(sockfd, buf, 8)); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); + + ASSERT_EQ(nullptr, GetSocketFromServer(0)); sockfd.reset(-1); StopServer(); @@ -1893,7 +2113,8 @@ TEST_F(RdmaTest, v3_server_rejects_invalid_pb_bytes) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); + usleep(100000); + Socket* s = GetSocketFromServer(0); ASSERT_TRUE(s != nullptr); // "RDM3" + pb_size = 8 + 8 bytes of 0xff (invalid protobuf body). @@ -1903,7 +2124,9 @@ TEST_F(RdmaTest, v3_server_rejects_invalid_pb_bytes) { memcpy(buf + 4, &pb_size_be, 4); memset(buf + 8, 0xff, 8); ASSERT_EQ(16, write(sockfd, buf, 16)); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); + + ASSERT_EQ(nullptr, GetSocketFromServer(0)); sockfd.reset(-1); StopServer(); @@ -1919,38 +2142,43 @@ TEST_F(RdmaTest, v3_server_invalid_sq_size_falls_back) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); + usleep(100000); + Socket* s = GetSocketFromServer(0); ASSERT_TRUE(s != nullptr); rdma::RdmaHello msg = MakeValidV3Hello(); msg.set_sq_size(0); // invalid: < MIN_QP_SIZE (16) std::string packet = MakeV3Packet(msg); - ASSERT_TRUE(WriteAll(sockfd, packet.data(), packet.size())); + ASSERT_EQ((ssize_t)packet.size(), + write(sockfd, packet.data(), packet.size())); + usleep(100000); // Server validated the hello as invalid -> _rdma_state = RDMA_OFF, - // but still proceeds to S_ACK_WAIT (sends its own reply hello). - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); - ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransportOf(s)->_rdma_state); + // but still proceeds to the common ACK_WAIT (sends its own reply hello). + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransport::Get(s)->_rdma_state); // Drain server's reply hello (content not asserted here; covered // by v3_server_hello_bytes_baseline). uint8_t reply_hdr[8]; ASSERT_EQ(8, read(sockfd, reply_hdr, 8)); ASSERT_EQ(0, memcmp(reply_hdr, "RDM3", 4)); - uint32_t reply_pb_size = butil::NetToHost32( - *reinterpret_cast(reply_hdr + 4)); + uint32_t reply_pb_size = + butil::NetToHost32(*reinterpret_cast(reply_hdr + 4)); std::string reply_body(reply_pb_size, '\0'); ASSERT_EQ((ssize_t)reply_pb_size, read(sockfd, &reply_body[0], reply_pb_size)); // Client ACK flags=0 -> server settles into FALLBACK_TCP. uint32_t flags = butil::HostToNet32(0); - ASSERT_EQ((ssize_t)sizeof(flags), - write(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + ASSERT_EQ((ssize_t)sizeof(flags), write(sockfd, &flags, sizeof(flags))); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -1961,16 +2189,14 @@ class EceFlagGuard { explicit EceFlagGuard(bool v) : _saved(rdma::FLAGS_rdma_ece) { rdma::FLAGS_rdma_ece = v; } - ~EceFlagGuard() { - rdma::FLAGS_rdma_ece = _saved; - } + ~EceFlagGuard() { rdma::FLAGS_rdma_ece = _saved; } + private: bool _saved; }; // Build a valid v3 hello that also carries an ECE block. -rdma::RdmaHello MakeValidV3HelloWithEce(uint32_t vendor_id, - uint32_t options, +rdma::RdmaHello MakeValidV3HelloWithEce(uint32_t vendor_id, uint32_t options, uint32_t comp_mask) { rdma::RdmaHello msg = MakeValidV3Hello(); rdma::RdmaEce* ece = msg.mutable_ece(); @@ -1986,17 +2212,18 @@ static void ReadServerV3Reply(int fd, rdma::RdmaHello* reply) { uint8_t reply_hdr[8]; ASSERT_EQ(8, read(fd, reply_hdr, 8)); ASSERT_EQ(0, memcmp(reply_hdr, "RDM3", 4)); - uint32_t reply_pb_size = butil::NetToHost32(*reinterpret_cast(reply_hdr + 4)); + uint32_t reply_pb_size = + butil::NetToHost32(*reinterpret_cast(reply_hdr + 4)); ASSERT_GT(reply_pb_size, 0u); ASSERT_LE(reply_pb_size, 4096u); std::string reply_body(reply_pb_size, '\0'); - ASSERT_EQ((ssize_t)reply_pb_size, - read(fd, &reply_body[0], reply_pb_size)); + ASSERT_EQ((ssize_t)reply_pb_size, read(fd, &reply_body[0], reply_pb_size)); ASSERT_TRUE(reply->ParseFromString(reply_body)); } // A client hello carrying ECE must not break the server handshake: with ECE -// enabled the server still parses the hello and advances to S_ACK_WAIT. +// enabled the server still parses the hello and advances to the common +// ACK_WAIT. TEST_F(RdmaTest, v3_server_accepts_client_hello_with_ece) { EceFlagGuard ece_flag_guard(true); StartServer(); @@ -2008,14 +2235,17 @@ TEST_F(RdmaTest, v3_server_accepts_client_hello_with_ece) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); + usleep(100000); + Socket* s = GetSocketFromServer(0); ASSERT_TRUE(s != nullptr); rdma::RdmaHello msg = MakeValidV3HelloWithEce(0x02c9, 0x1, 0x0); std::string packet = MakeV3Packet(msg); ASSERT_EQ((ssize_t)packet.size(), write(sockfd, packet.data(), packet.size())); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); + usleep(100000); + + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); rdma::RdmaHello reply; ReadServerV3Reply(sockfd, &reply); @@ -2023,10 +2253,13 @@ TEST_F(RdmaTest, v3_server_accepts_client_hello_with_ece) { // ACK flags=0 -> clean FALLBACK_TCP so the test ends without hardware. uint32_t flags = butil::HostToNet32(0); ASSERT_EQ((ssize_t)sizeof(flags), write(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -2043,30 +2276,33 @@ TEST_F(RdmaTest, v3_server_reply_has_no_ece_when_disabled) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); + usleep(100000); + Socket* s = GetSocketFromServer(0); ASSERT_TRUE(s != nullptr); rdma::RdmaHello msg = MakeValidV3HelloWithEce(0x02c9, 0x1, 0x0); std::string packet = MakeV3Packet(msg); - ASSERT_TRUE(WriteAll(sockfd, packet.data(), packet.size())); + ASSERT_EQ((ssize_t)packet.size(), + write(sockfd, packet.data(), packet.size())); + usleep(100000); - // Reading the reply in full doubles as the synchronization point. rdma::RdmaHello reply; ReadServerV3Reply(sockfd, &reply); EXPECT_FALSE(reply.has_ece()); uint32_t flags = butil::HostToNet32(0); - ASSERT_TRUE(WriteAll(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + ASSERT_EQ((ssize_t)sizeof(flags), write(sockfd, &flags, sizeof(flags))); + usleep(100000); sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); StopServer(); } // When ECE is enabled but there is no negotiated result (UT skips the real QP // bring-up, so the server never fills _outgoing_ece), the server reply must -// still NOT advertise ECE (FillLocalRdmaHello degrade branch #2 -> degrade-safe). +// still NOT advertise ECE (FillLocalRdmaHello degrade branch #2 -> +// degrade-safe). TEST_F(RdmaTest, v3_server_reply_has_no_ece_without_hw_negotiation) { EceFlagGuard ece_flag_guard(true); StartServer(); @@ -2078,24 +2314,26 @@ TEST_F(RdmaTest, v3_server_reply_has_no_ece_without_hw_negotiation) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); + usleep(100000); + Socket* s = GetSocketFromServer(0); ASSERT_TRUE(s != nullptr); rdma::RdmaHello msg = MakeValidV3HelloWithEce(0x02c9, 0x1, 0x0); std::string packet = MakeV3Packet(msg); - ASSERT_TRUE(WriteAll(sockfd, packet.data(), packet.size())); + ASSERT_EQ((ssize_t)packet.size(), + write(sockfd, packet.data(), packet.size())); + usleep(100000); - // Reading the reply in full doubles as the synchronization point. rdma::RdmaHello reply; ReadServerV3Reply(sockfd, &reply); EXPECT_FALSE(reply.has_ece()); uint32_t flags = butil::HostToNet32(0); - ASSERT_TRUE(WriteAll(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + ASSERT_EQ((ssize_t)sizeof(flags), write(sockfd, &flags, sizeof(flags))); + usleep(100000); sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); StopServer(); } @@ -2105,9 +2343,8 @@ class ResourceAllocFailGuard { : _saved(rdma::g_fail_resource_alloc_for_test) { rdma::g_fail_resource_alloc_for_test = v; } - ~ResourceAllocFailGuard() { - rdma::g_fail_resource_alloc_for_test = _saved; - } + ~ResourceAllocFailGuard() { rdma::g_fail_resource_alloc_for_test = _saved; } + private: bool _saved; }; @@ -2131,11 +2368,13 @@ TEST_F(RdmaTest, client_alloc_resource_fail_fallback_tcp) { req.set_sleep_us(200000); google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); - ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransportOf(s)->_rdma_state); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s.get())->handshake_phase()); + ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransport::Get(s)->_rdma_state); // The socket must not be failed, otherwise it can no longer carry TCP. ASSERT_FALSE(s->Failed()); @@ -2157,9 +2396,11 @@ TEST_F(RdmaTest, server_alloc_resource_fail_fallback_tcp) { butil::fd_guard sockfd(socket(AF_INET, SOCK_STREAM, 0)); ASSERT_TRUE(sockfd >= 0); ASSERT_EQ(0, connect(sockfd, (sockaddr*)&addr, sizeof(sockaddr))); - Socket* s = WaitForServerSocket(); + usleep(100000); // wait for server to handle the msg + Socket* s = GetSocketFromServer(0); ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); // Send a well-formed v2 hello: the negotiation succeeds // but the resource allocation does not. @@ -2178,18 +2419,22 @@ TEST_F(RdmaTest, server_alloc_resource_fail_fallback_tcp) { msg.Serialize(data + 4); ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, write(sockfd, data, rdma::HELLO_V2_MSG_LEN_MIN)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_ACK_WAIT, RdmaTransportOf(s)); - ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransportOf(s)->_rdma_state); + usleep(100000); + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); + ASSERT_EQ(RdmaTransport::RDMA_OFF, RdmaTransport::Get(s)->_rdma_state); ASSERT_FALSE(s->Failed()); // Ack without RDMA so that the server finishes the handshake in TCP mode. uint32_t flags = butil::HostToNet32(0); ASSERT_EQ(sizeof(flags), write(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); ASSERT_FALSE(s->Failed()); sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); StopServer(); } @@ -2213,9 +2458,11 @@ TEST_F(RdmaTest, try_global_disable_rdma) { req.set_sleep_us(200000); google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s.get())->handshake_phase()); bthread_id_join(cntl.call_id()); ASSERT_EQ(0, cntl.ErrorCode()); @@ -2312,13 +2559,7 @@ TEST_F(RdmaTest, channel_option_invalid) { ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); } -// Rounds, per-round RPC count and attachment sizes shared by the end-to-end -// tests below. One RPC per test leaves everything that only shows up on the -// second message untouched -- buffer reuse, an EOF read racing another writer -// of the same input stream, resource recycling. -static const int E2E_ROUND_NUM = 3; static const int E2E_RPC_NUM = 32; -static const size_t E2E_ATTACH_SIZE[] = { 0, 4096, 128 * 1024 }; static void ShutdownClientConnection(Controller& cntl) { SocketUniquePtr s; @@ -2342,7 +2583,7 @@ static int SendEchoRpcs(Channel& channel, int rpc_num, size_t attach_size, req[i].set_code(i + 1); if (attach_size > 0) { EXPECT_EQ(0, attach[i].resize( - attach_size, static_cast('a' + i % 26))); + attach_size, static_cast('a' + i % 26))); cntl[i].request_attachment().append(attach[i]); } ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], DoNothing()); @@ -2368,17 +2609,6 @@ static int SendEchoRpcs(Channel& channel, int rpc_num, size_t attach_size, return succeeded; } -static void SendEchoRpcsInRounds(Channel& channel) { - for (int round = 0; round < E2E_ROUND_NUM; ++round) { - for (size_t i = 0; i < arraysize(E2E_ATTACH_SIZE); ++i) { - ASSERT_EQ(E2E_RPC_NUM, - SendEchoRpcs(channel, E2E_RPC_NUM, E2E_ATTACH_SIZE[i])) - << "round=" << round - << " attach_size=" << E2E_ATTACH_SIZE[i]; - } - } -} - TEST_P(RdmaRpcTest, rdma_client_to_rdma_server) { if (!FLAGS_rdma_test_enable) { return; @@ -2390,10 +2620,18 @@ TEST_P(RdmaRpcTest, rdma_client_to_rdma_server) { ChannelOptions chan_options; chan_options.socket_mode = SOCKET_MODE_RDMA; chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 5000; + chan_options.timeout_ms = 500; chan_options.max_retry = 0; ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - ASSERT_NO_FATAL_FAILURE(SendEchoRpcsInRounds(channel)); + Controller cntl; + test::EchoRequest req; + test::EchoResponse res; + req.set_message(__FUNCTION__); + google::protobuf::Closure* done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + // usleep(100000); + bthread_id_join(cntl.call_id()); + ASSERT_EQ(0, cntl.ErrorCode()); StopServer(); } @@ -2404,10 +2642,18 @@ TEST_P(RdmaRpcTest, tcp_client_to_tcp_server) { Channel channel; ChannelOptions chan_options; chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 5000; + chan_options.timeout_ms = 500; chan_options.max_retry = 0; ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - ASSERT_NO_FATAL_FAILURE(SendEchoRpcsInRounds(channel)); + Controller cntl; + test::EchoRequest req; + test::EchoResponse res; + req.set_message(__FUNCTION__); + google::protobuf::Closure* done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); + bthread_id_join(cntl.call_id()); + ASSERT_EQ(0, cntl.ErrorCode()); StopServer(); } @@ -2418,10 +2664,18 @@ TEST_P(RdmaRpcTest, tcp_client_to_rdma_server) { Channel channel; ChannelOptions chan_options; chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 5000; + chan_options.timeout_ms = 500; chan_options.max_retry = 0; ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - ASSERT_NO_FATAL_FAILURE(SendEchoRpcsInRounds(channel)); + Controller cntl; + test::EchoRequest req; + test::EchoResponse res; + req.set_message(__FUNCTION__); + google::protobuf::Closure* done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); + bthread_id_join(cntl.call_id()); + ASSERT_EQ(0, cntl.ErrorCode()); StopServer(); } @@ -2433,10 +2687,18 @@ TEST_P(RdmaRpcTest, rdma_client_to_tcp_server) { ChannelOptions chan_options; chan_options.socket_mode = SOCKET_MODE_RDMA; chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 5000; + chan_options.timeout_ms = 500; chan_options.max_retry = 0; ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - ASSERT_NO_FATAL_FAILURE(SendEchoRpcsInRounds(channel)); + Controller cntl; + test::EchoRequest req; + test::EchoResponse res; + req.set_message(__FUNCTION__); + google::protobuf::Closure* done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); + bthread_id_join(cntl.call_id()); + ASSERT_FALSE(cntl.Failed()); StopServer(); } @@ -2453,13 +2715,12 @@ TEST_P(RdmaRpcTest, tcp_client_to_rdma_server_short_connection) { ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); for (int round = 0; round < 8; ++round) { ASSERT_EQ(E2E_RPC_NUM, SendEchoRpcs(channel, E2E_RPC_NUM, 4096)) - << "round=" << round; + << "round=" << round; } StopServer(); } -// Rounds of connection churn: a race needs attempts, not one well-timed shot. static const int CHURN_ROUND_NUM = 16; static const int CHURN_RPC_NUM = 64; static const size_t CHURN_ATTACH_SIZE = 32 * 1024; @@ -2507,10 +2768,12 @@ TEST_P(RdmaRpcTest, rdma_server_survives_connection_churn) { static const int RPC_NUM = 1024; void DumpRdmaEndpointInfo(Socket* client, Socket* server) { - std::cout << std::endl << "client:"; - static_cast(client->_transport.get())->_rdma_ep->DebugInfo(std::cout); - std::cout << std::endl << "server:"; - static_cast(server->_transport.get())->_rdma_ep->DebugInfo(std::cout); + std::cout << std::endl + << "client:"; + RdmaTransport::Get(client)->_rdma_ep->DebugInfo(std::cout); + std::cout << std::endl + << "server:"; + RdmaTransport::Get(server)->_rdma_ep->DebugInfo(std::cout); } TEST_P(RdmaRpcTest, send_rpcs_in_one_qp) { @@ -2584,8 +2847,8 @@ TEST_P(RdmaRpcTest, send_rpcs_in_one_qp) { Socket* m = GetSocketFromServer(0); DumpRdmaEndpointInfo(s.get(), m); } - ASSERT_TRUE(0 == cntl[i].ErrorCode() || - EOVERCROWDED == cntl[i].ErrorCode()) << "req[" << i << "] " << berror(cntl[i].ErrorCode()); + ASSERT_TRUE(0 == cntl[i].ErrorCode() || EOVERCROWDED == cntl[i].ErrorCode()) + << "req[" << i << "] " << berror(cntl[i].ErrorCode()); } SocketUniquePtr s; @@ -2638,7 +2901,8 @@ TEST_P(RdmaRpcTest, send_rpc_in_many_qp) { req[i].set_message(__FUNCTION__); cntl[i].request_attachment().append(attach); google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel[i]).Echo(&cntl[i], &req[i], &res[i], done); + ::test::EchoService::Stub(&channel[i]) + .Echo(&cntl[i], &req[i], &res[i], done); } for (int i = 0; i < RPC_NUM; ++i) { bthread_id_join(cntl[i].call_id()); @@ -2776,12 +3040,12 @@ TEST_P(RdmaRpcTest, server_stop_during_rpc) { for (int i = 0; i < RPC_NUM; ++i) { bthread_id_join(cntl[i].call_id()); - if (i == 0) StopServer(); + if (i == 0) + StopServer(); int error_code = cntl[i].ErrorCode(); - ASSERT_TRUE(error_code == 0 || - error_code == EEOF || - error_code == ELOGOFF || - error_code == EHOSTDOWN) << "req[" << i << "]: " << error_code; + ASSERT_TRUE(error_code == 0 || error_code == EEOF || + error_code == ELOGOFF || error_code == EHOSTDOWN) + << "req[" << i << "]: " << error_code; } } @@ -2818,10 +3082,9 @@ TEST_P(RdmaRpcTest, server_close_during_rpc) { for (int i = 0; i < RPC_NUM; ++i) { bthread_id_join(cntl[i].call_id()); int error_code = cntl[i].ErrorCode(); - ASSERT_TRUE(error_code == 0 || - error_code == EEOF || - error_code == EFAILEDSOCKET || - error_code == EHOSTDOWN) << "req[" << i << "]: " << error_code; + ASSERT_TRUE(error_code == 0 || error_code == EEOF || + error_code == EFAILEDSOCKET || error_code == EHOSTDOWN) + << "req[" << i << "]: " << error_code; } StopServer(); @@ -2859,9 +3122,9 @@ TEST_P(RdmaRpcTest, client_close_during_rpc) { for (int i = 0; i < RPC_NUM; ++i) { bthread_id_join(cntl[i].call_id()); int error_code = cntl[i].ErrorCode(); - ASSERT_TRUE(error_code == 0 || - error_code == ECLOSE || - error_code == EHOSTDOWN) << "req[" << i << "]: " << error_code; + ASSERT_TRUE(error_code == 0 || error_code == ECLOSE || + error_code == EHOSTDOWN) + << "req[" << i << "]: " << error_code; } StopServer(); @@ -2892,7 +3155,6 @@ TEST_P(RdmaRpcTest, rdma_client_close_during_rpc_repeatedly) { ASSERT_FALSE(HasFailure()) << "round=" << round; } - const int served = g_echo_served.load(butil::memory_order_relaxed) - served_before - CHURN_ROUND_NUM; LOG(INFO) << "server served " << served << " of " @@ -2930,10 +3192,10 @@ TEST_P(RdmaRpcTest, verbs_error_handling) { google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); + usleep(100000); // wait for rdma handshake complete + SocketUniquePtr s; ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - // The QP below only exists once the handshake is over. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, RdmaTransportOf(s)); ibv_send_wr wr; memset(&wr, 0, sizeof(wr)); ibv_sge sge; @@ -2944,7 +3206,7 @@ TEST_P(RdmaRpcTest, verbs_error_handling) { wr.sg_list = &sge; wr.num_sge = 1; ibv_send_wr* bad = nullptr; - auto rdma_transport = RdmaTransportOf(s); + auto rdma_transport = RdmaTransport::Get(s); ibv_post_send(rdma_transport->_rdma_ep->_resource->qp, &wr, &bad); bthread_id_join(cntl.call_id()); ASSERT_EQ(ERDMA, cntl.ErrorCode()); @@ -2967,9 +3229,8 @@ TEST_P(RdmaRpcTest, rdma_use_parallel_channel) { opts.socket_mode = SOCKET_MODE_RDMA; for (size_t i = 0; i < NCHANS; ++i) { ASSERT_EQ(0, subchans[i].Init(_naming_url.c_str(), "rR", &opts)); - ASSERT_EQ(0, channel.AddChannel( - &subchans[i], DOESNT_OWN_CHANNEL, - nullptr, nullptr)); + ASSERT_EQ(0, channel.AddChannel(&subchans[i], DOESNT_OWN_CHANNEL, nullptr, + nullptr)); } ASSERT_EQ(0, channel.Init(nullptr)); @@ -3015,7 +3276,7 @@ TEST_P(RdmaRpcTest, rdma_use_selective_channel) { StopServer(); } -static void MockFree(void* buf) { } +static void MockFree(void* buf) {} TEST_P(RdmaRpcTest, send_rpcs_with_user_defined_iobuf) { if (!FLAGS_rdma_test_enable) { @@ -3036,7 +3297,8 @@ TEST_P(RdmaRpcTest, send_rpcs_with_user_defined_iobuf) { test::EchoResponse res[RPC_NUM]; butil::IOBuf attach; - void* data = malloc(4096);; + void* data = malloc(4096); + ; attach.append_user_data(data, 4096, nullptr); req[0].set_message(__FUNCTION__); cntl[0].request_attachment().append(attach); @@ -3055,12 +3317,14 @@ TEST_P(RdmaRpcTest, send_rpcs_with_user_defined_iobuf) { memset(mr[2 * i], i % 100, 4096); lkey[2 * i] = rdma::RegisterMemoryForRdma(mr[2 * i], 4096); ASSERT_TRUE(lkey[2 * i] != 0); - cntl[i].request_attachment().append_user_data_with_meta(mr[2 * i] + i, 4096 - i, MockFree, lkey[2 * i]); + cntl[i].request_attachment().append_user_data_with_meta( + mr[2 * i] + i, 4096 - i, MockFree, lkey[2 * i]); mr[2 * i + 1] = (char*)malloc(4096); memset(mr[2 * i + 1], i % 100, 4096); lkey[2 * i + 1] = rdma::RegisterMemoryForRdma(mr[2 * i + 1], 4096); ASSERT_TRUE(lkey[2 * i + 1] != 0); - cntl[i].request_attachment().append_user_data_with_meta(mr[2 * i + 1] + i, 4096 - i, MockFree, lkey[2 * i + 1]); + cntl[i].request_attachment().append_user_data_with_meta( + mr[2 * i + 1] + i, 4096 - i, MockFree, lkey[2 * i + 1]); req[i].set_message(__FUNCTION__); google::protobuf::Closure* done = DoNothing(); ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); @@ -3126,12 +3390,10 @@ TEST_P(RdmaRpcTest, try_memory_pool_empty) { // The server always accepts both via magic-byte dispatch, so this // proves the upper-layer RPC paths behave identically under either // wire format. -INSTANTIATE_TEST_SUITE_P( - HandshakeVersion, RdmaRpcTest, - ::testing::Values(2, 3), - [](const ::testing::TestParamInfo& info) { - return std::string("v") + std::to_string(info.param); - }); +INSTANTIATE_TEST_SUITE_P(HandshakeVersion, RdmaRpcTest, ::testing::Values(2, 3), + [](const ::testing::TestParamInfo& info) { + return std::string("v") + std::to_string(info.param); + }); #endif // if BRPC_WITH_RDMA diff --git a/test/brpc_transport_handshake_unittest.cpp b/test/brpc_transport_handshake_unittest.cpp new file mode 100644 index 0000000000..2ed34445a1 --- /dev/null +++ b/test/brpc_transport_handshake_unittest.cpp @@ -0,0 +1,790 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include + +#include +#include +#include +#include +#include + +#include "bthread/bthread.h" +#include "butil/fd_guard.h" +#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" + +namespace brpc { +namespace handshake { + +class MemoryHandshakeIO : public HandshakeIO { +public: + explicit MemoryHandshakeIO(const std::string& input = std::string()) + : _input(input), _offset(0) {} + + int ReadExact(void* data, size_t len) override { + if (_input.size() - _offset < len) { + errno = EIO; + return -1; + } + memcpy(data, _input.data() + _offset, len); + _offset += len; + return 0; + } + + int WriteAll(const void* data, size_t len) override { + _output.append(static_cast(data), len); + return 0; + } + + int PushBack(const void* data, size_t len) override { + _pushed_back.append(static_cast(data), len); + return 0; + } + + const std::string& output() const { return _output; } + const std::string& pushed_back() const { return _pushed_back; } + +private: + std::string _input; + size_t _offset; + std::string _output; + std::string _pushed_back; +}; + +struct BlockingHandshakeRead { + SocketHandshakeIO* io; + int result; + int error; +}; + +static void* RunBlockingHandshakeRead(void* arg) { + BlockingHandshakeRead* read = + static_cast(arg); + char byte = 0; + read->result = read->io->ReadExact(&byte, sizeof(byte)); + read->error = errno; + 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, + FrameSpec::FIXED); +} + +class TestHandshakeProtocol : public HandshakeProtocol { +public: + explicit TestHandshakeProtocol(std::string* calls = NULL) + : _calls(calls), _with_extension(false), + _extension_size(0) {} + + int ProtocolVersion() const override { return 7; } + const FrameSpec& HelloFrameSpec() const override { + static const FrameSpec spec = FixedSpec("HS", 2, 4); + return spec; + } + const FrameSpec& AckFrameSpec() const override { + static const FrameSpec spec = FixedSpec(NULL, 0, 1); + return spec; + } + StepResult BuildHello(bool enabled, std::string* payload) override { + Append("build "); + *payload = enabled ? "LO" : "NO"; + return STEP_OK; + } + StepResult ParseHello(const std::string& payload) override { + Append("parse "); + return payload == "OK" ? STEP_OK : STEP_FALLBACK; + } + StepResult BuildAck(bool enabled, std::string* payload) override { + Append(enabled ? "ack1 " : "ack0 "); + *payload = enabled ? "1" : "0"; + return STEP_OK; + } + StepResult ParseAck( + const std::string& payload, bool* enabled) override { + Append("parse_ack "); + *enabled = payload == "1"; + return payload == "1" || payload == "0" ? STEP_OK : STEP_ERROR; + } + bool HasExtension() const override { return _with_extension; } + const FrameSpec& ExtensionFrameSpec() const override { + return _extension_spec; + } + StepResult BuildExtension( + bool enabled, std::string* payload) override { + *payload = enabled + ? std::string("E1", _extension_size) + : std::string("E0", _extension_size); + return STEP_OK; + } + StepResult ParseExtension(const std::string& payload) override { + return payload == std::string("E1", _extension_size) + ? STEP_OK : STEP_FALLBACK; + } + void EnableExtension(size_t size) { + _with_extension = true; + _extension_size = size; + _extension_spec = FixedSpec(NULL, 0, size); + } + +private: + void Append(const char* value) { + if (_calls != NULL) { + *_calls += value; + } + } + + std::string* _calls; + bool _with_extension; + size_t _extension_size; + FrameSpec _extension_spec; +}; + +class TestHandshakeTransport : public HandshakeTransport { +public: + explicit TestHandshakeTransport(std::string* calls = NULL) + : prepare_result(STEP_OK), negotiate_result(STEP_OK), + validate_result(STEP_OK), tcp_active(false), + high_speed_active(false), failed(false), _calls(calls) {} + + StepResult PrepareResources() override { + Append("prepare "); + return prepare_result; + } + StepResult NegotiateResources() override { + Append("negotiate "); + return negotiate_result; + } + void OnEstablished() override { + high_speed_active = true; + Append("activate"); + } + void OnFallback() override { tcp_active = true; } + void OnFailed() override { failed = true; } + StepResult ValidateEstablished() override { + Append("validate "); + return validate_result; + } + + StepResult prepare_result; + StepResult negotiate_result; + StepResult validate_result; + bool tcp_active; + bool high_speed_active; + bool failed; + +private: + void Append(const char* value) { + if (_calls != NULL) { + *_calls += value; + } + } + std::string* _calls; +}; + +static std::string MakeUBShmHello() { + std::string frame(64, '\0'); + memcpy(&frame[0], "UB", 2); + const uint16_t msg_len = butil::HostToNet16(64); + const uint16_t hello_ver = butil::HostToNet16(2); + const uint16_t impl_ver = butil::HostToNet16(1); + memcpy(&frame[2], &msg_len, sizeof(msg_len)); + memcpy(&frame[4], &hello_ver, sizeof(hello_ver)); + memcpy(&frame[6], &impl_ver, sizeof(impl_ver)); + return frame; +} + +TEST(TransportHandshakeTest, blocked_read_stops_after_socket_failure) { + int fds[2]; + ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds)); + butil::fd_guard peer_fd(fds[1]); + + SocketOptions options; + options.fd = fds[0]; + SocketId id; + ASSERT_EQ(0, Socket::Create(options, &id)); + + SocketUniquePtr socket; + ASSERT_EQ(0, Socket::Address(id, &socket)); + + SocketHandshakeIO io(socket.get()); + BlockingHandshakeRead read = {&io, 0, 0}; + bthread_t tid; + bthread_attr_t attr = BTHREAD_ATTR_NORMAL; + ASSERT_EQ(0, bthread_start_background( + &tid, &attr, RunBlockingHandshakeRead, &read)); + + bthread_usleep(10000); + ASSERT_EQ(0, socket->SetFailed( + EFAILEDSOCKET, "cancel blocked handshake read")); + ASSERT_EQ(0, bthread_join(tid, NULL)); + + EXPECT_EQ(-1, read.result); + EXPECT_EQ(EFAILEDSOCKET, read.error); +} + +TEST(HandshakeFrameTest, supports_two_byte_fixed_magic) { + const FrameSpec spec = FixedSpec("UB", 2, 6); + std::string frame; + ASSERT_EQ(FRAME_OK, FrameCodec::Encode(spec, "data", &frame)); + ASSERT_EQ("UBdata", frame); +} + +TEST(HandshakeFrameTest, encodes_u16_total_length) { + const FrameSpec spec( + "RDMA", 4, 8, 64, FrameSpec::U16_TOTAL_LENGTH); + std::string frame; + ASSERT_EQ(FRAME_OK, FrameCodec::Encode(spec, "xy", &frame)); + ASSERT_EQ(8UL, frame.size()); + uint16_t total_be = 0; + memcpy(&total_be, frame.data() + 4, sizeof(total_be)); + ASSERT_EQ(8, butil::NetToHost16(total_be)); + ASSERT_EQ("xy", frame.substr(6)); +} + +TEST(HandshakeFrameTest, encodes_u32_body_length) { + const FrameSpec spec( + "RDM3", 4, 9, 32, FrameSpec::U32_BODY_LENGTH); + std::string frame; + ASSERT_EQ(FRAME_OK, FrameCodec::Encode(spec, "abc", &frame)); + uint32_t body_be = 0; + memcpy(&body_be, frame.data() + 4, sizeof(body_be)); + ASSERT_EQ(3U, butil::NetToHost32(body_be)); + ASSERT_EQ("abc", frame.substr(8)); +} + +TEST(HandshakeFrameTest, buffered_partial_frame_is_not_consumed) { + const FrameSpec spec( + "RDMA", 4, 8, 64, FrameSpec::U16_TOTAL_LENGTH); + std::string frame; + ASSERT_EQ(FRAME_OK, FrameCodec::Encode(spec, "payload", &frame)); + butil::IOBuf source; + source.append(frame.data(), frame.size() - 1); + IOBufHandshakeInput input(&source); + std::string payload; + ASSERT_EQ(FRAME_NEED_MORE, + FrameCodec::ParseBufferedFrame(&input, spec, &payload)); + ASSERT_EQ(frame.size() - 1, source.size()); + + source.append(frame.data() + frame.size() - 1, 1); + ASSERT_EQ(FRAME_OK, + FrameCodec::ParseBufferedFrame(&input, spec, &payload)); + ASSERT_EQ("payload", payload); + ASSERT_TRUE(source.empty()); +} + +TEST(HandshakeFrameTest, buffered_magic_mismatch_is_not_consumed) { + const FrameSpec spec = FixedSpec("UB", 2, 6); + butil::IOBuf source; + source.append("XXdata", 6); + IOBufHandshakeInput input(&source); + std::string payload; + ASSERT_EQ(FRAME_NOT_MINE, + FrameCodec::ParseBufferedFrame(&input, spec, &payload)); + ASSERT_EQ(6UL, source.size()); +} + +TEST(HandshakeFrameTest, blocking_magic_mismatch_is_pushed_back) { + MemoryHandshakeIO io("XX"); + const FrameSpec spec = FixedSpec("UB", 2, 6); + std::string payload; + ASSERT_EQ(FRAME_NOT_MINE, + FrameCodec::ReadFrame(&io, spec, true, &payload)); + ASSERT_EQ("XX", io.pushed_back()); +} + +TEST(HandshakeFrameTest, rejects_lengths_outside_bounds) { + const FrameSpec spec( + "RDMA", 4, 8, 16, FrameSpec::U16_TOTAL_LENGTH); + std::string header("RDMA", 4); + const uint16_t total_be = butil::HostToNet16(17); + header.append(reinterpret_cast(&total_be), sizeof(total_be)); + butil::IOBuf source; + source.append(header); + IOBufHandshakeInput input(&source); + std::string payload; + ASSERT_EQ(FRAME_PROTOCOL_ERROR, + FrameCodec::ParseBufferedFrame(&input, spec, &payload)); + 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; + session.SetPhase(NEGOTIATING); + session.PublishFallback([&tcp_active]() { tcp_active = 1; }); + ASSERT_EQ(1, tcp_active); + ASSERT_EQ(FALLBACK_TCP, session.phase()); +} + +TEST(TransportHandshakeTest, client_runs_codec_and_resource_sequence) { + MemoryHandshakeIO io("HSOK"); + HandshakeSession session; + session.SetIOForTest(&io); + std::string calls; + + TestHandshakeProtocol protocol(&calls); + TestHandshakeTransport transport(&calls); + + ASSERT_EQ(STEP_OK, session.RunClient(&protocol, &transport)); + ASSERT_EQ("prepare build parse negotiate ack1 activate", calls); + ASSERT_EQ("HSLO1", io.output()); + ASSERT_EQ(ESTABLISHED, session.phase()); + ASSERT_EQ(7, session.protocol_version()); +} + +TEST(TransportHandshakeTest, client_exchanges_extension_before_ack) { + MemoryHandshakeIO io("HSOKE"); + HandshakeSession session; + session.SetIOForTest(&io); + TestHandshakeProtocol protocol; + protocol.EnableExtension(1); + TestHandshakeTransport transport; + + ASSERT_EQ(STEP_OK, session.RunClient(&protocol, &transport)); + EXPECT_EQ("HSLOE1", io.output()); + EXPECT_EQ(ESTABLISHED, session.phase()); +} + +TEST(TransportHandshakeTest, server_resumes_fragmented_extension_then_ack) { + MemoryHandshakeIO io; + HandshakeSession session; + session.SetIOForTest(&io); + butil::IOBuf source; + source.append("HSOK", 4); + IOBufHandshakeInput input(&source); + TestHandshakeProtocol protocol; + protocol.EnableExtension(2); + std::vector protocols(1, &protocol); + TestHandshakeTransport transport; + + ASSERT_EQ(STEP_NEED_MORE, + session.RunServer(protocols, &input, &transport, false)); + EXPECT_EQ(EXTENSION_WAIT, session.phase()); + EXPECT_EQ("HSLO", io.output()); + source.append("E", 1); + ASSERT_EQ(STEP_NEED_MORE, + session.RunServer(protocols, &input, &transport, false)); + EXPECT_EQ(EXTENSION_WAIT, session.phase()); + source.append("11", 2); // Remainder of extension, then ACK. + ASSERT_EQ(STEP_OK, + session.RunServer(protocols, &input, &transport, false)); + EXPECT_EQ("HSLOE1", io.output()); + EXPECT_TRUE(source.empty()); + EXPECT_EQ(ESTABLISHED, session.phase()); +} + +TEST(TransportHandshakeTest, client_resource_failure_falls_back_before_io) { + MemoryHandshakeIO io; + HandshakeSession session; + session.SetIOForTest(&io); + TestHandshakeProtocol protocol; + TestHandshakeTransport transport; + transport.prepare_result = STEP_FALLBACK; + transport.negotiate_result = STEP_ERROR; + + ASSERT_EQ(STEP_FALLBACK, session.RunClient(&protocol, &transport)); + ASSERT_TRUE(transport.tcp_active); + ASSERT_TRUE(io.output().empty()); + ASSERT_EQ(FALLBACK_TCP, session.phase()); +} + +TEST(TransportHandshakeTest, server_resumes_at_buffered_ack) { + MemoryHandshakeIO io; + HandshakeSession session; + session.SetIOForTest(&io); + butil::IOBuf source; + source.append("HSOK", 4); + IOBufHandshakeInput input(&source); + std::string calls; + + TestHandshakeProtocol protocol(&calls); + std::vector protocols(1, &protocol); + TestHandshakeTransport transport(&calls); + + ASSERT_EQ(STEP_NEED_MORE, + session.RunServer(protocols, &input, &transport, false)); + ASSERT_EQ(ACK_WAIT, session.phase()); + ASSERT_EQ("HSLO", io.output()); + ASSERT_TRUE(source.empty()); + + source.append("1", 1); + ASSERT_EQ(STEP_OK, + session.RunServer(protocols, &input, &transport, false)); + ASSERT_EQ("parse prepare negotiate build parse_ack validate activate", + calls); + ASSERT_EQ(ESTABLISHED, session.phase()); +} + +TEST(TransportHandshakeTest, server_resource_failure_falls_back_after_ack) { + MemoryHandshakeIO io; + HandshakeSession session; + session.SetIOForTest(&io); + butil::IOBuf source; + source.append("HSOK0", 5); + IOBufHandshakeInput input(&source); + TestHandshakeProtocol protocol; + std::vector protocols(1, &protocol); + TestHandshakeTransport transport; + transport.prepare_result = STEP_FALLBACK; + transport.negotiate_result = STEP_ERROR; + + ASSERT_EQ(STEP_FALLBACK, + session.RunServer(protocols, &input, &transport, false)); + ASSERT_EQ("HSNO", io.output()); + ASSERT_TRUE(source.empty()); + ASSERT_TRUE(transport.tcp_active); + ASSERT_EQ(FALLBACK_TCP, session.phase()); +} + +TEST(TransportHandshakeTest, server_falls_back_without_consuming_other_magic) { + MemoryHandshakeIO io; + HandshakeSession session; + session.SetIOForTest(&io); + butil::IOBuf source; + source.append("XXok", 4); + IOBufHandshakeInput input(&source); + TestHandshakeProtocol protocol; + std::vector protocols(1, &protocol); + TestHandshakeTransport transport; + transport.prepare_result = STEP_ERROR; + transport.negotiate_result = STEP_ERROR; + + ASSERT_EQ(STEP_FALLBACK, + session.RunServer(protocols, &input, &transport, true)); + ASSERT_TRUE(transport.tcp_active); + ASSERT_EQ(4UL, source.size()); + ASSERT_EQ(FALLBACK_TCP, session.phase()); + + ASSERT_EQ(STEP_NOT_MINE, + session.RunServer(protocols, &input, &transport, true)); + ASSERT_EQ(4UL, source.size()); + ASSERT_EQ(FALLBACK_TCP, session.phase()); +} + +TEST(TransportHandshakeTest, server_enters_hello_phase_after_magic_matches) { + MemoryHandshakeIO io; + HandshakeSession session; + session.SetIOForTest(&io); + butil::IOBuf source; + IOBufHandshakeInput input(&source); + + TestHandshakeProtocol protocol; + std::vector protocols(1, &protocol); + TestHandshakeTransport transport; + + source.append("H", 1); + ASSERT_EQ(STEP_NEED_MORE, + session.RunServer(protocols, &input, &transport, false)); + ASSERT_EQ(UNINITIALIZED, session.phase()); + + source.append("S", 1); + ASSERT_EQ(STEP_NEED_MORE, + session.RunServer(protocols, &input, &transport, false)); + ASSERT_EQ(HELLO_WAIT, session.phase()); + ASSERT_EQ(7, session.protocol_version()); + ASSERT_EQ(2UL, source.size()); +} + +TEST(TransportHandshakeTest, + impossible_magic_prefixes_try_other_protocols_without_consuming) { + int fds[2]; + ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds)); + butil::fd_guard peer_fd(fds[1]); + SocketOptions options; + options.fd = fds[0]; + SocketId id; + ASSERT_EQ(0, Socket::Create(options, &id)); + SocketUniquePtr socket; + ASSERT_EQ(0, Socket::Address(id, &socket)); + + const char* const invalid_prefixes[] = { + "X", "UX", "RX", "RDN", "RDMB", + }; + for (size_t i = 0; i < arraysize(invalid_prefixes); ++i) { + butil::IOBuf source; + source.append(invalid_prefixes[i]); + const size_t original_size = source.size(); + const ParseResult result = policy::ParseTransportHandshake( + &source, socket.get(), false, NULL); + ASSERT_FALSE(result.is_ok()); + EXPECT_EQ(PARSE_ERROR_TRY_OTHERS, result.error()) + << invalid_prefixes[i]; + EXPECT_EQ(original_size, source.size()) + << invalid_prefixes[i]; + EXPECT_EQ(UNINITIALIZED, + AdapterTransport::Get(socket.get())->handshake_phase()); + EXPECT_EQ(nullptr, socket->parsing_context()); + } + socket->SetFailed(); +} + +TEST(TransportHandshakeTest, + partial_rdma_magic_waits_for_more_data_without_consuming) { + int fds[2]; + ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds)); + butil::fd_guard peer_fd(fds[1]); + SocketOptions options; + options.fd = fds[0]; + SocketId id; + ASSERT_EQ(0, Socket::Create(options, &id)); + SocketUniquePtr socket; + ASSERT_EQ(0, Socket::Address(id, &socket)); + + const char* const partial_prefixes[] = { + "R", "RD", "RDM", + }; + for (size_t i = 0; i < arraysize(partial_prefixes); ++i) { + butil::IOBuf source; + source.append(partial_prefixes[i]); + const size_t original_size = source.size(); + const ParseResult result = policy::ParseTransportHandshake( + &source, socket.get(), false, NULL); + ASSERT_FALSE(result.is_ok()); + EXPECT_EQ(PARSE_ERROR_NOT_ENOUGH_DATA, result.error()) + << partial_prefixes[i]; + EXPECT_EQ(original_size, source.size()) + << partial_prefixes[i]; + EXPECT_EQ(UNINITIALIZED, + AdapterTransport::Get(socket.get())->handshake_phase()); + EXPECT_EQ(nullptr, socket->parsing_context()); + } + socket->SetFailed(); +} + +TEST(TransportHandshakeTest, + plain_tcp_server_incrementally_rejects_ubshm_upgrade) { + int fds[2]; + ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds)); + butil::fd_guard peer_fd(fds[1]); + SocketOptions options; + options.fd = fds[0]; + SocketId id; + ASSERT_EQ(0, Socket::Create(options, &id)); + SocketUniquePtr socket; + ASSERT_EQ(0, Socket::Address(id, &socket)); + + const std::string hello = MakeUBShmHello(); + butil::IOBuf source; + source.append(hello.data(), 1); + ParseResult result = policy::ParseTransportHandshake( + &source, socket.get(), false, NULL); + ASSERT_FALSE(result.is_ok()); + ASSERT_EQ(PARSE_ERROR_NOT_ENOUGH_DATA, result.error()); + ASSERT_EQ(1UL, source.size()); + ASSERT_EQ(UNINITIALIZED, + AdapterTransport::Get(socket.get())->handshake_phase()); + + source.append(hello.data() + 1, hello.size() - 1); + result = policy::ParseTransportHandshake( + &source, socket.get(), false, NULL); + ASSERT_FALSE(result.is_ok()); + ASSERT_EQ(PARSE_ERROR_NOT_ENOUGH_DATA, result.error()); + ASSERT_TRUE(source.empty()); + ASSERT_NE(nullptr, socket->parsing_context()); + + char reply[64]; + ASSERT_EQ(sizeof(reply), read(peer_fd, reply, sizeof(reply))); + EXPECT_EQ(0, memcmp(reply, "UB", 2)); + uint16_t reply_len = 0; + memcpy(&reply_len, reply + 2, sizeof(reply_len)); + EXPECT_EQ(64, butil::NetToHost16(reply_len)); + EXPECT_EQ(0, reply[4]); + EXPECT_EQ(0, reply[5]); + EXPECT_EQ(0, reply[6]); + EXPECT_EQ(0, reply[7]); + + const uint32_t ack = 0; + source.append(&ack, sizeof(ack)); + result = policy::ParseTransportHandshake( + &source, socket.get(), false, NULL); + ASSERT_FALSE(result.is_ok()); + ASSERT_EQ(PARSE_ERROR_TRY_OTHERS, result.error()); + ASSERT_TRUE(source.empty()); + ASSERT_EQ(FALLBACK_TCP, + AdapterTransport::Get(socket.get())->handshake_phase()); + ASSERT_EQ(nullptr, socket->parsing_context()); + socket->SetFailed(); +} + +TEST(TransportHandshakeTest, + plain_tcp_server_preserves_data_coalesced_behind_ubshm_ack) { + int fds[2]; + ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds)); + butil::fd_guard peer_fd(fds[1]); + SocketOptions options; + options.fd = fds[0]; + SocketId id; + ASSERT_EQ(0, Socket::Create(options, &id)); + SocketUniquePtr socket; + ASSERT_EQ(0, Socket::Address(id, &socket)); + + butil::IOBuf source; + source.append(MakeUBShmHello()); + const uint32_t ack = 0; + source.append(&ack, sizeof(ack)); + const char application_data[] = "application-data"; + source.append(application_data, sizeof(application_data) - 1); + + const ParseResult result = policy::ParseTransportHandshake( + &source, socket.get(), false, NULL); + ASSERT_FALSE(result.is_ok()); + ASSERT_EQ(PARSE_ERROR_TRY_OTHERS, result.error()); + ASSERT_EQ(sizeof(application_data) - 1, source.size()); + + char remaining[sizeof(application_data) - 1] = {}; + ASSERT_EQ(sizeof(remaining), + source.copy_to(remaining, sizeof(remaining))); + EXPECT_EQ(0, memcmp(application_data, remaining, sizeof(remaining))); + ASSERT_EQ(FALLBACK_TCP, + AdapterTransport::Get(socket.get())->handshake_phase()); + ASSERT_EQ(nullptr, socket->parsing_context()); + 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 295d3f1289..02ad6ecd6a 100644 --- a/test/brpc_ubring_unittest.cpp +++ b/test/brpc_ubring_unittest.cpp @@ -24,6 +24,7 @@ #if BRPC_WITH_UBRING #include "brpc/ubshm/common/common.h" +#include "brpc/handshake/ubshm_handshake.h" #include "brpc/ubshm/ub_endpoint.h" #include "brpc/ubshm/shm/shm_def.h" #include "brpc/ubshm/shm/shm_mgr.h" @@ -108,6 +109,26 @@ TEST(HelloFormatExtensionTest, deserialize_unknown_format) { EXPECT_EQ(0x1234, extension.format_id); } +TEST(UBShmHandshakeAdapterTest, rejects_unsupported_format_extension) { + brpc::ubring::UBShmHandshakeAdapter adapter; + std::string payload; + ASSERT_EQ(brpc::handshake::STEP_OK, + adapter.BuildExtension(true, &payload)); + EXPECT_EQ(std::string("\0\4\0\1", 4), payload); + EXPECT_EQ(brpc::handshake::STEP_OK, + adapter.ParseExtension(payload)); + + payload[3] = 2; + EXPECT_EQ(brpc::handshake::STEP_FALLBACK, + adapter.ParseExtension(payload)); + payload[3] = 0; + EXPECT_EQ(brpc::handshake::STEP_FALLBACK, + adapter.ParseExtension(payload)); + payload[1] = 3; + EXPECT_EQ(brpc::handshake::STEP_FALLBACK, + adapter.ParseExtension(payload)); +} + TEST_F(HelloMessageTest, serialize_deserialize_roundtrip) { msg.msg_len = 64; msg.hello_ver = 2; @@ -189,6 +210,84 @@ TEST_F(HelloMessageTest, toString_contains_fields) { EXPECT_NE(std::string::npos, s.find("UBRING_test")); } +TEST(UBShmHandshakeAdapterTest, codec_uses_v3_wire_format) { + brpc::ubring::UBShmHandshakeAdapter adapter; + char shm_name[SHM_MAX_NAME_BUFF_LEN] = {0}; + memcpy(shm_name, "UBRING_test_C", 14); + + std::string payload; + ASSERT_EQ(brpc::handshake::STEP_OK, + adapter.BuildHello(true, 4 * 1024 * 1024, + shm_name, &payload)); + brpc::ubring::HelloMessage decoded{}; + ASSERT_EQ(brpc::handshake::STEP_OK, + adapter.ParseHello(payload, &decoded)); + EXPECT_EQ(64, decoded.msg_len); + EXPECT_EQ(3, decoded.hello_ver); + EXPECT_EQ(1, decoded.impl_ver); + EXPECT_EQ(4 * 1024 * 1024, decoded.len); + EXPECT_EQ(0, memcmp(shm_name, decoded.shm_name, + SHM_MAX_NAME_BUFF_LEN)); + + std::string frame; + ASSERT_EQ(brpc::handshake::FRAME_OK, + brpc::handshake::FrameCodec::Encode( + adapter.HelloFrameSpec(), payload, &frame)); + ASSERT_EQ(64, frame.size()); + 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, rejects_unterminated_remote_name) { + brpc::ubring::HelloMessage message{}; + message.msg_len = 64; + message.hello_ver = 3; + 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; + ASSERT_EQ(brpc::handshake::STEP_OK, + adapter.BuildHello(false, 0, NULL, &payload)); + brpc::ubring::HelloMessage decoded{}; + EXPECT_EQ(brpc::handshake::STEP_FALLBACK, + adapter.ParseHello(payload, &decoded)); + EXPECT_EQ(64, decoded.msg_len); +} + TEST(UBRingConfigurationTest, time_flags_include_units_and_expected_defaults) { struct TimeFlagExpectation { const char* name; @@ -271,16 +370,16 @@ using brpc::ubring::UBShmEndpointTest; TEST_F(UBShmEndpointTest, construct_initial_state) { ASSERT_NE(nullptr, _ep); EXPECT_EQ(brpc::ubring::UBR_DATA_FORMAT_NONE, - _ep->_negotiated_data_format); + _ep->negotiated_data_format()); } TEST_F(UBShmEndpointTest, reset_clears_negotiated_data_format) { - _ep->_negotiated_data_format = brpc::ubring::UBR_DATA_FORMAT_LEGACY_64; + _ep->SetNegotiatedDataFormat(brpc::ubring::UBR_DATA_FORMAT_LEGACY_64); _ep->Reset(); EXPECT_EQ(brpc::ubring::UBR_DATA_FORMAT_NONE, - _ep->_negotiated_data_format); + _ep->negotiated_data_format()); } TEST_F(UBShmEndpointTest, allocate_client_resources_real_shm) {